proxy supports more methods
This commit is contained in:
@@ -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()
|
||||
try:
|
||||
result = asyncio.run(
|
||||
|
||||
@@ -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[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
|
||||
|
||||
@@ -165,6 +165,69 @@ class FakeAdapter:
|
||||
"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(
|
||||
self,
|
||||
connection: ConnectionConfig,
|
||||
|
||||
@@ -65,6 +65,37 @@ class BackendAdapter(Protocol):
|
||||
auth: AuthRecord | None,
|
||||
) -> 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(
|
||||
self,
|
||||
connection: ConnectionConfig,
|
||||
|
||||
@@ -96,6 +96,18 @@ class CombinedCatalog:
|
||||
result.extend(snapshot.prompts)
|
||||
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]:
|
||||
return {
|
||||
"nodes": [
|
||||
|
||||
@@ -4,14 +4,17 @@ from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from mcp import ClientResult
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.client.stdio import StdioServerParameters, stdio_client
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.types import CallToolResult as McpCallToolResult
|
||||
from mcp.types import ClientNotification, ClientRequest
|
||||
from mcp.types import ListPromptsResult, ListResourcesResult
|
||||
from mcp.types import ListToolsResult, Tool as McpTool
|
||||
from mcp.types import Prompt as McpPrompt
|
||||
from mcp.types import Resource as McpResource
|
||||
from pydantic import AnyUrl
|
||||
|
||||
from .adapters import (
|
||||
BackendAdapter,
|
||||
@@ -85,7 +88,6 @@ def _tool_result_to_call_result(result: McpCallToolResult) -> ToolCallResult:
|
||||
meta=result.meta or {},
|
||||
)
|
||||
|
||||
|
||||
class McpSdkAdapter(BackendAdapter):
|
||||
@asynccontextmanager
|
||||
async def _session(
|
||||
@@ -168,6 +170,53 @@ class McpSdkAdapter(BackendAdapter):
|
||||
"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(
|
||||
self,
|
||||
connection: ConnectionConfig,
|
||||
|
||||
+171
-1
@@ -12,7 +12,14 @@ from .catalog import CombinedCatalog, snapshot_from_specs
|
||||
from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name
|
||||
from .discovery import discover_connection_capabilities, specs_from_discovered_tools
|
||||
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
|
||||
|
||||
|
||||
@@ -103,6 +110,169 @@ class WfMcpService:
|
||||
snapshots[connection.id] = snapshot
|
||||
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(
|
||||
self,
|
||||
connection_id: str,
|
||||
|
||||
Reference in New Issue
Block a user