mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 10:38:32 +08:00
refactor: make session boundaries explicit for migration flows (#38379)
This commit is contained in:
parent
5b4ceacbe7
commit
070aed81d9
@ -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")
|
||||
|
||||
57
api/controllers/common/session.py
Normal file
57
api/controllers/common/session.py
Normal 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)
|
||||
@ -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:
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
)
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
174
api/tests/unit_tests/controllers/common/test_session.py
Normal file
174
api/tests/unit_tests/controllers/common/test_session.py
Normal 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"
|
||||
@ -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
|
||||
|
||||
@ -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]
|
||||
|
||||
@ -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"),
|
||||
|
||||
Loading…
Reference in New Issue
Block a user