refactor: route all mcp operations through runtime pool

This commit is contained in:
lda
2026-06-08 18:29:22 +07:00 Verified
parent 929103219e
commit 0a8cc4190f
14 changed files with 961 additions and 30 deletions
+42 -1
View File
@@ -11,7 +11,7 @@ from typing import Any, Generic, TypeVar
from mcp.client.session import ClientSession
from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool
from wf_sources_mcp.client import McpSourceClient, open_mcp_session
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.sdk import ToolCallResult
@@ -47,6 +47,10 @@ class PersistentSessionFactory:
get_prompt_callback=owner.get_prompt,
list_resources_callback=owner.list_resources,
list_prompts_callback=owner.list_prompts,
list_tools_callback=owner.list_tools,
get_connection_metadata_callback=owner.get_connection_metadata,
invoke_method_callback=owner.invoke_method,
send_notification_callback=owner.send_notification,
close_callback=owner.close,
)
@@ -180,6 +184,43 @@ class _SessionOwner:
run=lambda client: client.list_prompts(),
)
async def list_tools(self) -> list[DiscoveredTool]:
"""Submit tool listing through the generic owner-task operation queue."""
return await self.submit(
operation="list_tools",
run=lambda client: client.list_tools(),
)
async def get_connection_metadata(self) -> dict[str, Any]:
"""Return connection metadata computed locally without an upstream call."""
transport = self.connection.transport
return {
"server": self.connection.provider,
"transport": transport.kind if transport is not None else None,
}
async def invoke_method(
self,
method: str,
params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Submit a raw method invocation through the generic owner-task operation queue."""
return await self.submit(
operation="invoke_method",
run=lambda client: client.invoke_method(method, params),
)
async def send_notification(
self,
method: str,
params: dict[str, Any] | None = None,
) -> None:
"""Submit a notification send through the generic owner-task operation queue."""
await self.submit(
operation="send_notification",
run=lambda client: client.send_notification(method, params),
)
async def close(self) -> None:
"""Ask the owner task to close the MCP transport in its own scope."""
task = self._task
+37 -1
View File
@@ -7,7 +7,7 @@ from inspect import isawaitable
from typing import Any, cast
from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.sdk import ToolCallResult
@@ -119,6 +119,42 @@ class McpRuntimePool:
session = await self.get_session(connection, auth)
return await session.list_prompts()
async def list_tools(
self,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> list[DiscoveredTool]:
session = await self.get_session(connection, auth)
return await session.list_tools()
async def get_connection_metadata(
self,
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> dict[str, Any]:
session = await self.get_session(connection, auth)
return await session.get_connection_metadata()
async def invoke_method(
self,
connection: McpSourceConnection,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> dict[str, Any]:
session = await self.get_session(connection, auth)
return await session.invoke_method(method, params)
async def send_notification(
self,
connection: McpSourceConnection,
auth: AuthRecord | None,
method: str,
params: dict[str, Any] | None = None,
) -> None:
session = await self.get_session(connection, auth)
await session.send_notification(method, params)
async def close_connection(self, connection_id: str) -> None:
current = self._sessions.pop(connection_id, None)
if current is not None:
+67 -1
View File
@@ -8,7 +8,7 @@ from mcp.client.session import ClientSession
from pydantic import AnyUrl
from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.sdk import ToolCallResult
from wf_sources_mcp.sdk.converters import tool_result_to_call_result
@@ -21,6 +21,10 @@ RawPromptGetter = Callable[
]
RawResourceLister = Callable[[], Awaitable[list[DiscoveredResource]]]
RawPromptLister = Callable[[], Awaitable[list[DiscoveredPrompt]]]
RawToolLister = Callable[[], Awaitable[list[DiscoveredTool]]]
RawMetadataGetter = Callable[[], Awaitable[dict[str, Any]]]
RawMethodInvoker = Callable[[str, dict[str, Any] | None], Awaitable[dict[str, Any]]]
RawNotificationSender = Callable[[str, dict[str, Any] | None], Awaitable[None]]
@dataclass(slots=True)
@@ -41,6 +45,10 @@ class PersistentMcpSession:
get_prompt_callback: RawPromptGetter | None = None
list_resources_callback: RawResourceLister | None = None
list_prompts_callback: RawPromptLister | None = None
list_tools_callback: RawToolLister | None = None
get_connection_metadata_callback: RawMetadataGetter | None = None
invoke_method_callback: RawMethodInvoker | None = None
send_notification_callback: RawNotificationSender | None = None
close_callback: Callable[[], Awaitable[None]] | None = None
async def call_tool(
@@ -97,6 +105,64 @@ class PersistentMcpSession:
return [prompt_to_discovered(prompt) for prompt in result.prompts]
raise RuntimeError("persistent MCP session has no prompt list transport")
async def list_tools(self) -> list[DiscoveredTool]:
"""List MCP tools through the owner task or injected session."""
if self.list_tools_callback is not None:
return await self.list_tools_callback()
if self.client is not None:
from wf_sources_mcp.sdk.converters import tool_to_discovered
result = await self.client.list_tools()
return [tool_to_discovered(tool) for tool in result.tools]
raise RuntimeError("persistent MCP session has no tools list transport")
async def get_connection_metadata(self) -> dict[str, Any]:
"""Return connection metadata from callback or local connection info."""
if self.get_connection_metadata_callback is not None:
return await self.get_connection_metadata_callback()
transport = self.connection.transport
return {
"server": self.connection.provider,
"transport": transport.kind if transport is not None else None,
}
async def invoke_method(
self,
method: str,
params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Invoke a raw MCP method through the owner task or injected session."""
if self.invoke_method_callback is not None:
return await self.invoke_method_callback(method, params)
if self.client is not None:
from mcp import ClientResult
from mcp.types import ClientRequest
result = await self.client.send_request(
ClientRequest.model_validate({"method": method, "params": params}),
ClientResult,
)
return result.model_dump(by_alias=True, mode="json", exclude_none=True)
raise RuntimeError("persistent MCP session has no method invoke transport")
async def send_notification(
self,
method: str,
params: dict[str, Any] | None = None,
) -> None:
"""Send an MCP notification through the owner task or injected session."""
if self.send_notification_callback is not None:
await self.send_notification_callback(method, params)
return
if self.client is not None:
from mcp.types import ClientNotification
await self.client.send_notification(
ClientNotification.model_validate({"method": method, "params": params})
)
return
raise RuntimeError("persistent MCP session has no notification send transport")
async def close(self) -> None:
"""Close the transport/session stack owned by the runtime factory."""
if self.close_callback is not None:
+2
View File
@@ -10,6 +10,7 @@ from .converters import (
)
from .protocols import (
BackendAdapter,
McpSourceOperations,
PromptRuntime,
ResourceRuntime,
StatefulMcpRuntime,
@@ -21,6 +22,7 @@ from .protocols import (
__all__ = [
"BackendAdapter",
"McpSdkAdapter",
"McpSourceOperations",
"PromptRuntime",
"ResourceRuntime",
"StatefulMcpRuntime",
+10 -3
View File
@@ -17,7 +17,9 @@ class ToolCallResult:
meta: dict[str, Any] = field(default_factory=dict)
class BackendAdapter(Protocol):
class McpSourceOperations(Protocol):
"""Full MCP operation surface shared by one-shot adapters and persistent runtimes."""
async def list_tools(
self,
connection: McpSourceConnection,
@@ -82,6 +84,10 @@ class BackendAdapter(Protocol):
) -> ToolCallResult: ...
class BackendAdapter(McpSourceOperations, Protocol):
"""One-shot or adapter-style MCP operation executor."""
class ToolRuntime(Protocol):
"""Runtime boundary for executing MCP tools from workflow nodes."""
@@ -133,8 +139,8 @@ class PromptRuntime(Protocol):
) -> dict[str, Any]: ...
class StatefulMcpRuntime(ToolRuntime, ResourceRuntime, PromptRuntime, Protocol):
"""Stateful execution/read/list boundary for configured MCP sources.
class StatefulMcpRuntime(McpSourceOperations, Protocol):
"""Persistent MCP operation executor for configured sources.
Implementations keep source session state across calls. Catalog refresh may
still use one-shot adapters by policy.
@@ -143,6 +149,7 @@ class StatefulMcpRuntime(ToolRuntime, ResourceRuntime, PromptRuntime, Protocol):
__all__ = [
"BackendAdapter",
"McpSourceOperations",
"PromptRuntime",
"ResourceRuntime",
"StatefulMcpRuntime",
+14 -5
View File
@@ -47,11 +47,20 @@ _TRANSPORT_METADATA_KEYS = {
class LegacyConnectionConfigLike(Protocol):
"""Structural shape needed from legacy broker connection configs."""
id: str
server: str
account: str
enabled: bool
metadata: Mapping[str, object]
@property
def id(self) -> str: ...
@property
def server(self) -> str: ...
@property
def account(self) -> str: ...
@property
def enabled(self) -> bool: ...
@property
def metadata(self) -> Mapping[str, object]: ...
class McpSourceRegistryEntry(SourceRegistryBaseModel):