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

This commit is contained in:
Byron.wang
2026-07-04 04:43:01 +00:00
committed by GitHub
parent 5b4ceacbe7
commit 070aed81d9
10 changed files with 527 additions and 341 deletions
@@ -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
@@ -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"),