mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: make session boundaries explicit for migration flows (#38379)
This commit is contained in:
@@ -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"),
|
||||
|
||||
Reference in New Issue
Block a user