proxy supports more methods

This commit is contained in:
lda
2026-04-29 23:14:28 +07:00 Verified
parent 1e41a94fc6
commit 2e2ea468c2
7 changed files with 433 additions and 2 deletions
+19
View File
@@ -59,6 +59,25 @@ def test_mcp_sdk_adapter_lists_and_calls_stdio_tool() -> None:
} }
] ]
resource_result = asyncio.run(
service.read_resource("fixture.personal.resource.welcome")
)
prompt_result = asyncio.run(
service.render_prompt(
"fixture.personal.prompt.summarize",
arguments={"text": "hello"},
)
)
assert (
resource_result["contents"][0]["text"] == "Welcome from the fixture MCP server."
)
assert (
prompt_result["messages"][0]["content"]["text"]
== "Summarize this text:\n\nhello"
)
ping_result = asyncio.run(service.invoke_method("fixture.personal", "ping"))
assert ping_result == {}
adapter = McpSdkAdapter() adapter = McpSdkAdapter()
try: try:
result = asyncio.run( result = asyncio.run(
+87
View File
@@ -224,3 +224,90 @@ def test_service_records_tool_call_events() -> None:
] ]
assert tool_events[0].capability_id == "demo.personal.echo_tool" assert tool_events[0].capability_id == "demo.personal.echo_tool"
assert tool_events[1].payload["outcome"] == "ok" assert tool_events[1].payload["outcome"] == "ok"
def test_service_can_inspect_resources_and_prompts() -> None:
service = WfMcpService(store=FileStore(local_temp_root() / "inspect_store"))
service.register_connection(
ConnectionConfig(id="demo.personal", server="demo", account="personal")
)
service.register_adapter("demo", FakeAdapter())
asyncio.run(service.refresh_connection_catalog("demo.personal"))
resources = service.list_resources(connection_id="demo.personal")
prompts = service.list_prompts(connection_id="demo.personal")
assert [resource.qualified_name for resource in resources] == [
"demo.personal.resource.welcome"
]
assert [prompt.qualified_name for prompt in prompts] == [
"demo.personal.prompt.summarize"
]
resource = service.get_resource("demo.personal.resource.welcome")
prompt = service.get_prompt("demo.personal.prompt.summarize")
assert resource.uri == "demo://docs/welcome"
assert prompt.arguments[0]["name"] == "text"
def test_service_can_proxy_resource_reads_and_prompt_gets() -> None:
service = WfMcpService(store=FileStore(local_temp_root() / "proxy_store"))
service.register_connection(
ConnectionConfig(id="demo.personal", server="demo", account="personal")
)
service.register_adapter("demo", FakeAdapter())
asyncio.run(service.refresh_connection_catalog("demo.personal"))
resource_result = asyncio.run(
service.read_resource("demo.personal.resource.welcome")
)
prompt_result = asyncio.run(
service.render_prompt(
"demo.personal.prompt.summarize",
arguments={"text": "hello world"},
)
)
assert (
resource_result["contents"][0]["text"]
== "Welcome from the fake adapter resource."
)
assert (
prompt_result["messages"][0]["content"]["text"]
== "Summarize this text:\n\nhello world"
)
event_kinds = [event.kind for event in service.list_events()]
assert "resource_read_started" in event_kinds
assert "resource_read_completed" in event_kinds
assert "prompt_get_started" in event_kinds
assert "prompt_get_completed" in event_kinds
def test_service_can_invoke_raw_method_and_notification() -> None:
service = WfMcpService(store=FileStore(local_temp_root() / "raw_store"))
service.register_connection(
ConnectionConfig(id="demo.personal", server="demo", account="personal")
)
service.register_adapter("demo", FakeAdapter())
result = asyncio.run(
service.invoke_method("demo.personal", "demo.echo", params={"text": "hello"})
)
asyncio.run(
service.send_notification(
"demo.personal",
"notifications/progress",
params={"progress": 1},
)
)
assert result == {"echoed": "hello"}
event_kinds = [event.kind for event in service.list_events()]
assert "raw_method_started" in event_kinds
assert "raw_method_completed" in event_kinds
assert "raw_notification_started" in event_kinds
assert "raw_notification_completed" in event_kinds
+63
View File
@@ -165,6 +165,69 @@ class FakeAdapter:
"auth_scheme": auth.scheme if auth is not None else None, "auth_scheme": auth.scheme if auth is not None else None,
} }
async def read_resource(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
uri: str,
) -> dict[str, Any]:
if uri != "demo://docs/welcome":
raise KeyError(uri)
return {
"contents": [
{
"uri": uri,
"mimeType": "text/plain",
"text": "Welcome from the fake adapter resource.",
}
]
}
async def get_prompt(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
prompt_name: str,
arguments: dict[str, str] | None = None,
) -> dict[str, Any]:
if prompt_name != "prompt.summarize":
raise KeyError(prompt_name)
text = (arguments or {}).get("text", "")
return {
"description": "Summarize text",
"messages": [
{
"role": "user",
"content": {
"type": "text",
"text": f"Summarize this text:\n\n{text}",
},
}
],
}
async def invoke_method(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> dict[str, Any]:
if method == "ping":
return {}
if method == "demo.echo":
return {"echoed": (params or {}).get("text", "")}
raise KeyError(method)
async def send_notification(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> None:
return None
async def call_tool( async def call_tool(
self, self,
connection: ConnectionConfig, connection: ConnectionConfig,
+31
View File
@@ -65,6 +65,37 @@ class BackendAdapter(Protocol):
auth: AuthRecord | None, auth: AuthRecord | None,
) -> dict[str, Any]: ... ) -> dict[str, Any]: ...
async def read_resource(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
uri: str,
) -> dict[str, Any]: ...
async def get_prompt(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
prompt_name: str,
arguments: dict[str, str] | None = None,
) -> dict[str, Any]: ...
async def invoke_method(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> dict[str, Any]: ...
async def send_notification(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> None: ...
async def call_tool( async def call_tool(
self, self,
connection: ConnectionConfig, connection: ConnectionConfig,
+12
View File
@@ -96,6 +96,18 @@ class CombinedCatalog:
result.extend(snapshot.prompts) result.extend(snapshot.prompts)
return sorted(result, key=lambda entry: entry.qualified_name) return sorted(result, key=lambda entry: entry.qualified_name)
def find_resource(self, qualified_name: str) -> CatalogResourceEntry | None:
for entry in self.resource_entries():
if entry.qualified_name == qualified_name:
return entry
return None
def find_prompt(self, qualified_name: str) -> CatalogPromptEntry | None:
for entry in self.prompt_entries():
if entry.qualified_name == qualified_name:
return entry
return None
def as_payload(self) -> dict[str, Any]: def as_payload(self) -> dict[str, Any]:
return { return {
"nodes": [ "nodes": [
+50 -1
View File
@@ -4,14 +4,17 @@ from contextlib import asynccontextmanager
from typing import Any from typing import Any
import httpx import httpx
from mcp import ClientResult
from mcp.client.session import ClientSession from mcp.client.session import ClientSession
from mcp.client.stdio import StdioServerParameters, stdio_client from mcp.client.stdio import StdioServerParameters, stdio_client
from mcp.client.streamable_http import streamable_http_client from mcp.client.streamable_http import streamable_http_client
from mcp.types import CallToolResult as McpCallToolResult from mcp.types import CallToolResult as McpCallToolResult
from mcp.types import ClientNotification, ClientRequest
from mcp.types import ListPromptsResult, ListResourcesResult from mcp.types import ListPromptsResult, ListResourcesResult
from mcp.types import ListToolsResult, Tool as McpTool from mcp.types import ListToolsResult, Tool as McpTool
from mcp.types import Prompt as McpPrompt from mcp.types import Prompt as McpPrompt
from mcp.types import Resource as McpResource from mcp.types import Resource as McpResource
from pydantic import AnyUrl
from .adapters import ( from .adapters import (
BackendAdapter, BackendAdapter,
@@ -85,7 +88,6 @@ def _tool_result_to_call_result(result: McpCallToolResult) -> ToolCallResult:
meta=result.meta or {}, meta=result.meta or {},
) )
class McpSdkAdapter(BackendAdapter): class McpSdkAdapter(BackendAdapter):
@asynccontextmanager @asynccontextmanager
async def _session( async def _session(
@@ -168,6 +170,53 @@ class McpSdkAdapter(BackendAdapter):
"transport": connection.metadata.get("transport", "stdio"), "transport": connection.metadata.get("transport", "stdio"),
} }
async def read_resource(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
uri: str,
) -> dict[str, Any]:
async with self._session(connection, auth) as session:
result = await session.read_resource(AnyUrl(uri))
return result.model_dump(by_alias=True, mode="json", exclude_none=True)
async def get_prompt(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
prompt_name: str,
arguments: dict[str, str] | None = None,
) -> dict[str, Any]:
async with self._session(connection, auth) as session:
result = await session.get_prompt(prompt_name, arguments)
return result.model_dump(by_alias=True, mode="json", exclude_none=True)
async def invoke_method(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> dict[str, Any]:
async with self._session(connection, auth) as session:
result = await session.send_request(
ClientRequest.model_validate({"method": method, "params": params}),
ClientResult,
)
return result.model_dump(by_alias=True, mode="json", exclude_none=True)
async def send_notification(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> None:
async with self._session(connection, auth) as session:
await session.send_notification(
ClientNotification.model_validate({"method": method, "params": params})
)
async def call_tool( async def call_tool(
self, self,
connection: ConnectionConfig, connection: ConnectionConfig,
+171 -1
View File
@@ -12,7 +12,14 @@ from .catalog import CombinedCatalog, snapshot_from_specs
from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name
from .discovery import discover_connection_capabilities, specs_from_discovered_tools from .discovery import discover_connection_capabilities, specs_from_discovered_tools
from .events import McpEvent, make_event from .events import McpEvent, make_event
from .models import AuthRecord, CatalogSnapshot, ConnectionConfig, RawWorkflowPlan from .models import (
AuthRecord,
CatalogPromptEntry,
CatalogResourceEntry,
CatalogSnapshot,
ConnectionConfig,
RawWorkflowPlan,
)
from .store import Store from .store import Store
@@ -103,6 +110,169 @@ class WfMcpService:
snapshots[connection.id] = snapshot snapshots[connection.id] = snapshot
return CombinedCatalog(snapshots=snapshots) return CombinedCatalog(snapshots=snapshots)
def get_connection_snapshot(self, connection_id: str) -> CatalogSnapshot | None:
self.connections.get(connection_id)
return self.store.load_catalog(connection_id)
def list_resources(
self,
*,
connection_id: str | None = None,
) -> list[CatalogResourceEntry]:
if connection_id is None:
return self.get_catalog().resource_entries()
snapshot = self.get_connection_snapshot(connection_id)
if snapshot is None:
return []
return sorted(snapshot.resources, key=lambda entry: entry.qualified_name)
def list_prompts(
self,
*,
connection_id: str | None = None,
) -> list[CatalogPromptEntry]:
if connection_id is None:
return self.get_catalog().prompt_entries()
snapshot = self.get_connection_snapshot(connection_id)
if snapshot is None:
return []
return sorted(snapshot.prompts, key=lambda entry: entry.qualified_name)
def get_resource(self, qualified_name: str) -> CatalogResourceEntry:
entry = self.get_catalog().find_resource(qualified_name)
if entry is None:
raise KeyError(f"unknown resource {qualified_name!r}")
return entry
def get_prompt(self, qualified_name: str) -> CatalogPromptEntry:
entry = self.get_catalog().find_prompt(qualified_name)
if entry is None:
raise KeyError(f"unknown prompt {qualified_name!r}")
return entry
async def read_resource(self, qualified_name: str) -> dict[str, Any]:
resource = self.get_resource(qualified_name)
connection = self.connections.get(resource.connection_id)
adapter = self.adapters.get(connection.server)
if adapter is None:
raise KeyError(f"no adapter registered for server {connection.server!r}")
auth = self.load_auth(resource.connection_id)
self._record_event(
make_event(
"resource_read_started",
connection_id=resource.connection_id,
capability_id=qualified_name,
payload={"uri": resource.uri},
)
)
result = await adapter.read_resource(connection, auth, resource.uri)
self._record_event(
make_event(
"resource_read_completed",
connection_id=resource.connection_id,
capability_id=qualified_name,
payload={"uri": resource.uri},
)
)
return result
async def invoke_method(
self,
connection_id: str,
method: str,
*,
params: dict[str, Any] | None = None,
) -> dict[str, Any]:
connection = self.connections.get(connection_id)
adapter = self.adapters.get(connection.server)
if adapter is None:
raise KeyError(f"no adapter registered for server {connection.server!r}")
auth = self.load_auth(connection_id)
self._record_event(
make_event(
"raw_method_started",
connection_id=connection_id,
capability_id=method,
payload={"params": params or {}},
)
)
result = await adapter.invoke_method(connection, auth, method, params)
self._record_event(
make_event(
"raw_method_completed",
connection_id=connection_id,
capability_id=method,
payload={"result_keys": sorted(result.keys())},
)
)
return result
async def send_notification(
self,
connection_id: str,
method: str,
*,
params: dict[str, Any] | None = None,
) -> None:
connection = self.connections.get(connection_id)
adapter = self.adapters.get(connection.server)
if adapter is None:
raise KeyError(f"no adapter registered for server {connection.server!r}")
auth = self.load_auth(connection_id)
self._record_event(
make_event(
"raw_notification_started",
connection_id=connection_id,
capability_id=method,
payload={"params": params or {}},
)
)
await adapter.send_notification(connection, auth, method, params)
self._record_event(
make_event(
"raw_notification_completed",
connection_id=connection_id,
capability_id=method,
payload={},
)
)
async def render_prompt(
self,
qualified_name: str,
*,
arguments: dict[str, str] | None = None,
) -> dict[str, Any]:
prompt = self.get_prompt(qualified_name)
connection = self.connections.get(prompt.connection_id)
adapter = self.adapters.get(connection.server)
if adapter is None:
raise KeyError(f"no adapter registered for server {connection.server!r}")
auth = self.load_auth(prompt.connection_id)
self._record_event(
make_event(
"prompt_get_started",
connection_id=prompt.connection_id,
capability_id=qualified_name,
payload={"argument_keys": sorted((arguments or {}).keys())},
)
)
result = await adapter.get_prompt(
connection,
auth,
prompt.local_name,
arguments,
)
self._record_event(
make_event(
"prompt_get_completed",
connection_id=prompt.connection_id,
capability_id=qualified_name,
payload={"argument_keys": sorted((arguments or {}).keys())},
)
)
return result
async def refresh_connection_catalog( async def refresh_connection_catalog(
self, self,
connection_id: str, connection_id: str,