From 070aed81d96bee03d72b29542dac8c48d1d0a314 Mon Sep 17 00:00:00 2001 From: "Byron.wang" Date: Sat, 4 Jul 2026 12:43:01 +0800 Subject: [PATCH] refactor: make session boundaries explicit for migration flows (#38379) --- api/commands/data_migration.py | 31 ++-- api/controllers/common/session.py | 57 ++++++ api/controllers/console/app/wraps.py | 58 +----- api/services/data_migration/export_service.py | 36 ++-- api/services/data_migration/import_service.py | 113 +++++++----- .../commands/test_data_migration_commands.py | 114 ++++++++++++ .../controllers/common/test_session.py | 174 ++++++++++++++++++ .../controllers/console/app/test_wraps.py | 152 +-------------- .../data_migration/test_export_service.py | 12 +- .../data_migration/test_import_service.py | 121 ++++++------ 10 files changed, 527 insertions(+), 341 deletions(-) create mode 100644 api/controllers/common/session.py create mode 100644 api/tests/unit_tests/controllers/common/test_session.py diff --git a/api/commands/data_migration.py b/api/commands/data_migration.py index 52836af08a0..bd56c41ea44 100644 --- a/api/commands/data_migration.py +++ b/api/commands/data_migration.py @@ -10,6 +10,7 @@ import click import sqlalchemy as sa import yaml +from core.db.session_factory import session_factory from extensions.ext_database import db from models import Tenant from models.model import App @@ -106,7 +107,8 @@ def export_migration_data(input_file: str | None, output_file: str | None, overw assert output_file is not None raw_config = _load_json_object(input_file, "Export config") selection = ExportConfigParser().parse(raw_config) - result = MigrationExportService().export(selection) + with session_factory.create_session() as session: + result = MigrationExportService().export(session, selection) MigrationPackageService().save_package(result.package, output_file, overwrite=overwrite) click.echo(click.style(f"Output written to {output_file}", fg="green")) _render_report(result.report_items, context=_with_output_path(result.report_context, output_file)) @@ -153,19 +155,21 @@ def import_migration_data( _require_options(("--input", input_file)) assert input_file is not None package = MigrationPackageService().load_package(input_file) - result = MigrationImportService().import_package( - ImportRequest( - package=package, - cli_target_tenant=target_tenant, - operator_email=operator_email, - options_override=_build_options_override( - package.metadata.import_options, - id_strategy=id_strategy, - conflict_strategy=conflict_strategy, - create_app_api_token_on_import=create_app_api_token_on_import, + with session_factory.create_session() as session: + result = MigrationImportService().import_package( + session, + ImportRequest( + package=package, + cli_target_tenant=target_tenant, + operator_email=operator_email, + options_override=_build_options_override( + package.metadata.import_options, + id_strategy=id_strategy, + conflict_strategy=conflict_strategy, + create_app_api_token_on_import=create_app_api_token_on_import, + ), ), ) - ) _render_report(result.report_items, context=result.report_context) except MigrationDataError as exc: raise click.ClickException(str(exc)) from exc @@ -248,7 +252,8 @@ def migration_data_wizard() -> None: conflict_strategy=conflict_strategy, output_file=output_file, ) - result = MigrationExportService().export(selection) + with session_factory.create_session() as session: + result = MigrationExportService().export(session, selection) MigrationPackageService().save_package(result.package, output_file, overwrite=overwrite) click.echo(click.style(f"Output written to {output_file}", fg="green")) _print_wizard_step("Report") diff --git a/api/controllers/common/session.py b/api/controllers/common/session.py new file mode 100644 index 00000000000..fac2bec6767 --- /dev/null +++ b/api/controllers/common/session.py @@ -0,0 +1,57 @@ +"""Controller session decorators. + +`with_session` is an HTTP controller helper: it opens one SQLAlchemy session +for a Resource handler and injects it as the first argument after `self`. +Handlers use a transaction by default so migrated write paths keep +commit/rollback handling; pure read handlers may opt out with `write=False`. +""" + +from collections.abc import Callable +from functools import wraps +from typing import Concatenate, overload + +from sqlalchemy.orm import Session + +from core.db.session_factory import session_factory + + +@overload +def with_session[T, **P, R]( + view: Callable[Concatenate[T, Session, P], R], + *, + write: bool = True, +) -> Callable[Concatenate[T, P], R]: ... + + +@overload +def with_session[T, **P, R]( + view: None = None, + *, + write: bool = True, +) -> Callable[[Callable[Concatenate[T, Session, P], R]], Callable[Concatenate[T, P], R]]: ... + + +def with_session[T, **P, R]( + view: Callable[Concatenate[T, Session, P], R] | None = None, + *, + write: bool = True, +) -> ( + Callable[Concatenate[T, P], R] | Callable[[Callable[Concatenate[T, Session, P], R]], Callable[Concatenate[T, P], R]] +): + """Inject a request-scoped session, using a transaction only for write handlers.""" + + def decorator(view: Callable[Concatenate[T, Session, P], R]) -> Callable[Concatenate[T, P], R]: + @wraps(view) + def wrapper(self: T, *args: P.args, **kwargs: P.kwargs) -> R: + if write: + with session_factory.get_session_maker().begin() as session: + return view(self, session, *args, **kwargs) + + with session_factory.create_session() as session: + return view(self, session, *args, **kwargs) + + return wrapper + + if view is None: + return decorator + return decorator(view) diff --git a/api/controllers/console/app/wraps.py b/api/controllers/console/app/wraps.py index 8cb06533346..3d79273685f 100644 --- a/api/controllers/console/app/wraps.py +++ b/api/controllers/console/app/wraps.py @@ -1,26 +1,26 @@ """Controller decorators for console app resources. -`with_session` opens one SQLAlchemy session for a request handler and injects it -as the first argument after `self`. Handlers use a transaction by default so -migrated write paths keep commit/rollback handling; pure read handlers may opt -out with `write=False`. App-loading decorators prefer that injected session when -present, while still supporting existing handlers that have not been migrated -yet and still rely on Flask-SQLAlchemy's scoped `db.session`. +App-loading decorators prefer a session injected by +`controllers.common.session.with_session` when present, while still supporting +existing handlers that have not been migrated yet and still rely on +Flask-SQLAlchemy's scoped `db.session`. """ from collections.abc import Callable from functools import wraps -from typing import Concatenate, cast, overload +from typing import cast, overload from sqlalchemy import select from sqlalchemy.orm import Session +from controllers.common.session import with_session from controllers.console.app.error import AppNotFoundError -from core.db.session_factory import session_factory from extensions.ext_database import db from libs.login import current_account_with_tenant from models import App, AppMode +__all__ = ["get_app_model", "get_app_model_with_trial", "with_session"] + def _load_app_model(session: Session, app_id: str) -> App | None: """Load the tenant-scoped app row with the request session owned by `with_session`.""" @@ -45,48 +45,6 @@ def _load_app_model_with_trial(app_id: str) -> App | None: return app_model -@overload -def with_session[T, **P, R]( - view: Callable[Concatenate[T, Session, P], R], - *, - write: bool = True, -) -> Callable[Concatenate[T, P], R]: ... - - -@overload -def with_session[T, **P, R]( - view: None = None, - *, - write: bool = True, -) -> Callable[[Callable[Concatenate[T, Session, P], R]], Callable[Concatenate[T, P], R]]: ... - - -def with_session[T, **P, R]( - view: Callable[Concatenate[T, Session, P], R] | None = None, - *, - write: bool = True, -) -> ( - Callable[Concatenate[T, P], R] | Callable[[Callable[Concatenate[T, Session, P], R]], Callable[Concatenate[T, P], R]] -): - """Inject a request-scoped session, using a transaction only for write handlers.""" - - def decorator(view: Callable[Concatenate[T, Session, P], R]) -> Callable[Concatenate[T, P], R]: - @wraps(view) - def wrapper(self: T, *args: P.args, **kwargs: P.kwargs) -> R: - if write: - with session_factory.get_session_maker().begin() as session: - return view(self, session, *args, **kwargs) - - with session_factory.create_session() as session: - return view(self, session, *args, **kwargs) - - return wrapper - - if view is None: - return decorator - return decorator(view) - - def _get_injected_session(args: tuple[object, ...]) -> Session | None: """Return the request session inserted by `with_session`, if this handler has been migrated.""" if len(args) < 2: diff --git a/api/services/data_migration/export_service.py b/api/services/data_migration/export_service.py index 645a719e81d..f5d214d230b 100644 --- a/api/services/data_migration/export_service.py +++ b/api/services/data_migration/export_service.py @@ -6,9 +6,9 @@ from uuid import UUID import sqlalchemy as sa import yaml +from sqlalchemy.orm import Session from core.tools.tool_manager import ToolManager -from extensions.ext_database import db from graphon.model_runtime.utils.encoders import jsonable_encoder from models import Account, Tenant from models.account import TenantAccountJoin @@ -120,8 +120,8 @@ class MigrationExportService: self.package_service = package_service or MigrationPackageService() self.dependency_discovery_service = dependency_discovery_service or DependencyDiscoveryService() - def export(self, selection: ExportSelection) -> ExportResult: - tenant = self._get_tenant(selection) + def export(self, session: Session, selection: ExportSelection) -> ExportResult: + tenant = self._get_tenant(session, selection) package = self.package_service.build_empty_package( source_tenant_id=tenant.id, source_tenant_name=tenant.name, @@ -131,7 +131,7 @@ class MigrationExportService: report_items: list[ResourceReportItem] = [] discovered_dependencies: list[DiscoveredDependency] = [] - apps = self._selected_apps(tenant.id, selection) + apps = self._selected_apps(session, tenant.id, selection) exported_app_ids = {app.id for app in apps} for app in apps: dsl_content = AppDslService.export_dsl(app_model=app, include_secret=selection.include_secrets) @@ -157,6 +157,7 @@ class MigrationExportService: report_items=report_items, ) self._export_workflow_tools( + session, tenant, self._provider_ids( selection.additional_workflow_tools, discovered_dependencies, DependencyKind.WORKFLOW_TOOL @@ -167,6 +168,7 @@ class MigrationExportService: report_items=report_items, ) self._export_mcp_tools( + session, tenant_id=tenant.id, provider_ids=self._provider_ids( selection.additional_mcp_tools, @@ -193,9 +195,9 @@ class MigrationExportService: ), ) - def _get_tenant(self, selection: ExportSelection) -> Tenant: + def _get_tenant(self, session: Session, selection: ExportSelection) -> Tenant: if selection.source_tenant_id: - tenant = db.session.get(Tenant, selection.source_tenant_id) + tenant = session.get(Tenant, selection.source_tenant_id) if tenant is None: raise MigrationDataError(f"Source tenant not found: {selection.source_tenant_id}") if tenant.name != selection.source_tenant_name: @@ -203,7 +205,7 @@ class MigrationExportService: f"Source tenant id/name mismatch: {selection.source_tenant_id} / {selection.source_tenant_name}" ) return tenant - tenants = list(db.session.scalars(sa.select(Tenant).where(Tenant.name == selection.source_tenant_name)).all()) + tenants = list(session.scalars(sa.select(Tenant).where(Tenant.name == selection.source_tenant_name)).all()) if not tenants: raise MigrationDataError(f"Source tenant not found: {selection.source_tenant_name}") if len(tenants) > 1: @@ -212,13 +214,13 @@ class MigrationExportService: ) return tenants[0] - def _selected_apps(self, tenant_id: str, selection: ExportSelection) -> list[App]: + def _selected_apps(self, session: Session, tenant_id: str, selection: ExportSelection) -> list[App]: query = sa.select(App).where(App.tenant_id == tenant_id, App.mode.in_(SUPPORTED_APP_MODES)) if not selection.export_all_apps: if not selection.app_ids: return [] query = query.where(App.id.in_(selection.app_ids)) - apps = list(db.session.scalars(query).all()) + apps = list(session.scalars(query).all()) if not selection.export_all_apps and len(apps) != len(set(selection.app_ids)): found_ids = {app.id for app in apps} missing_ids = [app_id for app_id in selection.app_ids if app_id not in found_ids] @@ -265,6 +267,7 @@ class MigrationExportService: def _export_workflow_tools( self, + session: Session, tenant: Tenant, provider_ids: Iterable[str], *, @@ -276,7 +279,7 @@ class MigrationExportService: provider_ids = self._dedupe(provider_ids) if not provider_ids: return - owner = self._get_tenant_owner(tenant.id) + owner = self._get_tenant_owner(session, tenant.id) if owner is None: for provider_id in provider_ids: report_items.append( @@ -306,7 +309,7 @@ class MigrationExportService: exported_workflow_tools.append(tool_info) if tool_info.get("app_id") not in exported_app_ids: workflow_app_id = str(tool_info.get("app_id") or "") - workflow_app = db.session.get(App, workflow_app_id) if workflow_app_id else None + workflow_app = session.get(App, workflow_app_id) if workflow_app_id else None self._record_dependency_metadata( [ DiscoveredDependency( @@ -327,8 +330,8 @@ class MigrationExportService: ResourceReportItem(ResourceType.WORKFLOW_TOOL, provider_id, provider_id, "unresolved", str(exc)) ) - def _get_tenant_owner(self, tenant_id: str) -> Account | None: - return db.session.scalar( + def _get_tenant_owner(self, session: Session, tenant_id: str) -> Account | None: + return session.scalar( sa.select(Account) .join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id) .where(TenantAccountJoin.tenant_id == tenant_id, TenantAccountJoin.role == "owner") @@ -338,6 +341,7 @@ class MigrationExportService: def _export_mcp_tools( self, + session: Session, *, tenant_id: str, provider_ids: Iterable[str], @@ -355,7 +359,7 @@ class MigrationExportService: ) continue try: - provider = self._get_mcp_provider(tenant_id, provider_id) + provider = self._get_mcp_provider(session, tenant_id, provider_id) exported_mcp_tools.append(self._serialize_mcp_provider(provider)) report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider.name, "exported")) except Exception as exc: @@ -363,11 +367,11 @@ class MigrationExportService: ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider_id, "unresolved", str(exc)) ) - def _get_mcp_provider(self, tenant_id: str, provider_id: str) -> MCPToolProvider: + def _get_mcp_provider(self, session: Session, tenant_id: str, provider_id: str) -> MCPToolProvider: predicates = [MCPToolProvider.server_identifier == provider_id] if self._is_uuid_string(provider_id): predicates.append(MCPToolProvider.id == provider_id) - provider = db.session.scalar( + provider = session.scalar( sa.select(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant_id, sa.or_(*predicates)) ) if provider is None: diff --git a/api/services/data_migration/import_service.py b/api/services/data_migration/import_service.py index 3e0d1919e5e..3eb251bbaef 100644 --- a/api/services/data_migration/import_service.py +++ b/api/services/data_migration/import_service.py @@ -82,24 +82,24 @@ class ImportTargetResolver: "Target tenant must be provided by --target-tenant, import config, or package metadata." ) - def resolve(self, request: ImportRequest) -> ImportTarget: + def resolve(self, session: Session, request: ImportRequest) -> ImportTarget: target_tenant_name = self.select_target_tenant_name(request) package_target = request.package.metadata.target_tenant or {} if request.cli_target_tenant or request.config_target_tenant: - tenant = self._resolve_tenant_by_id_or_name(target_tenant_name) + tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name) elif package_target.get("id") and self._is_uuid(package_target["id"]): - tenant = db.session.get(Tenant, package_target["id"]) + tenant = session.get(Tenant, package_target["id"]) if tenant is not None and package_target.get("name") and tenant.name != package_target.get("name"): raise MigrationDataError( f"Target tenant id/name mismatch: {package_target['id']} / {package_target['name']}" ) else: - tenant = self._resolve_tenant_by_id_or_name(target_tenant_name) + tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name) if tenant is None: raise MigrationDataError(f"Target tenant not found: {target_tenant_name}") account_query = ( - db.session.query(Account) + session.query(Account) .join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id) .filter(TenantAccountJoin.tenant_id == tenant.id) ) @@ -123,12 +123,12 @@ class ImportTargetResolver: operator_email=account.email, ) - def _resolve_tenant_by_id_or_name(self, value: str) -> Tenant | None: + def _resolve_tenant_by_id_or_name(self, session: Session, value: str) -> Tenant | None: if self._is_uuid(value): - tenant = db.session.get(Tenant, value) + tenant = session.get(Tenant, value) if tenant is not None: return tenant - tenants = list(db.session.scalars(sa.select(Tenant).where(Tenant.name == value)).all()) + tenants = list(session.scalars(sa.select(Tenant).where(Tenant.name == value)).all()) if len(tenants) > 1: raise MigrationDataError(f"Target tenant name is ambiguous; use target_tenant.id: {value}") return tenants[0] if tenants else None @@ -149,8 +149,8 @@ class MigrationImportService: def __init__(self, *, target_resolver: ImportTargetResolver | None = None) -> None: self.target_resolver = target_resolver or ImportTargetResolver() - def import_package(self, request: ImportRequest) -> ImportResult: - target = self.target_resolver.resolve(request) + def import_package(self, session: Session, request: ImportRequest) -> ImportResult: + target = self.target_resolver.resolve(session, request) options = request.options_override or request.package.metadata.import_options report_items = [ ResourceReportItem( @@ -165,6 +165,7 @@ class MigrationImportService: id_mapping_details: list[ResourceIdMapping] = [] self._import_api_tools( + session, request.package, target, options, @@ -173,12 +174,13 @@ class MigrationImportService: id_mapping_details, self._source_api_provider_ids_by_name(request.package), ) - self._import_mcp_tools(request.package, target, options, report_items, id_mapping, id_mapping_details) - self._preflight_dependency_only_mcp(request.package, target, report_items) + self._import_mcp_tools(session, request.package, target, options, report_items, id_mapping, id_mapping_details) + self._preflight_dependency_only_mcp(session, request.package, target, report_items) workflow_tool_app_ids = self._workflow_tool_source_app_ids(request.package) imported_workflow_ids: set[str] = set() if workflow_tool_app_ids: self._import_workflows( + session, request.package, target, options, @@ -188,8 +190,11 @@ class MigrationImportService: imported_workflow_ids=imported_workflow_ids, only_app_ids=workflow_tool_app_ids, ) - self._import_workflow_tools(request.package, target, options, id_mapping, id_mapping_details, report_items) + self._import_workflow_tools( + session, request.package, target, options, id_mapping, id_mapping_details, report_items + ) self._import_workflows( + session, request.package, target, options, @@ -213,6 +218,7 @@ class MigrationImportService: def _import_workflows( self, + session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -223,8 +229,8 @@ class MigrationImportService: only_app_ids: set[str] | None = None, skip_app_ids: set[str] | None = None, ) -> None: - account = db.session.get(Account, target.operator_id) - tenant = db.session.get(Tenant, target.tenant_id) + account = session.get(Account, target.operator_id) + tenant = session.get(Tenant, target.tenant_id) if account is None: raise MigrationDataError(f"Operator account not found: {target.operator_id}") if tenant is None: @@ -242,7 +248,7 @@ class MigrationImportService: id_mapping, ) existing_app = ( - self._find_existing_app(app_id, target.tenant_id) + self._find_existing_app(session, app_id, target.tenant_id) if options.id_strategy == IdStrategy.PRESERVE_ID else None ) @@ -264,6 +270,7 @@ class MigrationImportService: continue imported_app_id = self._import_workflow_app( + session=session, account=account, workflow_data=workflow_data, dsl_content=dsl_content, @@ -283,7 +290,7 @@ class MigrationImportService: if imported_workflow_ids is not None: imported_workflow_ids.add(app_id) if options.create_app_api_token_on_import: - self._create_or_reuse_app_api_token(imported_app_id, target.tenant_id) + self._create_or_reuse_app_api_token(session, imported_app_id, target.tenant_id) report_items.append( ResourceReportItem( ResourceType.WORKFLOW, @@ -304,6 +311,7 @@ class MigrationImportService: def _import_workflow_app( self, *, + session: Session, account: Account, workflow_data: dict[str, object], dsl_content: str, @@ -311,7 +319,7 @@ class MigrationImportService: existing_app: App | None, options: ImportOptions, ) -> str: - import_service = AppDslService(cast(Session, db.session)) + import_service = AppDslService(session) if existing_app is not None: import_result = import_service.import_app( account=account, @@ -332,7 +340,7 @@ class MigrationImportService: raise MigrationDataError(f"Workflow import failed: {error}") if import_result.app_id is None: raise MigrationDataError(f"Workflow import did not return an app id: {workflow_data.get('name')}") - db.session.commit() + session.commit() return import_result.app_id def _rewrite_workflow_dsl_provider_ids(self, dsl_content: str, id_mapping: dict[str, str]) -> str: @@ -400,13 +408,13 @@ class MigrationImportService: def _should_preserve_source_app_id(self, options: ImportOptions) -> bool: return options.id_strategy == IdStrategy.PRESERVE_ID - def _find_existing_app(self, app_id: str | None, tenant_id: str) -> App | None: + def _find_existing_app(self, session: Session, app_id: str | None, tenant_id: str) -> App | None: if not self._is_uuid_string(app_id): return None - return db.session.scalar(sa.select(App).where(App.id == app_id, App.tenant_id == tenant_id)) + return session.scalar(sa.select(App).where(App.id == app_id, App.tenant_id == tenant_id)) - def _create_or_reuse_app_api_token(self, app_id: str, tenant_id: str) -> None: - existing = db.session.scalar( + def _create_or_reuse_app_api_token(self, session: Session, app_id: str, tenant_id: str) -> None: + existing = session.scalar( sa.select(ApiToken).where( ApiToken.type == ApiTokenType.APP, ApiToken.app_id == app_id, @@ -420,11 +428,12 @@ class MigrationImportService: api_token.tenant_id = tenant_id api_token.token = ApiToken.generate_api_key("app", 24) api_token.type = ApiTokenType.APP - db.session.add(api_token) - db.session.commit() + session.add(api_token) + session.commit() def _import_api_tools( self, + session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -436,7 +445,7 @@ class MigrationImportService: for tool_data in package.tools: provider_name = self._required_string(tool_data, "provider_name", "api_tool") schema = self._required_string(tool_data, "schema", "api_tool") - existing = db.session.scalar( + existing = session.scalar( sa.select(ApiToolProvider).where( ApiToolProvider.tenant_id == target.tenant_id, ApiToolProvider.name == provider_name, @@ -501,7 +510,7 @@ class MigrationImportService: icon=icon, ) status = "created" - target_provider = self._find_api_tool_provider(target.tenant_id, provider_name) + target_provider = self._find_api_tool_provider(session, target.tenant_id, provider_name) if target_provider is not None: self._record_id_mappings( id_mapping, @@ -513,8 +522,8 @@ class MigrationImportService: ) report_items.append(ResourceReportItem(ResourceType.API_TOOL, provider_name, provider_name, status)) - def _find_api_tool_provider(self, tenant_id: str, provider_name: str) -> ApiToolProvider | None: - return db.session.scalar( + def _find_api_tool_provider(self, session: Session, tenant_id: str, provider_name: str) -> ApiToolProvider | None: + return session.scalar( sa.select(ApiToolProvider).where( ApiToolProvider.tenant_id == tenant_id, ApiToolProvider.name == provider_name, @@ -549,6 +558,7 @@ class MigrationImportService: def _import_workflow_tools( self, + session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -558,13 +568,13 @@ class MigrationImportService: ) -> None: if not package.workflow_tools: return - account = db.session.get(Account, target.operator_id) + account = session.get(Account, target.operator_id) if account is None: raise MigrationDataError(f"Operator account not found: {target.operator_id}") for workflow_tool_data in package.workflow_tools: app_id = self._optional_string(workflow_tool_data.get("app_id")) resolved_app_id = id_mapping.get(app_id or "", app_id) - if not resolved_app_id or self._find_existing_app(resolved_app_id, target.tenant_id) is None: + if not resolved_app_id or self._find_existing_app(session, resolved_app_id, target.tenant_id) is None: report_items.append( ResourceReportItem( ResourceType.WORKFLOW_TOOL, @@ -576,7 +586,7 @@ class MigrationImportService: ) continue try: - self._ensure_workflow_app_is_published(target, account, resolved_app_id) + self._ensure_workflow_app_is_published(session, target, account, resolved_app_id) except Exception as exc: report_items.append( ResourceReportItem( @@ -592,7 +602,7 @@ class MigrationImportService: tool_name = self._required_string(workflow_tool_data, "name", "workflow_tool") lookup_workflow_tool_id = workflow_tool_id if options.id_strategy == IdStrategy.PRESERVE_ID else None existing = self._find_existing_workflow_tool( - target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id + session, target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id ) if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL: raise MigrationDataError(f"Workflow tool already exists and conflict_strategy=fail: {tool_name}") @@ -659,7 +669,7 @@ class MigrationImportService: ) status = "created" target_provider = self._find_existing_workflow_tool( - target.tenant_id, import_id or None, tool_name, resolved_app_id + session, target.tenant_id, import_id or None, tool_name, resolved_app_id ) if target_provider is None: raise MigrationDataError(f"Workflow tool was not created: {tool_name}") @@ -675,8 +685,10 @@ class MigrationImportService: ) report_items.append(ResourceReportItem(ResourceType.WORKFLOW_TOOL, identifier, tool_name, status)) - def _ensure_workflow_app_is_published(self, target: ImportTarget, account: Account, app_id: str) -> None: - app = self._find_existing_app(app_id, target.tenant_id) + def _ensure_workflow_app_is_published( + self, session: Session, target: ImportTarget, account: Account, app_id: str + ) -> None: + app = self._find_existing_app(session, app_id, target.tenant_id) if app is None: raise MigrationDataError(f"Referenced workflow app was not found in target tenant: {app_id}") if app.workflow_id: @@ -702,6 +714,7 @@ class MigrationImportService: def _import_mcp_tools( self, + session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -714,7 +727,7 @@ class MigrationImportService: server_identifier = self._required_string(mcp_data, "server_identifier", "mcp_tool") provider_id = self._optional_string(mcp_data.get("id")) lookup_provider_id = provider_id if options.id_strategy == IdStrategy.PRESERVE_ID else None - existing = self._find_existing_mcp_tool(target.tenant_id, lookup_provider_id, server_identifier) + existing = self._find_existing_mcp_tool(session, target.tenant_id, lookup_provider_id, server_identifier) if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL: raise MigrationDataError(f"MCP tool already exists and conflict_strategy=fail: {name}") if existing is not None and options.conflict_strategy == ConflictStrategy.SKIP: @@ -730,7 +743,7 @@ class MigrationImportService: report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, existing.id, name, "skipped")) continue - service = MCPToolManageService(session=cast(Session, db.session)) + service = MCPToolManageService(session=session) configuration = MCPConfiguration.model_validate(mcp_data.get("configuration") or {}) authentication = ( MCPAuthentication.model_validate(mcp_data["authentication"]) if mcp_data.get("authentication") else None @@ -752,7 +765,7 @@ class MigrationImportService: # stored mode (update_provider now defaults to OFF when omitted). identity_mode=IdentityMode(existing.identity_mode), ) - db.session.commit() + session.commit() status = "updated" identifier = existing.id provider = existing @@ -770,14 +783,16 @@ class MigrationImportService: configuration=configuration, authentication=authentication, ) - created_provider = self._find_existing_mcp_tool(target.tenant_id, lookup_provider_id, server_identifier) + created_provider = self._find_existing_mcp_tool( + session, target.tenant_id, lookup_provider_id, server_identifier + ) if created_provider is None: raise MigrationDataError(f"MCP provider was not created: {name}") status = "created" provider = created_provider identifier = provider.id self._restore_mcp_provider_tools(provider, mcp_data) - db.session.commit() + session.commit() if provider_id: self._record_id_mappings( id_mapping, @@ -797,12 +812,12 @@ class MigrationImportService: provider.authed = True def _find_existing_mcp_tool( - self, tenant_id: str, provider_id: str | None, server_identifier: str + self, session: Session, tenant_id: str, provider_id: str | None, server_identifier: str ) -> MCPToolProvider | None: predicates = [MCPToolProvider.server_identifier == server_identifier] if self._is_uuid_string(provider_id): predicates.append(MCPToolProvider.id == provider_id) - return db.session.scalar( + return session.scalar( sa.select(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant_id, or_(*predicates)).limit(1) ) @@ -816,26 +831,26 @@ class MigrationImportService: return True def _find_existing_workflow_tool( - self, tenant_id: str, workflow_tool_id: str | None, tool_name: str, app_id: str + self, session: Session, tenant_id: str, workflow_tool_id: str | None, tool_name: str, app_id: str ) -> WorkflowToolProvider | None: predicates = [WorkflowToolProvider.name == tool_name, WorkflowToolProvider.app_id == app_id] if self._is_uuid_string(workflow_tool_id): predicates.append(WorkflowToolProvider.id == workflow_tool_id) - return db.session.scalar( + return session.scalar( sa.select(WorkflowToolProvider) .where(WorkflowToolProvider.tenant_id == tenant_id, or_(*predicates)) .limit(1) ) def _preflight_dependency_only_mcp( - self, package: MigrationPackage, target: ImportTarget, report_items: list[ResourceReportItem] + self, session: Session, package: MigrationPackage, target: ImportTarget, report_items: list[ResourceReportItem] ) -> None: for dependency in package.dependencies: if dependency.get("kind") != DependencyKind.MCP_TOOL.value: continue provider_id = str(dependency.get("provider_id", dependency.get("id", ""))) provider_name = self._optional_string(dependency.get("provider_name") or dependency.get("name")) - existing = self._find_dependency_only_mcp_provider(target.tenant_id, provider_id, provider_name) + existing = self._find_dependency_only_mcp_provider(session, target.tenant_id, provider_id, provider_name) report_name = f"mcp_tool {provider_name or getattr(existing, 'name', None) or provider_id}" if existing is not None: report_items.append( @@ -864,12 +879,12 @@ class MigrationImportService: ) def _find_dependency_only_mcp_provider( - self, tenant_id: str, provider_id: str, provider_name: str | None + self, session: Session, tenant_id: str, provider_id: str, provider_name: str | None ) -> MCPToolProvider | None: predicates = [MCPToolProvider.server_identifier == provider_id] if self._is_uuid_string(provider_id): predicates.append(MCPToolProvider.id == provider_id) - return db.session.scalar( + return session.scalar( sa.select(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant_id, or_(*predicates)).limit(1) ) diff --git a/api/tests/unit_tests/commands/test_data_migration_commands.py b/api/tests/unit_tests/commands/test_data_migration_commands.py index 82aabacc5c7..b7f92f3291a 100644 --- a/api/tests/unit_tests/commands/test_data_migration_commands.py +++ b/api/tests/unit_tests/commands/test_data_migration_commands.py @@ -1,13 +1,41 @@ import json +from pathlib import Path from click.testing import CliRunner +from commands import data_migration from commands.data_migration import ( ID_STRATEGY_CHOICES, export_migration_data, export_migration_data_template, import_migration_data, ) +from services.data_migration.entities import ( + ConflictStrategy, + ExportResult, + ImportOptions, + ImportResult, + MigrationPackage, + ReportContext, +) + + +class FakeSessionContext: + session: object + entered: bool + exited: bool + + def __init__(self, session: object) -> None: + self.session = session + self.entered = False + self.exited = False + + def __enter__(self) -> object: + self.entered = True + return self.session + + def __exit__(self, *_args: object) -> None: + self.exited = True def test_export_command_requires_input_and_output(): @@ -69,3 +97,89 @@ def test_export_template_command_requires_overwrite_for_existing_output(tmp_path assert result.exit_code != 0 assert "already exists" in result.output + + +def test_export_command_uses_cli_owned_session(monkeypatch, tmp_path: Path): + session = object() + session_context = FakeSessionContext(session) + captured: dict[str, object] = {} + input_file = tmp_path / "export-config.json" + output_file = tmp_path / "migration-package.json" + input_file.write_text(json.dumps({"source_tenant": {"name": "source"}, "apps": {"all": True}})) + package = MigrationPackage.from_mapping({"metadata": {"version": "1", "source_scope": "single"}}) + + class FakeMigrationExportService: + def export(self, export_session, selection): + captured["session"] = export_session + captured["selection"] = selection + return ExportResult(package=package, report_items=[], report_context=ReportContext()) + + class FakeMigrationPackageService: + def save_package(self, package_to_save, path, *, overwrite): + captured["package"] = package_to_save + captured["path"] = path + captured["overwrite"] = overwrite + + monkeypatch.setattr(data_migration.session_factory, "create_session", lambda: session_context) + monkeypatch.setattr(data_migration, "MigrationExportService", FakeMigrationExportService) + monkeypatch.setattr(data_migration, "MigrationPackageService", FakeMigrationPackageService) + + result = CliRunner().invoke( + export_migration_data, + ["--input", str(input_file), "--output", str(output_file)], + ) + + assert result.exit_code == 0 + assert captured["session"] is session + assert captured["package"] is package + assert captured["path"] == str(output_file) + assert captured["overwrite"] is False + assert session_context.entered + assert session_context.exited + + +def test_import_command_uses_cli_owned_session(monkeypatch, tmp_path: Path): + session = object() + session_context = FakeSessionContext(session) + captured: dict[str, object] = {} + input_file = tmp_path / "migration-package.json" + input_file.write_text("{}") + package = MigrationPackage.from_mapping( + { + "metadata": { + "version": "1", + "source_scope": "single", + "target_tenant": {"name": "target"}, + "import_options": {"conflict_strategy": "fail"}, + } + } + ) + + class FakeMigrationImportService: + def import_package(self, import_session, request): + captured["session"] = import_session + captured["request"] = request + return ImportResult(report_items=[], report_context=ReportContext(target_tenant="target")) + + class FakeMigrationPackageService: + def load_package(self, path): + captured["path"] = path + return package + + monkeypatch.setattr(data_migration.session_factory, "create_session", lambda: session_context) + monkeypatch.setattr(data_migration, "MigrationImportService", FakeMigrationImportService) + monkeypatch.setattr(data_migration, "MigrationPackageService", FakeMigrationPackageService) + + result = CliRunner().invoke( + import_migration_data, + ["--input", str(input_file), "--conflict-strategy", "skip"], + ) + + assert result.exit_code == 0 + assert captured["session"] is session + assert captured["path"] == str(input_file) + request = captured["request"] + assert request.package is package + assert request.options_override == ImportOptions(conflict_strategy=ConflictStrategy.SKIP) + assert session_context.entered + assert session_context.exited diff --git a/api/tests/unit_tests/controllers/common/test_session.py b/api/tests/unit_tests/controllers/common/test_session.py new file mode 100644 index 00000000000..05b96059572 --- /dev/null +++ b/api/tests/unit_tests/controllers/common/test_session.py @@ -0,0 +1,174 @@ +from __future__ import annotations + +import pytest + +from controllers.common import session as session_module + + +class FakeSession: + committed: bool + rolled_back: bool + closed: bool + + def __init__(self) -> None: + self.committed = False + self.rolled_back = False + self.closed = False + + def commit(self) -> None: + self.committed = True + + def rollback(self) -> None: + self.rolled_back = True + + +class FakeSessionBegin: + session: FakeSession + entered: bool + exited: bool + exc_type: object | None + + def __init__(self, session: FakeSession) -> None: + self.session = session + self.entered = False + self.exited = False + self.exc_type = None + + def __enter__(self) -> FakeSession: + self.entered = True + return self.session + + def __exit__(self, exc_type: object | None, *_args: object) -> None: + self.exited = True + self.exc_type = exc_type + if exc_type is None: + self.session.commit() + else: + self.session.rollback() + self.session.closed = True + + +class FakeSessionContext: + session: FakeSession + entered: bool + exited: bool + exc_type: object | None + + def __init__(self, session: FakeSession) -> None: + self.session = session + self.entered = False + self.exited = False + self.exc_type = None + + def __enter__(self) -> FakeSession: + self.entered = True + return self.session + + def __exit__(self, exc_type: object | None, *_args: object) -> None: + self.exited = True + self.exc_type = exc_type + self.session.closed = True + + +class FakeSessionMaker: + begin_context: FakeSessionBegin + + def __init__(self, session: FakeSession) -> None: + self.begin_context = FakeSessionBegin(session) + + def begin(self) -> FakeSessionBegin: + return self.begin_context + + +def test_with_session_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) -> None: + session = FakeSession() + session_maker = FakeSessionMaker(session) + monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + + class Handler: + @session_module.with_session(write=True) + def post(self, injected_session): + assert injected_session is session + return "ok" + + assert Handler().post() == "ok" + + assert session.closed + assert session.committed + assert not session.rolled_back + assert session_maker.begin_context.entered + assert session_maker.begin_context.exited + assert session_maker.begin_context.exc_type is None + + +def test_with_session_default_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) -> None: + session = FakeSession() + session_maker = FakeSessionMaker(session) + monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + + class Handler: + @session_module.with_session + def post(self, injected_session): + assert injected_session is session + return "ok" + + assert Handler().post() == "ok" + assert session.committed + assert not session.rolled_back + + +def test_with_session_write_rolls_back_on_error(monkeypatch: pytest.MonkeyPatch) -> None: + session = FakeSession() + session_maker = FakeSessionMaker(session) + monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + + class Handler: + @session_module.with_session(write=True) + def get(self, _session): + raise RuntimeError("boom") + + with pytest.raises(RuntimeError, match="boom"): + Handler().get() + + assert session.closed + assert not session.committed + assert session.rolled_back + assert session_maker.begin_context.entered + assert session_maker.begin_context.exited + assert session_maker.begin_context.exc_type is RuntimeError + + +def test_with_session_read_mode_does_not_commit(monkeypatch: pytest.MonkeyPatch) -> None: + session = FakeSession() + session_context = FakeSessionContext(session) + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) + + class Handler: + @session_module.with_session(write=False) + def get(self, injected_session): + assert injected_session is session + return "ok" + + assert Handler().get() == "ok" + + assert session.closed + assert not session.committed + assert not session.rolled_back + assert session_context.entered + assert session_context.exited + assert session_context.exc_type is None + + +def test_with_session_preserves_wrapped_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + session = FakeSession() + session_maker = FakeSessionMaker(session) + monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + + class Handler: + @session_module.with_session + def get(self, _session): + """handler docs""" + return "ok" + + assert Handler.get.__name__ == "get" + assert Handler.get.__doc__ == "handler docs" diff --git a/api/tests/unit_tests/controllers/console/app/test_wraps.py b/api/tests/unit_tests/controllers/console/app/test_wraps.py index d46d22c5a21..4060d101c80 100644 --- a/api/tests/unit_tests/controllers/console/app/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/app/test_wraps.py @@ -4,6 +4,7 @@ from types import SimpleNamespace import pytest +from controllers.common.session import with_session from controllers.console.app import wraps as wraps_module from controllers.console.app.error import AppNotFoundError from models.model import AppMode @@ -11,16 +12,10 @@ from models.model import AppMode class FakeSession: app_model: object | None - committed: bool - rolled_back: bool - closed: bool scalar_called: bool def __init__(self, app_model: object | None = None) -> None: self.app_model = app_model - self.committed = False - self.rolled_back = False - self.closed = False self.scalar_called = False def scalar(self, *_args: object, **_kwargs: object) -> object | None: @@ -28,68 +23,10 @@ class FakeSession: return self.app_model def commit(self) -> None: - self.committed = True + pass def rollback(self) -> None: - self.rolled_back = True - - -class FakeSessionBegin: - session: FakeSession - entered: bool - exited: bool - exc_type: object | None - - def __init__(self, session: FakeSession) -> None: - self.session = session - self.entered = False - self.exited = False - self.exc_type = None - - def __enter__(self) -> FakeSession: - self.entered = True - return self.session - - def __exit__(self, exc_type: object | None, *_args: object) -> None: - self.exited = True - self.exc_type = exc_type - if exc_type is None: - self.session.commit() - else: - self.session.rollback() - self.session.closed = True - - -class FakeSessionContext: - session: FakeSession - entered: bool - exited: bool - exc_type: object | None - - def __init__(self, session: FakeSession) -> None: - self.session = session - self.entered = False - self.exited = False - self.exc_type = None - - def __enter__(self) -> FakeSession: - self.entered = True - return self.session - - def __exit__(self, exc_type: object | None, *_args: object) -> None: - self.exited = True - self.exc_type = exc_type - self.session.closed = True - - -class FakeSessionMaker: - begin_context: FakeSessionBegin - - def __init__(self, session: FakeSession) -> None: - self.begin_context = FakeSessionBegin(session) - - def begin(self) -> FakeSessionBegin: - return self.begin_context + pass def test_get_app_model_injects_model(monkeypatch: pytest.MonkeyPatch) -> None: @@ -126,11 +63,13 @@ def test_get_app_model_requires_app_id() -> None: handler() -def test_with_session_defaults_to_write_session_for_get_app_model(monkeypatch: pytest.MonkeyPatch) -> None: +def test_wraps_with_session_reexports_common_session_decorator() -> None: + assert wraps_module.with_session is with_session + + +def test_get_app_model_prefers_injected_session(monkeypatch: pytest.MonkeyPatch) -> None: app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1") session = FakeSession(app_model) - session_maker = FakeSessionMaker(session) - monkeypatch.setattr(wraps_module.session_factory, "get_session_maker", lambda: session_maker) monkeypatch.setattr(wraps_module, "current_account_with_tenant", lambda: (None, "t1")) monkeypatch.setattr( wraps_module.db, @@ -139,80 +78,9 @@ def test_with_session_defaults_to_write_session_for_get_app_model(monkeypatch: p ) class Handler: - @wraps_module.with_session @wraps_module.get_app_model - def get(self, injected_session, app_model): - assert injected_session is session + def get(self, _injected_session, app_model): return app_model.id - assert Handler().get(app_id="app-1") == "app-1" + assert Handler().get(session, app_id="app-1") == "app-1" assert session.scalar_called - assert session.committed - assert not session.rolled_back - assert session.closed - assert session_maker.begin_context.entered - assert session_maker.begin_context.exited - assert session_maker.begin_context.exc_type is None - - -def test_with_session_read_mode_does_not_commit(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_context = FakeSessionContext(session) - monkeypatch.setattr(wraps_module.session_factory, "create_session", lambda: session_context) - - class Handler: - @wraps_module.with_session(write=False) - def get(self, injected_session): - assert injected_session is session - return "ok" - - assert Handler().get() == "ok" - - assert session.closed - assert not session.committed - assert not session.rolled_back - assert session_context.entered - assert session_context.exited - assert session_context.exc_type is None - - -def test_with_session_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_maker = FakeSessionMaker(session) - monkeypatch.setattr(wraps_module.session_factory, "get_session_maker", lambda: session_maker) - - class Handler: - @wraps_module.with_session(write=True) - def post(self, injected_session): - assert injected_session is session - return "ok" - - assert Handler().post() == "ok" - - assert session.closed - assert session.committed - assert not session.rolled_back - assert session_maker.begin_context.entered - assert session_maker.begin_context.exited - assert session_maker.begin_context.exc_type is None - - -def test_with_session_write_rolls_back_on_error(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_maker = FakeSessionMaker(session) - monkeypatch.setattr(wraps_module.session_factory, "get_session_maker", lambda: session_maker) - - class Handler: - @wraps_module.with_session(write=True) - def get(self, _session): - raise RuntimeError("boom") - - with pytest.raises(RuntimeError, match="boom"): - Handler().get() - - assert session.closed - assert not session.committed - assert session.rolled_back - assert session_maker.begin_context.entered - assert session_maker.begin_context.exited - assert session_maker.begin_context.exc_type is RuntimeError diff --git a/api/tests/unit_tests/services/data_migration/test_export_service.py b/api/tests/unit_tests/services/data_migration/test_export_service.py index fcecd328f31..f5480ff52af 100644 --- a/api/tests/unit_tests/services/data_migration/test_export_service.py +++ b/api/tests/unit_tests/services/data_migration/test_export_service.py @@ -126,6 +126,7 @@ def test_secret_free_mcp_dependencies_are_dependency_only(): report_items = [] service._export_mcp_tools( + object(), tenant_id="tenant-1", provider_ids=["mcp-1"], include_secrets=False, @@ -147,16 +148,15 @@ def test_secret_free_mcp_dependencies_are_dependency_only(): assert report_items[0].name == "mcp_tool mcp-1" -def test_get_mcp_provider_does_not_compare_non_uuid_identifier_to_uuid_id(monkeypatch): +def test_get_mcp_provider_does_not_compare_non_uuid_identifier_to_uuid_id(): statements = [] - def capture_scalar(statement): - statements.append(str(statement)) - - monkeypatch.setattr("services.data_migration.export_service.db.session.scalar", capture_scalar) + class StubSession: + def scalar(self, statement): + statements.append(str(statement)) with pytest.raises(MigrationDataError, match="MCP provider not found"): - MigrationExportService()._get_mcp_provider("tenant-1", "my-test-mcp") + MigrationExportService()._get_mcp_provider(StubSession(), "tenant-1", "my-test-mcp") assert len(statements) == 1 assert "tool_mcp_providers.id =" not in statements[0] diff --git a/api/tests/unit_tests/services/data_migration/test_import_service.py b/api/tests/unit_tests/services/data_migration/test_import_service.py index bd130cdd78b..2b11d575ed6 100644 --- a/api/tests/unit_tests/services/data_migration/test_import_service.py +++ b/api/tests/unit_tests/services/data_migration/test_import_service.py @@ -92,12 +92,8 @@ def test_package_target_tenant_id_ignores_invalid_uuid(monkeypatch): return EmptyResult() - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - with pytest.raises(MigrationDataError, match="Target tenant not found"): - ImportTargetResolver().resolve(ImportRequest(package=package)) + ImportTargetResolver().resolve(StubSession(), ImportRequest(package=package)) def test_options_override_replaces_package_defaults(): @@ -117,7 +113,7 @@ def test_options_override_replaces_package_defaults(): captured_options: list[ImportOptions] = [] class StubResolver(ImportTargetResolver): - def resolve(self, request: ImportRequest) -> ImportTarget: + def resolve(self, session, request: ImportRequest) -> ImportTarget: return ImportTarget( tenant_id="tenant-1", tenant_name="target", @@ -128,6 +124,7 @@ def test_options_override_replaces_package_defaults(): class CapturingImportService(MigrationImportService): def _import_workflows( self, + session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -140,7 +137,7 @@ def test_options_override_replaces_package_defaults(): override = ImportOptions(create_app_api_token_on_import=False, conflict_strategy=ConflictStrategy.SKIP) CapturingImportService(target_resolver=StubResolver()).import_package( - ImportRequest(package=package, options_override=override) + object(), ImportRequest(package=package, options_override=override) ) assert captured_options == [override] @@ -153,47 +150,37 @@ def test_only_preserve_id_strategy_reuses_source_app_id(): assert service._should_preserve_source_app_id(ImportOptions(id_strategy=IdStrategy.GENERATE_NEW_ID)) is False -def test_find_existing_app_ignores_invalid_uuid(monkeypatch): +def test_find_existing_app_ignores_invalid_uuid(): class StubSession: def scalar(self, statement): raise AssertionError("invalid UUID should not be queried against App.id") - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - - assert MigrationImportService()._find_existing_app("not-a-uuid", "tenant-1") is None + assert MigrationImportService()._find_existing_app(StubSession(), "not-a-uuid", "tenant-1") is None -def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(monkeypatch): +def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_existing_workflow_tool("tenant-1", "not-a-uuid", "tool-name", "app-id") + MigrationImportService()._find_existing_workflow_tool( + StubSession(), "tenant-1", "not-a-uuid", "tool-name", "app-id" + ) where_clause = str(captured[0].whereclause) assert f"{WorkflowToolProvider.__tablename__}.id" not in where_clause -def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(monkeypatch): +def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_existing_mcp_tool("tenant-1", "my-test-mcp", "my-test-mcp") + MigrationImportService()._find_existing_mcp_tool(StubSession(), "tenant-1", "my-test-mcp", "my-test-mcp") where_clause = str(captured[0].whereclause) assert f"{MCPToolProvider.__tablename__}.id" not in where_clause @@ -224,10 +211,10 @@ def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction( from services.data_migration import import_service - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "AppDslService", StubAppDslService) imported_app_id = MigrationImportService()._import_workflow_app( + session=StubSession(), account=object(), workflow_data={"name": "main_chatflow"}, dsl_content="app:\n mode: workflow\n", @@ -330,20 +317,19 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch return account class PublishingImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id): + def _find_existing_app(self, session, app_id, tenant_id): return object() - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): if ("created", app_id) in events: return type("WorkflowToolProvider", (), {"id": workflow_tool_id or "created-workflow-tool-id"})() return None - def _ensure_workflow_app_is_published(self, target, account, app_id): + def _ensure_workflow_app_is_published(self, session, target, account, app_id): events.append(("published", app_id)) from services.data_migration import import_service - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.WorkflowToolManageService, "create_workflow_tool", @@ -351,6 +337,7 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch ) PublishingImportService()._import_workflow_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -397,18 +384,17 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP return account class StrategyImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id): + def _find_existing_app(self, session, app_id, tenant_id): return object() - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): return target_provider if created_kwargs else None - def _ensure_workflow_app_is_published(self, target, account, app_id): + def _ensure_workflow_app_is_published(self, session, target, account, app_id): return None from services.data_migration import import_service - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.WorkflowToolManageService, "create_workflow_tool", @@ -416,6 +402,7 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP ) StrategyImportService()._import_workflow_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -462,20 +449,17 @@ def test_workflow_tool_skip_records_id_mapping(monkeypatch): return account class SkipImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id): + def _find_existing_app(self, session, app_id, tenant_id): return object() - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): return existing_provider - def _ensure_workflow_app_is_published(self, target, account, app_id): + def _ensure_workflow_app_is_published(self, session, target, account, app_id): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - SkipImportService()._import_workflow_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -511,18 +495,22 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str report_items = [] class ExistingApiImportService(MigrationImportService): - def _find_api_tool_provider(self, tenant_id, provider_name): + def _find_api_tool_provider(self, session, tenant_id, provider_name): + return target_provider + + class StubSession: + def scalar(self, statement): return target_provider from services.data_migration import import_service - monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: target_provider) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) monkeypatch.setattr(import_service.ApiToolManageService, "update_api_tool_provider", lambda **kwargs: None) ExistingApiImportService()._import_api_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -561,18 +549,18 @@ def test_api_tool_create_records_id_mapping(monkeypatch): return None class CreatedApiImportService(MigrationImportService): - def _find_api_tool_provider(self, tenant_id, provider_name): + def _find_api_tool_provider(self, session, tenant_id, provider_name): return target_provider from services.data_migration import import_service - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) monkeypatch.setattr(import_service.ApiToolManageService, "create_api_tool_provider", lambda **kwargs: None) CreatedApiImportService()._import_api_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -617,10 +605,10 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): from services.data_migration import import_service - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) MigrationImportService()._import_mcp_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -665,7 +653,7 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str return None class ExistingMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier): + def _find_existing_mcp_tool(self, session, tenant_id, provider_id, server_identifier): return provider class StubMCPToolManageService: @@ -677,10 +665,10 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str from services.data_migration import import_service - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) ExistingMCPImportService()._import_mcp_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -727,7 +715,7 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): return None class CreatedMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier): + def _find_existing_mcp_tool(self, session, tenant_id, provider_id, server_identifier): return provider if provider_created else None class StubMCPToolManageService: @@ -740,10 +728,10 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): from services.data_migration import import_service - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) CreatedMCPImportService()._import_mcp_tools( + StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -812,11 +800,12 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work } ) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: None) + class StubSession: + def scalar(self, statement): + return None MigrationImportService()._preflight_dependency_only_mcp( + StubSession(), package, ImportTarget( tenant_id="tenant-1", @@ -839,18 +828,15 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work ] -def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(monkeypatch): +def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - MigrationImportService()._find_dependency_only_mcp_provider( + StubSession(), "tenant-1", "my-test-mcp-server", "my-test-mcp", @@ -874,11 +860,12 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp {"id": "target-provider-id", "name": "my-test-mcp", "server_identifier": "my-test-mcp-server"}, )() - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: provider) + class StubSession: + def scalar(self, statement): + return provider MigrationImportService()._preflight_dependency_only_mcp( + StubSession(), package, ImportTarget( tenant_id="tenant-1", @@ -904,7 +891,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): events = [] class StubResolver(ImportTargetResolver): - def resolve(self, request): + def resolve(self, session, request): return ImportTarget( tenant_id="tenant-1", tenant_name="target", @@ -915,6 +902,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): class OrderedImportService(MigrationImportService): def _import_api_tools( self, + session, package, target, options, @@ -927,6 +915,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): def _import_workflows( self, + session, package, target, options, @@ -951,10 +940,12 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): if imported_workflow_ids is not None: imported_workflow_ids.add(app_id) - def _import_workflow_tools(self, package, target, options, id_mapping, id_mapping_details, report_items): + def _import_workflow_tools( + self, session, package, target, options, id_mapping, id_mapping_details, report_items + ): events.append(("workflow_tool", package.workflow_tools[0]["id"])) - def _import_mcp_tools(self, package, target, options, report_items, id_mapping, id_mapping_details): + def _import_mcp_tools(self, session, package, target, options, report_items, id_mapping, id_mapping_details): events.append(("mcp_tools", "imported")) package = MigrationPackage.from_mapping( @@ -968,7 +959,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): } ) - OrderedImportService(target_resolver=StubResolver()).import_package(ImportRequest(package=package)) + OrderedImportService(target_resolver=StubResolver()).import_package(object(), ImportRequest(package=package)) assert events == [ ("api_tools", "imported"),