start handling the rest of MCP

its a BIG THING god damn.
This commit is contained in:
lda
2026-04-29 21:39:27 +07:00 Verified
parent 6bff79c304
commit c2681a6976
12 changed files with 492 additions and 227 deletions
+13 -1
View File
@@ -1,9 +1,17 @@
from .adapters import BackendAdapter, DiscoveredTool, ToolCallResult
from .adapters import (
BackendAdapter,
DiscoveredPrompt,
DiscoveredResource,
DiscoveredTool,
ToolCallResult,
)
from .catalog import CombinedCatalog
from .connections import ConnectionRegistry, parse_connection_id, qualify_node_name
from .models import (
AuthRecord,
CatalogNodeEntry,
CatalogPromptEntry,
CatalogResourceEntry,
CatalogSnapshot,
ConnectionConfig,
RawWorkflowPlan,
@@ -17,10 +25,14 @@ __all__ = [
"AuthRecord",
"BackendAdapter",
"CatalogNodeEntry",
"CatalogPromptEntry",
"CatalogResourceEntry",
"CatalogSnapshot",
"CombinedCatalog",
"ConnectionConfig",
"ConnectionRegistry",
"DiscoveredPrompt",
"DiscoveredResource",
"DiscoveredTool",
"FileStore",
"McpSdkAdapter",
+35
View File
@@ -16,6 +16,23 @@ class DiscoveredTool:
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class DiscoveredResource:
uri: str
name: str
description: str | None
mime_type: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class DiscoveredPrompt:
name: str
description: str | None
arguments: list[dict[str, Any]] = field(default_factory=list)
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class ToolCallResult:
outcome: str
@@ -30,6 +47,24 @@ class BackendAdapter(Protocol):
auth: AuthRecord | None,
) -> list[DiscoveredTool]: ...
async def list_resources(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> list[DiscoveredResource]: ...
async def list_prompts(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> list[DiscoveredPrompt]: ...
async def get_connection_metadata(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> dict[str, Any]: ...
async def call_tool(
self,
connection: ConnectionConfig,
+84 -2
View File
@@ -5,14 +5,23 @@ from typing import Any
from wf_authoring import NodeCatalog, NodeSpec
from .adapters import DiscoveredPrompt, DiscoveredResource
from .connections import qualify_node_name
from .models import CatalogNodeEntry, CatalogSnapshot
from .models import (
CatalogNodeEntry,
CatalogPromptEntry,
CatalogResourceEntry,
CatalogSnapshot,
)
def snapshot_from_specs(
connection_id: str,
*,
specs: dict[str, NodeSpec[Any, Any]],
resources: list[DiscoveredResource] | None = None,
prompts: list[DiscoveredPrompt] | None = None,
metadata: dict[str, Any] | None = None,
fetched_at_epoch_ms: int,
max_age_seconds: int,
) -> CatalogSnapshot:
@@ -31,11 +40,37 @@ def snapshot_from_specs(
)
for entry in catalog.entries()
]
resource_entries = [
CatalogResourceEntry(
qualified_name=qualify_node_name(connection_id, resource.name),
connection_id=connection_id,
local_name=resource.name,
uri=resource.uri,
description=resource.description,
mime_type=resource.mime_type,
metadata=resource.metadata,
)
for resource in resources or []
]
prompt_entries = [
CatalogPromptEntry(
qualified_name=qualify_node_name(connection_id, prompt.name),
connection_id=connection_id,
local_name=prompt.name,
description=prompt.description,
arguments=prompt.arguments,
metadata=prompt.metadata,
)
for prompt in prompts or []
]
return CatalogSnapshot(
connection_id=connection_id,
fetched_at_epoch_ms=fetched_at_epoch_ms,
max_age_seconds=max_age_seconds,
nodes=nodes,
resources=resource_entries,
prompts=prompt_entries,
metadata=metadata or {},
)
@@ -49,6 +84,18 @@ class CombinedCatalog:
result.extend(snapshot.nodes)
return sorted(result, key=lambda entry: entry.qualified_name)
def resource_entries(self) -> list[CatalogResourceEntry]:
result: list[CatalogResourceEntry] = []
for snapshot in self.snapshots.values():
result.extend(snapshot.resources)
return sorted(result, key=lambda entry: entry.qualified_name)
def prompt_entries(self) -> list[CatalogPromptEntry]:
result: list[CatalogPromptEntry] = []
for snapshot in self.snapshots.values():
result.extend(snapshot.prompts)
return sorted(result, key=lambda entry: entry.qualified_name)
def as_payload(self) -> dict[str, Any]:
return {
"nodes": [
@@ -62,5 +109,40 @@ class CombinedCatalog:
"output_schema": entry.output_schema,
}
for entry in self.entries()
]
],
"resources": [
{
"qualified_name": entry.qualified_name,
"connection_id": entry.connection_id,
"local_name": entry.local_name,
"uri": entry.uri,
"description": entry.description,
"mime_type": entry.mime_type,
"metadata": entry.metadata,
}
for entry in self.resource_entries()
],
"prompts": [
{
"qualified_name": entry.qualified_name,
"connection_id": entry.connection_id,
"local_name": entry.local_name,
"description": entry.description,
"arguments": entry.arguments,
"metadata": entry.metadata,
}
for entry in self.prompt_entries()
],
"connections": [
{
"connection_id": snapshot.connection_id,
"fetched_at_epoch_ms": snapshot.fetched_at_epoch_ms,
"max_age_seconds": snapshot.max_age_seconds,
"metadata": snapshot.metadata,
}
for snapshot in sorted(
self.snapshots.values(),
key=lambda snapshot: snapshot.connection_id,
)
],
}
+60 -1
View File
@@ -8,9 +8,18 @@ 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 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 .adapters import BackendAdapter, DiscoveredTool, ToolCallResult
from .adapters import (
BackendAdapter,
DiscoveredPrompt,
DiscoveredResource,
DiscoveredTool,
ToolCallResult,
)
from .models import AuthRecord, ConnectionConfig
@@ -39,6 +48,28 @@ def _tool_to_discovered(tool: McpTool) -> DiscoveredTool:
)
def _resource_to_discovered(resource: McpResource) -> DiscoveredResource:
return DiscoveredResource(
uri=str(resource.uri),
name=str(resource.uri),
description=resource.description,
mime_type=resource.mimeType,
metadata=resource.model_dump(by_alias=True),
)
def _prompt_to_discovered(prompt: McpPrompt) -> DiscoveredPrompt:
arguments = [
argument.model_dump(by_alias=True) for argument in prompt.arguments or []
]
return DiscoveredPrompt(
name=prompt.name,
description=prompt.description,
arguments=arguments,
metadata=prompt.model_dump(by_alias=True),
)
def _tool_result_to_call_result(result: McpCallToolResult) -> ToolCallResult:
if result.structuredContent is not None:
output = result.structuredContent
@@ -107,6 +138,34 @@ class McpSdkAdapter(BackendAdapter):
result: ListToolsResult = await session.list_tools()
return [_tool_to_discovered(tool) for tool in result.tools]
async def list_resources(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> list[DiscoveredResource]:
async with self._session(connection, auth) as session:
result: ListResourcesResult = await session.list_resources()
return [_resource_to_discovered(resource) for resource in result.resources]
async def list_prompts(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> list[DiscoveredPrompt]:
async with self._session(connection, auth) as session:
result: ListPromptsResult = await session.list_prompts()
return [_prompt_to_discovered(prompt) for prompt in result.prompts]
async def get_connection_metadata(
self,
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> dict[str, Any]:
return {
"server": connection.server,
"transport": connection.metadata.get("transport", "stdio"),
}
async def call_tool(
self,
connection: ConnectionConfig,
+27
View File
@@ -31,12 +31,36 @@ class CatalogNodeEntry:
output_schema: dict[str, Any]
@dataclass(slots=True)
class CatalogResourceEntry:
qualified_name: str
connection_id: str
local_name: str
uri: str
description: str | None
mime_type: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class CatalogPromptEntry:
qualified_name: str
connection_id: str
local_name: str
description: str | None
arguments: list[dict[str, Any]] = field(default_factory=list)
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class CatalogSnapshot:
connection_id: str
fetched_at_epoch_ms: int
max_age_seconds: int
nodes: list[CatalogNodeEntry] = field(default_factory=list)
resources: list[CatalogResourceEntry] = field(default_factory=list)
prompts: list[CatalogPromptEntry] = field(default_factory=list)
metadata: dict[str, Any] = field(default_factory=dict)
def is_stale(self, now_epoch_ms: int) -> bool:
age_ms = now_epoch_ms - self.fetched_at_epoch_ms
@@ -60,4 +84,7 @@ def dump_catalog_snapshot(snapshot: CatalogSnapshot) -> dict[str, Any]:
"fetched_at_epoch_ms": snapshot.fetched_at_epoch_ms,
"max_age_seconds": snapshot.max_age_seconds,
"nodes": [asdict(node) for node in snapshot.nodes],
"resources": [asdict(resource) for resource in snapshot.resources],
"prompts": [asdict(prompt) for prompt in snapshot.prompts],
"metadata": snapshot.metadata,
}
+13
View File
@@ -93,6 +93,9 @@ class WfMcpService:
auth = self.load_auth(connection_id)
tools = await adapter.list_tools(connection, auth)
resources = await adapter.list_resources(connection, auth)
prompts = await adapter.list_prompts(connection, auth)
metadata = await adapter.get_connection_metadata(connection, auth)
specs = [
wrap_discovered_tool(
connection=connection,
@@ -107,6 +110,16 @@ class WfMcpService:
*specs,
max_age_seconds=max_age_seconds,
)
snapshot = snapshot_from_specs(
connection_id,
specs=self.specs_by_connection.get(connection_id, {}),
resources=resources,
prompts=prompts,
metadata=metadata,
fetched_at_epoch_ms=int(time.time() * 1000),
max_age_seconds=max_age_seconds or self.default_catalog_max_age_seconds,
)
self.store.save_catalog(snapshot)
def compile_plan(self, plan: RawWorkflowPlan) -> Workflow:
node_defs: dict[str, Any] = {}
+25 -10
View File
@@ -1,10 +1,16 @@
from __future__ import annotations
import json
from dataclasses import asdict
from pathlib import Path
from .models import AuthRecord, CatalogNodeEntry, CatalogSnapshot
from .models import (
AuthRecord,
CatalogNodeEntry,
CatalogPromptEntry,
CatalogResourceEntry,
CatalogSnapshot,
dump_catalog_snapshot,
)
class Store:
@@ -44,7 +50,14 @@ class FileStore(Store):
def save_auth(self, record: AuthRecord) -> None:
self._auth_path(record.connection_id).write_text(
json.dumps(asdict(record), indent=2),
json.dumps(
{
"connection_id": record.connection_id,
"scheme": record.scheme,
"payload": record.payload,
},
indent=2,
),
encoding="utf-8",
)
@@ -56,14 +69,8 @@ class FileStore(Store):
return AuthRecord(**data)
def save_catalog(self, snapshot: CatalogSnapshot) -> None:
payload = {
"connection_id": snapshot.connection_id,
"fetched_at_epoch_ms": snapshot.fetched_at_epoch_ms,
"max_age_seconds": snapshot.max_age_seconds,
"nodes": [asdict(node) for node in snapshot.nodes],
}
self._catalog_path(snapshot.connection_id).write_text(
json.dumps(payload, indent=2),
json.dumps(dump_catalog_snapshot(snapshot), indent=2),
encoding="utf-8",
)
@@ -77,4 +84,12 @@ class FileStore(Store):
fetched_at_epoch_ms=data["fetched_at_epoch_ms"],
max_age_seconds=data["max_age_seconds"],
nodes=[CatalogNodeEntry(**node) for node in data.get("nodes", [])],
resources=[
CatalogResourceEntry(**resource)
for resource in data.get("resources", [])
],
prompts=[
CatalogPromptEntry(**prompt) for prompt in data.get("prompts", [])
],
metadata=data.get("metadata", {}),
)