refactor: make session boundaries explicit for migration flows (#38379)

This commit is contained in:
Byron.wang 2026-07-04 12:43:01 +08:00 committed by GitHub
parent 5b4ceacbe7
commit 070aed81d9
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
10 changed files with 527 additions and 341 deletions

View File

@ -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")

View File

@ -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)

View File

@ -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:

View File

@ -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:

View File

@ -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)
)

View File

@ -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

View File

@ -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"

View File

@ -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

View File

@ -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]

View File

@ -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"),