Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions api/controllers/console/auth/error.py
Original file line number Diff line number Diff line change
Expand Up @@ -155,3 +155,21 @@ class MemberNotInTenantError(BaseHTTPException):
error_code = "member_not_in_tenant"
description = "The member is not in the workspace."
code = 400


class MFARequiredError(BaseHTTPException):
error_code = "mfa_required"
description = "Multi-factor authentication is required."
code = 401


class MFATokenRequiredError(BaseHTTPException):
error_code = "mfa_token_invalid"
description = "The MFA token is invalid or expired."
code = 401


class MFASetupRequiredError(BaseHTTPException):
error_code = "mfa_setup_required"
description = "MFA setup is required to complete this action."
code = 400
18 changes: 18 additions & 0 deletions api/controllers/console/auth/login.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from services.errors.account import AccountRegisterError
from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError
from services.feature_service import FeatureService
from services.mfa_service import MFAService


class LoginApi(Resource):
Expand All @@ -46,6 +47,8 @@ def post(self):
parser.add_argument("password", type=str, required=True, location="json")
parser.add_argument("remember_me", type=bool, required=False, default=False, location="json")
parser.add_argument("invite_token", type=str, required=False, default=None, location="json")
parser.add_argument("mfa_code", type=str, required=False, default=None, location="json")
parser.add_argument("is_backup_code", type=bool, required=False, default=False, location="json")
args = parser.parse_args()

if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(args["email"]):
Expand All @@ -72,7 +75,22 @@ def post(self):
raise AccountBannedError()
except services.errors.account.AccountPasswordError:
AccountService.add_login_error_rate_limit(args["email"])
# Perform dummy MFA check to prevent timing attacks
# This ensures similar response time regardless of authentication failure
if args.get("mfa_code"):
import time

time.sleep(0.05) # Simulate MFA verification delay
raise AuthenticationFailedError()

# Check MFA requirement
if MFAService.is_mfa_required(account):
if not args["mfa_code"]:
return {"result": "fail", "code": "mfa_required"}

if not MFAService.authenticate_with_mfa(account, args["mfa_code"]):
return {"result": "fail", "code": "mfa_token_invalid", "data": "The MFA token is invalid or expired."}

# SELF_HOSTED only have one workspace
tenants = TenantService.get_join_tenants(account)
if len(tenants) == 0:
Expand Down
140 changes: 140 additions & 0 deletions api/controllers/console/auth/mfa.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
import logging
from typing import cast

import flask_login
from flask_restx import Resource, reqparse

from controllers.console.wraps import account_initialization_required
from libs.login import login_required
from models.account import Account
from services.mfa_service import MFAService


class MFASetupInitApi(Resource):
@login_required
@account_initialization_required
def get(self):
"""Initialize MFA setup - generate secret and QR code (GET method for compatibility)."""
# Call post method directly with proper context
account = cast(Account, flask_login.current_user)

try:
mfa_status = MFAService.get_mfa_status(account)
if mfa_status["enabled"]:
return {"error": "MFA is already enabled"}, 400

setup_data = MFAService.generate_mfa_setup_data(account)
return {"secret": setup_data["secret"], "qr_code": setup_data["qr_code"]}
except Exception:
# Log the actual error for debugging
logging.exception("MFA setup error")
# Return generic error message to avoid exposing internal details
return {"error": "Failed to setup MFA. Please try again."}, 500

@login_required
@account_initialization_required
def post(self): # type: ignore
"""Initialize MFA setup - generate secret and QR code."""
account = cast(Account, flask_login.current_user)

try:
mfa_status = MFAService.get_mfa_status(account)
if mfa_status["enabled"]:
return {"error": "MFA is already enabled"}, 400

setup_data = MFAService.generate_mfa_setup_data(account)
return {"secret": setup_data["secret"], "qr_code": setup_data["qr_code"]}
except Exception:
# Log the actual error for debugging
logging.exception("MFA setup error")
# Return generic error message to avoid exposing internal details
return {"error": "Failed to setup MFA. Please try again."}, 500


class MFASetupCompleteApi(Resource):
@login_required
@account_initialization_required
def post(self):
"""Complete MFA setup with TOTP verification."""
parser = reqparse.RequestParser()
parser.add_argument("totp_token", type=str, required=True, help="TOTP token is required")
args = parser.parse_args()

account = cast(Account, flask_login.current_user)

try:
result = MFAService.setup_mfa(account, args["totp_token"])
return {
"message": "MFA setup completed successfully",
"backup_codes": result["backup_codes"],
"setup_at": result["setup_at"].isoformat(),
}
except ValueError as e:
return {"error": str(e)}, 400
except Exception as e:
return {"error": str(e)}, 500


class MFADisableApi(Resource):
@login_required
@account_initialization_required
def post(self):
"""Disable MFA with password verification."""
parser = reqparse.RequestParser()
parser.add_argument("password", type=str, required=True, help="Password is required")
args = parser.parse_args()

account = cast(Account, flask_login.current_user)

try:
mfa_status = MFAService.get_mfa_status(account)
if not mfa_status["enabled"]:
return {"error": "MFA is not enabled"}, 400

if MFAService.disable_mfa(account, args["password"]):
return {"message": "MFA disabled successfully"}
else:
return {"error": "Invalid password"}, 400
except Exception as e:
return {"error": str(e)}, 500


class MFAStatusApi(Resource):
@login_required
@account_initialization_required
def get(self):
"""Get current MFA status."""
account = cast(Account, flask_login.current_user)

try:
status = MFAService.get_mfa_status(account)
return status
except Exception as e:
return {"error": str(e)}, 500


class MFAVerifyApi(Resource):
def post(self):
"""Verify MFA token during login (public endpoint)."""
parser = reqparse.RequestParser()
parser.add_argument("email", type=str, required=True, help="Email is required")
parser.add_argument("mfa_token", type=str, required=True, help="MFA token is required")
args = parser.parse_args()

from models.engine import db

account = db.session.query(Account).filter_by(email=args["email"]).first()

if not account:
return {"error": "Account not found"}, 404

if not MFAService.is_mfa_required(account):
return {"error": "MFA not required for this account"}, 400

try:
if MFAService.authenticate_with_mfa(account, args["mfa_token"]):
return {"message": "MFA verification successful"}
else:
return {"error": "Invalid MFA token"}, 400
except Exception as e:
return {"error": str(e)}, 500
8 changes: 8 additions & 0 deletions api/controllers/console/workspace/account.py
Original file line number Diff line number Diff line change
Expand Up @@ -583,3 +583,11 @@ def post(self):
api.add_resource(CheckEmailUnique, "/account/change-email/check-email-unique")
# api.add_resource(AccountEmailApi, '/account/email')
# api.add_resource(AccountEmailVerifyApi, '/account/email-verify')

# MFA endpoints
from controllers.console.auth.mfa import MFADisableApi, MFASetupCompleteApi, MFASetupInitApi, MFAStatusApi

api.add_resource(MFAStatusApi, "/account/mfa/status")
api.add_resource(MFASetupInitApi, "/account/mfa/setup")
api.add_resource(MFASetupCompleteApi, "/account/mfa/setup/complete")
api.add_resource(MFADisableApi, "/account/mfa/disable")
Original file line number Diff line number Diff line change
@@ -0,0 +1,43 @@
"""add account mfa settings table

Revision ID: c699fb438ab7
Revises: 68519ad5cd18
Create Date: 2025-09-22 13:50:00.000000

"""
from alembic import op
import models as models
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql

# revision identifiers, used by Alembic.
revision = 'c699fb438ab7'
down_revision = '68519ad5cd18'
branch_labels = None
depends_on = None


def upgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.create_table('account_mfa_settings',
sa.Column('id', models.types.StringUUID(), server_default=sa.text('uuid_generate_v4()'), nullable=False),
sa.Column('account_id', models.types.StringUUID(), nullable=False),
sa.Column('enabled', sa.Boolean(), server_default=sa.text('false'), nullable=False),
sa.Column('secret', sa.Text(), nullable=True),
sa.Column('backup_codes', sa.Text(), nullable=True),
sa.Column('setup_at', sa.DateTime(), nullable=True),
sa.Column('created_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False),
sa.Column('updated_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False),
sa.ForeignKeyConstraint(['account_id'], ['accounts.id'], ),
sa.PrimaryKeyConstraint('id', name='account_mfa_settings_pkey'),
sa.UniqueConstraint('account_id', name='unique_account_mfa_settings')
)
op.create_index('account_mfa_settings_account_id_idx', 'account_mfa_settings', ['account_id'], unique=False)
# ### end Alembic commands ###


def downgrade():
# ### commands auto generated by Alembic - please adjust! ###
op.drop_index('account_mfa_settings_account_id_idx', table_name='account_mfa_settings')
op.drop_table('account_mfa_settings')
# ### end Alembic commands ###
2 changes: 2 additions & 0 deletions api/models/__init__.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from .account import (
Account,
AccountIntegrate,
AccountMFASettings,
AccountStatus,
InvitationCode,
Tenant,
Expand Down Expand Up @@ -97,6 +98,7 @@
"APIBasedExtensionPoint",
"Account",
"AccountIntegrate",
"AccountMFASettings",
"AccountStatus",
"ApiRequest",
"ApiToken",
Expand Down
23 changes: 22 additions & 1 deletion api/models/account.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import sqlalchemy as sa
from flask_login import UserMixin # type: ignore[import-untyped]
from sqlalchemy import DateTime, String, func, select
from sqlalchemy.orm import Mapped, Session, mapped_column, reconstructor
from sqlalchemy.orm import Mapped, Session, backref, mapped_column, reconstructor, relationship
from typing_extensions import deprecated

from models.base import Base
Expand Down Expand Up @@ -361,3 +361,24 @@ class UpgradeMode(enum.StrEnum):
include_plugins: Mapped[list[str]] = mapped_column(sa.ARRAY(String(255)), nullable=False) # plugin_id (author/name)
created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, server_default=func.current_timestamp())
updated_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, server_default=func.current_timestamp())


class AccountMFASettings(Base):
__tablename__ = "account_mfa_settings"
__table_args__ = (
sa.PrimaryKeyConstraint("id", name="account_mfa_settings_pkey"),
sa.UniqueConstraint("account_id", name="unique_account_mfa_settings"),
sa.Index("account_mfa_settings_account_id_idx", "account_id"),
)

id: Mapped[str] = mapped_column(StringUUID, server_default=sa.text("uuid_generate_v4()"))
account_id: Mapped[str] = mapped_column(StringUUID, sa.ForeignKey("accounts.id"), nullable=False)
enabled: Mapped[bool] = mapped_column(sa.Boolean, nullable=False, server_default=sa.text("false"))
secret: Mapped[str | None] = mapped_column(sa.Text, nullable=True)
backup_codes: Mapped[str | None] = mapped_column(sa.Text, nullable=True)
setup_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, server_default=func.current_timestamp())
updated_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, server_default=func.current_timestamp())

# Relationship
account = relationship("Account", backref=backref("mfa_settings", uselist=False, cascade="all, delete-orphan"))
2 changes: 2 additions & 0 deletions api/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -68,10 +68,12 @@ dependencies = [
"pydantic-extra-types~=2.10.3",
"pydantic-settings~=2.9.1",
"pyjwt~=2.10.1",
"pyotp~=2.9.0",
"pypdfium2==4.30.0",
"python-docx~=1.1.0",
"python-dotenv==1.0.1",
"pyyaml~=6.0.1",
"qrcode~=7.4.2",
"readabilipy~=0.3.0",
"redis[hiredis]~=6.1.0",
"resend~=2.9.0",
Expand Down
Loading