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
+20 -16
View File
@@ -6,9 +6,9 @@ from uuid import UUID
import sqlalchemy as sa
import yaml
from sqlalchemy.orm import Session
from core.tools.tool_manager import ToolManager
from extensions.ext_database import db
from graphon.model_runtime.utils.encoders import jsonable_encoder
from models import Account, Tenant
from models.account import TenantAccountJoin
@@ -120,8 +120,8 @@ class MigrationExportService:
self.package_service = package_service or MigrationPackageService()
self.dependency_discovery_service = dependency_discovery_service or DependencyDiscoveryService()
def export(self, selection: ExportSelection) -> ExportResult:
tenant = self._get_tenant(selection)
def export(self, session: Session, selection: ExportSelection) -> ExportResult:
tenant = self._get_tenant(session, selection)
package = self.package_service.build_empty_package(
source_tenant_id=tenant.id,
source_tenant_name=tenant.name,
@@ -131,7 +131,7 @@ class MigrationExportService:
report_items: list[ResourceReportItem] = []
discovered_dependencies: list[DiscoveredDependency] = []
apps = self._selected_apps(tenant.id, selection)
apps = self._selected_apps(session, tenant.id, selection)
exported_app_ids = {app.id for app in apps}
for app in apps:
dsl_content = AppDslService.export_dsl(app_model=app, include_secret=selection.include_secrets)
@@ -157,6 +157,7 @@ class MigrationExportService:
report_items=report_items,
)
self._export_workflow_tools(
session,
tenant,
self._provider_ids(
selection.additional_workflow_tools, discovered_dependencies, DependencyKind.WORKFLOW_TOOL
@@ -167,6 +168,7 @@ class MigrationExportService:
report_items=report_items,
)
self._export_mcp_tools(
session,
tenant_id=tenant.id,
provider_ids=self._provider_ids(
selection.additional_mcp_tools,
@@ -193,9 +195,9 @@ class MigrationExportService:
),
)
def _get_tenant(self, selection: ExportSelection) -> Tenant:
def _get_tenant(self, session: Session, selection: ExportSelection) -> Tenant:
if selection.source_tenant_id:
tenant = db.session.get(Tenant, selection.source_tenant_id)
tenant = session.get(Tenant, selection.source_tenant_id)
if tenant is None:
raise MigrationDataError(f"Source tenant not found: {selection.source_tenant_id}")
if tenant.name != selection.source_tenant_name:
@@ -203,7 +205,7 @@ class MigrationExportService:
f"Source tenant id/name mismatch: {selection.source_tenant_id} / {selection.source_tenant_name}"
)
return tenant
tenants = list(db.session.scalars(sa.select(Tenant).where(Tenant.name == selection.source_tenant_name)).all())
tenants = list(session.scalars(sa.select(Tenant).where(Tenant.name == selection.source_tenant_name)).all())
if not tenants:
raise MigrationDataError(f"Source tenant not found: {selection.source_tenant_name}")
if len(tenants) > 1:
@@ -212,13 +214,13 @@ class MigrationExportService:
)
return tenants[0]
def _selected_apps(self, tenant_id: str, selection: ExportSelection) -> list[App]:
def _selected_apps(self, session: Session, tenant_id: str, selection: ExportSelection) -> list[App]:
query = sa.select(App).where(App.tenant_id == tenant_id, App.mode.in_(SUPPORTED_APP_MODES))
if not selection.export_all_apps:
if not selection.app_ids:
return []
query = query.where(App.id.in_(selection.app_ids))
apps = list(db.session.scalars(query).all())
apps = list(session.scalars(query).all())
if not selection.export_all_apps and len(apps) != len(set(selection.app_ids)):
found_ids = {app.id for app in apps}
missing_ids = [app_id for app_id in selection.app_ids if app_id not in found_ids]
@@ -265,6 +267,7 @@ class MigrationExportService:
def _export_workflow_tools(
self,
session: Session,
tenant: Tenant,
provider_ids: Iterable[str],
*,
@@ -276,7 +279,7 @@ class MigrationExportService:
provider_ids = self._dedupe(provider_ids)
if not provider_ids:
return
owner = self._get_tenant_owner(tenant.id)
owner = self._get_tenant_owner(session, tenant.id)
if owner is None:
for provider_id in provider_ids:
report_items.append(
@@ -306,7 +309,7 @@ class MigrationExportService:
exported_workflow_tools.append(tool_info)
if tool_info.get("app_id") not in exported_app_ids:
workflow_app_id = str(tool_info.get("app_id") or "")
workflow_app = db.session.get(App, workflow_app_id) if workflow_app_id else None
workflow_app = session.get(App, workflow_app_id) if workflow_app_id else None
self._record_dependency_metadata(
[
DiscoveredDependency(
@@ -327,8 +330,8 @@ class MigrationExportService:
ResourceReportItem(ResourceType.WORKFLOW_TOOL, provider_id, provider_id, "unresolved", str(exc))
)
def _get_tenant_owner(self, tenant_id: str) -> Account | None:
return db.session.scalar(
def _get_tenant_owner(self, session: Session, tenant_id: str) -> Account | None:
return session.scalar(
sa.select(Account)
.join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id)
.where(TenantAccountJoin.tenant_id == tenant_id, TenantAccountJoin.role == "owner")
@@ -338,6 +341,7 @@ class MigrationExportService:
def _export_mcp_tools(
self,
session: Session,
*,
tenant_id: str,
provider_ids: Iterable[str],
@@ -355,7 +359,7 @@ class MigrationExportService:
)
continue
try:
provider = self._get_mcp_provider(tenant_id, provider_id)
provider = self._get_mcp_provider(session, tenant_id, provider_id)
exported_mcp_tools.append(self._serialize_mcp_provider(provider))
report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider.name, "exported"))
except Exception as exc:
@@ -363,11 +367,11 @@ class MigrationExportService:
ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider_id, "unresolved", str(exc))
)
def _get_mcp_provider(self, tenant_id: str, provider_id: str) -> MCPToolProvider:
def _get_mcp_provider(self, session: Session, tenant_id: str, provider_id: str) -> MCPToolProvider:
predicates = [MCPToolProvider.server_identifier == provider_id]
if self._is_uuid_string(provider_id):
predicates.append(MCPToolProvider.id == provider_id)
provider = db.session.scalar(
provider = session.scalar(
sa.select(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant_id, sa.or_(*predicates))
)
if provider is None:
+64 -49
View File
@@ -82,24 +82,24 @@ class ImportTargetResolver:
"Target tenant must be provided by --target-tenant, import config, or package metadata."
)
def resolve(self, request: ImportRequest) -> ImportTarget:
def resolve(self, session: Session, request: ImportRequest) -> ImportTarget:
target_tenant_name = self.select_target_tenant_name(request)
package_target = request.package.metadata.target_tenant or {}
if request.cli_target_tenant or request.config_target_tenant:
tenant = self._resolve_tenant_by_id_or_name(target_tenant_name)
tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name)
elif package_target.get("id") and self._is_uuid(package_target["id"]):
tenant = db.session.get(Tenant, package_target["id"])
tenant = session.get(Tenant, package_target["id"])
if tenant is not None and package_target.get("name") and tenant.name != package_target.get("name"):
raise MigrationDataError(
f"Target tenant id/name mismatch: {package_target['id']} / {package_target['name']}"
)
else:
tenant = self._resolve_tenant_by_id_or_name(target_tenant_name)
tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name)
if tenant is None:
raise MigrationDataError(f"Target tenant not found: {target_tenant_name}")
account_query = (
db.session.query(Account)
session.query(Account)
.join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id)
.filter(TenantAccountJoin.tenant_id == tenant.id)
)
@@ -123,12 +123,12 @@ class ImportTargetResolver:
operator_email=account.email,
)
def _resolve_tenant_by_id_or_name(self, value: str) -> Tenant | None:
def _resolve_tenant_by_id_or_name(self, session: Session, value: str) -> Tenant | None:
if self._is_uuid(value):
tenant = db.session.get(Tenant, value)
tenant = session.get(Tenant, value)
if tenant is not None:
return tenant
tenants = list(db.session.scalars(sa.select(Tenant).where(Tenant.name == value)).all())
tenants = list(session.scalars(sa.select(Tenant).where(Tenant.name == value)).all())
if len(tenants) > 1:
raise MigrationDataError(f"Target tenant name is ambiguous; use target_tenant.id: {value}")
return tenants[0] if tenants else None
@@ -149,8 +149,8 @@ class MigrationImportService:
def __init__(self, *, target_resolver: ImportTargetResolver | None = None) -> None:
self.target_resolver = target_resolver or ImportTargetResolver()
def import_package(self, request: ImportRequest) -> ImportResult:
target = self.target_resolver.resolve(request)
def import_package(self, session: Session, request: ImportRequest) -> ImportResult:
target = self.target_resolver.resolve(session, request)
options = request.options_override or request.package.metadata.import_options
report_items = [
ResourceReportItem(
@@ -165,6 +165,7 @@ class MigrationImportService:
id_mapping_details: list[ResourceIdMapping] = []
self._import_api_tools(
session,
request.package,
target,
options,
@@ -173,12 +174,13 @@ class MigrationImportService:
id_mapping_details,
self._source_api_provider_ids_by_name(request.package),
)
self._import_mcp_tools(request.package, target, options, report_items, id_mapping, id_mapping_details)
self._preflight_dependency_only_mcp(request.package, target, report_items)
self._import_mcp_tools(session, request.package, target, options, report_items, id_mapping, id_mapping_details)
self._preflight_dependency_only_mcp(session, request.package, target, report_items)
workflow_tool_app_ids = self._workflow_tool_source_app_ids(request.package)
imported_workflow_ids: set[str] = set()
if workflow_tool_app_ids:
self._import_workflows(
session,
request.package,
target,
options,
@@ -188,8 +190,11 @@ class MigrationImportService:
imported_workflow_ids=imported_workflow_ids,
only_app_ids=workflow_tool_app_ids,
)
self._import_workflow_tools(request.package, target, options, id_mapping, id_mapping_details, report_items)
self._import_workflow_tools(
session, request.package, target, options, id_mapping, id_mapping_details, report_items
)
self._import_workflows(
session,
request.package,
target,
options,
@@ -213,6 +218,7 @@ class MigrationImportService:
def _import_workflows(
self,
session: Session,
package: MigrationPackage,
target: ImportTarget,
options: ImportOptions,
@@ -223,8 +229,8 @@ class MigrationImportService:
only_app_ids: set[str] | None = None,
skip_app_ids: set[str] | None = None,
) -> None:
account = db.session.get(Account, target.operator_id)
tenant = db.session.get(Tenant, target.tenant_id)
account = session.get(Account, target.operator_id)
tenant = session.get(Tenant, target.tenant_id)
if account is None:
raise MigrationDataError(f"Operator account not found: {target.operator_id}")
if tenant is None:
@@ -242,7 +248,7 @@ class MigrationImportService:
id_mapping,
)
existing_app = (
self._find_existing_app(app_id, target.tenant_id)
self._find_existing_app(session, app_id, target.tenant_id)
if options.id_strategy == IdStrategy.PRESERVE_ID
else None
)
@@ -264,6 +270,7 @@ class MigrationImportService:
continue
imported_app_id = self._import_workflow_app(
session=session,
account=account,
workflow_data=workflow_data,
dsl_content=dsl_content,
@@ -283,7 +290,7 @@ class MigrationImportService:
if imported_workflow_ids is not None:
imported_workflow_ids.add(app_id)
if options.create_app_api_token_on_import:
self._create_or_reuse_app_api_token(imported_app_id, target.tenant_id)
self._create_or_reuse_app_api_token(session, imported_app_id, target.tenant_id)
report_items.append(
ResourceReportItem(
ResourceType.WORKFLOW,
@@ -304,6 +311,7 @@ class MigrationImportService:
def _import_workflow_app(
self,
*,
session: Session,
account: Account,
workflow_data: dict[str, object],
dsl_content: str,
@@ -311,7 +319,7 @@ class MigrationImportService:
existing_app: App | None,
options: ImportOptions,
) -> str:
import_service = AppDslService(cast(Session, db.session))
import_service = AppDslService(session)
if existing_app is not None:
import_result = import_service.import_app(
account=account,
@@ -332,7 +340,7 @@ class MigrationImportService:
raise MigrationDataError(f"Workflow import failed: {error}")
if import_result.app_id is None:
raise MigrationDataError(f"Workflow import did not return an app id: {workflow_data.get('name')}")
db.session.commit()
session.commit()
return import_result.app_id
def _rewrite_workflow_dsl_provider_ids(self, dsl_content: str, id_mapping: dict[str, str]) -> str:
@@ -400,13 +408,13 @@ class MigrationImportService:
def _should_preserve_source_app_id(self, options: ImportOptions) -> bool:
return options.id_strategy == IdStrategy.PRESERVE_ID
def _find_existing_app(self, app_id: str | None, tenant_id: str) -> App | None:
def _find_existing_app(self, session: Session, app_id: str | None, tenant_id: str) -> App | None:
if not self._is_uuid_string(app_id):
return None
return db.session.scalar(sa.select(App).where(App.id == app_id, App.tenant_id == tenant_id))
return session.scalar(sa.select(App).where(App.id == app_id, App.tenant_id == tenant_id))
def _create_or_reuse_app_api_token(self, app_id: str, tenant_id: str) -> None:
existing = db.session.scalar(
def _create_or_reuse_app_api_token(self, session: Session, app_id: str, tenant_id: str) -> None:
existing = session.scalar(
sa.select(ApiToken).where(
ApiToken.type == ApiTokenType.APP,
ApiToken.app_id == app_id,
@@ -420,11 +428,12 @@ class MigrationImportService:
api_token.tenant_id = tenant_id
api_token.token = ApiToken.generate_api_key("app", 24)
api_token.type = ApiTokenType.APP
db.session.add(api_token)
db.session.commit()
session.add(api_token)
session.commit()
def _import_api_tools(
self,
session: Session,
package: MigrationPackage,
target: ImportTarget,
options: ImportOptions,
@@ -436,7 +445,7 @@ class MigrationImportService:
for tool_data in package.tools:
provider_name = self._required_string(tool_data, "provider_name", "api_tool")
schema = self._required_string(tool_data, "schema", "api_tool")
existing = db.session.scalar(
existing = session.scalar(
sa.select(ApiToolProvider).where(
ApiToolProvider.tenant_id == target.tenant_id,
ApiToolProvider.name == provider_name,
@@ -501,7 +510,7 @@ class MigrationImportService:
icon=icon,
)
status = "created"
target_provider = self._find_api_tool_provider(target.tenant_id, provider_name)
target_provider = self._find_api_tool_provider(session, target.tenant_id, provider_name)
if target_provider is not None:
self._record_id_mappings(
id_mapping,
@@ -513,8 +522,8 @@ class MigrationImportService:
)
report_items.append(ResourceReportItem(ResourceType.API_TOOL, provider_name, provider_name, status))
def _find_api_tool_provider(self, tenant_id: str, provider_name: str) -> ApiToolProvider | None:
return db.session.scalar(
def _find_api_tool_provider(self, session: Session, tenant_id: str, provider_name: str) -> ApiToolProvider | None:
return session.scalar(
sa.select(ApiToolProvider).where(
ApiToolProvider.tenant_id == tenant_id,
ApiToolProvider.name == provider_name,
@@ -549,6 +558,7 @@ class MigrationImportService:
def _import_workflow_tools(
self,
session: Session,
package: MigrationPackage,
target: ImportTarget,
options: ImportOptions,
@@ -558,13 +568,13 @@ class MigrationImportService:
) -> None:
if not package.workflow_tools:
return
account = db.session.get(Account, target.operator_id)
account = session.get(Account, target.operator_id)
if account is None:
raise MigrationDataError(f"Operator account not found: {target.operator_id}")
for workflow_tool_data in package.workflow_tools:
app_id = self._optional_string(workflow_tool_data.get("app_id"))
resolved_app_id = id_mapping.get(app_id or "", app_id)
if not resolved_app_id or self._find_existing_app(resolved_app_id, target.tenant_id) is None:
if not resolved_app_id or self._find_existing_app(session, resolved_app_id, target.tenant_id) is None:
report_items.append(
ResourceReportItem(
ResourceType.WORKFLOW_TOOL,
@@ -576,7 +586,7 @@ class MigrationImportService:
)
continue
try:
self._ensure_workflow_app_is_published(target, account, resolved_app_id)
self._ensure_workflow_app_is_published(session, target, account, resolved_app_id)
except Exception as exc:
report_items.append(
ResourceReportItem(
@@ -592,7 +602,7 @@ class MigrationImportService:
tool_name = self._required_string(workflow_tool_data, "name", "workflow_tool")
lookup_workflow_tool_id = workflow_tool_id if options.id_strategy == IdStrategy.PRESERVE_ID else None
existing = self._find_existing_workflow_tool(
target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id
session, target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id
)
if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL:
raise MigrationDataError(f"Workflow tool already exists and conflict_strategy=fail: {tool_name}")
@@ -659,7 +669,7 @@ class MigrationImportService:
)
status = "created"
target_provider = self._find_existing_workflow_tool(
target.tenant_id, import_id or None, tool_name, resolved_app_id
session, target.tenant_id, import_id or None, tool_name, resolved_app_id
)
if target_provider is None:
raise MigrationDataError(f"Workflow tool was not created: {tool_name}")
@@ -675,8 +685,10 @@ class MigrationImportService:
)
report_items.append(ResourceReportItem(ResourceType.WORKFLOW_TOOL, identifier, tool_name, status))
def _ensure_workflow_app_is_published(self, target: ImportTarget, account: Account, app_id: str) -> None:
app = self._find_existing_app(app_id, target.tenant_id)
def _ensure_workflow_app_is_published(
self, session: Session, target: ImportTarget, account: Account, app_id: str
) -> None:
app = self._find_existing_app(session, app_id, target.tenant_id)
if app is None:
raise MigrationDataError(f"Referenced workflow app was not found in target tenant: {app_id}")
if app.workflow_id:
@@ -702,6 +714,7 @@ class MigrationImportService:
def _import_mcp_tools(
self,
session: Session,
package: MigrationPackage,
target: ImportTarget,
options: ImportOptions,
@@ -714,7 +727,7 @@ class MigrationImportService:
server_identifier = self._required_string(mcp_data, "server_identifier", "mcp_tool")
provider_id = self._optional_string(mcp_data.get("id"))
lookup_provider_id = provider_id if options.id_strategy == IdStrategy.PRESERVE_ID else None
existing = self._find_existing_mcp_tool(target.tenant_id, lookup_provider_id, server_identifier)
existing = self._find_existing_mcp_tool(session, target.tenant_id, lookup_provider_id, server_identifier)
if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL:
raise MigrationDataError(f"MCP tool already exists and conflict_strategy=fail: {name}")
if existing is not None and options.conflict_strategy == ConflictStrategy.SKIP:
@@ -730,7 +743,7 @@ class MigrationImportService:
report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, existing.id, name, "skipped"))
continue
service = MCPToolManageService(session=cast(Session, db.session))
service = MCPToolManageService(session=session)
configuration = MCPConfiguration.model_validate(mcp_data.get("configuration") or {})
authentication = (
MCPAuthentication.model_validate(mcp_data["authentication"]) if mcp_data.get("authentication") else None
@@ -752,7 +765,7 @@ class MigrationImportService:
# stored mode (update_provider now defaults to OFF when omitted).
identity_mode=IdentityMode(existing.identity_mode),
)
db.session.commit()
session.commit()
status = "updated"
identifier = existing.id
provider = existing
@@ -770,14 +783,16 @@ class MigrationImportService:
configuration=configuration,
authentication=authentication,
)
created_provider = self._find_existing_mcp_tool(target.tenant_id, lookup_provider_id, server_identifier)
created_provider = self._find_existing_mcp_tool(
session, target.tenant_id, lookup_provider_id, server_identifier
)
if created_provider is None:
raise MigrationDataError(f"MCP provider was not created: {name}")
status = "created"
provider = created_provider
identifier = provider.id
self._restore_mcp_provider_tools(provider, mcp_data)
db.session.commit()
session.commit()
if provider_id:
self._record_id_mappings(
id_mapping,
@@ -797,12 +812,12 @@ class MigrationImportService:
provider.authed = True
def _find_existing_mcp_tool(
self, tenant_id: str, provider_id: str | None, server_identifier: str
self, session: Session, tenant_id: str, provider_id: str | None, server_identifier: str
) -> MCPToolProvider | None:
predicates = [MCPToolProvider.server_identifier == server_identifier]
if self._is_uuid_string(provider_id):
predicates.append(MCPToolProvider.id == provider_id)
return db.session.scalar(
return session.scalar(
sa.select(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant_id, or_(*predicates)).limit(1)
)
@@ -816,26 +831,26 @@ class MigrationImportService:
return True
def _find_existing_workflow_tool(
self, tenant_id: str, workflow_tool_id: str | None, tool_name: str, app_id: str
self, session: Session, tenant_id: str, workflow_tool_id: str | None, tool_name: str, app_id: str
) -> WorkflowToolProvider | None:
predicates = [WorkflowToolProvider.name == tool_name, WorkflowToolProvider.app_id == app_id]
if self._is_uuid_string(workflow_tool_id):
predicates.append(WorkflowToolProvider.id == workflow_tool_id)
return db.session.scalar(
return session.scalar(
sa.select(WorkflowToolProvider)
.where(WorkflowToolProvider.tenant_id == tenant_id, or_(*predicates))
.limit(1)
)
def _preflight_dependency_only_mcp(
self, package: MigrationPackage, target: ImportTarget, report_items: list[ResourceReportItem]
self, session: Session, package: MigrationPackage, target: ImportTarget, report_items: list[ResourceReportItem]
) -> None:
for dependency in package.dependencies:
if dependency.get("kind") != DependencyKind.MCP_TOOL.value:
continue
provider_id = str(dependency.get("provider_id", dependency.get("id", "")))
provider_name = self._optional_string(dependency.get("provider_name") or dependency.get("name"))
existing = self._find_dependency_only_mcp_provider(target.tenant_id, provider_id, provider_name)
existing = self._find_dependency_only_mcp_provider(session, target.tenant_id, provider_id, provider_name)
report_name = f"mcp_tool {provider_name or getattr(existing, 'name', None) or provider_id}"
if existing is not None:
report_items.append(
@@ -864,12 +879,12 @@ class MigrationImportService:
)
def _find_dependency_only_mcp_provider(
self, tenant_id: str, provider_id: str, provider_name: str | None
self, session: Session, tenant_id: str, provider_id: str, provider_name: str | None
) -> MCPToolProvider | None:
predicates = [MCPToolProvider.server_identifier == provider_id]
if self._is_uuid_string(provider_id):
predicates.append(MCPToolProvider.id == provider_id)
return db.session.scalar(
return session.scalar(
sa.select(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant_id, or_(*predicates)).limit(1)
)