fix: load mcp auth at tool call time

This commit is contained in:
lda
2026-06-13 07:07:39 +07:00 Verified
parent 38913cfdcb
commit 4cab4a7d5d
5 changed files with 89 additions and 1 deletions
+2
View File
@@ -34,6 +34,7 @@ def specs_from_discovered_tools(
*,
connection: ConnectionConfig,
auth: AuthRecord | None,
auth_loader: Callable[[], AuthRecord | None] | None = None,
executor: ToolExecutor,
tools: list[DiscoveredTool],
emit_event: Callable[[McpEvent], None] | None = None,
@@ -50,6 +51,7 @@ def specs_from_discovered_tools(
return source_specs_from_discovered_tools(
connection=source_connection,
auth=auth,
auth_loader=auth_loader,
executor=executor,
tools=tools,
emit_event=emit_tool_event if emit_event is not None else None,
@@ -273,6 +273,7 @@ class UpstreamTransportService:
specs = specs_from_discovered_tools(
connection=connection,
auth=auth,
auth_loader=lambda: self.load_connection_auth(connection),
executor=self.tool_executor_for(connection),
tools=capabilities.tools,
emit_event=self.event_sink,
+2
View File
@@ -85,6 +85,7 @@ def specs_from_discovered_tools(
*,
connection: McpSourceConnection,
auth: AuthRecord | None,
auth_loader: Callable[[], AuthRecord | None] | None = None,
executor: ToolExecutor,
tools: list[DiscoveredTool],
emit_event: ToolWrapperEventSink | None = None,
@@ -93,6 +94,7 @@ def specs_from_discovered_tools(
wrap_discovered_tool(
connection=connection,
auth=auth,
auth_loader=auth_loader,
executor=executor,
tool=tool,
emit_event=emit_event,
+4 -1
View File
@@ -1,5 +1,7 @@
from __future__ import annotations
from collections.abc import Callable
from pydantic import BaseModel
from wf_authoring import NodeReturn, NodeSpec
@@ -20,6 +22,7 @@ def wrap_discovered_tool(
*,
connection: McpSourceConnection,
auth: AuthRecord | None,
auth_loader: Callable[[], AuthRecord | None] | None = None,
executor: ToolExecutor,
tool: DiscoveredTool,
emit_event: ToolWrapperEventSink | None = None,
@@ -47,7 +50,7 @@ def wrap_discovered_tool(
)
result = await executor.call_tool(
connection=connection,
auth=auth,
auth=auth_loader() if auth_loader is not None else auth,
tool_name=tool.name,
# Pydantic fills absent optional fields with None, but strict MCP
# servers such as Playwright distinguish omitted from explicit null.
@@ -1,8 +1,11 @@
from __future__ import annotations
from pathlib import Path
from typing import Any
from wf_artifacts import WorkflowDeployment
from wf_authoring import build_async_registry
from wf_core import RuntimeContext
from wf_mcp.broker import WfMcpService
from wf_mcp.broker.service.source_catalog import SourceCatalogService
from wf_mcp.broker.service.upstream_transport import UpstreamTransportService
@@ -12,6 +15,7 @@ from wf_mcp.models import AuthRecord, CatalogSnapshot, ConnectionConfig
from wf_mcp.storage import FileAuthStore, FileCatalogStore, FileStore
from wf_platform import CapabilityBuckets, CapabilitySource, SourcePermissions
from wf_sources_mcp.catalog import DiscoveredTool
from wf_sources_mcp.sdk import ToolCallResult
from ..test_support import FakeAdapter, local_temp_root
from ..workflow_surface.conftest import echo_artifact
@@ -147,6 +151,82 @@ async def test_upstream_transport_refreshes_catalog_directly() -> None:
assert "catalog_refresh_completed" in [event.kind for event in events]
async def test_refreshed_tool_specs_load_auth_at_call_time(tmp_path: Path) -> None:
class RecordingAuthAdapter(FakeAdapter):
def __init__(self) -> None:
self.seen_auth_payloads: list[dict[str, Any]] = []
async def call_tool(
self,
connection,
auth,
tool_name: str,
payload: dict[str, Any],
) -> ToolCallResult:
if auth is not None:
self.seen_auth_payloads.append(dict(auth.payload))
return await super().call_tool(connection, auth, tool_name, payload)
events: list[McpEvent] = []
store = FileStore(tmp_path / "refreshed_spec_auth")
connections = ConnectionRegistry()
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata={**_fake_transport_metadata(), "auth_ref": "demo.creds"},
)
connections.register(connection)
transport = UpstreamTransportService(
auth_store=store,
catalog_store=store,
event_sink=events.append,
)
adapter = RecordingAuthAdapter()
transport.register_adapter("demo", adapter)
transport.save_auth(
AuthRecord(
connection_id="demo.creds",
scheme="bearer",
payload={"token": "old"},
)
)
source_catalog = SourceCatalogService(
store=store,
connection_lookup=connections.get,
connection_list_enabled=connections.list_enabled,
connection_list_all=connections.list_all,
tool_executor_for=transport.tool_executor_for,
load_auth=transport.load_connection_auth,
emit_event=events.append,
)
source_catalog.hydrate_connection_source_from_snapshot(connection)
await transport.refresh_connection_catalog(
connection,
source_catalog=source_catalog,
record_catalog_change_events=lambda source_id, snapshot, reason: None,
)
transport.save_auth(
AuthRecord(
connection_id="demo.creds",
scheme="bearer",
payload={"token": "new"},
)
)
spec = source_catalog.get_qualified_spec("demo.personal.echo_tool")
handler = build_async_registry(spec)[spec.name]
result = await handler(
{"text": "hello"},
RuntimeContext(current_node_id="echo"),
)
assert result["outcome"] == "ok"
assert result["output"]["echoed"] == "hello"
assert adapter.seen_auth_payloads == [{"token": "new"}]
async def test_upstream_transport_live_diagnostics_report_missing_connection() -> None:
transport = UpstreamTransportService(
auth_store=FileStore(local_temp_root() / "upstream_live_missing"),