From 21f3d487aefaa6124573e1d5cbae1679aabcd20f Mon Sep 17 00:00:00 2001 From: Byron Wang Date: Tue, 11 Aug 2026 15:41:17 +0800 Subject: [PATCH] refactor(api): decouple token manager from account model --- api/extensions/ext_application_services.py | 2 +- api/libs/helper.py | 9 +++----- api/services/account_change_email_adapters.py | 15 ++----------- api/services/account_deletion_adapters.py | 15 +++---------- api/services/account_service.py | 22 ++++++++++++++----- api/services/webapp_auth_service.py | 5 ++++- .../unit_tests/libs/test_token_manager.py | 11 ++++------ .../test_account_change_email_adapters.py | 4 +--- .../test_account_deletion_adapters.py | 5 ++--- 9 files changed, 37 insertions(+), 51 deletions(-) diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index 4b9fc479db2..737a4017a18 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -118,7 +118,7 @@ def build_application_services( verification_lockout_duration=dify_config.CHANGE_EMAIL_LOCKOUT_DURATION, ), email_policy=BillingAccountEmailPolicyGateway( - billing_enabled=dify_config.BILLING_ENABLED, + billing_enabled=deployment_edition == DeploymentEdition.CLOUD, ), ), deletion=AccountDeletionService( diff --git a/api/libs/helper.py b/api/libs/helper.py index 7610b834d31..98a593cbe34 100644 --- a/api/libs/helper.py +++ b/api/libs/helper.py @@ -500,16 +500,13 @@ class TokenManager: def generate_token( cls, token_type: str, - account: "Account | None" = None, + account_id: str | None = None, email: str | None = None, additional_data: dict[str, Any] | None = None, ) -> str: - if account is None and email is None: + if account_id is None and email is None: raise ValueError("Account or email must be provided") - account_id = account.id if account else None - account_email = email if email is not None else account.email if account else None - if account_id: old_token = cls._get_current_token_for_account(account_id, token_type) if old_token: @@ -518,7 +515,7 @@ class TokenManager: cls.revoke_token(old_token, token_type) token = str(uuid.uuid4()) - token_data = {"account_id": account_id, "email": account_email, "token_type": token_type} + token_data = {"account_id": account_id, "email": email, "token_type": token_type} if additional_data: token_data.update(additional_data) diff --git a/api/services/account_change_email_adapters.py b/api/services/account_change_email_adapters.py index 942a69de171..3ccde9a5200 100644 --- a/api/services/account_change_email_adapters.py +++ b/api/services/account_change_email_adapters.py @@ -1,8 +1,7 @@ """Infrastructure adapters for the account change-email application service.""" import secrets -from dataclasses import dataclass -from typing import TYPE_CHECKING, cast, override +from typing import override from pydantic import TypeAdapter, ValidationError from redis import RedisError @@ -35,18 +34,9 @@ from services.entities.auth_entities import ( ) from tasks.mail_change_mail_task import send_change_mail_completed_notification_task, send_change_mail_task -if TYPE_CHECKING: - from models.account import Account - _token_adapter: TypeAdapter[ChangeEmailTokenData] = TypeAdapter(ChangeEmailTokenData) -@dataclass(frozen=True, slots=True) -class _TokenAccount: - id: str - email: str - - class TokenManagerChangeEmailTokenGateway(ChangeEmailTokenGateway): @override def get(self, token: str) -> AccountChangeEmailTokenData | None: @@ -75,9 +65,8 @@ class TokenManagerChangeEmailTokenGateway(ChangeEmailTokenGateway): @override def issue(self, token_data: AccountChangeEmailTokenData) -> str: - account = cast("Account", _TokenAccount(id=token_data.account_id, email=token_data.email)) return TokenManager.generate_token( - account=account, + account_id=token_data.account_id, email=token_data.email, token_type="change_email", additional_data={ diff --git a/api/services/account_deletion_adapters.py b/api/services/account_deletion_adapters.py index cf22354ab9d..bcf705c5924 100644 --- a/api/services/account_deletion_adapters.py +++ b/api/services/account_deletion_adapters.py @@ -2,8 +2,7 @@ import secrets from collections.abc import Sequence -from dataclasses import dataclass -from typing import TYPE_CHECKING, cast, override +from typing import override from libs.helper import RateLimiter, TokenManager from services.account_errors import AccountDeletionRateLimitError @@ -18,22 +17,14 @@ from services.entities.account_entities import AccountDeletionChallenge from tasks.delete_account_task import delete_account_task from tasks.mail_account_deletion_task import send_account_deletion_verification_code -if TYPE_CHECKING: - from models.account import Account - - -@dataclass(frozen=True, slots=True) -class _TokenAccount: - id: str - email: str - class TokenManagerAccountDeletionVerificationGateway(AccountDeletionVerificationGateway): @override def create(self, *, account_id: str, email: str) -> AccountDeletionChallenge: code = "".join(str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6)) token = TokenManager.generate_token( - account=cast("Account", _TokenAccount(id=account_id, email=email)), + account_id=account_id, + email=email, token_type="account_deletion", additional_data={"code": code}, ) diff --git a/api/services/account_service.py b/api/services/account_service.py index 3f460275e20..68e7f2b1dc0 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -541,7 +541,10 @@ class AccountService: def generate_account_deletion_verification_code(account: Account) -> tuple[str, str]: code = "".join([str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6)]) token = TokenManager.generate_token( - account=account, token_type="account_deletion", additional_data={"code": code} + account_id=account.id, + email=account.email, + token_type="account_deletion", + additional_data={"code": code}, ) return token, code @@ -917,7 +920,10 @@ class AccountService: code = "".join([str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6)]) additional_data["code"] = code token = TokenManager.generate_token( - account=account, email=email, token_type="reset_password", additional_data=additional_data + account_id=account.id if account else None, + email=email, + token_type="reset_password", + additional_data=additional_data, ) return code, token @@ -941,7 +947,7 @@ class AccountService: account: Account, ) -> str: token = TokenManager.generate_token( - account=account, + account_id=account.id, email=token_data.email, token_type="change_email", additional_data=token_data.to_token_manager_payload(), @@ -960,7 +966,10 @@ class AccountService: code = "".join([str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6)]) additional_data["code"] = code token = TokenManager.generate_token( - account=account, email=email, token_type="owner_transfer", additional_data=additional_data + account_id=account.id if account else None, + email=email, + token_type="owner_transfer", + additional_data=additional_data, ) return code, token @@ -1020,7 +1029,10 @@ class AccountService: code = "".join([str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6)]) token = TokenManager.generate_token( - account=account, email=email, token_type="email_code_login", additional_data={"code": code} + account_id=account.id if account else None, + email=email, + token_type="email_code_login", + additional_data={"code": code}, ) send_email_code_login_mail_task.delay( language=language, diff --git a/api/services/webapp_auth_service.py b/api/services/webapp_auth_service.py index 3373a5ddf66..d05f04f7dd4 100644 --- a/api/services/webapp_auth_service.py +++ b/api/services/webapp_auth_service.py @@ -74,7 +74,10 @@ class WebAppAuthService: code = "".join([str(secrets.randbelow(exclusive_upper_bound=10)) for _ in range(6)]) token = TokenManager.generate_token( - account=account, email=email, token_type="email_code_login", additional_data={"code": code} + account_id=account.id if account else None, + email=email, + token_type="email_code_login", + additional_data={"code": code}, ) send_email_code_login_mail_task.delay( language=language, diff --git a/api/tests/unit_tests/libs/test_token_manager.py b/api/tests/unit_tests/libs/test_token_manager.py index bbe8a7e30bc..35d28efbbda 100644 --- a/api/tests/unit_tests/libs/test_token_manager.py +++ b/api/tests/unit_tests/libs/test_token_manager.py @@ -61,19 +61,16 @@ def test_token_manager_roundtrip_preserves_untyped_metadata_keys(monkeypatch: py assert data.get("custom_marker") == "preserve-me" -def test_token_manager_roundtrip_uses_explicit_email_with_account(monkeypatch: pytest.MonkeyPatch) -> None: - """When both `account` and `email` are supplied, the token should bind the - stable `account_id` from the account and the target email from the explicit - email argument. +def test_token_manager_roundtrip_uses_explicit_account_id(monkeypatch: pytest.MonkeyPatch) -> None: + """When both `account_id` and `email` are supplied, the token should bind + both primitive values without depending on an account model. """ storage: dict[str, str] = {} monkeypatch.setattr(helper_module, "redis_client", _build_fake_redis(storage)) - account = SimpleNamespace(id="acc-1", email="old@example.com") - token = TokenManager.generate_token( - account=account, + account_id="acc-1", email="new@example.com", token_type="change_email", additional_data={ diff --git a/api/tests/unit_tests/services/test_account_change_email_adapters.py b/api/tests/unit_tests/services/test_account_change_email_adapters.py index 4c8970bb38c..f0b2d65c8a9 100644 --- a/api/tests/unit_tests/services/test_account_change_email_adapters.py +++ b/api/tests/unit_tests/services/test_account_change_email_adapters.py @@ -39,9 +39,7 @@ def test_token_gateway_issues_account_bound_state() -> None: ) as generate_token: assert gateway.issue(token_data) == "token" - token_account = generate_token.call_args.kwargs["account"] - assert token_account.id == "account-1" - assert token_account.email == "new@example.com" + assert generate_token.call_args.kwargs["account_id"] == "account-1" assert generate_token.call_args.kwargs["email"] == "new@example.com" assert generate_token.call_args.kwargs["additional_data"] == { "old_email": "old@example.com", diff --git a/api/tests/unit_tests/services/test_account_deletion_adapters.py b/api/tests/unit_tests/services/test_account_deletion_adapters.py index 6ced40b1dd2..880b5e8a094 100644 --- a/api/tests/unit_tests/services/test_account_deletion_adapters.py +++ b/api/tests/unit_tests/services/test_account_deletion_adapters.py @@ -35,9 +35,8 @@ def test_verification_gateway_creates_six_digit_account_bound_challenge() -> Non assert challenge.token == "token" assert challenge.code == "123456" - token_account = generate_token.call_args.kwargs["account"] - assert token_account.id == "account-1" - assert token_account.email == "account@example.com" + assert generate_token.call_args.kwargs["account_id"] == "account-1" + assert generate_token.call_args.kwargs["email"] == "account@example.com" assert generate_token.call_args.kwargs["additional_data"] == {"code": "123456"}