feat: add neutral auth store boundary

This commit is contained in:
lda
2026-06-06 10:18:16 +07:00 Verified
parent 71facde1a2
commit 34726433a8
15 changed files with 456 additions and 41 deletions
+1 -1
View File
@@ -91,7 +91,7 @@ class WfMcpService:
connection_list_enabled=self.connection_service.list_enabled,
connection_list_all=self.connection_service.list_all,
tool_executor_for=self.upstream.tool_executor_for,
load_auth=self.upstream.load_auth,
load_auth=self.upstream.load_connection_auth,
emit_event=self.events.record_event,
default_catalog_max_age_seconds=self.default_catalog_max_age_seconds,
)
+2 -2
View File
@@ -36,7 +36,7 @@ from .specs import get_qualified_spec, qualify_spec
ConnectionLookup = Callable[[str], ConnectionConfig]
ConnectionList = Callable[[], list[ConnectionConfig]]
ToolExecutorLookup = Callable[[ConnectionConfig], ToolExecutor]
AuthLoader = Callable[[str], AuthRecord | None]
AuthLoader = Callable[[ConnectionConfig], AuthRecord | None]
EventEmitter = Callable[[McpEvent], None]
@@ -270,7 +270,7 @@ class SourceCatalogService:
async def invoke_tool(payload: BaseModel) -> NodeReturn[BaseModel]:
connection = self.connection_lookup(entry.connection_id)
auth = self.load_auth(entry.connection_id)
auth = self.load_auth(connection)
result = await self.tool_executor_for(connection).call_tool(
connection,
auth,
@@ -63,6 +63,19 @@ class UpstreamTransportService:
def load_auth(self, connection_id: str) -> AuthRecord | None:
return self.store.load_auth(connection_id)
def load_connection_auth(self, connection: ConnectionConfig) -> AuthRecord | None:
"""Resolve auth for a connection, preferring explicit source auth_ref.
Legacy MCP auth records are keyed by connection id. New source registry
and neutral config entries carry `auth_ref`; keep both paths until the
old compatibility surface has no callers.
"""
auth_ref = connection.metadata.get("auth_ref")
if isinstance(auth_ref, str):
return self.load_auth(auth_ref)
return self.load_auth(connection.id)
def tool_executor_for(self, connection: ConnectionConfig) -> ToolExecutor:
"""Return the executor used by generated workflow NodeSpecs.
@@ -81,7 +94,7 @@ class UpstreamTransportService:
uri: str,
) -> dict[str, Any]:
adapter = require_adapter(connection, self.adapters)
auth = self.load_auth(connection.id)
auth = self.load_connection_auth(connection)
self.event_sink(
make_event(
"resource_read_started",
@@ -109,7 +122,7 @@ class UpstreamTransportService:
arguments: dict[str, str] | None = None,
) -> dict[str, Any]:
adapter = require_adapter(connection, self.adapters)
auth = self.load_auth(connection.id)
auth = self.load_connection_auth(connection)
self.event_sink(
make_event(
"prompt_get_started",
@@ -137,7 +150,7 @@ class UpstreamTransportService:
params: dict[str, Any] | None = None,
) -> dict[str, Any]:
adapter = require_adapter(connection, self.adapters)
auth = self.load_auth(connection.id)
auth = self.load_connection_auth(connection)
self.event_sink(
make_event(
"raw_method_started",
@@ -165,7 +178,7 @@ class UpstreamTransportService:
params: dict[str, Any] | None = None,
) -> None:
adapter = require_adapter(connection, self.adapters)
auth = self.load_auth(connection.id)
auth = self.load_connection_auth(connection)
self.event_sink(
make_event(
"raw_notification_started",
@@ -193,7 +206,7 @@ class UpstreamTransportService:
default_catalog_max_age_seconds: int = 300,
record_catalog_change_events: Callable[[str, CatalogSnapshot, str], None],
) -> None:
auth = self.load_auth(connection.id)
auth = self.load_connection_auth(connection)
self.event_sink(
make_event(
"catalog_refresh_started",
@@ -295,7 +308,7 @@ class UpstreamTransportService:
continue
try:
adapter = require_adapter(connection, self.adapters)
auth = self.load_auth(source_id)
auth = self.load_connection_auth(connection)
await asyncio.wait_for(
adapter.list_tools(connection, auth),
timeout=LIVE_SOURCE_CHECK_TIMEOUT_SECONDS,