diff --git a/api/tests/unit_tests/services/data_migration/test_import_service.py b/api/tests/unit_tests/services/data_migration/test_import_service.py index 10460aba470..02a897a8f20 100644 --- a/api/tests/unit_tests/services/data_migration/test_import_service.py +++ b/api/tests/unit_tests/services/data_migration/test_import_service.py @@ -1,7 +1,15 @@ +from dataclasses import dataclass +from typing import cast + import pytest import yaml +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session -from models.tools import MCPToolProvider, WorkflowToolProvider +from core.tools.entities.tool_entities import ApiProviderSchemaType +from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole +from models.model import App, AppMode +from models.tools import ApiToolProvider, MCPToolProvider, WorkflowToolProvider from services.app_dsl_service import Import from services.data_migration import import_service from services.data_migration.entities import ( @@ -19,6 +27,142 @@ from services.data_migration.import_service import ImportRequest, ImportTargetRe from services.entities.dsl_entities import ImportStatus +@dataclass(frozen=True) +class Database: + """Typed database binding used by import code that still reads ``db.engine``.""" + + engine: Engine + session: Session + + +@pytest.fixture +def database(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Database: + database = Database(engine=cast(Engine, sqlite_session.get_bind()), session=sqlite_session) + monkeypatch.setattr(import_service, "db", database) + return database + + +def _persist_tenant_account( + session: Session, + *, + tenant_id: str = "tenant-1", + tenant_name: str = "target", + account_id: str = "account-1", +) -> tuple[Tenant, Account]: + tenant = Tenant(name=tenant_name) + tenant.id = tenant_id + account = Account(name="Owner", email="owner@example.com") + account.id = account_id + join = TenantAccountJoin( + tenant_id=tenant_id, + account_id=account_id, + current=True, + role=TenantAccountRole.OWNER, + ) + session.add_all([tenant, account, join]) + session.commit() + return tenant, account + + +def _persist_app( + session: Session, + *, + app_id: str, + tenant_id: str = "tenant-1", +) -> App: + app = App( + id=app_id, + tenant_id=tenant_id, + name=f"App {app_id}", + description="Migration fixture", + mode=AppMode.WORKFLOW, + icon_type=None, + icon="", + icon_background=None, + workflow_id=None, + enable_site=False, + enable_api=False, + max_active_requests=None, + created_by="account-1", + maintainer="account-1", + ) + session.add(app) + session.commit() + return app + + +def _persist_workflow_provider( + session: Session, + *, + provider_id: str, + app_id: str, + name: str = "embedded_workflow_as_tool", + tenant_id: str = "tenant-1", +) -> WorkflowToolProvider: + provider = WorkflowToolProvider( + name=name, + label=name, + icon="{}", + app_id=app_id, + version="", + user_id="account-1", + tenant_id=tenant_id, + description="", + parameter_configuration="[]", + ) + provider.id = provider_id + session.add(provider) + session.commit() + return provider + + +def _persist_api_provider( + session: Session, + *, + provider_id: str, + name: str = "weather", + tenant_id: str = "tenant-1", +) -> ApiToolProvider: + provider = ApiToolProvider( + name=name, + icon="{}", + schema="openapi: 3.0.0", + schema_type_str=ApiProviderSchemaType.OPENAPI, + user_id="account-1", + tenant_id=tenant_id, + description="", + tools_str="[]", + credentials_str="{}", + ) + provider.id = provider_id + session.add(provider) + session.commit() + return provider + + +def _persist_mcp_provider( + session: Session, + *, + provider_id: str, + server_identifier: str = "my-test-mcp", + name: str = "my-test-mcp", + tenant_id: str = "tenant-1", +) -> MCPToolProvider: + provider = MCPToolProvider( + name=name, + server_identifier=server_identifier, + server_url="http://localhost:3000/mcp", + server_url_hash=f"hash-{provider_id}", + icon=None, + tenant_id=tenant_id, + user_id="account-1", + ) + provider.id = provider_id + session.add(provider) + session.commit() + return provider + + def test_target_tenant_precedence_cli_then_config_then_package(): package = MigrationPackage.from_mapping( { @@ -77,31 +221,18 @@ def test_target_tenant_name_is_not_treated_as_uuid(): assert resolver._is_uuid("49a99e46-bc2c-4885-91fa-47615f6192b5") is True -def test_package_target_tenant_id_ignores_invalid_uuid(monkeypatch): +def test_package_target_tenant_id_ignores_invalid_uuid(database: Database): package = MigrationPackage.from_mapping( {"metadata": {"version": "1", "source_scope": "single", "target_tenant": {"id": "not-a-uuid"}}} ) - class StubSession: - def get(self, model, identifier): - raise AssertionError("invalid UUID should not be passed to session.get") - - def scalars(self, statement): - class EmptyResult: - def all(self): - return [] - - 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), session=import_service.db.session) + ImportTargetResolver().resolve(ImportRequest(package=package), session=database.session) + + assert database.session.query(Tenant).count() == 0 -def test_options_override_replaces_package_defaults(): +def test_options_override_replaces_package_defaults(database: Database): package = MigrationPackage.from_mapping( { "metadata": { @@ -142,7 +273,7 @@ def test_options_override_replaces_package_defaults(): CapturingImportService(target_resolver=StubResolver()).import_package( ImportRequest(package=package, options_override=override), - session=import_service.db.session, + session=database.session, ) assert captured_options == [override] @@ -155,74 +286,55 @@ 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): - 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", session=import_service.db.session) is None - ) +def test_find_existing_app_ignores_invalid_uuid(database: Database): + _persist_app(database.session, app_id="app-other", tenant_id="tenant-2") + assert MigrationImportService()._find_existing_app("not-a-uuid", "tenant-1", session=database.session) is None -def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(monkeypatch): - captured = [] +def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(database: Database): + provider = _persist_workflow_provider(database.session, provider_id="provider-1", app_id="app-id", name="tool-name") + statements = [] - class StubSession: - def scalar(self, statement): - captured.append(statement) + def capture_statement(orm_execute_state) -> None: + statements.append(orm_execute_state.statement) - from services.data_migration import import_service + event.listen(database.session, "do_orm_execute", capture_statement) + try: + result = MigrationImportService()._find_existing_workflow_tool( + "tenant-1", "not-a-uuid", "tool-name", "app-id", session=database.session + ) + finally: + event.remove(database.session, "do_orm_execute", capture_statement) - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_existing_workflow_tool( - "tenant-1", "not-a-uuid", "tool-name", "app-id", session=import_service.db.session - ) - - where_clause = str(captured[0].whereclause) + assert result is provider + where_clause = str(statements[0].whereclause) assert f"{WorkflowToolProvider.__tablename__}.id" not in where_clause -def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(monkeypatch): - captured = [] +def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(database: Database): + provider = _persist_mcp_provider(database.session, provider_id="provider-1") + statements = [] - class StubSession: - def scalar(self, statement): - captured.append(statement) + def capture_statement(orm_execute_state) -> None: + statements.append(orm_execute_state.statement) - from services.data_migration import import_service + event.listen(database.session, "do_orm_execute", capture_statement) + try: + result = MigrationImportService()._find_existing_mcp_tool( + "tenant-1", "my-test-mcp", "my-test-mcp", session=database.session + ) + finally: + event.remove(database.session, "do_orm_execute", capture_statement) - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_existing_mcp_tool( - "tenant-1", "my-test-mcp", "my-test-mcp", session=import_service.db.session - ) - - where_clause = str(captured[0].whereclause) + assert result is provider + where_clause = str(statements[0].whereclause) assert f"{MCPToolProvider.__tablename__}.id" not in where_clause assert f"{MCPToolProvider.__tablename__}.name" not in where_clause -def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction(monkeypatch): - class FailingNestedTransaction: - def __enter__(self): - raise AssertionError("nested transaction should not be opened") - - def __exit__(self, exc_type, exc_value, traceback): - return False - - class StubSession: - def begin_nested(self): - return FailingNestedTransaction() - - def commit(self): - return None - +def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction( + monkeypatch: pytest.MonkeyPatch, database: Database +): class StubAppDslService: def __init__(self, session): self.session = session @@ -230,22 +342,30 @@ def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction( def import_app(self, **kwargs): return Import(id="import-id", status=ImportStatus.COMPLETED, app_id="imported-app-id") - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "AppDslService", StubAppDslService) + nested_transactions = [] - imported_app_id = MigrationImportService()._import_workflow_app( - account=object(), - workflow_data={"name": "main_chatflow"}, - dsl_content="app:\n mode: workflow\n", - app_id="source-app-id", - existing_app=None, - options=ImportOptions(id_strategy=IdStrategy.PRESERVE_ID), - session=import_service.db.session, - ) + def capture_transaction(_session, transaction) -> None: + if transaction.nested: + nested_transactions.append(transaction) + + event.listen(database.session, "after_transaction_create", capture_transaction) + + try: + imported_app_id = MigrationImportService()._import_workflow_app( + account=object(), + workflow_data={"name": "main_chatflow"}, + dsl_content="app:\n mode: workflow\n", + app_id="source-app-id", + existing_app=None, + options=ImportOptions(id_strategy=IdStrategy.PRESERVE_ID), + session=database.session, + ) + finally: + event.remove(database.session, "after_transaction_create", capture_transaction) assert imported_app_id == "imported-app-id" + assert nested_transactions == [] def test_rewrite_workflow_dsl_replaces_tool_provider_ids(): @@ -330,211 +450,149 @@ def test_source_api_provider_ids_are_discovered_from_workflow_dsl(): assert MigrationImportService()._source_api_provider_ids_by_name(package) == {"weather": {"source-api-provider-id"}} -def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch): +def test_workflow_tool_import_publishes_referenced_app_before_create( + monkeypatch: pytest.MonkeyPatch, database: Database +): events = [] - account = type("Account", (), {"id": "account-1"})() - - class StubSession: - def get(self, model, identifier): - return account + _persist_tenant_account(database.session) + app_id = "00000000-0000-0000-0000-000000000001" + _persist_app(database.session, app_id=app_id) class PublishingImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id, session): - return object() - - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): - 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, session): 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", - lambda **kwargs: events.append(("created", kwargs["workflow_app_id"])), - ) + def create_workflow_tool(**kwargs) -> None: + events.append(("created", kwargs["workflow_app_id"])) + _persist_workflow_provider( + database.session, + provider_id="00000000-0000-0000-0000-000000000010", + app_id=kwargs["workflow_app_id"], + ) + monkeypatch.setattr(import_service.WorkflowToolManageService, "create_workflow_tool", create_workflow_tool) PublishingImportService()._import_workflow_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, - "workflow_tools": [ - { - "id": "workflow-tool-1", - "name": "embedded_workflow_as_tool", - "app_id": "workflow-app-1", - } - ], + "workflow_tools": [{"id": "workflow-tool-1", "name": "embedded_workflow_as_tool", "app_id": app_id}], } ), - ImportTarget( - tenant_id="tenant-1", - tenant_name="target", - operator_id="account-1", - operator_email="owner@example.com", - ), + ImportTarget("tenant-1", "target", "account-1", "owner@example.com"), ImportOptions(), {}, [], [], - session=import_service.db.session, + session=database.session, ) - assert events == [("published", "workflow-app-1"), ("created", "workflow-app-1")] + assert events == [("published", app_id), ("created", app_id)] -@pytest.mark.parametrize( - ("id_strategy", "expected_import_id"), - [ - (IdStrategy.PRESERVE_ID, "source-workflow-tool-id"), - (IdStrategy.GENERATE_NEW_ID, ""), - ], -) -def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyPatch, id_strategy, expected_import_id): +@pytest.mark.parametrize("id_strategy", [IdStrategy.PRESERVE_ID, IdStrategy.GENERATE_NEW_ID]) +def test_workflow_tool_import_id_follows_id_strategy( + monkeypatch: pytest.MonkeyPatch, database: Database, id_strategy: IdStrategy +): created_kwargs = [] - target_provider = type("WorkflowToolProvider", (), {"id": "target-workflow-tool-id"})() - account = type("Account", (), {"id": "account-1"})() - id_mapping = {"source-app-id": "target-app-id"} + _persist_tenant_account(database.session) + source_app_id = "00000000-0000-0000-0000-000000000021" + target_app_id = "00000000-0000-0000-0000-000000000022" + source_provider_id = "00000000-0000-0000-0000-000000000020" + generated_provider_id = "00000000-0000-0000-0000-000000000023" + _persist_app(database.session, app_id=target_app_id) + id_mapping = {source_app_id: target_app_id} id_mapping_details = [] - class StubSession: - def get(self, model, identifier): - return account - class StrategyImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id, session): - return object() - - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): - return target_provider if created_kwargs else None - def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - monkeypatch.setattr( - import_service.WorkflowToolManageService, - "create_workflow_tool", - lambda **kwargs: created_kwargs.append(kwargs), - ) + def create_workflow_tool(**kwargs) -> None: + created_kwargs.append(kwargs) + _persist_workflow_provider( + database.session, + provider_id=kwargs["import_id"] or generated_provider_id, + app_id=kwargs["workflow_app_id"], + ) + monkeypatch.setattr(import_service.WorkflowToolManageService, "create_workflow_tool", create_workflow_tool) StrategyImportService()._import_workflow_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, "workflow_tools": [ - { - "id": "source-workflow-tool-id", - "name": "embedded_workflow_as_tool", - "app_id": "source-app-id", - } + {"id": source_provider_id, "name": "embedded_workflow_as_tool", "app_id": source_app_id} ], } ), - ImportTarget( - tenant_id="tenant-1", - tenant_name="target", - operator_id="account-1", - operator_email="owner@example.com", - ), + ImportTarget("tenant-1", "target", "account-1", "owner@example.com"), ImportOptions(id_strategy=id_strategy), id_mapping, id_mapping_details, [], - session=import_service.db.session, + session=database.session, ) + expected_import_id = source_provider_id if id_strategy == IdStrategy.PRESERVE_ID else "" + target_id = expected_import_id or generated_provider_id assert created_kwargs[0]["import_id"] == expected_import_id - assert id_mapping["source-workflow-tool-id"] == "target-workflow-tool-id" + assert id_mapping[source_provider_id] == target_id assert id_mapping_details == [ - ResourceIdMapping( - ResourceType.WORKFLOW_TOOL, - "embedded_workflow_as_tool", - "source-workflow-tool-id", - "target-workflow-tool-id", - ) + ResourceIdMapping(ResourceType.WORKFLOW_TOOL, "embedded_workflow_as_tool", source_provider_id, target_id) ] -def test_workflow_tool_skip_records_id_mapping(monkeypatch): - account = type("Account", (), {"id": "account-1"})() - existing_provider = type("WorkflowToolProvider", (), {"id": "existing-workflow-tool-id"})() - id_mapping = {"source-app-id": "target-app-id"} - - class StubSession: - def get(self, model, identifier): - return account +def test_workflow_tool_skip_records_id_mapping(database: Database): + _persist_tenant_account(database.session) + source_app_id = "00000000-0000-0000-0000-000000000031" + target_app_id = "00000000-0000-0000-0000-000000000032" + source_provider_id = "00000000-0000-0000-0000-000000000034" + _persist_app(database.session, app_id=target_app_id) + existing = _persist_workflow_provider( + database.session, provider_id="00000000-0000-0000-0000-000000000033", app_id=target_app_id + ) + id_mapping = {source_app_id: target_app_id} class SkipImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id, session): - return object() - - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): - return existing_provider - def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - SkipImportService()._import_workflow_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, "workflow_tools": [ - { - "id": "source-workflow-tool-id", - "name": "embedded_workflow_as_tool", - "app_id": "source-app-id", - } + {"id": source_provider_id, "name": "embedded_workflow_as_tool", "app_id": source_app_id} ], } ), - ImportTarget( - tenant_id="tenant-1", - tenant_name="target", - operator_id="account-1", - operator_email="owner@example.com", - ), + ImportTarget("tenant-1", "target", "account-1", "owner@example.com"), ImportOptions(conflict_strategy=ConflictStrategy.SKIP, id_strategy=IdStrategy.GENERATE_NEW_ID), id_mapping, [], [], - session=import_service.db.session, + session=database.session, ) - assert id_mapping["source-workflow-tool-id"] == "existing-workflow-tool-id" + assert id_mapping[source_provider_id] == existing.id @pytest.mark.parametrize("conflict_strategy", [ConflictStrategy.SKIP, ConflictStrategy.UPDATE]) -def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_strategy): - target_provider = type("ApiToolProvider", (), {"id": "target-api-provider-id", "name": "weather"})() +def test_api_tool_existing_provider_records_id_mapping( + monkeypatch: pytest.MonkeyPatch, database: Database, conflict_strategy: ConflictStrategy +): + target_provider = _persist_api_provider(database.session, provider_id="target-api-provider-id") + _persist_api_provider(database.session, provider_id="other-tenant-provider-id", tenant_id="tenant-2") id_mapping = {} id_mapping_details = [] report_items = [] - class ExistingApiImportService(MigrationImportService): - def _find_api_tool_provider(self, tenant_id, provider_name, session): - 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( + MigrationImportService()._import_api_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -552,7 +610,7 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str id_mapping, id_mapping_details, {"weather": {"source-api-provider-id-from-dsl"}}, - session=import_service.db.session, + session=database.session, ) assert id_mapping == { @@ -565,27 +623,19 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str ) -def test_api_tool_create_records_id_mapping(monkeypatch): - target_provider = type("ApiToolProvider", (), {"id": "target-api-provider-id", "name": "weather"})() +def test_api_tool_create_records_id_mapping(monkeypatch: pytest.MonkeyPatch, database: Database): id_mapping = {} - class StubSession: - def scalar(self, statement): - return None - - class CreatedApiImportService(MigrationImportService): - def _find_api_tool_provider(self, tenant_id, provider_name, session): - 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( + def create_api_tool_provider(**_kwargs) -> None: + _persist_api_provider(database.session, provider_id="target-api-provider-id") + + monkeypatch.setattr(import_service.ApiToolManageService, "create_api_tool_provider", create_api_tool_provider) + + MigrationImportService()._import_api_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -603,25 +653,16 @@ def test_api_tool_create_records_id_mapping(monkeypatch): id_mapping, [], {}, - session=import_service.db.session, + session=database.session, ) assert id_mapping["source-api-provider-id"] == "target-api-provider-id" -def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): - provider = type( - "Provider", (), {"id": "target-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"} - )() +def test_mcp_tool_import_restores_exported_tool_list(monkeypatch: pytest.MonkeyPatch, database: Database): + provider = _persist_mcp_provider(database.session, provider_id="target-provider-id") report_items = [] - class StubSession: - def scalar(self, statement): - return provider - - def commit(self): - return None - class StubMCPToolManageService: def __init__(self, session): self.session = session @@ -629,9 +670,6 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): def update_provider(self, **kwargs): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) MigrationImportService()._import_mcp_tools( @@ -660,29 +698,29 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): report_items, {}, [], - session=import_service.db.session, + session=database.session, ) + database.session.refresh(provider) assert provider.tools == '[{"name": "echo"}]' assert provider.authed is True @pytest.mark.parametrize("conflict_strategy", [ConflictStrategy.SKIP, ConflictStrategy.UPDATE]) -def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_strategy): - provider = type( - "Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"} - )() +def test_mcp_tool_existing_provider_records_id_mapping( + monkeypatch: pytest.MonkeyPatch, database: Database, conflict_strategy: ConflictStrategy +): + provider = _persist_mcp_provider(database.session, provider_id="target-mcp-provider-id") + _persist_mcp_provider( + database.session, + provider_id="other-mcp-provider-id", + server_identifier="other-mcp", + name="other-mcp", + tenant_id="tenant-2", + ) id_mapping = {} id_mapping_details = [] - class StubSession: - def commit(self): - return None - - class ExistingMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): - return provider - class StubMCPToolManageService: def __init__(self, session): self.session = session @@ -690,12 +728,9 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str def update_provider(self, **kwargs): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) - ExistingMCPImportService()._import_mcp_tools( + MigrationImportService()._import_mcp_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -721,7 +756,7 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str [], id_mapping, id_mapping_details, - session=import_service.db.session, + session=database.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" @@ -731,35 +766,25 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str ] -def test_mcp_tool_create_records_id_mapping(monkeypatch): - provider = type( - "Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"} - )() +def test_mcp_tool_create_records_id_mapping(monkeypatch: pytest.MonkeyPatch, database: Database): id_mapping = {} - provider_created = False - - class StubSession: - def commit(self): - return None - - class CreatedMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): - return provider if provider_created else None class StubMCPToolManageService: def __init__(self, session): self.session = session def create_provider(self, **kwargs): - nonlocal provider_created - provider_created = True + _persist_mcp_provider( + database.session, + provider_id="target-mcp-provider-id", + server_identifier=kwargs["server_identifier"], + name=kwargs["name"], + tenant_id=kwargs["tenant_id"], + ) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) - CreatedMCPImportService()._import_mcp_tools( + MigrationImportService()._import_mcp_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -784,13 +809,13 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): [], id_mapping, [], - session=import_service.db.session, + session=database.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" -def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_workflow_context(monkeypatch): +def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_workflow_context(database: Database): report_items = [] package = MigrationPackage.from_mapping( { @@ -829,10 +854,6 @@ 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) - MigrationImportService()._preflight_dependency_only_mcp( package, ImportTarget( @@ -842,7 +863,7 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work operator_email="owner@example.com", ), report_items, - session=import_service.db.session, + session=database.session, ) assert report_items == [ @@ -857,29 +878,30 @@ 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): - captured = [] +def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(database: Database): + provider = _persist_mcp_provider(database.session, provider_id="provider-1", server_identifier="my-test-mcp-server") + statements = [] - class StubSession: - def scalar(self, statement): - captured.append(statement) + def capture_statement(orm_execute_state) -> None: + statements.append(orm_execute_state.statement) - from services.data_migration import import_service + event.listen(database.session, "do_orm_execute", capture_statement) + try: + result = MigrationImportService()._find_dependency_only_mcp_provider( + "tenant-1", + "my-test-mcp-server", + "my-test-mcp", + session=database.session, + ) + finally: + event.remove(database.session, "do_orm_execute", capture_statement) - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_dependency_only_mcp_provider( - "tenant-1", - "my-test-mcp-server", - "my-test-mcp", - session=import_service.db.session, - ) - - where_clause = str(captured[0].whereclause) + assert result is provider + where_clause = str(statements[0].whereclause) assert f"{MCPToolProvider.__tablename__}.id" not in where_clause -def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeypatch): +def test_dependency_only_mcp_preflight_reports_available_target_provider(database: Database): report_items = [] package = MigrationPackage.from_mapping( { @@ -887,15 +909,18 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp "dependencies": [{"kind": "mcp_tool", "provider_id": "my-test-mcp-server"}], } ) - provider = type( - "Provider", - (), - {"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) + _persist_mcp_provider( + database.session, + provider_id="target-provider-id", + server_identifier="my-test-mcp-server", + ) + _persist_mcp_provider( + database.session, + provider_id="other-provider-id", + server_identifier="other-server", + name="other-provider", + tenant_id="tenant-2", + ) MigrationImportService()._preflight_dependency_only_mcp( package, @@ -906,7 +931,7 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp operator_email="owner@example.com", ), report_items, - session=import_service.db.session, + session=database.session, ) assert report_items == [ @@ -920,7 +945,7 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp ] -def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): +def test_import_package_imports_workflow_tool_provider_apps_before_consumers(database: Database): events = [] class StubResolver(ImportTargetResolver): @@ -996,7 +1021,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): ) OrderedImportService(target_resolver=StubResolver()).import_package( - ImportRequest(package=package), session=import_service.db.session + ImportRequest(package=package), session=database.session ) assert events == [