mirror of
https://github.com/langgenius/dify.git
synced 2026-09-28 06:13:22 +08:00
refactor(api): move app-scoped end users behind application services (#41398)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 ""
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -216,6 +216,9 @@ class EndUserType(StrEnum):
|
||||
TRIGGER = "trigger"
|
||||
|
||||
|
||||
DEFAULT_END_USER_SESSION_ID = "DEFAULT-USER"
|
||||
|
||||
|
||||
class DocumentDocType(StrEnum):
|
||||
"""Document doc_type classification"""
|
||||
|
||||
|
||||
@@ -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__ = (
|
||||
|
||||
@@ -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 |
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
+45
-23
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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))]
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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")
|
||||
@@ -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 == []
|
||||
|
||||
@@ -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
|
||||
@@ -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")]
|
||||
@@ -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",
|
||||
)
|
||||
]
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user