Files
lda-wf/tests/wf_mcp/service/test_upstream_transport.py
T
2026-07-30 01:27:46 +07:00

748 lines
24 KiB
Python

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
from wf_mcp.connections import ConnectionRegistry
from wf_mcp.events import McpEvent
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
def _fake_transport_metadata() -> dict[str, object]:
return {"transport": "stdio", "command": "fake-mcp-server"}
def _transport(root: Path) -> UpstreamTransportService:
events: list[McpEvent] = []
return UpstreamTransportService(
auth_store=FileStore(root),
catalog_store=FileStore(root),
event_sink=events.append,
)
def test_upstream_transport_registers_adapter() -> None:
events: list[McpEvent] = []
transport = UpstreamTransportService(
auth_store=FileStore(local_temp_root() / "upstream_adapter"),
catalog_store=FileStore(local_temp_root() / "upstream_adapter"),
event_sink=events.append,
)
adapter = FakeAdapter()
transport.register_adapter("demo", adapter)
assert transport.adapters["demo"] is adapter
def test_upstream_transport_saves_and_loads_auth_with_event() -> None:
events: list[McpEvent] = []
transport = UpstreamTransportService(
auth_store=FileStore(local_temp_root() / "upstream_auth"),
catalog_store=FileStore(local_temp_root() / "upstream_auth"),
event_sink=events.append,
)
record = AuthRecord(connection_id="demo.personal", scheme="bearer")
transport.save_auth(record)
loaded = transport.load_auth("demo.personal")
assert loaded is not None
assert loaded.connection_id == "demo.personal"
assert events[-1].kind == "auth_saved"
assert events[-1].connection_id == "demo.personal"
def test_wfmcpservice_uses_upstream_transport_for_adapters_and_auth() -> None:
service = WfMcpService(store=FileStore(local_temp_root() / "service_upstream"))
adapter = FakeAdapter()
service.register_adapter("demo", adapter)
service.save_auth(AuthRecord(connection_id="demo.personal", scheme="bearer"))
assert service.upstream.adapters["demo"] is adapter
assert service.adapters is service.upstream.adapters
assert service.load_auth("demo.personal") is not None
assert service.list_events()[-1].kind == "auth_saved"
async def test_upstream_transport_invokes_raw_method_and_records_events() -> None:
events: list[McpEvent] = []
connections = ConnectionRegistry()
connections.register(
ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
)
transport = UpstreamTransportService(
auth_store=FileStore(local_temp_root() / "upstream_raw_method"),
catalog_store=FileStore(local_temp_root() / "upstream_raw_method"),
event_sink=events.append,
)
transport.register_adapter("demo", FakeAdapter())
result = await transport.invoke_method(
connections.get("demo.personal"),
"demo.echo",
params={"text": "hello"},
)
assert result["echoed"] == "hello"
assert [event.kind for event in events] == [
"raw_method_started",
"raw_method_completed",
]
async def test_upstream_transport_refreshes_catalog_directly() -> None:
events: list[McpEvent] = []
store = FileStore(local_temp_root() / "upstream_refresh")
connections = ConnectionRegistry()
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
connections.register(connection)
transport = UpstreamTransportService(
auth_store=store,
catalog_store=store,
event_sink=events.append,
)
transport.register_adapter("demo", FakeAdapter())
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,
)
snapshot = store.load_catalog("demo.personal")
assert snapshot is not None
assert len(snapshot.nodes) >= 1
assert "catalog_refresh_started" in [event.kind for event in events]
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"),
catalog_store=FileStore(local_temp_root() / "upstream_live_missing"),
event_sink=lambda event: None,
)
def _raise_missing_connection(connection_id: str) -> ConnectionConfig:
raise KeyError(connection_id)
source_catalog = SourceCatalogService(
store=transport.catalog_store,
connection_lookup=_raise_missing_connection,
connection_list_enabled=list,
connection_list_all=list,
tool_executor_for=transport.tool_executor_for,
load_auth=transport.load_connection_auth,
emit_event=lambda event: None,
)
source_catalog.register_capability_source(
CapabilitySource(
id="demo.personal",
kind="connection",
permissions=SourcePermissions(calls_upstream=True),
capabilities=CapabilityBuckets(),
)
)
artifact = echo_artifact()
deployment = WorkflowDeployment(
id="echo.personal",
artifact_id="echo",
artifact_version=1,
bindings=[{"logical_source": "demo", "concrete_source": "demo.personal"}],
)
diagnostics = await transport.deployment_diagnostics(
deployment=deployment,
artifacts=[artifact],
source_catalog=source_catalog,
)
assert diagnostics[0].code == "source_unreachable"
assert diagnostics[0].bound_source == "demo.personal"
def test_upstream_load_connection_auth_prefers_auth_ref(tmp_path: Path) -> None:
service = _transport(tmp_path)
service.save_auth(
AuthRecord(
connection_id="github.creds",
scheme="bearer",
payload={"token": "secret"},
)
)
service.save_auth(
AuthRecord(
connection_id="github.work",
scheme="bearer",
payload={"token": "wrong"},
)
)
connection = ConnectionConfig(
id="github.work",
server="github",
account="work",
metadata={**_fake_transport_metadata(), "auth_ref": "github.creds"},
)
assert service.load_connection_auth(connection) == AuthRecord(
connection_id="github.creds",
scheme="bearer",
payload={"token": "secret"},
)
def test_upstream_load_connection_auth_falls_back_to_connection_id(
tmp_path: Path,
) -> None:
service = _transport(tmp_path)
service.save_auth(
AuthRecord(
connection_id="github.work",
scheme="bearer",
payload={"token": "legacy"},
)
)
connection = ConnectionConfig(
id="github.work",
server="github",
account="work",
)
assert service.load_connection_auth(connection) == AuthRecord(
connection_id="github.work",
scheme="bearer",
payload={"token": "legacy"},
)
def test_upstream_load_connection_auth_ignores_non_string_auth_ref(
tmp_path: Path,
) -> None:
service = _transport(tmp_path)
service.save_auth(
AuthRecord(
connection_id="github.work",
scheme="bearer",
payload={"token": "legacy"},
)
)
connection = ConnectionConfig(
id="github.work",
server="github",
account="work",
metadata={"auth_ref": 123},
)
assert service.load_connection_auth(connection) == AuthRecord(
connection_id="github.work",
scheme="bearer",
payload={"token": "legacy"},
)
async def test_upstream_transport_live_diagnostics_report_missing_auth_ref(
tmp_path: Path,
) -> None:
events: list[McpEvent] = []
store = FileStore(tmp_path)
connections = ConnectionRegistry()
connection = ConnectionConfig(
id="github.work",
server="github",
account="work",
metadata={**_fake_transport_metadata(), "auth_ref": "github.creds"},
)
connections.register(connection)
transport = UpstreamTransportService(
auth_store=store,
catalog_store=store,
event_sink=events.append,
)
transport.register_adapter("demo", FakeAdapter())
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.register_capability_source(
CapabilitySource(
id="github.work",
kind="connection",
permissions=SourcePermissions(calls_upstream=True),
capabilities=CapabilityBuckets(),
)
)
artifact = echo_artifact()
deployment = WorkflowDeployment(
id="echo.personal",
artifact_id="echo",
artifact_version=1,
bindings=[{"logical_source": "demo", "concrete_source": "github.work"}],
)
diagnostics = await transport.deployment_diagnostics(
deployment=deployment,
artifacts=[artifact],
source_catalog=source_catalog,
)
assert diagnostics[0].code == "auth_not_found"
assert diagnostics[0].bound_source == "github.work"
assert "github.creds" in diagnostics[0].message
def test_upstream_transport_uses_separate_auth_and_catalog_stores(
tmp_path: Path,
) -> None:
auth_store = FileAuthStore(tmp_path / "auth")
catalog_store = FileCatalogStore(tmp_path / "catalog")
events = []
transport = UpstreamTransportService(
auth_store=auth_store,
catalog_store=catalog_store,
event_sink=events.append,
)
record = AuthRecord(connection_id="demo.personal", scheme="bearer")
transport.save_auth(record)
snapshot = CatalogSnapshot(
connection_id="demo.personal",
fetched_at_epoch_ms=1,
max_age_seconds=300,
nodes=[],
resources=[],
prompts=[],
metadata={},
)
transport.catalog_store.save_catalog(snapshot)
assert (tmp_path / "auth" / "auth" / "demo.personal.json").exists()
assert (tmp_path / "catalog" / "catalog" / "demo.personal.json").exists()
class _StatefulRuntime:
def __init__(self) -> None:
self.resources: list[tuple[str, str]] = []
self.prompts: list[tuple[str, str, dict[str, str] | None]] = []
self.tools_called: list[tuple[str, str]] = []
self.methods_invoked: list[tuple[str, str, str, dict[str, object] | None]] = []
self.notifications_sent: list[
tuple[str, str, str, dict[str, object] | None]
] = []
async def call_tool(self, connection, auth, tool_name, payload):
raise AssertionError("not used by these tests")
async def read_resource(self, connection, auth, uri: str):
self.resources.append((connection.id, uri))
return {"contents": [{"uri": uri, "text": "stateful resource"}]}
async def get_prompt(
self,
connection,
auth,
prompt_name: str,
arguments: dict[str, str] | None = None,
):
self.prompts.append((connection.id, prompt_name, arguments))
return {
"messages": [
{
"role": "user",
"content": {"type": "text", "text": "stateful prompt"},
}
]
}
async def list_tools(self, connection, auth):
self.tools_called.append((connection.id, connection.provider))
return [
DiscoveredTool(
name="stateful_tool",
title="Stateful Tool",
description="A stateful tool",
input_schema={"type": "object"},
output_schema={"type": "object"},
)
]
async def list_resources(self, connection, auth):
return []
async def list_prompts(self, connection, auth):
return []
async def get_connection_metadata(self, connection, auth):
return {"server": connection.provider, "transport": "stdio"}
async def invoke_method(self, connection, auth, method, params=None):
self.methods_invoked.append(
(connection.id, connection.provider, method, params)
)
return {"echoed": (params or {}).get("text", "")}
async def send_notification(self, connection, auth, method, params=None):
self.notifications_sent.append(
(connection.id, connection.provider, method, params)
)
class _ExplodingContentAdapter(FakeAdapter):
async def read_resource(self, connection, auth, uri):
raise AssertionError("adapter read_resource should not be used")
async def get_prompt(self, connection, auth, prompt_name, arguments=None):
raise AssertionError("adapter get_prompt should not be used")
async def test_upstream_transport_prefers_stateful_runtime_for_resource_reads(
tmp_path: Path,
) -> None:
events: list[McpEvent] = []
runtime = _StatefulRuntime()
transport = UpstreamTransportService(
auth_store=FileStore(tmp_path),
catalog_store=FileStore(tmp_path),
event_sink=events.append,
stateful_runtime=runtime,
)
transport.register_adapter("demo", _ExplodingContentAdapter())
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
result = await transport.read_resource(
connection,
"demo.personal.resource.welcome",
"fixture://docs/welcome",
)
assert result["contents"][0]["text"] == "stateful resource"
assert runtime.resources == [("demo.personal", "fixture://docs/welcome")]
assert [event.kind for event in events] == [
"resource_read_started",
"resource_read_completed",
]
async def test_upstream_transport_prefers_stateful_runtime_for_prompts(
tmp_path: Path,
) -> None:
events: list[McpEvent] = []
runtime = _StatefulRuntime()
transport = UpstreamTransportService(
auth_store=FileStore(tmp_path),
catalog_store=FileStore(tmp_path),
event_sink=events.append,
stateful_runtime=runtime,
)
transport.register_adapter("demo", _ExplodingContentAdapter())
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
result = await transport.render_prompt(
connection,
"demo.personal.prompt.summarize",
"prompt.summarize",
{"text": "hello"},
)
assert result["messages"][0]["content"]["text"] == "stateful prompt"
assert runtime.prompts == [("demo.personal", "prompt.summarize", {"text": "hello"})]
assert [event.kind for event in events] == [
"prompt_get_started",
"prompt_get_completed",
]
async def test_upstream_transport_prefers_stateful_runtime_for_invoke_method(
tmp_path: Path,
) -> None:
events: list[McpEvent] = []
runtime = _StatefulRuntime()
transport = UpstreamTransportService(
auth_store=FileStore(tmp_path),
catalog_store=FileStore(tmp_path),
event_sink=events.append,
stateful_runtime=runtime,
)
transport.register_adapter("demo", _ExplodingContentAdapter())
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
result = await transport.invoke_method(
connection,
"demo.echo",
params={"text": "hello"},
)
assert result["echoed"] == "hello"
assert runtime.methods_invoked == [
("demo.personal", "demo", "demo.echo", {"text": "hello"})
]
assert [event.kind for event in events] == [
"raw_method_started",
"raw_method_completed",
]
async def test_upstream_transport_prefers_stateful_runtime_for_send_notification(
tmp_path: Path,
) -> None:
events: list[McpEvent] = []
runtime = _StatefulRuntime()
transport = UpstreamTransportService(
auth_store=FileStore(tmp_path),
catalog_store=FileStore(tmp_path),
event_sink=events.append,
stateful_runtime=runtime,
)
transport.register_adapter("demo", _ExplodingContentAdapter())
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
await transport.send_notification(
connection,
"notifications/test",
params={"data": "value"},
)
assert runtime.notifications_sent == [
("demo.personal", "demo", "notifications/test", {"data": "value"})
]
assert [event.kind for event in events] == [
"raw_notification_started",
"raw_notification_completed",
]
async def test_upstream_transport_prefers_stateful_runtime_for_catalog_refresh(
tmp_path: Path,
) -> None:
events: list[McpEvent] = []
store = FileStore(tmp_path)
connections = ConnectionRegistry()
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
connections.register(connection)
runtime = _StatefulRuntime()
transport = UpstreamTransportService(
auth_store=store,
catalog_store=store,
event_sink=events.append,
stateful_runtime=runtime,
)
transport.register_adapter("demo", _ExplodingContentAdapter())
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,
)
assert runtime.tools_called == [("demo.personal", "demo")]
assert "catalog_refresh_started" in [event.kind for event in events]
assert "catalog_refresh_completed" in [event.kind for event in events]
class _ExplodingAdapterForDiagnostics(FakeAdapter):
async def list_tools(self, connection, auth):
raise AssertionError("adapter list_tools should not be used in diagnostics")
async def test_upstream_transport_prefers_stateful_runtime_for_deployment_diagnostics(
tmp_path: Path,
) -> None:
runtime = _StatefulRuntime()
connections = ConnectionRegistry()
connection = ConnectionConfig(
id="demo.personal",
server="demo",
account="personal",
metadata=_fake_transport_metadata(),
)
connections.register(connection)
transport = UpstreamTransportService(
auth_store=FileStore(tmp_path),
catalog_store=FileStore(tmp_path),
event_sink=lambda event: None,
stateful_runtime=runtime,
)
transport.register_adapter("demo", _ExplodingAdapterForDiagnostics())
source_catalog = SourceCatalogService(
store=transport.catalog_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=lambda event: None,
)
source_catalog.register_capability_source(
CapabilitySource(
id="demo.personal",
kind="connection",
permissions=SourcePermissions(calls_upstream=True),
capabilities=CapabilityBuckets(),
)
)
artifact = echo_artifact()
deployment = WorkflowDeployment(
id="echo.personal",
artifact_id="echo",
artifact_version=1,
bindings=[{"logical_source": "demo", "concrete_source": "demo.personal"}],
)
diagnostics = await transport.deployment_diagnostics(
deployment=deployment,
artifacts=[artifact],
source_catalog=source_catalog,
)
assert diagnostics == []
assert runtime.tools_called == [("demo.personal", "demo")]