diff --git a/api/.importlinter b/api/.importlinter index ea5b463ff11..6d693295d6a 100644 --- a/api/.importlinter +++ b/api/.importlinter @@ -264,6 +264,37 @@ forbidden_modules = sqlalchemy werkzeug +[importlinter:contract:end-user-query-service-boundary] +name = App-scoped end user query service is framework and persistence neutral +type = forbidden +source_modules = + services.app_scoped_end_user_query_service + services.entities.app_scoped_end_user_entities +forbidden_modules = + configs + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + +[importlinter:contract:end-user-command-service-boundary] +name = App-scoped end user command service is framework and persistence neutral +type = forbidden +source_modules = + services.app_scoped_end_user_service +forbidden_modules = + configs + controllers + extensions + flask + models.model + repositories + sqlalchemy + werkzeug + [importlinter:contract:app-tracing-config-service-boundary] name = App tracing configuration service is framework and persistence neutral type = forbidden diff --git a/api/controllers/inner_api/plugin/plugin.py b/api/controllers/inner_api/plugin/plugin.py index 221887f73c7..65103cee9bd 100644 --- a/api/controllers/inner_api/plugin/plugin.py +++ b/api/controllers/inner_api/plugin/plugin.py @@ -33,6 +33,7 @@ from core.plugin.entities.request import ( ) from core.tools.entities.tool_entities import ToolProviderType from core.tools.signature import bind_file_uri, get_signed_file_uri_for_plugin +from extensions.ext_application_services import application_services from extensions.ext_database import db from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.helper import length_prefixed_response @@ -351,6 +352,7 @@ class PluginInvokeAppApi(Resource): stream=payload.response_mode == "streaming", inputs=payload.inputs, files=payload.files, + end_users=application_services().app_scoped_end_users.commands, ) return length_prefixed_response(0xF, PluginAppBackwardsInvocation.convert_to_event_stream(response)) diff --git a/api/controllers/inner_api/plugin/wraps.py b/api/controllers/inner_api/plugin/wraps.py index 3b58d82bdb8..9f0fc259430 100644 --- a/api/controllers/inner_api/plugin/wraps.py +++ b/api/controllers/inner_api/plugin/wraps.py @@ -10,8 +10,8 @@ from sqlalchemy.orm import sessionmaker from extensions.ext_database import db from libs.login import current_user from models.account import Tenant -from models.enums import EndUserType -from models.model import DefaultEndUserSessionID, EndUser +from models.enums import DEFAULT_END_USER_SESSION_ID, EndUserType +from models.model import EndUser class TenantUserPayload(BaseModel): @@ -30,8 +30,8 @@ def get_user(tenant_id: str, user_id: str | None) -> EndUser: context. """ if not user_id: - user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID - is_anonymous = user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID + user_id = DEFAULT_END_USER_SESSION_ID + is_anonymous = user_id == DEFAULT_END_USER_SESSION_ID try: with sessionmaker(db.engine, expire_on_commit=False).begin() as session: user_model = None @@ -102,7 +102,7 @@ def get_user_tenant[**P, R](view_func: Callable[P, R]) -> Callable[P, R]: raise ValueError("tenant_id is required") if not user_id: - user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID + user_id = DEFAULT_END_USER_SESSION_ID tenant_model = db.session.get(Tenant, tenant_id) diff --git a/api/controllers/openapi/auth/subjects.py b/api/controllers/openapi/auth/subjects.py index 2bac031e0bb..ba50cbaa16b 100644 --- a/api/controllers/openapi/auth/subjects.py +++ b/api/controllers/openapi/auth/subjects.py @@ -15,12 +15,12 @@ from werkzeug.exceptions import Unauthorized from controllers.openapi.auth.context import Context from controllers.openapi.auth.data import ExternalIdentity from controllers.openapi.auth.loaders import load_app, load_workspace, route_has_app +from extensions.ext_application_services import application_services from libs.oauth_bearer import AuthContext, Scope, SubjectType from models.account import Account from models.enums import CreatorUserRole, EndUserType from models.model import EndUser from services.account_service import AccountService -from services.end_user_service import EndUserService from services.enterprise.enterprise_service import WebAppAccessMode _SUBJECT_CLASSES: dict[SubjectType, type[Subject]] = {} @@ -123,7 +123,7 @@ class ExternalSsoSubject(Subject): identity = self.external_identity if identity is None: raise Unauthorized("missing context for external user resolution") - return EndUserService.get_or_create_end_user_by_type( + return application_services().app_scoped_end_users.commands.get_or_create_end_user_by_type( EndUserType.OPENAPI, tenant_id=str(load_workspace(ctx).id), app_id=str(load_app(ctx).id), diff --git a/api/controllers/service_api/end_user/end_user.py b/api/controllers/service_api/end_user/end_user.py index 607ed12e5b6..723c915a145 100644 --- a/api/controllers/service_api/end_user/end_user.py +++ b/api/controllers/service_api/end_user/end_user.py @@ -6,9 +6,11 @@ from controllers.common.schema import register_response_schema_models from controllers.service_api import service_api_ns from controllers.service_api.end_user.error import EndUserNotFoundError from controllers.service_api.wraps import validate_app_token +from extensions.ext_application_services import application_services from fields.end_user_fields import EndUserDetail +from machinery.context import ServiceApiRequestContext from models.model import App -from services.end_user_service import EndUserService +from services.app_scoped_end_user_query_service import AppScopedEndUserNotFoundError register_response_schema_models(service_api_ns, EndUserDetail) @@ -48,10 +50,13 @@ class EndUserApi(Resource): cross-tenant/app access when an end-user ID is known. """ - end_user = EndUserService.get_end_user_by_id( - tenant_id=app_model.tenant_id, app_id=app_model.id, end_user_id=str(end_user_id) + request_context = ServiceApiRequestContext( + tenant_id=app_model.tenant_id, + app_id=app_model.id, ) - if end_user is None: - raise EndUserNotFoundError() + try: + end_user = application_services().app_scoped_end_users.queries.get_by_id(request_context, str(end_user_id)) + except AppScopedEndUserNotFoundError as error: + raise EndUserNotFoundError() from error return EndUserDetail.model_validate(end_user).model_dump(mode="json") diff --git a/api/controllers/service_api/wraps.py b/api/controllers/service_api/wraps.py index a28e731cd72..1b6192928f0 100644 --- a/api/controllers/service_api/wraps.py +++ b/api/controllers/service_api/wraps.py @@ -32,7 +32,6 @@ from models.dataset import Dataset, RateLimitLog from models.model import ApiToken, App from services import dataset_api_key_service from services.api_token_service import ApiTokenCache, fetch_token_with_single_flight, record_token_usage -from services.end_user_service import EndUserService from services.feature_service import FeatureService logger = logging.getLogger(__name__) @@ -147,7 +146,11 @@ def validate_app_token[**P, R]( if user_id: user_id = str(user_id) - end_user = EndUserService.get_or_create_end_user(app_model, user_id) + end_user = application_services().app_scoped_end_users.commands.get_or_create_end_user( + app_model.tenant_id, + app_model.id, + user_id, + ) kwargs["end_user"] = end_user # Set EndUser as current logged-in user for flask_login.current_user diff --git a/api/controllers/trigger/webhook.py b/api/controllers/trigger/webhook.py index 7715090b967..4b5e4cfc9b8 100644 --- a/api/controllers/trigger/webhook.py +++ b/api/controllers/trigger/webhook.py @@ -8,6 +8,7 @@ from controllers.trigger import bp from core.trigger.debug.event_bus import TriggerDebugEventBus from core.trigger.debug.events import WebhookDebugEvent, build_webhook_pool_key from enums import QuotaType +from extensions.ext_application_services import application_services from services.errors.app import QuotaExceededError from services.trigger.webhook_service import RawWebhookDataDict, WebhookService @@ -70,7 +71,12 @@ def handle_webhook(webhook_id: str): return jsonify({"error": "Bad Request", "message": error}), 400 # Process webhook call (send to Celery) - WebhookService.trigger_workflow_execution(webhook_trigger, webhook_data, workflow) + WebhookService.trigger_workflow_execution( + webhook_trigger, + webhook_data, + workflow, + end_users=application_services().app_scoped_end_users.commands, + ) # Return configured response response_data, status_code = WebhookService.generate_webhook_response(node_config) diff --git a/api/core/plugin/backwards_invocation/app.py b/api/core/plugin/backwards_invocation/app.py index 4165b585d1c..66a1c547eea 100644 --- a/api/core/plugin/backwards_invocation/app.py +++ b/api/core/plugin/backwards_invocation/app.py @@ -1,6 +1,6 @@ import uuid from collections.abc import Generator, Mapping -from typing import Any, cast +from typing import Any, Protocol, cast from sqlalchemy import select from sqlalchemy.orm import Session @@ -26,7 +26,15 @@ from models.model import ( load_annotation_reply_config, ) from models.workflow import Workflow -from services.end_user_service import EndUserService + + +class AppScopedEndUserProvisioner(Protocol): + def get_or_create_end_user( + self, + tenant_id: str, + app_id: str, + user_id: str | None = None, + ) -> EndUser: ... class PluginAppBackwardsInvocation(BaseBackwardsInvocation): @@ -70,19 +78,20 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): inputs: Mapping, files: list[dict], session: Session, + end_users: AppScopedEndUserProvisioner, ) -> Generator[Mapping | str, None, None] | Mapping: """ invoke app """ app = cls._get_app(app_id, tenant_id) if not user_id: - user = EndUserService.get_or_create_end_user(app) + user = end_users.get_or_create_end_user(app.tenant_id, app.id, None) else: try: user = cls._get_user(user_id, app) except ValueError: # Plugins such as WeCom Bot pass external sender IDs rather than EndUser UUIDs. - user = EndUserService.get_or_create_end_user(app, user_id=user_id) + user = end_users.get_or_create_end_user(app.tenant_id, app.id, user_id) conversation_id = conversation_id or "" diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index 8169e98c18f..bc635a9c2f9 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -29,6 +29,7 @@ from libs.helper import RateLimiter from libs.oauth import GitHubOAuth, GoogleOAuth from libs.oauth_bearer import invalidate_oauth_token_cache from libs.passport import PassportService +from models.model import EndUser from repositories.account_activation_repository import SQLAlchemyAccountActivationRepository from repositories.account_integration_repository import SQLAlchemyAccountIntegrationRepository from repositories.account_oauth_repository import ( @@ -40,6 +41,7 @@ from repositories.account_oauth_repository import ( from repositories.account_repository import SQLAlchemyAccountRepository from repositories.app_definition_query_repository import AppDefinitionQueryRepository from repositories.app_preview_query_repository import AppPreviewQueryRepository +from repositories.app_scoped_end_user_repository import AppScopedEndUserRepo from repositories.app_site_command_repository import AppSiteCommandRepository from repositories.app_statistic_query_repository import AppStatisticQueryRepository from repositories.app_tracing_config_repository import SQLAlchemyAppTracingConfigRepository @@ -142,6 +144,8 @@ from services.app_definition_query_service import AppDefinitionQueryService from services.app_preview_details_adapters import AppPreviewDetailsRuntime from services.app_preview_details_service import AppPreviewDetails from services.app_preview_query_service import AppPreviewQueryService +from services.app_scoped_end_user_query_service import AppScopedEndUserQueryService +from services.app_scoped_end_user_service import AppScopedEndUserService from services.app_site_service import AppSiteService from services.app_statistic_query import AppStatisticQuery from services.app_task_service import AppTaskControlService @@ -261,6 +265,12 @@ class AccountServices: profile: AccountProfileService +@dataclass(frozen=True, slots=True) +class AppScopedEndUserServices: + commands: AppScopedEndUserService[EndUser] + queries: AppScopedEndUserQueryService + + @dataclass(frozen=True, slots=True) class ApplicationServices: accounts: AccountServices @@ -275,6 +285,7 @@ class ApplicationServices: compliance_downloads: ComplianceDownloadService data_source_api_key_auth: DataSourceApiKeyAuthService data_source_oauth: Mapping[str, DataSourceOAuthService] + app_scoped_end_users: AppScopedEndUserServices webapp_access: WebAppAccessQueryService web_app_runtime: WebAppRuntimeQueryService explore_banner_queries: ExploreBannerQueryService @@ -462,6 +473,7 @@ def build_application_services( trial_enabled=trial_app_enabled, ) workspace_query_repository = WorkspaceQueryRepository(session_factory=database_client) + app_scoped_end_user_repository = AppScopedEndUserRepo(session_factory=database_client) file_service = FileService(session_factory=database_client) remote_file_service = RemoteFileService(files=file_service) passwords = DefaultAccountPasswordHasher() @@ -666,6 +678,10 @@ def build_application_services( encryptor=TenantApiKeyAuthCredentialEncryptor(), ), data_source_oauth=_build_data_source_oauth_services(database_client=database_client), + app_scoped_end_users=AppScopedEndUserServices( + commands=AppScopedEndUserService(end_users=app_scoped_end_user_repository), + queries=AppScopedEndUserQueryService(end_users=app_scoped_end_user_repository), + ), webapp_access=WebAppAccessQueryService( access=WebAppAccessQueryRepository(session_factory=database_client), webapp_auth_enabled=SystemFeatureService.is_webapp_auth_enabled(deployment_edition=deployment_edition), diff --git a/api/fields/end_user_fields.py b/api/fields/end_user_fields.py index 646004ad103..c58bc2b13d0 100644 --- a/api/fields/end_user_fields.py +++ b/api/fields/end_user_fields.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime from typing import Annotated -from pydantic import Field, WithJsonSchema +from pydantic import WithJsonSchema from fields.base import ResponseModel @@ -18,12 +18,7 @@ class SimpleEndUser(ResponseModel): class EndUserDetail(ResponseModel): - """Full EndUser record for API responses. - - Note: The SQLAlchemy model defines an `is_anonymous` property for Flask-Login semantics - (always False). The database column is exposed as `_is_anonymous`, so this DTO maps - `is_anonymous` from `_is_anonymous` to return the stored value. - """ + """Full end-user detail returned by the Service API.""" id: UUIDString tenant_id: UUIDString @@ -31,7 +26,7 @@ class EndUserDetail(ResponseModel): type: str external_user_id: str | None = None name: str | None = None - is_anonymous: bool = Field(validation_alias="_is_anonymous") + is_anonymous: bool session_id: str created_at: datetime updated_at: datetime diff --git a/api/machinery/context.py b/api/machinery/context.py index 4facf6caf9b..d64b8882190 100644 --- a/api/machinery/context.py +++ b/api/machinery/context.py @@ -1,5 +1,6 @@ """Stable values passed from API admission into application services.""" +from dataclasses import dataclass from typing import NamedTuple @@ -10,6 +11,14 @@ class RequestContext(NamedTuple): active_workspace_id: str +@dataclass(frozen=True, slots=True, kw_only=True) +class ServiceApiRequestContext: + """Stable app scope admitted for a Service API request.""" + + tenant_id: str + app_id: str + + class AccountRequestContext(NamedTuple): """Stable identity for account-scoped use cases that do not require a workspace.""" diff --git a/api/models/enums.py b/api/models/enums.py index 0274823ff93..b38e372ab8a 100644 --- a/api/models/enums.py +++ b/api/models/enums.py @@ -216,6 +216,9 @@ class EndUserType(StrEnum): TRIGGER = "trigger" +DEFAULT_END_USER_SESSION_ID = "DEFAULT-USER" + + class DocumentDocType(StrEnum): """Document doc_type classification""" diff --git a/api/models/model.py b/api/models/model.py index 491c5549fee..3280f46fa41 100644 --- a/api/models/model.py +++ b/api/models/model.py @@ -2068,14 +2068,6 @@ class OperationLog(TypeBase): ) -class DefaultEndUserSessionID(StrEnum): - """ - End User Session ID enum. - """ - - DEFAULT_SESSION_ID = "DEFAULT-USER" - - class EndUser(Base, UserMixin): __tablename__ = "end_users" __table_args__ = ( diff --git a/api/openapi/markdown/service-openapi.md b/api/openapi/markdown/service-openapi.md index f5c26d3719a..9fd290e93dc 100644 --- a/api/openapi/markdown/service-openapi.md +++ b/api/openapi/markdown/service-openapi.md @@ -3277,11 +3277,7 @@ Request payload for bulk downloading documents as a zip archive. #### EndUserDetail -Full EndUser record for API responses. - -Note: The SQLAlchemy model defines an `is_anonymous` property for Flask-Login semantics -(always False). The database column is exposed as `_is_anonymous`, so this DTO maps -`is_anonymous` from `_is_anonymous` to return the stored value. +Full end-user detail returned by the Service API. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | diff --git a/api/repositories/app_scoped_end_user_repository.py b/api/repositories/app_scoped_end_user_repository.py new file mode 100644 index 00000000000..3e5dcb7af1c --- /dev/null +++ b/api/repositories/app_scoped_end_user_repository.py @@ -0,0 +1,173 @@ +"""SQLAlchemy persistence adapter for app-scoped end users.""" + +from collections.abc import Sequence +from typing import override + +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.enums import EndUserType +from models.model import EndUser +from services.app_scoped_end_user_query_service import AppScopedEndUserQuery +from services.app_scoped_end_user_service import AppScopedEndUserRepository +from services.entities.app_scoped_end_user_entities import ( + AppScopedEndUserRecord, + NewAppScopedEndUser, + StoredAppScopedEndUser, +) + + +class AppScopedEndUserRepo(AppScopedEndUserQuery, AppScopedEndUserRepository[EndUser]): + def __init__(self, *, session_factory: sessionmaker[Session]) -> None: + self._session_factory = session_factory + + @override + def find_by_id(self, *, tenant_id: str, app_id: str, end_user_id: str) -> AppScopedEndUserRecord | None: + statement = ( + select( + EndUser.id, + EndUser.tenant_id, + EndUser.app_id, + EndUser.type, + EndUser.external_user_id, + EndUser.name, + EndUser._is_anonymous, + EndUser.session_id, + EndUser.created_at, + EndUser.updated_at, + ) + .where( + EndUser.id == end_user_id, + EndUser.tenant_id == tenant_id, + EndUser.app_id == app_id, + ) + .limit(1) + ) + + with self._session_factory() as session: + row = session.execute(statement).one_or_none() + + if row is None: + return None + + ( + record_id, + record_tenant_id, + record_app_id, + end_user_type, + external_user_id, + name, + is_anonymous, + session_id, + created_at, + updated_at, + ) = row + return AppScopedEndUserRecord( + id=record_id, + tenant_id=record_tenant_id, + app_id=record_app_id, + type=end_user_type.value, + external_user_id=external_user_id, + name=name, + is_anonymous=is_anonymous, + session_id=session_id, + created_at=created_at, + updated_at=updated_at, + ) + + @override + def find_by_session( + self, + *, + tenant_id: str, + app_id: str, + user_id: str, + ) -> Sequence[StoredAppScopedEndUser[EndUser]]: + with self._session_factory() as session: + end_users = session.scalars( + select(EndUser).where( + EndUser.tenant_id == tenant_id, + EndUser.app_id == app_id, + EndUser.session_id == user_id, + ) + ).all() + return [self._stored(end_user) for end_user in end_users] + + @override + def find_by_apps( + self, + *, + tenant_id: str, + app_ids: Sequence[str], + user_id: str, + type: str, + ) -> Sequence[StoredAppScopedEndUser[EndUser]]: + with self._session_factory() as session: + end_users = session.scalars( + select(EndUser).where( + EndUser.tenant_id == tenant_id, + EndUser.app_id.in_(app_ids), + EndUser.session_id == user_id, + EndUser.type == EndUserType(type), + ) + ).all() + return [self._stored(end_user) for end_user in end_users] + + @override + def create(self, command: NewAppScopedEndUser) -> StoredAppScopedEndUser[EndUser]: + end_user = EndUser( + tenant_id=command.tenant_id, + app_id=command.app_id, + type=EndUserType(command.type), + is_anonymous=command.is_anonymous, + session_id=command.session_id, + external_user_id=command.external_user_id, + ) + with self._session_factory.begin() as session: + session.add(end_user) + session.flush() + return self._stored(end_user) + + @override + def create_batch( + self, + commands: Sequence[NewAppScopedEndUser], + ) -> Sequence[StoredAppScopedEndUser[EndUser]]: + end_users = [ + EndUser( + tenant_id=command.tenant_id, + app_id=command.app_id, + type=EndUserType(command.type), + is_anonymous=command.is_anonymous, + session_id=command.session_id, + external_user_id=command.external_user_id, + ) + for command in commands + ] + with self._session_factory.begin() as session: + session.add_all(end_users) + if end_users: + session.flush() + return [self._stored(end_user) for end_user in end_users] + + @override + def update_type(self, end_user_id: str, type: str) -> StoredAppScopedEndUser[EndUser]: + with self._session_factory.begin() as session: + end_user = session.get(EndUser, end_user_id) + if end_user is None: + raise RuntimeError(f"End user {end_user_id} disappeared before it could be updated") + end_user.type = EndUserType(type) + return self._stored(end_user) + + @staticmethod + def _stored(end_user: EndUser) -> StoredAppScopedEndUser[EndUser]: + if end_user.id is None: + raise RuntimeError("App-scoped end user must be flushed before it is returned") + if end_user.app_id is None: + raise RuntimeError("App-scoped end user must have an app ID") + return StoredAppScopedEndUser( + id=end_user.id, + app_id=end_user.app_id, + type=end_user.type.value, + value=end_user, + ) diff --git a/api/services/app_scoped_end_user_query_service.py b/api/services/app_scoped_end_user_query_service.py new file mode 100644 index 00000000000..ab41d8a2ea5 --- /dev/null +++ b/api/services/app_scoped_end_user_query_service.py @@ -0,0 +1,29 @@ +"""Application boundary for retrieving app-scoped Service API end users.""" + +from typing import Protocol + +from machinery.context import ServiceApiRequestContext +from services.entities.app_scoped_end_user_entities import AppScopedEndUserRecord + + +class AppScopedEndUserQuery(Protocol): + def find_by_id(self, *, tenant_id: str, app_id: str, end_user_id: str) -> AppScopedEndUserRecord | None: ... + + +class AppScopedEndUserNotFoundError(Exception): + """Raised when an end user is not visible to the admitted app.""" + + +class AppScopedEndUserQueryService: + def __init__(self, *, end_users: AppScopedEndUserQuery) -> None: + self._app_scoped_end_users = end_users + + def get_by_id(self, context: ServiceApiRequestContext, end_user_id: str) -> AppScopedEndUserRecord: + end_user = self._app_scoped_end_users.find_by_id( + tenant_id=context.tenant_id, + app_id=context.app_id, + end_user_id=end_user_id, + ) + if end_user is None: + raise AppScopedEndUserNotFoundError(end_user_id) + return end_user diff --git a/api/services/app_scoped_end_user_service.py b/api/services/app_scoped_end_user_service.py new file mode 100644 index 00000000000..6dd26bd4362 --- /dev/null +++ b/api/services/app_scoped_end_user_service.py @@ -0,0 +1,135 @@ +import logging +from collections.abc import Mapping, Sequence +from typing import Protocol + +from models.enums import DEFAULT_END_USER_SESSION_ID, EndUserType +from services.entities.app_scoped_end_user_entities import NewAppScopedEndUser, StoredAppScopedEndUser + +logger = logging.getLogger(__name__) + + +class AppScopedEndUserRepository[T](Protocol): + def find_by_session( + self, + *, + tenant_id: str, + app_id: str, + user_id: str, + ) -> Sequence[StoredAppScopedEndUser[T]]: ... + + def find_by_apps( + self, + *, + tenant_id: str, + app_ids: Sequence[str], + user_id: str, + type: str, + ) -> Sequence[StoredAppScopedEndUser[T]]: ... + + def create(self, command: NewAppScopedEndUser) -> StoredAppScopedEndUser[T]: ... + + def create_batch(self, commands: Sequence[NewAppScopedEndUser]) -> Sequence[StoredAppScopedEndUser[T]]: ... + + def update_type(self, end_user_id: str, type: str) -> StoredAppScopedEndUser[T]: ... + + +class AppScopedEndUserService[T]: + """Application service for provisioning app-scoped end users.""" + + def __init__(self, *, end_users: AppScopedEndUserRepository[T]) -> None: + self._app_scoped_end_users = end_users + + def get_or_create_end_user( + self, + tenant_id: str, + app_id: str, + user_id: str | None = None, + ) -> T: + return self.get_or_create_end_user_by_type( + EndUserType.SERVICE_API, + tenant_id=tenant_id, + app_id=app_id, + user_id=user_id, + ) + + def get_or_create_end_user_by_type( + self, + type: EndUserType, + tenant_id: str, + app_id: str, + user_id: str | None = None, + ) -> T: + normalized_user_id = user_id or DEFAULT_END_USER_SESSION_ID + candidates = self._app_scoped_end_users.find_by_session( + tenant_id=tenant_id, + app_id=app_id, + user_id=normalized_user_id, + ) + # An AppDeploy row is never a legacy row this service may upgrade. FileGrantService + # reads those rows by type, so retyping one would hide it and strand its files. + candidates = [candidate for candidate in candidates if candidate.type != EndUserType.APP_DEPLOY.value] + end_user = next((candidate for candidate in candidates if candidate.type == type.value), None) + end_user = end_user or next(iter(candidates), None) + + if end_user is None: + return self._app_scoped_end_users.create( + NewAppScopedEndUser( + tenant_id=tenant_id, + app_id=app_id, + type=type.value, + is_anonymous=normalized_user_id == DEFAULT_END_USER_SESSION_ID, + session_id=normalized_user_id, + external_user_id=normalized_user_id, + ) + ).value + + if end_user.type != type.value: + logger.info( + "Upgrading legacy EndUser %s from type=%s to %s for session_id=%s", + end_user.id, + end_user.type, + type.value, + normalized_user_id, + ) + end_user = self._app_scoped_end_users.update_type(end_user.id, type.value) + + return end_user.value + + def create_end_user_batch( + self, + type: EndUserType, + tenant_id: str, + app_ids: list[str], + user_id: str, + ) -> Mapping[str, T]: + normalized_user_id = user_id or DEFAULT_END_USER_SESSION_ID + unique_app_ids = list(dict.fromkeys(app_ids)) + if not unique_app_ids: + return {} + + existing_end_users = self._app_scoped_end_users.find_by_apps( + tenant_id=tenant_id, + app_ids=unique_app_ids, + user_id=normalized_user_id, + type=type.value, + ) + result: dict[str, T] = {} + for end_user in existing_end_users: + result.setdefault(end_user.app_id, end_user.value) + + commands = [ + NewAppScopedEndUser( + tenant_id=tenant_id, + app_id=app_id, + type=type.value, + is_anonymous=normalized_user_id == DEFAULT_END_USER_SESSION_ID, + session_id=normalized_user_id, + external_user_id=normalized_user_id, + ) + for app_id in unique_app_ids + if app_id not in result + ] + for end_user in self._app_scoped_end_users.create_batch(commands): + result[end_user.app_id] = end_user.value + + return result diff --git a/api/services/end_user_service.py b/api/services/end_user_service.py deleted file mode 100644 index cd55a3dba41..00000000000 --- a/api/services/end_user_service.py +++ /dev/null @@ -1,188 +0,0 @@ -import logging -from collections.abc import Mapping - -from sqlalchemy import case, select -from sqlalchemy.orm import sessionmaker - -from extensions.ext_database import db -from models.enums import EndUserType -from models.model import App, DefaultEndUserSessionID, EndUser - -logger = logging.getLogger(__name__) - - -class EndUserService: - """ - Service for managing end users. - """ - - @classmethod - def get_end_user_by_id(cls, *, tenant_id: str, app_id: str, end_user_id: str) -> EndUser | None: - """Get an end user by primary key. - - This is scoped to the provided tenant and app to prevent cross-tenant/app access - when an end-user ID is known. - """ - - with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: - return session.scalar( - select(EndUser) - .where( - EndUser.id == end_user_id, - EndUser.tenant_id == tenant_id, - EndUser.app_id == app_id, - ) - .limit(1) - ) - - @classmethod - def get_or_create_end_user(cls, app_model: App, user_id: str | None = None) -> EndUser: - """ - Get or create an end user for a given app. - """ - - return cls.get_or_create_end_user_by_type(EndUserType.SERVICE_API, app_model.tenant_id, app_model.id, user_id) - - @classmethod - def get_or_create_end_user_by_type( - cls, type: EndUserType, tenant_id: str, app_id: str, user_id: str | None = None - ) -> EndUser: - """ - Get or create an end user for a given app and type. - """ - - if not user_id: - user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID - - with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: - # Query with ORDER BY to prioritize exact type matches while maintaining backward compatibility - # This single query approach is more efficient than separate queries - end_user = session.scalar( - select(EndUser) - .where( - EndUser.tenant_id == tenant_id, - EndUser.app_id == app_id, - EndUser.session_id == user_id, - # An AppDeploy row is never a legacy row this could upgrade: - # the type was added after the split, and FileGrantService - # reads its rows by type. Retyping one here would hide it - # from that read and strand the files it owns. - EndUser.type != EndUserType.APP_DEPLOY, - ) - .order_by( - # Prioritize records with matching type (0 = match, 1 = no match) - case((EndUser.type == type, 0), else_=1) - ) - .limit(1) - ) - - if end_user: - # If found a legacy end user with different type, update it for future consistency - if end_user.type != type: - logger.info( - "Upgrading legacy EndUser %s from type=%s to %s for session_id=%s", - end_user.id, - end_user.type, - type, - user_id, - ) - end_user.type = type - else: - # Create new end user if none exists - end_user = EndUser( - tenant_id=tenant_id, - app_id=app_id, - type=type, - is_anonymous=user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID, - session_id=user_id, - external_user_id=user_id, - ) - session.add(end_user) - - return end_user - - @classmethod - def create_end_user_batch( - cls, type: EndUserType, tenant_id: str, app_ids: list[str], user_id: str - ) -> Mapping[str, EndUser]: - """Create end users in batch. - - Creates end users in batch for the specified tenant and application IDs in O(1) time. - - This batch creation is necessary because trigger subscriptions can span multiple applications, - and trigger events may be dispatched to multiple applications simultaneously. - - For each app_id in app_ids, check if an `EndUser` with the given - `user_id` (as session_id/external_user_id) already exists for the - tenant/app and type `type`. If it exists, return it; otherwise, - create it. Operates with minimal DB I/O by querying and inserting in - batches. - - Returns a mapping of `app_id -> EndUser`. - """ - - # Normalize user_id to default if empty - if not user_id: - user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID - - # Deduplicate app_ids while preserving input order - seen: set[str] = set() - unique_app_ids: list[str] = [] - for app_id in app_ids: - if app_id not in seen: - seen.add(app_id) - unique_app_ids.append(app_id) - - # Result is a simple app_id -> EndUser mapping - result: dict[str, EndUser] = {} - if not unique_app_ids: - return result - - with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: - # Fetch existing end users for all target apps in a single query - existing_end_users: list[EndUser] = list( - session.scalars( - select(EndUser).where( - EndUser.tenant_id == tenant_id, - EndUser.app_id.in_(unique_app_ids), - EndUser.session_id == user_id, - EndUser.type == type, - ) - ).all() - ) - - found_app_ids: set[str] = set() - for eu in existing_end_users: - # If duplicates exist due to weak DB constraints, prefer the first - if eu.app_id is None: - continue - if eu.app_id not in result: - result[eu.app_id] = eu - found_app_ids.add(eu.app_id) - - # Determine which apps still need an EndUser created - missing_app_ids = [app_id for app_id in unique_app_ids if app_id not in found_app_ids] - - if missing_app_ids: - new_end_users: list[EndUser] = [] - is_anonymous = user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID - for app_id in missing_app_ids: - new_end_users.append( - EndUser( - tenant_id=tenant_id, - app_id=app_id, - type=type, - is_anonymous=is_anonymous, - session_id=user_id, - external_user_id=user_id, - ) - ) - - session.add_all(new_end_users) - - for eu in new_end_users: - if eu.app_id is None: - continue - result[eu.app_id] = eu - - return result diff --git a/api/services/entities/app_scoped_end_user_entities.py b/api/services/entities/app_scoped_end_user_entities.py new file mode 100644 index 00000000000..bb475ddf576 --- /dev/null +++ b/api/services/entities/app_scoped_end_user_entities.py @@ -0,0 +1,38 @@ +"""Framework-neutral contracts for app-scoped end-user use cases.""" + +from dataclasses import dataclass +from datetime import datetime +from typing import NamedTuple + + +class AppScopedEndUserRecord(NamedTuple): + id: str + tenant_id: str + app_id: str + type: str + external_user_id: str | None + name: str | None + is_anonymous: bool + session_id: str + created_at: datetime + updated_at: datetime + + +@dataclass(frozen=True, slots=True) +class StoredAppScopedEndUser[T]: + """Persistence metadata paired with an opaque caller-facing entity.""" + + id: str + app_id: str + type: str + value: T + + +@dataclass(frozen=True, slots=True) +class NewAppScopedEndUser: + tenant_id: str + app_id: str + type: str + is_anonymous: bool + session_id: str + external_user_id: str diff --git a/api/services/trigger/webhook_service.py b/api/services/trigger/webhook_service.py index 382fd1ed3d7..7381e9b3e36 100644 --- a/api/services/trigger/webhook_service.py +++ b/api/services/trigger/webhook_service.py @@ -3,7 +3,7 @@ import logging import mimetypes import secrets from collections.abc import Callable, Mapping, Sequence -from typing import Any, NotRequired, TypedDict +from typing import Any, NotRequired, Protocol, TypedDict import orjson from flask import request @@ -31,11 +31,10 @@ from graphon.entities.graph_config import NodeConfigDict from graphon.file import FileTransferMethod from graphon.variables.types import ArrayValidation, SegmentType from models.enums import AppTriggerStatus, AppTriggerType, EndUserType -from models.model import App +from models.model import App, EndUser from models.trigger import AppTrigger, WorkflowWebhookTrigger from models.workflow import Workflow from services.async_workflow_service import AsyncWorkflowService -from services.end_user_service import EndUserService from services.errors.app import QuotaExceededError from services.quota_service import QuotaService from services.trigger.app_trigger_service import AppTriggerService @@ -71,6 +70,16 @@ class WorkflowInputsDict(TypedDict): webhook_body: dict[str, Any] +class WebhookEndUserProvisioner(Protocol): + def get_or_create_end_user_by_type( + self, + type: EndUserType, + tenant_id: str, + app_id: str, + user_id: str | None = None, + ) -> EndUser: ... + + class WebhookService: """Service for handling webhook operations.""" @@ -792,7 +801,12 @@ class WebhookService: @classmethod def trigger_workflow_execution( - cls, webhook_trigger: WorkflowWebhookTrigger, webhook_data: RawWebhookDataDict, workflow: Workflow + cls, + webhook_trigger: WorkflowWebhookTrigger, + webhook_data: RawWebhookDataDict, + workflow: Workflow, + *, + end_users: WebhookEndUserProvisioner, ) -> None: """Trigger workflow execution via AsyncWorkflowService. @@ -817,7 +831,7 @@ class WebhookService: tenant_id=webhook_trigger.tenant_id, ) - end_user = EndUserService.get_or_create_end_user_by_type( + end_user = end_users.get_or_create_end_user_by_type( type=EndUserType.TRIGGER, tenant_id=webhook_trigger.tenant_id, app_id=webhook_trigger.app_id, diff --git a/api/tasks/trigger_processing_tasks.py b/api/tasks/trigger_processing_tasks.py index 59e2f4e364c..241452045b7 100644 --- a/api/tasks/trigger_processing_tasks.py +++ b/api/tasks/trigger_processing_tasks.py @@ -9,7 +9,7 @@ import json import logging from collections.abc import Mapping, Sequence from datetime import UTC, datetime -from typing import Any +from typing import Any, Protocol from celery import shared_task from sqlalchemy import select @@ -27,6 +27,7 @@ from core.trigger.provider import PluginTriggerProviderController from core.trigger.trigger_manager import TriggerManager from core.workflow.nodes.trigger_plugin.entities import TriggerEventNodeData from enums import QuotaType +from extensions.ext_application_services import application_services from graphon.enums import WorkflowExecutionStatus from models.enums import ( AppTriggerType, @@ -40,7 +41,6 @@ from models.provider_ids import TriggerProviderID from models.trigger import TriggerSubscription, WorkflowPluginTrigger, WorkflowTriggerLog from models.workflow import Workflow, WorkflowAppLog, WorkflowAppLogCreatedFrom, WorkflowRun from services.async_workflow_service import AsyncWorkflowService -from services.end_user_service import EndUserService from services.errors.app import QuotaExceededError from services.quota_service import QuotaService, unlimited from services.trigger.app_trigger_service import AppTriggerService @@ -56,6 +56,16 @@ logger = logging.getLogger(__name__) TRIGGER_QUEUE = "triggered_workflow_dispatcher" +class TriggerEndUserProvisioner(Protocol): + def create_end_user_batch( + self, + type: EndUserType, + tenant_id: str, + app_ids: list[str], + user_id: str, + ) -> Mapping[str, EndUser]: ... + + def dispatch_trigger_debug_event( events: list[str], user_id: str, @@ -234,6 +244,8 @@ def dispatch_triggered_workflow( subscription: TriggerSubscription, event_name: str, request_id: str, + *, + end_users: TriggerEndUserProvisioner, ) -> int: """Process triggered workflows. @@ -266,7 +278,7 @@ def dispatch_triggered_workflow( with session_factory.create_session() as session: workflows: Mapping[str, Workflow] = _get_published_workflows_by_app_ids(session, subscribers) - end_users: Mapping[str, EndUser] = EndUserService.create_end_user_batch( + end_users_by_app: Mapping[str, EndUser] = end_users.create_end_user_batch( type=EndUserType.TRIGGER, tenant_id=subscription.tenant_id, app_ids=[plugin_trigger.app_id for plugin_trigger in subscribers], @@ -333,7 +345,7 @@ def dispatch_triggered_workflow( error_message = e.to_user_friendly_error(plugin_name=trigger_entity.identity.name) try: - end_user = end_users.get(plugin_trigger.app_id) + end_user = end_users_by_app.get(plugin_trigger.app_id) _record_trigger_failure_log( session=session, workflow=workflow, @@ -384,7 +396,7 @@ def dispatch_triggered_workflow( # Trigger async workflow try: - end_user = end_users.get(plugin_trigger.app_id) + end_user = end_users_by_app.get(plugin_trigger.app_id) if not end_user: raise ValueError(f"End user not found for app {plugin_trigger.app_id}") @@ -412,6 +424,8 @@ def dispatch_triggered_workflows( events: list[str], subscription: TriggerSubscription, request_id: str, + *, + end_users: TriggerEndUserProvisioner, ) -> int: dispatched_count = 0 for event_name in events: @@ -421,6 +435,7 @@ def dispatch_triggered_workflows( subscription=subscription, event_name=event_name, request_id=request_id, + end_users=end_users, ) except Exception: logger.exception( @@ -494,6 +509,7 @@ def dispatch_triggered_workflows_async( events=events, subscription=subscription, request_id=request_id, + end_users=application_services().app_scoped_end_users.commands, ) debug_dispatched = dispatch_trigger_debug_event( diff --git a/api/tests/test_containers_integration_tests/services/test_end_user_service.py b/api/tests/test_containers_integration_tests/services/test_end_user_service.py index b6104f94f21..2864f135ca1 100644 --- a/api/tests/test_containers_integration_tests/services/test_end_user_service.py +++ b/api/tests/test_containers_integration_tests/services/test_end_user_service.py @@ -6,11 +6,15 @@ from uuid import uuid4 import pytest from sqlalchemy.orm import Session +from extensions.ext_application_services import application_services from models import TenantAccountRole from models.account import Account, Tenant, TenantAccountJoin -from models.enums import EndUserType -from models.model import App, DefaultEndUserSessionID, EndUser -from services.end_user_service import EndUserService +from models.enums import DEFAULT_END_USER_SESSION_ID, EndUserType +from models.model import App, EndUser + + +def _service(): + return application_services().app_scoped_end_users.commands class TestEndUserServiceFactory: @@ -90,7 +94,7 @@ class TestEndUserServiceFactory: class TestEndUserServiceGetOrCreateEndUser: """ - Unit tests for EndUserService.get_or_create_end_user method. + Unit tests for _service().get_or_create_end_user method. This test suite covers: - Creating new end users @@ -113,7 +117,7 @@ class TestEndUserServiceGetOrCreateEndUser: user_id = "custom-user-123" # Act - result = EndUserService.get_or_create_end_user(app_model=app, user_id=user_id) + result = _service().get_or_create_end_user(tenant_id=app.tenant_id, app_id=app.id, user_id=user_id) # Assert assert result.tenant_id == app.tenant_id @@ -130,10 +134,10 @@ class TestEndUserServiceGetOrCreateEndUser: app = factory.create_app_and_account(db_session_with_containers) # Act - result = EndUserService.get_or_create_end_user(app_model=app, user_id=None) + result = _service().get_or_create_end_user(tenant_id=app.tenant_id, app_id=app.id, user_id=None) # Assert - assert result.session_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID + assert result.session_id == DEFAULT_END_USER_SESSION_ID # Verify _is_anonymous is set correctly (property always returns False) assert result._is_anonymous is True @@ -151,7 +155,7 @@ class TestEndUserServiceGetOrCreateEndUser: ) # Act - result = EndUserService.get_or_create_end_user(app_model=app, user_id=user_id) + result = _service().get_or_create_end_user(tenant_id=app.tenant_id, app_id=app.id, user_id=user_id) # Assert assert result.id == existing_user.id @@ -159,7 +163,7 @@ class TestEndUserServiceGetOrCreateEndUser: class TestEndUserServiceGetOrCreateEndUserByType: """ - Unit tests for EndUserService.get_or_create_end_user_by_type method. + Unit tests for _service().get_or_create_end_user_by_type method. This test suite covers: - Creating end users with different EndUserType values @@ -184,7 +188,7 @@ class TestEndUserServiceGetOrCreateEndUserByType: user_id = "user-789" # Act - result = EndUserService.get_or_create_end_user_by_type( + result = _service().get_or_create_end_user_by_type( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_id=app_id, @@ -208,7 +212,7 @@ class TestEndUserServiceGetOrCreateEndUserByType: user_id = "user-789" # Act - result = EndUserService.get_or_create_end_user_by_type( + result = _service().get_or_create_end_user_by_type( type=EndUserType.BROWSER, tenant_id=tenant_id, app_id=app_id, @@ -236,9 +240,9 @@ class TestEndUserServiceGetOrCreateEndUserByType: session_id=user_id, invoke_type=EndUserType.SERVICE_API, ) - with caplog.at_level(logging.INFO, logger="services.end_user_service"): + with caplog.at_level(logging.INFO, logger="services.app_scoped_end_user_service"): # Act - Request with different type - result = EndUserService.get_or_create_end_user_by_type( + result = _service().get_or_create_end_user_by_type( type=EndUserType.BROWSER, tenant_id=tenant_id, app_id=app_id, @@ -251,7 +255,7 @@ class TestEndUserServiceGetOrCreateEndUserByType: matching_logs = [ record for record in caplog.records - if record.name == "services.end_user_service" + if record.name == "services.app_scoped_end_user_service" and record.levelno == logging.INFO and "Upgrading legacy EndUser" in record.message ] @@ -277,8 +281,8 @@ class TestEndUserServiceGetOrCreateEndUserByType: ) # Act - Request with same type - with caplog.at_level(logging.INFO, logger="services.end_user_service"): - result = EndUserService.get_or_create_end_user_by_type( + with caplog.at_level(logging.INFO, logger="services.app_scoped_end_user_service"): + result = _service().get_or_create_end_user_by_type( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_id=app_id, @@ -301,7 +305,7 @@ class TestEndUserServiceGetOrCreateEndUserByType: app_id = app.id # Act - result = EndUserService.get_or_create_end_user_by_type( + result = _service().get_or_create_end_user_by_type( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_id=app_id, @@ -309,10 +313,10 @@ class TestEndUserServiceGetOrCreateEndUserByType: ) # Assert - assert result.session_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID + assert result.session_id == DEFAULT_END_USER_SESSION_ID # Verify _is_anonymous is set correctly (property always returns False) assert result._is_anonymous is True - assert result.external_user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID + assert result.external_user_id == DEFAULT_END_USER_SESSION_ID def test_query_ordering_prioritizes_matching_type( self, db_session_with_containers: Session, factory: TestEndUserServiceFactory @@ -340,7 +344,7 @@ class TestEndUserServiceGetOrCreateEndUserByType: ) # Act - result = EndUserService.get_or_create_end_user_by_type( + result = _service().get_or_create_end_user_by_type( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_id=app_id, @@ -362,7 +366,7 @@ class TestEndUserServiceGetOrCreateEndUserByType: user_id = "custom-external-id" # Act - result = EndUserService.get_or_create_end_user_by_type( + result = _service().get_or_create_end_user_by_type( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_id=app_id, @@ -393,7 +397,7 @@ class TestEndUserServiceGetOrCreateEndUserByType: user_id = f"user-{uuid4()}" # Act - result = EndUserService.get_or_create_end_user_by_type( + result = _service().get_or_create_end_user_by_type( type=invoke_type, tenant_id=tenant_id, app_id=app_id, @@ -404,51 +408,8 @@ class TestEndUserServiceGetOrCreateEndUserByType: assert result.type == invoke_type -class TestEndUserServiceGetEndUserById: - """Unit tests for EndUserService.get_end_user_by_id.""" - - @pytest.fixture - def factory(self): - """Provide test data factory.""" - return TestEndUserServiceFactory() - - def test_get_end_user_by_id_returns_end_user( - self, db_session_with_containers: Session, factory: TestEndUserServiceFactory - ): - app = factory.create_app_and_account(db_session_with_containers) - existing_user = factory.create_end_user( - db_session_with_containers, - tenant_id=app.tenant_id, - app_id=app.id, - session_id=f"session-{uuid4()}", - invoke_type=EndUserType.SERVICE_API, - ) - - result = EndUserService.get_end_user_by_id( - tenant_id=app.tenant_id, - app_id=app.id, - end_user_id=existing_user.id, - ) - - assert result is not None - assert result.id == existing_user.id - - def test_get_end_user_by_id_returns_none( - self, db_session_with_containers: Session, factory: TestEndUserServiceFactory - ): - app = factory.create_app_and_account(db_session_with_containers) - - result = EndUserService.get_end_user_by_id( - tenant_id=app.tenant_id, - app_id=app.id, - end_user_id=str(uuid4()), - ) - - assert result is None - - class TestEndUserServiceCreateBatch: - """Integration tests for EndUserService.create_end_user_batch.""" + """Integration tests for _service().create_end_user_batch.""" @pytest.fixture def factory(self): @@ -486,7 +447,7 @@ class TestEndUserServiceCreateBatch: return tenant_id, all_apps def test_create_batch_empty_app_ids(self, db_session_with_containers: Session): - result = EndUserService.create_end_user_batch( + result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=str(uuid4()), app_ids=[], user_id="user-1" ) assert result == {} @@ -498,7 +459,7 @@ class TestEndUserServiceCreateBatch: app_ids = [a.id for a in apps] user_id = f"user-{uuid4()}" - result = EndUserService.create_end_user_batch( + result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_ids=app_ids, user_id=user_id ) @@ -514,13 +475,13 @@ class TestEndUserServiceCreateBatch: tenant_id, apps = self._create_multiple_apps(db_session_with_containers, factory, count=2) app_ids = [a.id for a in apps] - result = EndUserService.create_end_user_batch( + result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_ids=app_ids, user_id="" ) assert len(result) == 2 for end_user in result.values(): - assert end_user.session_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID + assert end_user.session_id == DEFAULT_END_USER_SESSION_ID assert end_user._is_anonymous is True def test_create_batch_deduplicate_app_ids( @@ -530,7 +491,7 @@ class TestEndUserServiceCreateBatch: app_ids = [apps[0].id, apps[1].id, apps[0].id, apps[1].id] user_id = f"user-{uuid4()}" - result = EndUserService.create_end_user_batch( + result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_ids=app_ids, user_id=user_id ) @@ -544,12 +505,12 @@ class TestEndUserServiceCreateBatch: user_id = f"user-{uuid4()}" # Create batch first time - first_result = EndUserService.create_end_user_batch( + first_result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_ids=app_ids, user_id=user_id ) # Create batch second time — should return existing users - second_result = EndUserService.create_end_user_batch( + second_result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_ids=app_ids, user_id=user_id ) @@ -564,7 +525,7 @@ class TestEndUserServiceCreateBatch: user_id = f"user-{uuid4()}" # Create for first 2 apps - first_result = EndUserService.create_end_user_batch( + first_result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_ids=[apps[0].id, apps[1].id], @@ -572,7 +533,7 @@ class TestEndUserServiceCreateBatch: ) # Create for all 3 apps — should reuse first 2, create 3rd - all_result = EndUserService.create_end_user_batch( + all_result = _service().create_end_user_batch( type=EndUserType.SERVICE_API, tenant_id=tenant_id, app_ids=[a.id for a in apps], @@ -594,7 +555,7 @@ class TestEndUserServiceCreateBatch: tenant_id, apps = self._create_multiple_apps(db_session_with_containers, factory, count=1) user_id = f"user-{uuid4()}" - result = EndUserService.create_end_user_batch( + result = _service().create_end_user_batch( type=invoke_type, tenant_id=tenant_id, app_ids=[apps[0].id], user_id=user_id ) diff --git a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py index 432a483b8c0..22fb7da062e 100644 --- a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py +++ b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py @@ -3,7 +3,6 @@ from __future__ import annotations import json import logging from datetime import datetime -from types import SimpleNamespace from unittest.mock import MagicMock, patch from uuid import uuid4 @@ -15,14 +14,42 @@ from sqlalchemy.orm import Session from core.trigger.constants import TRIGGER_WEBHOOK_NODE_TYPE from enums import QuotaType from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole -from models.enums import AppTriggerStatus, AppTriggerType -from models.model import App +from models.enums import AppTriggerStatus, AppTriggerType, EndUserType +from models.model import App, EndUser from models.trigger import AppTrigger, WorkflowWebhookTrigger from models.workflow import Workflow from services.errors.app import QuotaExceededError from services.trigger.webhook_service import WebhookService +class _EndUserServiceStub: + def __init__(self, result: EndUser | Exception) -> None: + self._result = result + self.calls: list[tuple[EndUserType, str, str, str | None]] = [] + + def get_or_create_end_user_by_type( + self, + type: EndUserType, + tenant_id: str, + app_id: str, + user_id: str | None = None, + ) -> EndUser: + self.calls.append((type, tenant_id, app_id, user_id)) + if isinstance(self._result, Exception): + raise self._result + return self._result + + +def _end_user(*, tenant_id: str, app_id: str) -> EndUser: + return EndUser( + id=str(uuid4()), + tenant_id=tenant_id, + app_id=app_id, + type=EndUserType.TRIGGER, + session_id="trigger-session", + ) + + class WebhookServiceRelationshipFactory: @staticmethod def create_account_and_tenant(db_session_with_containers: Session) -> tuple[Account, Tenant]: @@ -329,23 +356,24 @@ class TestWebhookServiceTriggerExecutionWithContainers: db_session_with_containers, app=app, account=account, node_id="node-1" ) - end_user = SimpleNamespace(id=str(uuid4())) + end_user = _end_user(tenant_id=tenant.id, app_id=app.id) webhook_data = {"body": {"value": 1}, "headers": {}, "query_params": {}, "files": {}, "method": "POST"} quota_charge = MagicMock() with ( - patch( - "services.trigger.webhook_service.EndUserService.get_or_create_end_user_by_type", - return_value=end_user, - ), patch( "services.trigger.webhook_service.QuotaService.reserve", return_value=quota_charge, ) as mock_reserve, patch("services.trigger.webhook_service.AsyncWorkflowService.trigger_workflow_async") as mock_trigger, ): - WebhookService.trigger_workflow_execution(webhook_trigger, webhook_data, workflow) + WebhookService.trigger_workflow_execution( + webhook_trigger, + webhook_data, + workflow, + end_users=_EndUserServiceStub(end_user), + ) mock_reserve.assert_called_once() reserve_args = mock_reserve.call_args.args @@ -374,10 +402,6 @@ class TestWebhookServiceTriggerExecutionWithContainers: ) with ( - patch( - "services.trigger.webhook_service.EndUserService.get_or_create_end_user_by_type", - return_value=SimpleNamespace(id=str(uuid4())), - ), patch( "services.trigger.webhook_service.QuotaService.reserve", side_effect=QuotaExceededError(feature="trigger", tenant_id=tenant.id, required=1), @@ -391,6 +415,7 @@ class TestWebhookServiceTriggerExecutionWithContainers: webhook_trigger, {"body": {}, "headers": {}, "query_params": {}, "files": {}, "method": "POST"}, workflow, + end_users=_EndUserServiceStub(_end_user(tenant_id=tenant.id, app_id=app.id)), ) mock_mark_rate_limited.assert_called_once_with(tenant.id) @@ -413,16 +438,13 @@ class TestWebhookServiceTriggerExecutionWithContainers: ) caplog.set_level(logging.ERROR, logger="services.trigger.webhook_service") - with patch( - "services.trigger.webhook_service.EndUserService.get_or_create_end_user_by_type", - side_effect=RuntimeError("boom"), - ): - with pytest.raises(RuntimeError, match="boom"): - WebhookService.trigger_workflow_execution( - webhook_trigger, - {"body": {}, "headers": {}, "query_params": {}, "files": {}, "method": "POST"}, - workflow, - ) + with pytest.raises(RuntimeError, match="boom"): + WebhookService.trigger_workflow_execution( + webhook_trigger, + {"body": {}, "headers": {}, "query_params": {}, "files": {}, "method": "POST"}, + workflow, + end_users=_EndUserServiceStub(RuntimeError("boom")), + ) assert caplog.messages.count(f"Failed to trigger workflow for webhook {webhook_trigger.webhook_id}") == 1 diff --git a/api/tests/unit_tests/controllers/inner_api/app/test_file_grants.py b/api/tests/unit_tests/controllers/inner_api/app/test_file_grants.py index 795fce195d9..49350be56cd 100644 --- a/api/tests/unit_tests/controllers/inner_api/app/test_file_grants.py +++ b/api/tests/unit_tests/controllers/inner_api/app/test_file_grants.py @@ -29,8 +29,8 @@ from libs.datetime_utils import naive_utc_now from models.enums import CreatorUserRole, EndUserType from models.model import App, EndUser, UploadFile from models.tools import ToolFile -from services import end_user_service -from services.end_user_service import EndUserService +from repositories.app_scoped_end_user_repository import AppScopedEndUserRepo +from services.app_scoped_end_user_service import AppScopedEndUserService from services.file_grant_gateways import FILE_GRANT_AUDIENCE from services.file_grant_service import ( MAX_RUN_GRANT_TTL_SECONDS, @@ -437,9 +437,8 @@ def test_mint_reports_optional_files_item_by_item(app: Flask, sqlite_session: Se @pytest.mark.usefixtures("seeded_app") def test_end_user_service_never_retypes_an_app_deploy_row( - sqlite_engine: Engine, sqlite_session: Session, - monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], ) -> None: """Retyping would hide the row from the grant read and strand its files.""" @@ -456,9 +455,10 @@ def test_end_user_service_never_retypes_an_app_deploy_row( sqlite_session.add(owner) sqlite_session.commit() owner_id = owner.id - monkeypatch.setattr(end_user_service, "db", SimpleNamespace(engine=sqlite_engine)) - EndUserService.get_or_create_end_user_by_type(EndUserType.SERVICE_API, TENANT_ID, APP_ID, session_id) + AppScopedEndUserService( + end_users=AppScopedEndUserRepo(session_factory=sqlite_session_factory) + ).get_or_create_end_user_by_type(EndUserType.SERVICE_API, TENANT_ID, APP_ID, session_id) sqlite_session.expire_all() persisted_owner = sqlite_session.get(EndUser, owner_id) @@ -472,7 +472,7 @@ SKIPPED_DIRECTORY_NAMES = frozenset({".git", ".venv", "__pycache__", "migrations def test_app_deploy_end_users_have_exactly_one_writer() -> None: """``end_users`` has no unique constraint, so a second writer would fork identities. - ``end_user_service`` names the type only to exclude it from the legacy retype; + ``app_scoped_end_user_service`` names the type only to exclude it from the legacy retype; the behavioural guard above is what holds that exclusion in place. """ @@ -491,5 +491,5 @@ def test_app_deploy_end_users_have_exactly_one_writer() -> None: assert referencing_modules == { "repositories/file_grant_repository.py", - "services/end_user_service.py", + "services/app_scoped_end_user_service.py", } diff --git a/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin_wraps.py b/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin_wraps.py index 374c2e7ea9c..44482d46611 100644 --- a/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin_wraps.py +++ b/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin_wraps.py @@ -22,8 +22,8 @@ from controllers.inner_api.plugin.wraps import ( ) from models.account import Tenant from models.base import TypeBase -from models.enums import EndUserType -from models.model import DefaultEndUserSessionID, EndUser +from models.enums import DEFAULT_END_USER_SESSION_ID, EndUserType +from models.model import EndUser @pytest.fixture @@ -201,7 +201,7 @@ class TestGetUser: _persist_end_user( sqlite_plugin_engine, user_id="default-user-id", - session_id=DefaultEndUserSessionID.DEFAULT_SESSION_ID, + session_id=DEFAULT_END_USER_SESSION_ID, is_anonymous=True, ) @@ -209,7 +209,7 @@ class TestGetUser: result = get_user("tenant123", None) assert result.id == "default-user-id" - assert result.session_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID + assert result.session_id == DEFAULT_END_USER_SESSION_ID def test_should_raise_error_on_database_exception(self, sqlite_plugin_engine: Engine, app: Flask): """Test raising ValueError when database operation fails""" @@ -301,7 +301,7 @@ class TestGetUserTenant: _persist_end_user( sqlite_plugin_engine, user_id="default-user-id", - session_id=DefaultEndUserSessionID.DEFAULT_SESSION_ID, + session_id=DEFAULT_END_USER_SESSION_ID, is_anonymous=True, ) @@ -312,7 +312,7 @@ class TestGetUserTenant: assert result["tenant"].id == "tenant123" assert result["user"].id == "default-user-id" - assert result["user"].session_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID + assert result["user"].session_id == DEFAULT_END_USER_SESSION_ID class PluginTestPayload: diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_pipelines.py b/api/tests/unit_tests/controllers/openapi/auth/test_pipelines.py index 298bb01aea5..8b306667ce8 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_pipelines.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_pipelines.py @@ -33,7 +33,6 @@ from enums import DeploymentEdition from libs.oauth_bearer import AuthContext, try_get_auth_ctx from services.account_service import AccountService, TenantService from services.app_service import AppService -from services.end_user_service import EndUserService from services.enterprise.enterprise_service import WebAppAccessMode from ._world import ( @@ -235,7 +234,7 @@ def test_a_refused_sso_request_never_creates_an_end_user( config_overrides(DEPLOYMENT_EDITION=DeploymentEdition.ENTERPRISE) persist(sqlite_session, make_app(enable_api=enable_api), make_tenant()) monkeypatch.setattr(MOUNT, never_reached) - monkeypatch.setattr(EndUserService, "get_or_create_end_user_by_type", never_reached) + monkeypatch.setattr("controllers.openapi.auth.subjects.application_services", never_reached) subject = sso_subject() ctx = make_ctx(sqlite_session, subject, app_id=APP_ID) diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_subjects.py b/api/tests/unit_tests/controllers/openapi/auth/test_subjects.py index ce73594e2cc..a3bd212dae6 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_subjects.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_subjects.py @@ -1,9 +1,10 @@ from __future__ import annotations -from unittest.mock import patch +from types import SimpleNamespace import pytest -from sqlalchemy.orm import Session +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import Unauthorized from controllers.openapi.auth.subjects import AccountSubject, ExternalSsoSubject @@ -11,6 +12,8 @@ from libs.oauth_bearer import TokenType from models import Account, EndUser, TenantAccountJoin from models.account import TenantAccountRole from models.enums import EndUserType +from repositories.app_scoped_end_user_repository import AppScopedEndUserRepo +from services.app_scoped_end_user_service import AppScopedEndUserService from ._world import ( ACCOUNT_ID, @@ -76,7 +79,12 @@ class TestAccountResolveCaller: class TestExternalSsoResolveCaller: - def test_resolves_the_end_user_against_the_apps_workspace(self, sqlite_session: Session) -> None: + def test_resolves_the_end_user_against_the_apps_workspace( + self, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + monkeypatch: pytest.MonkeyPatch, + ) -> None: """It loads both itself. Nothing before it on an SSO route needs a workspace, so a subject that expected one to be there already would resolve an end user against nothing. @@ -84,26 +92,23 @@ class TestExternalSsoResolveCaller: persist(sqlite_session, make_app(), make_tenant()) subject = ExternalSsoSubject(make_auth(TokenType.OAUTH_EXTERNAL_SSO)) ctx = make_ctx(sqlite_session, subject, app_id=APP_ID) - end_user = EndUser( - tenant_id=TENANT_ID, - app_id=APP_ID, - type=EndUserType.OPENAPI, - is_anonymous=False, - session_id=SSO_EMAIL, + commands = AppScopedEndUserService( + end_users=AppScopedEndUserRepo(session_factory=sqlite_session_factory), ) + services = SimpleNamespace(app_scoped_end_users=SimpleNamespace(commands=commands)) + monkeypatch.setattr("controllers.openapi.auth.subjects.application_services", lambda: services) - with patch( - "controllers.openapi.auth.subjects.EndUserService.get_or_create_end_user_by_type", - return_value=end_user, - ) as get_or_create: - assert subject.resolve_caller(ctx, sqlite_session) is end_user + caller = subject.resolve_caller(ctx, sqlite_session) - get_or_create.assert_called_once_with( - EndUserType.OPENAPI, - tenant_id=TENANT_ID, - app_id=APP_ID, - user_id=SSO_EMAIL, - ) + assert isinstance(caller, EndUser) + with sqlite_session_factory() as observer: + persisted = observer.scalar(select(EndUser).where(EndUser.session_id == SSO_EMAIL)) + assert persisted is not None + assert persisted.id == caller.id + assert persisted.tenant_id == TENANT_ID + assert persisted.app_id == APP_ID + assert persisted.type == EndUserType.OPENAPI + assert persisted.external_user_id == SSO_EMAIL def test_rejects_a_token_without_an_external_identity(self, sqlite_session: Session) -> None: subject = ExternalSsoSubject(make_auth(TokenType.OAUTH_EXTERNAL_SSO, subject_email=None)) diff --git a/api/tests/unit_tests/controllers/openapi/test_auth_matrix.py b/api/tests/unit_tests/controllers/openapi/test_auth_matrix.py index e95d696b217..7fd823c714e 100644 --- a/api/tests/unit_tests/controllers/openapi/test_auth_matrix.py +++ b/api/tests/unit_tests/controllers/openapi/test_auth_matrix.py @@ -104,7 +104,6 @@ from models.enums import EndUserType from models.model import App, EndUser from models.oauth import OAuthAccessToken from services.account_service import AccountService -from services.end_user_service import EndUserService from services.enterprise.enterprise_service import EnterpriseService from services.entities.feature_entities import LicenseStatus, SystemFeatureModel from services.rbac_resource_service import RBACResourceService @@ -1108,13 +1107,8 @@ def _run_case( return_value={"data": [], "total": 0, "hasMore": False}, ) ) - stack.enter_context( - patch.object( - EndUserService, - "get_or_create_end_user_by_type", - side_effect=_end_user, - ) - ) + services = stack.enter_context(patch("controllers.openapi.auth.subjects.application_services")) + services.return_value.app_scoped_end_users.commands.get_or_create_end_user_by_type.side_effect = _end_user stack.enter_context(patch.object(RBACResourceService, "get_app_agent_binding", return_value=None)) stack.enter_context(patch.object(RBACResourceService, "get_app_maintainer", return_value=None)) stack.enter_context( diff --git a/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py b/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py index 30449f4b6b3..7177b256aae 100644 --- a/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py +++ b/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py @@ -1,66 +1,104 @@ from datetime import UTC, datetime +from inspect import unwrap from uuid import UUID, uuid4 import pytest -from pytest_mock import MockerFixture +from controllers.service_api.end_user import end_user as controller_module from controllers.service_api.end_user.end_user import EndUserApi from controllers.service_api.end_user.error import EndUserNotFoundError -from models.enums import EndUserType -from models.model import App, EndUser +from machinery.context import ServiceApiRequestContext +from models.model import App +from services.app_scoped_end_user_query_service import AppScopedEndUserNotFoundError +from services.entities.app_scoped_end_user_entities import AppScopedEndUserRecord + + +def _request_context(*, tenant_id: str = "workspace-1", app_id: str = "app-1") -> ServiceApiRequestContext: + return ServiceApiRequestContext( + tenant_id=tenant_id, + app_id=app_id, + ) + + +class EndUserQueryServiceStub: + def __init__(self, result: AppScopedEndUserRecord | Exception) -> None: + self._result = result + self.calls: list[tuple[ServiceApiRequestContext, str]] = [] + + def get_by_id(self, context: ServiceApiRequestContext, end_user_id: str) -> AppScopedEndUserRecord: + self.calls.append((context, end_user_id)) + if isinstance(self._result, Exception): + raise self._result + return self._result + + +class ApplicationServicesStub: + def __init__(self, end_user_queries: EndUserQueryServiceStub) -> None: + self.app_scoped_end_users = AppScopedEndUserServicesStub(end_user_queries) + + +class AppScopedEndUserServicesStub: + def __init__(self, queries: EndUserQueryServiceStub) -> None: + self.queries = queries class TestEndUserApi: - @pytest.fixture - def resource(self) -> EndUserApi: - return EndUserApi() - - @pytest.fixture - def app_model(self) -> App: - app = App( + def test_get_end_user_returns_all_attributes(self, monkeypatch: pytest.MonkeyPatch) -> None: + end_user = AppScopedEndUserRecord( id=str(uuid4()), tenant_id=str(uuid4()), - ) - return app - - def test_get_end_user_returns_all_attributes( - self, mocker: MockerFixture, resource: EndUserApi, app_model: App - ) -> None: - end_user = EndUser( - id=str(uuid4()), - tenant_id=app_model.tenant_id, - app_id=app_model.id, - type=EndUserType.SERVICE_API, + app_id=str(uuid4()), + type="service-api", external_user_id="external-123", name="Alice", - _is_anonymous=True, + is_anonymous=True, session_id="session-xyz", created_at=datetime(2024, 1, 1, tzinfo=UTC), updated_at=datetime(2024, 1, 2, tzinfo=UTC), ) - - get_end_user_by_id = mocker.patch( - "controllers.service_api.end_user.end_user.EndUserService.get_end_user_by_id", return_value=end_user + context = _request_context(tenant_id=end_user.tenant_id, app_id=end_user.app_id) + service = EndUserQueryServiceStub(end_user) + monkeypatch.setattr( + controller_module, + "application_services", + lambda: ApplicationServicesStub(service), ) - result = EndUserApi.get.__wrapped__(resource, app_model=app_model, end_user_id=UUID(end_user.id)) - - get_end_user_by_id.assert_called_once_with( - tenant_id=app_model.tenant_id, app_id=app_model.id, end_user_id=end_user.id + result = unwrap(EndUserApi.get)( + EndUserApi(), + app_model=App(id=context.app_id, tenant_id=context.tenant_id), + end_user_id=UUID(end_user.id), ) - assert result["id"] == end_user.id - assert result["tenant_id"] == end_user.tenant_id - assert result["app_id"] == end_user.app_id - assert result["type"] == end_user.type - assert result["external_user_id"] == end_user.external_user_id - assert result["name"] == end_user.name - assert result["is_anonymous"] is True - assert result["session_id"] == end_user.session_id - assert result["created_at"].startswith("2024-01-01T00:00:00") - assert result["updated_at"].startswith("2024-01-02T00:00:00") - def test_get_end_user_not_found(self, mocker: MockerFixture, resource: EndUserApi, app_model: App) -> None: - mocker.patch("controllers.service_api.end_user.end_user.EndUserService.get_end_user_by_id", return_value=None) + assert service.calls == [(context, end_user.id)] + assert result == { + "id": end_user.id, + "tenant_id": end_user.tenant_id, + "app_id": end_user.app_id, + "type": end_user.type, + "external_user_id": end_user.external_user_id, + "name": end_user.name, + "is_anonymous": True, + "session_id": end_user.session_id, + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-02T00:00:00Z", + } + + def test_get_end_user_maps_application_not_found_error(self, monkeypatch: pytest.MonkeyPatch) -> None: + context = _request_context() + service = EndUserQueryServiceStub(AppScopedEndUserNotFoundError()) + monkeypatch.setattr( + controller_module, + "application_services", + lambda: ApplicationServicesStub(service), + ) + end_user_id = uuid4() with pytest.raises(EndUserNotFoundError): - EndUserApi.get.__wrapped__(resource, app_model=app_model, end_user_id=uuid4()) + unwrap(EndUserApi.get)( + EndUserApi(), + app_model=App(id=context.app_id, tenant_id=context.tenant_id), + end_user_id=end_user_id, + ) + + assert service.calls == [(context, str(end_user_id))] diff --git a/api/tests/unit_tests/controllers/service_api/test_wraps.py b/api/tests/unit_tests/controllers/service_api/test_wraps.py index 997017fddaa..9fe59a388da 100644 --- a/api/tests/unit_tests/controllers/service_api/test_wraps.py +++ b/api/tests/unit_tests/controllers/service_api/test_wraps.py @@ -28,11 +28,47 @@ from enums import CloudPlan, DeploymentEdition from models import Account, Tenant, TenantAccountJoin from models.account import TenantAccountRole from models.dataset import Dataset, RateLimitLog -from models.enums import ApiTokenType -from models.model import ApiToken, App, AppMode, DatasetApiTokenBinding, IconType +from models.enums import ApiTokenType, EndUserType +from models.model import ApiToken, App, AppMode, DatasetApiTokenBinding, EndUser, IconType from tests.unit_tests.config_override import config_overrides_context +class _RecordingEndUserCommands: + def __init__(self, result: EndUser) -> None: + self._result = result + self.calls: list[tuple[str, str, str | None]] = [] + + def get_or_create_end_user(self, tenant_id: str, app_id: str, user_id: str | None = None) -> EndUser: + self.calls.append((tenant_id, app_id, user_id)) + return self._result + + +class _AppScopedEndUserServicesStub: + def __init__(self, commands: _RecordingEndUserCommands) -> None: + self.commands = commands + + +class _ApplicationServicesStub: + def __init__(self, commands: _RecordingEndUserCommands) -> None: + self.app_scoped_end_users = _AppScopedEndUserServicesStub(commands) + + +class _RecordingLoginManager: + def __init__(self) -> None: + self.users: list[EndUser] = [] + + def _update_request_context_with_user(self, user: EndUser) -> None: + self.users.append(user) + + +class _RecordingSignal: + def __init__(self) -> None: + self.calls: list[tuple[object, EndUser]] = [] + + def send(self, sender: object, *, user: EndUser) -> None: + self.calls.append((sender, user)) + + def _configure_current_app_mock(mock_current_app): mock_current_app.login_manager = Mock() mock_current_app._get_current_object = Mock(return_value=Mock()) @@ -223,6 +259,82 @@ class TestValidateAppToken: assert result["app_id"] == app_model.id assert account.current_tenant_id == tenant.id + @pytest.mark.parametrize( + ("fetch_from", "method", "path", "request_kwargs", "expected_user_id"), + [ + pytest.param(WhereisUserArg.QUERY, "GET", "/?user=query-user", {}, "query-user", id="query"), + pytest.param( + WhereisUserArg.JSON, + "POST", + "/", + {"json": {"user": "json-user"}}, + "json-user", + id="json", + ), + pytest.param( + WhereisUserArg.FORM, + "POST", + "/", + {"data": {"user": "form-user"}}, + "form-user", + id="form", + ), + ], + ) + @patch("controllers.service_api.wraps.validate_and_get_api_token") + @pytest.mark.parametrize("sqlite_session", [(App, Tenant)], indirect=True) + def test_fetch_user_arg_resolves_and_injects_end_user( + self, + mock_validate_token, + app: Flask, + sqlite_session: Session, + monkeypatch: pytest.MonkeyPatch, + fetch_from: WhereisUserArg, + method: str, + path: str, + request_kwargs: dict[str, object], + expected_user_id: str, + ) -> None: + tenant = Tenant(name="Workspace") + app_model = _app_model(tenant_id=tenant.id) + sqlite_session.add_all([tenant, app_model]) + sqlite_session.commit() + mock_validate_token.return_value = _api_token( + tenant_id=tenant.id, + app_id=app_model.id, + token_type=ApiTokenType.APP, + ) + + end_user = EndUser( + id=str(uuid.uuid4()), + tenant_id=tenant.id, + app_id=app_model.id, + type=EndUserType.SERVICE_API, + session_id=expected_user_id, + ) + commands = _RecordingEndUserCommands(end_user) + login_manager = _RecordingLoginManager() + signal = _RecordingSignal() + app.login_manager = login_manager # type: ignore[attr-defined] + monkeypatch.setattr(wraps_module, "application_services", lambda: _ApplicationServicesStub(commands)) + monkeypatch.setattr(wraps_module, "user_logged_in", signal) + + @validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=fetch_from, required=True)) + def protected_view(*, app_model: App, end_user: EndUser) -> tuple[App, EndUser]: + return app_model, end_user + + with ( + app.test_request_context(path, method=method, **request_kwargs), + patch("controllers.service_api.wraps.db.session", _session_proxy(sqlite_session)), + ): + injected_app, injected_end_user = protected_view() + + assert injected_app is app_model + assert injected_end_user is end_user + assert commands.calls == [(tenant.id, app_model.id, expected_user_id)] + assert login_manager.users == [end_user] + assert signal.calls == [(app, end_user)] + @patch("controllers.service_api.wraps.validate_and_get_api_token") @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_app_not_found_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session): diff --git a/api/tests/unit_tests/controllers/trigger/test_webhook.py b/api/tests/unit_tests/controllers/trigger/test_webhook.py index 087922554ae..ac1cfaf1ebb 100644 --- a/api/tests/unit_tests/controllers/trigger/test_webhook.py +++ b/api/tests/unit_tests/controllers/trigger/test_webhook.py @@ -1,4 +1,3 @@ -import types from unittest.mock import patch import pytest @@ -11,13 +10,25 @@ from services.errors.app import QuotaExceededError from tests.unit_tests.model_factories import make_workflow +class _RequestStub: + method = "POST" + headers = {"x-test": "1"} + args = {"a": "b"} + + +class _AppScopedEndUserServicesStub: + def __init__(self, commands: object) -> None: + self.commands = commands + + +class _ApplicationServicesStub: + def __init__(self, commands: object) -> None: + self.app_scoped_end_users = _AppScopedEndUserServicesStub(commands) + + @pytest.fixture(autouse=True) def mock_request(): - module.request = types.SimpleNamespace( - method="POST", - headers={"x-test": "1"}, - args={"a": "b"}, - ) + module.request = _RequestStub() @pytest.fixture(autouse=True) @@ -25,6 +36,13 @@ def mock_jsonify(): module.jsonify = lambda payload: payload +@pytest.fixture(autouse=True) +def end_user_commands(monkeypatch: pytest.MonkeyPatch) -> object: + commands = object() + monkeypatch.setattr(module, "application_services", lambda: _ApplicationServicesStub(commands)) + return commands + + def _webhook_trigger() -> WorkflowWebhookTrigger: return WorkflowWebhookTrigger( webhook_id="wh-1", @@ -76,6 +94,7 @@ class TestHandleWebhook: mock_trigger, mock_extract, mock_get, + end_user_commands, ): mock_get.return_value = (_webhook_trigger(), _workflow(), "node_config") mock_extract.return_value = {"input": "x"} @@ -86,6 +105,7 @@ class TestHandleWebhook: assert status == 200 assert response["ok"] is True mock_trigger.assert_called_once() + assert mock_trigger.call_args.kwargs["end_users"] is end_user_commands @patch.object(module.WebhookService, "get_webhook_trigger_and_workflow") @patch.object(module.WebhookService, "extract_and_validate_webhook_data", side_effect=ValueError("bad")) diff --git a/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py b/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py index adfd4084b3e..3734cf1f30b 100644 --- a/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py +++ b/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py @@ -26,6 +26,23 @@ class _DatabaseWithEngine: self.engine = engine +class _EndUserProvisioner: + def __init__(self) -> None: + self.result: EndUser | None = None + self.calls: list[tuple[str, str, str | None]] = [] + + def get_or_create_end_user( + self, + tenant_id: str, + app_id: str, + user_id: str | None = None, + ) -> EndUser: + self.calls.append((tenant_id, app_id, user_id)) + if self.result is None: + raise AssertionError("unexpected end-user provisioning") + return self.result + + def _app( *, app_id: str = "app-1", @@ -104,6 +121,7 @@ class TestPluginAppBackwardsInvocation: self.session = sqlite_session self.session_factory = sqlite_session_factory self.sqlite_engine = sqlite_engine + self.end_users = _EndUserProvisioner() mocker.patch("core.plugin.backwards_invocation.app.create_session", side_effect=sqlite_session_factory) def test_fetch_app_info_workflow_path(self, mocker: MockerFixture): @@ -167,6 +185,7 @@ class TestPluginAppBackwardsInvocation: inputs={"x": 1}, files=[], session=self.session, + end_users=self.end_users, ) assert result == {"routed": True} @@ -178,10 +197,7 @@ class TestPluginAppBackwardsInvocation: workflow = _workflow() mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow) - get_or_create = mocker.patch( - "core.plugin.backwards_invocation.app.EndUserService.get_or_create_end_user", - return_value=end_user, - ) + self.end_users.result = end_user route = mocker.patch.object(PluginAppBackwardsInvocation, "invoke_workflow_app", return_value={"ok": True}) result = PluginAppBackwardsInvocation.invoke_app( @@ -194,10 +210,11 @@ class TestPluginAppBackwardsInvocation: inputs={}, files=[], session=self.session, + end_users=self.end_users, ) assert result == {"ok": True} - get_or_create.assert_called_once_with(app) + assert self.end_users.calls == [(app.tenant_id, app.id, None)] assert route.call_args.args[1] is workflow assert route.call_args.args[2] is end_user @@ -216,6 +233,7 @@ class TestPluginAppBackwardsInvocation: inputs={}, files=[], session=self.session, + end_users=self.end_users, ) def test_invoke_app_unexpected_mode_raises(self, mocker: MockerFixture): @@ -237,6 +255,7 @@ class TestPluginAppBackwardsInvocation: inputs={}, files=[], session=self.session, + end_users=self.end_users, ) @pytest.mark.parametrize( @@ -374,6 +393,7 @@ class TestPluginAppBackwardsInvocation: inputs={}, files=[], session=self.session, + end_users=self.end_users, ) def test_invoke_completion_app(self, mocker: MockerFixture): @@ -510,10 +530,7 @@ class TestPluginAppBackwardsInvocation: mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow) mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", side_effect=ValueError("user not found")) - get_or_create = mocker.patch( - "core.plugin.backwards_invocation.app.EndUserService.get_or_create_end_user", - return_value=end_user, - ) + self.end_users.result = end_user route = mocker.patch.object(PluginAppBackwardsInvocation, "invoke_workflow_app", return_value={"ok": True}) result = PluginAppBackwardsInvocation.invoke_app( @@ -526,10 +543,11 @@ class TestPluginAppBackwardsInvocation: inputs={}, files=[], session=self.session, + end_users=self.end_users, ) assert result == {"ok": True} - get_or_create.assert_called_once_with(app, user_id="wecom-sender-1") + assert self.end_users.calls == [(app.tenant_id, app.id, "wecom-sender-1")] assert route.call_args.args[2] is end_user def test_get_app_returns_app(self): diff --git a/api/tests/unit_tests/extensions/test_ext_application_services.py b/api/tests/unit_tests/extensions/test_ext_application_services.py index 0e502ec7a9f..52dc3184035 100644 --- a/api/tests/unit_tests/extensions/test_ext_application_services.py +++ b/api/tests/unit_tests/extensions/test_ext_application_services.py @@ -31,6 +31,7 @@ from repositories.account_oauth_repository import ( RegisterServiceOAuthInvitationGateway, ) from repositories.account_repository import SQLAlchemyAccountRepository +from repositories.app_scoped_end_user_repository import AppScopedEndUserRepo from repositories.app_site_command_repository import AppSiteCommandRepository from repositories.app_statistic_query_repository import AppStatisticQueryRepository from repositories.app_tracing_config_repository import SQLAlchemyAppTracingConfigRepository @@ -67,6 +68,8 @@ from services.account_oauth_adapters import ( ) from services.app_generate_service import AppGenerateService from services.app_preview_query_service import AppPreviewRef, AppPreviewUnavailableError +from services.app_scoped_end_user_query_service import AppScopedEndUserQueryService +from services.app_scoped_end_user_service import AppScopedEndUserService from services.app_site_service import AppSiteService from services.app_tracing_config_gateway import OpsTraceManagerGateway from services.app_tracing_config_service import AppTracingConfigService @@ -163,6 +166,12 @@ def test_init_app_registers_services_for_the_current_app( services = ext_application_services.application_services() assert services is app.extensions["application_services"] assert services.init_validation.is_validated(session_validated=False) is False + assert isinstance(services.app_scoped_end_users.commands, AppScopedEndUserService) + assert isinstance(services.app_scoped_end_users.queries, AppScopedEndUserQueryService) + repository = services.app_scoped_end_users.queries._app_scoped_end_users + assert isinstance(repository, AppScopedEndUserRepo) + assert services.app_scoped_end_users.commands._app_scoped_end_users is repository + assert repository._session_factory is sqlite_session_factory assert isinstance(services.workflow_statistics, WorkflowStatisticQueryService) diff --git a/api/tests/unit_tests/machinery/test_service_api_context.py b/api/tests/unit_tests/machinery/test_service_api_context.py new file mode 100644 index 00000000000..dc17aa7890c --- /dev/null +++ b/api/tests/unit_tests/machinery/test_service_api_context.py @@ -0,0 +1,12 @@ +from machinery.context import ServiceApiRequestContext + + +def test_service_api_request_context_contains_only_app_scope() -> None: + context = ServiceApiRequestContext( + tenant_id="tenant-1", + app_id="app-1", + ) + + assert context.tenant_id == "tenant-1" + assert context.app_id == "app-1" + assert not hasattr(context, "account_id") diff --git a/api/tests/unit_tests/models/test_end_user_type.py b/api/tests/unit_tests/models/test_end_user_type.py index d31852caccb..313a628a4b5 100644 --- a/api/tests/unit_tests/models/test_end_user_type.py +++ b/api/tests/unit_tests/models/test_end_user_type.py @@ -7,7 +7,7 @@ import pytest from models.enums import EndUserType from models.model import EndUser from models.types import EnumText -from services.end_user_service import EndUserService +from services.app_scoped_end_user_service import AppScopedEndUserService API_ROOT = Path(__file__).resolve().parents[3] @@ -40,13 +40,15 @@ def test_end_user_type_still_rejects_unknown_values(): def test_end_user_service_creation_methods_accept_end_user_type(): - assert inspect.signature(EndUserService.get_or_create_end_user_by_type).parameters["type"].annotation is EndUserType - assert inspect.signature(EndUserService.create_end_user_batch).parameters["type"].annotation is EndUserType + get_or_create_type = inspect.signature(AppScopedEndUserService.get_or_create_end_user_by_type).parameters["type"] + assert get_or_create_type.annotation is EndUserType + assert inspect.signature(AppScopedEndUserService.create_end_user_batch).parameters["type"].annotation is EndUserType def test_end_user_service_callers_pass_end_user_type(): violations: list[str] = [] method_names = {"get_or_create_end_user_by_type", "create_end_user_batch"} + checked_calls = 0 for source_path in API_ROOT.rglob("*.py"): if "tests" in source_path.parts or ".venv" in source_path.parts: @@ -58,8 +60,7 @@ def test_end_user_service_callers_pass_end_user_type(): continue if not isinstance(node.func, ast.Attribute) or node.func.attr not in method_names: continue - if not isinstance(node.func.value, ast.Name) or node.func.value.id != "EndUserService": - continue + checked_calls += 1 type_arg = next((keyword.value for keyword in node.keywords if keyword.arg == "type"), None) if type_arg is None and node.args: @@ -72,6 +73,7 @@ def test_end_user_service_callers_pass_end_user_type(): ): violations.append(f"{source_path.relative_to(API_ROOT)}:{node.lineno}") + assert checked_calls > 0 assert violations == [] @@ -105,12 +107,10 @@ def test_production_end_user_constructors_use_end_user_type_enum(): and isinstance(value.value, ast.Name) and value.value.id == "EndUserType" ) - uses_end_user_service_type_parameter = ( - source_path.relative_to(API_ROOT) == Path("services/end_user_service.py") - and isinstance(value, ast.Name) - and value.id == "type" + uses_end_user_type_conversion = ( + isinstance(value, ast.Call) and isinstance(value.func, ast.Name) and value.func.id == "EndUserType" ) - if not (uses_end_user_type_member or uses_end_user_service_type_parameter): + if not (uses_end_user_type_member or uses_end_user_type_conversion): violations.append(f"{source_path.relative_to(API_ROOT)}:{node.lineno}") assert violations == [] diff --git a/api/tests/unit_tests/repositories/test_end_user_repository.py b/api/tests/unit_tests/repositories/test_end_user_repository.py new file mode 100644 index 00000000000..c990a886d38 --- /dev/null +++ b/api/tests/unit_tests/repositories/test_end_user_repository.py @@ -0,0 +1,147 @@ +from datetime import datetime + +from sqlalchemy.orm import Session, sessionmaker + +from models.enums import EndUserType +from models.model import EndUser +from repositories.app_scoped_end_user_repository import AppScopedEndUserRepo +from services.app_scoped_end_user_service import AppScopedEndUserService +from services.entities.app_scoped_end_user_entities import AppScopedEndUserRecord, NewAppScopedEndUser + +_END_USER_ID = "11111111-1111-1111-1111-111111111111" +_TENANT_ID = "22222222-2222-2222-2222-222222222222" +_APP_ID = "33333333-3333-3333-3333-333333333333" + + +def _persist_end_user(session: Session) -> None: + timestamp = datetime(2026, 1, 1) + session.add( + EndUser( + id=_END_USER_ID, + tenant_id=_TENANT_ID, + app_id=_APP_ID, + type=EndUserType.SERVICE_API, + external_user_id="external-1", + name="Alice", + is_anonymous=True, + session_id="session-1", + created_at=timestamp, + updated_at=timestamp, + ) + ) + session.commit() + + +def test_find_by_id_returns_detached_read_contract( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + _persist_end_user(sqlite_session) + + result = AppScopedEndUserRepo(session_factory=sqlite_session_factory).find_by_id( + tenant_id=_TENANT_ID, + app_id=_APP_ID, + end_user_id=_END_USER_ID, + ) + + assert result == AppScopedEndUserRecord( + id=_END_USER_ID, + tenant_id=_TENANT_ID, + app_id=_APP_ID, + type=EndUserType.SERVICE_API.value, + external_user_id="external-1", + name="Alice", + is_anonymous=True, + session_id="session-1", + created_at=datetime(2026, 1, 1), + updated_at=datetime(2026, 1, 1), + ) + + +def test_find_by_id_scopes_reads_to_tenant_and_app( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + _persist_end_user(sqlite_session) + repository = AppScopedEndUserRepo(session_factory=sqlite_session_factory) + + assert ( + repository.find_by_id( + tenant_id="44444444-4444-4444-4444-444444444444", + app_id=_APP_ID, + end_user_id=_END_USER_ID, + ) + is None + ) + assert ( + repository.find_by_id( + tenant_id=_TENANT_ID, + app_id="55555555-5555-5555-5555-555555555555", + end_user_id=_END_USER_ID, + ) + is None + ) + + +def test_create_flushes_before_returning_stored_metadata( + sqlite_session_factory: sessionmaker[Session], +) -> None: + repository = AppScopedEndUserRepo(session_factory=sqlite_session_factory) + + stored = repository.create( + NewAppScopedEndUser( + tenant_id=_TENANT_ID, + app_id=_APP_ID, + type=EndUserType.SERVICE_API.value, + is_anonymous=False, + session_id="session-1", + external_user_id="session-1", + ) + ) + + assert stored.id + assert stored.value.id == stored.id + + +def test_get_or_create_reuses_and_upgrades_an_existing_end_user( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + _persist_end_user(sqlite_session) + repository = AppScopedEndUserRepo(session_factory=sqlite_session_factory) + + result = AppScopedEndUserService(end_users=repository).get_or_create_end_user_by_type( + EndUserType.OPENAPI, + tenant_id=_TENANT_ID, + app_id=_APP_ID, + user_id="session-1", + ) + + assert result.id == _END_USER_ID + assert result.type == EndUserType.OPENAPI + with sqlite_session_factory() as observer: + persisted = observer.get(EndUser, _END_USER_ID) + assert persisted is not None + assert persisted.type == EndUserType.OPENAPI + + +def test_get_or_create_batch_reuses_existing_users_and_creates_missing_ones( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + _persist_end_user(sqlite_session) + repository = AppScopedEndUserRepo(session_factory=sqlite_session_factory) + second_app_id = "66666666-6666-6666-6666-666666666666" + + result = AppScopedEndUserService(end_users=repository).create_end_user_batch( + EndUserType.SERVICE_API, + tenant_id=_TENANT_ID, + app_ids=[_APP_ID, second_app_id], + user_id="session-1", + ) + + assert result[_APP_ID].id == _END_USER_ID + assert result[second_app_id].id + assert result[second_app_id].app_id == second_app_id + assert result[second_app_id].external_user_id == "session-1" + assert result[second_app_id]._is_anonymous is False diff --git a/api/tests/unit_tests/services/test_end_user_query_service.py b/api/tests/unit_tests/services/test_end_user_query_service.py new file mode 100644 index 00000000000..2d2b256a2ab --- /dev/null +++ b/api/tests/unit_tests/services/test_end_user_query_service.py @@ -0,0 +1,61 @@ +from datetime import datetime + +import pytest + +from machinery.context import ServiceApiRequestContext +from services.app_scoped_end_user_query_service import AppScopedEndUserNotFoundError, AppScopedEndUserQueryService +from services.entities.app_scoped_end_user_entities import AppScopedEndUserRecord + + +def _context() -> ServiceApiRequestContext: + return ServiceApiRequestContext( + tenant_id="workspace-1", + app_id="app-1", + ) + + +def _record() -> AppScopedEndUserRecord: + timestamp = datetime(2026, 1, 1) + return AppScopedEndUserRecord( + id="end-user-1", + tenant_id="workspace-1", + app_id="app-1", + type="service-api", + external_user_id="external-1", + name="Alice", + is_anonymous=False, + session_id="session-1", + created_at=timestamp, + updated_at=timestamp, + ) + + +class RecordingEndUserQuery: + def __init__(self, result: AppScopedEndUserRecord | None) -> None: + self._result = result + self.calls: list[tuple[str, str, str]] = [] + + def find_by_id(self, *, tenant_id: str, app_id: str, end_user_id: str) -> AppScopedEndUserRecord | None: + self.calls.append((tenant_id, app_id, end_user_id)) + return self._result + + +def test_get_scopes_query_to_admitted_workspace_and_app() -> None: + record = _record() + query = RecordingEndUserQuery(record) + service = AppScopedEndUserQueryService(end_users=query) + + result = service.get_by_id(_context(), "end-user-1") + + assert result == record + assert query.calls == [("workspace-1", "app-1", "end-user-1")] + + +def test_get_raises_framework_neutral_not_found_error() -> None: + query = RecordingEndUserQuery(None) + service = AppScopedEndUserQueryService(end_users=query) + + with pytest.raises(AppScopedEndUserNotFoundError): + service.get_by_id(_context(), "missing") + + assert query.calls == [("workspace-1", "app-1", "missing")] diff --git a/api/tests/unit_tests/services/test_end_user_service.py b/api/tests/unit_tests/services/test_end_user_service.py new file mode 100644 index 00000000000..f97c65638c6 --- /dev/null +++ b/api/tests/unit_tests/services/test_end_user_service.py @@ -0,0 +1,194 @@ +from collections.abc import Sequence +from typing import override + +from models.enums import DEFAULT_END_USER_SESSION_ID, EndUserType +from services.app_scoped_end_user_service import AppScopedEndUserRepository, AppScopedEndUserService +from services.entities.app_scoped_end_user_entities import NewAppScopedEndUser, StoredAppScopedEndUser + + +class RecordingEndUserRepository(AppScopedEndUserRepository[str]): + def __init__( + self, + *, + session_candidates: Sequence[StoredAppScopedEndUser[str]] = (), + app_candidates: Sequence[StoredAppScopedEndUser[str]] = (), + ) -> None: + self.session_candidates = list(session_candidates) + self.app_candidates = list(app_candidates) + self.find_session_calls: list[tuple[str, str, str]] = [] + self.find_apps_calls: list[tuple[str, list[str], str, str]] = [] + self.created: list[NewAppScopedEndUser] = [] + self.updated_types: list[tuple[str, str]] = [] + + @override + def find_by_session( + self, + *, + tenant_id: str, + app_id: str, + user_id: str, + ) -> Sequence[StoredAppScopedEndUser[str]]: + self.find_session_calls.append((tenant_id, app_id, user_id)) + return self.session_candidates + + @override + def find_by_apps( + self, + *, + tenant_id: str, + app_ids: Sequence[str], + user_id: str, + type: str, + ) -> Sequence[StoredAppScopedEndUser[str]]: + self.find_apps_calls.append((tenant_id, list(app_ids), user_id, type)) + return self.app_candidates + + @override + def create(self, command: NewAppScopedEndUser) -> StoredAppScopedEndUser[str]: + self.created.append(command) + return StoredAppScopedEndUser( + id=f"new-{command.app_id}", + app_id=command.app_id, + type=command.type, + value="created", + ) + + @override + def create_batch( + self, + commands: Sequence[NewAppScopedEndUser], + ) -> Sequence[StoredAppScopedEndUser[str]]: + self.created.extend(commands) + return [ + StoredAppScopedEndUser( + id=f"new-{command.app_id}", + app_id=command.app_id, + type=command.type, + value=command.app_id, + ) + for command in commands + ] + + @override + def update_type(self, end_user_id: str, type: str) -> StoredAppScopedEndUser[str]: + self.updated_types.append((end_user_id, type)) + candidate = next(candidate for candidate in self.session_candidates if candidate.id == end_user_id) + return StoredAppScopedEndUser(id=candidate.id, app_id=candidate.app_id, type=type, value=candidate.value) + + +def _stored(*, id: str, app_id: str = "app-1", type: EndUserType, value: str) -> StoredAppScopedEndUser[str]: + return StoredAppScopedEndUser(id=id, app_id=app_id, type=type.value, value=value) + + +def _service(repository: RecordingEndUserRepository) -> AppScopedEndUserService[str]: + return AppScopedEndUserService(end_users=repository) + + +def test_get_or_create_prioritizes_matching_type_in_service() -> None: + repository = RecordingEndUserRepository( + session_candidates=[ + _stored(id="legacy", type=EndUserType.BROWSER, value="legacy"), + _stored(id="matching", type=EndUserType.OPENAPI, value="matching"), + ] + ) + + result = _service(repository).get_or_create_end_user_by_type( + EndUserType.OPENAPI, + tenant_id="tenant-1", + app_id="app-1", + user_id="user-1", + ) + + assert result == "matching" + assert repository.updated_types == [] + assert repository.created == [] + + +def test_get_or_create_upgrades_legacy_type() -> None: + repository = RecordingEndUserRepository( + session_candidates=[_stored(id="legacy", type=EndUserType.BROWSER, value="legacy")] + ) + + result = _service(repository).get_or_create_end_user_by_type( + EndUserType.SERVICE_API, + tenant_id="tenant-1", + app_id="app-1", + user_id="user-1", + ) + + assert result == "legacy" + assert repository.updated_types == [("legacy", EndUserType.SERVICE_API.value)] + + +def test_get_or_create_never_retypes_an_app_deploy_end_user() -> None: + repository = RecordingEndUserRepository( + session_candidates=[_stored(id="app-deploy", type=EndUserType.APP_DEPLOY, value="app-deploy")] + ) + + result = _service(repository).get_or_create_end_user_by_type( + EndUserType.SERVICE_API, + tenant_id="tenant-1", + app_id="app-1", + user_id="user-1", + ) + + assert result == "created" + assert repository.updated_types == [] + assert repository.created == [ + NewAppScopedEndUser( + tenant_id="tenant-1", + app_id="app-1", + type=EndUserType.SERVICE_API.value, + is_anonymous=False, + session_id="user-1", + external_user_id="user-1", + ) + ] + + +def test_get_or_create_builds_anonymous_creation_command() -> None: + repository = RecordingEndUserRepository() + + result = _service(repository).get_or_create_end_user( + tenant_id="tenant-1", + app_id="app-1", + ) + + assert result == "created" + assert repository.find_session_calls == [("tenant-1", "app-1", DEFAULT_END_USER_SESSION_ID)] + assert repository.created == [ + NewAppScopedEndUser( + tenant_id="tenant-1", + app_id="app-1", + type=EndUserType.SERVICE_API.value, + is_anonymous=True, + session_id=DEFAULT_END_USER_SESSION_ID, + external_user_id=DEFAULT_END_USER_SESSION_ID, + ) + ] + + +def test_create_batch_deduplicates_apps_and_creates_only_missing_users() -> None: + repository = RecordingEndUserRepository( + app_candidates=[_stored(id="existing", app_id="app-1", type=EndUserType.TRIGGER, value="existing")] + ) + + result = _service(repository).create_end_user_batch( + EndUserType.TRIGGER, + tenant_id="tenant-1", + app_ids=["app-1", "app-2", "app-1"], + user_id="user-1", + ) + + assert result == {"app-1": "existing", "app-2": "app-2"} + assert repository.find_apps_calls == [("tenant-1", ["app-1", "app-2"], "user-1", EndUserType.TRIGGER.value)] + assert repository.created == [ + NewAppScopedEndUser( + tenant_id="tenant-1", + app_id="app-2", + type=EndUserType.TRIGGER.value, + is_anonymous=False, + session_id="user-1", + external_user_id="user-1", + ) + ] diff --git a/api/tests/unit_tests/services/test_webhook_service.py b/api/tests/unit_tests/services/test_webhook_service.py index 301997252cf..8d3257d4905 100644 --- a/api/tests/unit_tests/services/test_webhook_service.py +++ b/api/tests/unit_tests/services/test_webhook_service.py @@ -21,6 +21,24 @@ from services.trigger.webhook_service import WebhookService from tests.unit_tests.model_factories import make_app, make_end_user, make_workflow +class _EndUserServiceStub: + def __init__(self, result: EndUser | Exception) -> None: + self._result = result + self.calls: list[tuple[EndUserType, str, str, str | None]] = [] + + def get_or_create_end_user_by_type( + self, + type: EndUserType, + tenant_id: str, + app_id: str, + user_id: str | None = None, + ) -> EndUser: + self.calls.append((type, tenant_id, app_id, user_id)) + if isinstance(self._result, Exception): + raise self._result + return self._result + + def _webhook_trigger( *, webhook_id: str = "webhook-123", @@ -229,10 +247,6 @@ class TestWebhookServiceUnit: caplog.set_level(logging.INFO) with ( - patch( - "services.trigger.webhook_service.EndUserService.get_or_create_end_user_by_type", - return_value=_end_user(), - ), patch("services.trigger.webhook_service.QuotaService.reserve", return_value=quota_charge), patch( "services.trigger.webhook_service.AsyncWorkflowService.trigger_workflow_async", @@ -244,6 +258,7 @@ class TestWebhookServiceUnit: webhook_trigger, {"body": {}, "headers": {}, "query_params": {}, "files": {}, "method": "POST"}, workflow, + end_users=_EndUserServiceStub(_end_user()), ) assert exc_info.value is quota_error @@ -276,18 +291,18 @@ class TestWebhookServiceUnit: } with ( - patch.object( - webhook_service_module.EndUserService, - "get_or_create_end_user_by_type", - return_value=end_user, - ), patch.object(webhook_service_module.QuotaService, "reserve", return_value=quota_charge), patch.object( webhook_service_module.AsyncWorkflowService, "trigger_workflow_async", ) as mock_trigger, ): - WebhookService.trigger_workflow_execution(webhook_trigger, webhook_data, workflow) + WebhookService.trigger_workflow_execution( + webhook_trigger, + webhook_data, + workflow, + end_users=_EndUserServiceStub(end_user), + ) call_session = mock_trigger.call_args.kwargs["session"] assert call_session.get_bind() is sqlite_engine @@ -299,13 +314,13 @@ class TestWebhookServiceUnit: workflow = _workflow() webhook_data = {"method": "POST", "headers": {}, "query_params": {}, "body": {}, "files": {}} - with patch.object( - webhook_service_module.EndUserService, - "get_or_create_end_user_by_type", - side_effect=ValueError("Failed to create end user"), - ): - with pytest.raises(ValueError, match="Failed to create end user"): - WebhookService.trigger_workflow_execution(webhook_trigger, webhook_data, workflow) + with pytest.raises(ValueError, match="Failed to create end user"): + WebhookService.trigger_workflow_execution( + webhook_trigger, + webhook_data, + workflow, + end_users=_EndUserServiceStub(ValueError("Failed to create end user")), + ) def test_extract_webhook_data_json(self): """Test webhook data extraction from JSON request.""" diff --git a/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py b/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py index ecb320b5a9d..2e96619116f 100644 --- a/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py +++ b/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from unittest.mock import MagicMock, patch import pytest @@ -42,6 +43,22 @@ def _end_user() -> EndUser: ) +class _EndUserServiceStub: + def __init__(self) -> None: + self.result: Mapping[str, EndUser] = {} + self.calls: list[tuple[EndUserType, str, list[str], str]] = [] + + def create_end_user_batch( + self, + type: EndUserType, + tenant_id: str, + app_ids: list[str], + user_id: str, + ) -> Mapping[str, EndUser]: + self.calls.append((type, tenant_id, app_ids, user_id)) + return self.result + + class TestDispatchTriggeredWorkflow: """Unit tests covering branch behaviours of ``dispatch_triggered_workflow``. @@ -97,7 +114,7 @@ class TestDispatchTriggeredWorkflow: Defaults are configured so the code flow can reach the final async trigger block (line ~385); each test overrides specific handles - (``get_workflows``, ``reserve``, ``create_end_user_batch``, ...) to + (``get_workflows``, ``reserve``, ``end_users``, ...) to drive the path it targets. """ invoke_response = MagicMock() @@ -105,6 +122,7 @@ class TestDispatchTriggeredWorkflow: invoke_response.variables = {} quota_charge = MagicMock() + end_users = _EndUserServiceStub() with ( patch.object( @@ -141,11 +159,6 @@ class TestDispatchTriggeredWorkflow: trigger_processing_tasks_module, "_get_published_workflows_by_app_ids", ) as get_workflows, - patch.object( - trigger_processing_tasks_module.EndUserService, - "create_end_user_batch", - return_value={}, - ) as create_end_user_batch, patch.object( trigger_processing_tasks_module.QuotaService, "reserve", @@ -167,7 +180,7 @@ class TestDispatchTriggeredWorkflow: "mark_rate_limited": mark_rate_limited, "invoke_trigger_event": invoke_trigger_event, "invoke_response": invoke_response, - "create_end_user_batch": create_end_user_batch, + "end_users": end_users, "trigger_workflow_async": trigger_workflow_async, } @@ -180,6 +193,7 @@ class TestDispatchTriggeredWorkflow: subscription=subscription, event_name="test_event", request_id="request-123", + end_users=dispatch_mocks["end_users"], ) assert dispatched == 0 @@ -200,6 +214,7 @@ class TestDispatchTriggeredWorkflow: subscription=subscription, event_name="test_event", request_id="request-123", + end_users=dispatch_mocks["end_users"], ) assert dispatched == 0 @@ -215,13 +230,14 @@ class TestDispatchTriggeredWorkflow: dispatch_mocks["get_workflows"].return_value = {plugin_trigger.app_id: workflow} end_user = _end_user() - dispatch_mocks["create_end_user_batch"].return_value = {plugin_trigger.app_id: end_user} + dispatch_mocks["end_users"].result = {plugin_trigger.app_id: end_user} dispatched = dispatch_triggered_workflow( user_id="user-123", subscription=subscription, event_name="test_event", request_id="request-123", + end_users=dispatch_mocks["end_users"], ) assert dispatched == 1 diff --git a/packages/contracts/generated/api/service/zod.gen.ts b/packages/contracts/generated/api/service/zod.gen.ts index c04a4650041..9b2bcb161b9 100644 --- a/packages/contracts/generated/api/service/zod.gen.ts +++ b/packages/contracts/generated/api/service/zod.gen.ts @@ -1060,11 +1060,7 @@ export const zDocumentStatusListResponse = z.object({ /** * EndUserDetail * - * Full EndUser record for API responses. - * - * Note: The SQLAlchemy model defines an `is_anonymous` property for Flask-Login semantics - * (always False). The database column is exposed as `_is_anonymous`, so this DTO maps - * `is_anonymous` from `_is_anonymous` to return the stored value. + * Full end-user detail returned by the Service API. */ export const zEndUserDetail = z.object({ app_id: z.uuid().nullish(),