stateful proxy
This commit is contained in:
@@ -8,6 +8,19 @@
|
||||
|
||||
**Tech Stack:** Python 3.14, FastMCP, MCP Python SDK, `wf_core`, `wf_authoring`, `wf_mcp`, pytest, basedpyright, ruff.
|
||||
|
||||
## Implementation Status
|
||||
|
||||
- Tasks 1-3 are implemented: generated MCP workflow NodeSpecs now depend on
|
||||
the `ToolExecutor` protocol instead of directly baking in the one-shot SDK
|
||||
adapter.
|
||||
- Tasks 4-7 are implemented: `McpRuntimePool`, `PersistentMcpSession`, and
|
||||
`PersistentSessionFactory` exist, and config-built services use the runtime
|
||||
pool for generated workflow node execution while discovery/catalog refreshes
|
||||
still use short-lived SDK adapter sessions.
|
||||
- Remaining work starts at hiding/removing unsafe raw public `call_tool`
|
||||
surfaces and then renaming the legacy `transparent_proxy` package to the
|
||||
clearer proxy/provider-layer package.
|
||||
|
||||
---
|
||||
|
||||
## Problem Statement
|
||||
|
||||
@@ -7,6 +7,7 @@ from wf_artifacts import FileDraftWorkspaceStore, FileWorkflowArtifactStore
|
||||
|
||||
from ..control import BrokerConfigFile
|
||||
from ..models import BrokerConfig
|
||||
from ..runtime import McpRuntimePool, PersistentSessionFactory
|
||||
from ..sdk import McpSdkAdapter
|
||||
from ..storage import FileStore
|
||||
from .service import WfMcpService
|
||||
@@ -21,10 +22,15 @@ def load_broker_config(path: str | Path) -> BrokerConfig:
|
||||
|
||||
def build_service_from_config(config: BrokerConfig) -> WfMcpService:
|
||||
"""Create a broker service with SDK adapters for configured connections."""
|
||||
runtime_factory = PersistentSessionFactory()
|
||||
service = WfMcpService(
|
||||
store=FileStore(config.store_root),
|
||||
artifact_store=FileWorkflowArtifactStore(config.store_root),
|
||||
draft_workspace_store=FileDraftWorkspaceStore(config.store_root),
|
||||
# Discovery can use short-lived SDK sessions. Workflow execution needs
|
||||
# a persistent runtime so stateful MCP servers keep session/page state
|
||||
# across sequential workflow nodes.
|
||||
tool_executor=McpRuntimePool(runtime_factory.create),
|
||||
)
|
||||
for connection in config.connections:
|
||||
service.register_connection(connection)
|
||||
|
||||
@@ -41,6 +41,7 @@ from ...models import (
|
||||
ConnectionConfig,
|
||||
RawWorkflowPlan,
|
||||
)
|
||||
from ...runtime import ToolExecutor
|
||||
from ...sdk import BackendAdapter
|
||||
from ...shared.errors import error_payload
|
||||
from ...shared.names import RESERVED_CONNECTION_IDS
|
||||
@@ -72,6 +73,7 @@ class WfMcpService:
|
||||
include_builtin_specs: bool = True
|
||||
artifact_store: WorkflowArtifactStore | None = None
|
||||
draft_workspace_store: DraftWorkspaceStore | None = None
|
||||
tool_executor: ToolExecutor | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Install broker-local system specs when enabled."""
|
||||
@@ -130,6 +132,17 @@ class WfMcpService:
|
||||
def register_adapter(self, server: str, adapter: BackendAdapter) -> None:
|
||||
self.adapters[server] = adapter
|
||||
|
||||
def _tool_executor_for(self, connection: ConnectionConfig) -> ToolExecutor:
|
||||
"""Return the executor used by generated workflow NodeSpecs.
|
||||
|
||||
Discovery still uses the short-lived adapter path. Generated workflow
|
||||
nodes use this executor hook so config-built services can swap in a
|
||||
persistent runtime pool for stateful MCP servers.
|
||||
"""
|
||||
if self.tool_executor is not None:
|
||||
return self.tool_executor
|
||||
return require_adapter(connection, self.adapters)
|
||||
|
||||
def save_auth(self, record: AuthRecord) -> None:
|
||||
self.store.save_auth(record)
|
||||
self._record_event(
|
||||
@@ -606,7 +619,7 @@ class WfMcpService:
|
||||
specs = specs_from_discovered_tools(
|
||||
connection=connection,
|
||||
auth=auth,
|
||||
executor=adapter,
|
||||
executor=self._tool_executor_for(connection),
|
||||
tools=capabilities.tools,
|
||||
emit_event=self._record_event,
|
||||
)
|
||||
|
||||
@@ -1,3 +1,12 @@
|
||||
from .factory import PersistentSessionFactory
|
||||
from .pool import McpRuntimePool, connection_runtime_fingerprint
|
||||
from .protocols import ToolExecutor
|
||||
from .session import PersistentMcpSession
|
||||
|
||||
__all__ = ["ToolExecutor"]
|
||||
__all__ = [
|
||||
"McpRuntimePool",
|
||||
"PersistentMcpSession",
|
||||
"PersistentSessionFactory",
|
||||
"ToolExecutor",
|
||||
"connection_runtime_fingerprint",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import AsyncExitStack
|
||||
from dataclasses import dataclass
|
||||
|
||||
import httpx
|
||||
from mcp.client.session import ClientSession
|
||||
from mcp.client.stdio import StdioServerParameters, stdio_client
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
|
||||
from ..models import AuthRecord, ConnectionConfig
|
||||
from .session import PersistentMcpSession
|
||||
|
||||
|
||||
def _auth_headers(auth: AuthRecord | None) -> dict[str, str]:
|
||||
if auth is None:
|
||||
return {}
|
||||
headers = dict(auth.payload.get("headers", {}))
|
||||
token = auth.payload.get("token")
|
||||
if isinstance(token, str) and "Authorization" not in headers:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
return headers
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PersistentSessionFactory:
|
||||
"""Create initialized persistent MCP sessions for configured connections."""
|
||||
|
||||
async def create(
|
||||
self,
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None,
|
||||
) -> PersistentMcpSession:
|
||||
stack = AsyncExitStack()
|
||||
try:
|
||||
session = await self._create_with_stack(stack, connection, auth)
|
||||
except BaseException:
|
||||
await stack.aclose()
|
||||
raise
|
||||
return PersistentMcpSession(
|
||||
connection=connection,
|
||||
auth=auth,
|
||||
client=session,
|
||||
close_callback=stack.aclose,
|
||||
)
|
||||
|
||||
async def _create_with_stack(
|
||||
self,
|
||||
stack: AsyncExitStack,
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None,
|
||||
) -> ClientSession:
|
||||
transport = connection.metadata.get("transport", "stdio")
|
||||
if transport == "stdio":
|
||||
env = connection.metadata.get("env")
|
||||
if auth is not None:
|
||||
auth_env = auth.payload.get("env")
|
||||
if isinstance(auth_env, dict):
|
||||
env = {**(env or {}), **auth_env}
|
||||
params = StdioServerParameters(
|
||||
command=connection.metadata["command"],
|
||||
args=list(connection.metadata.get("args", [])),
|
||||
env=env,
|
||||
cwd=connection.metadata.get("cwd"),
|
||||
)
|
||||
read_stream, write_stream = await stack.enter_async_context(
|
||||
stdio_client(params)
|
||||
)
|
||||
session = await stack.enter_async_context(
|
||||
ClientSession(read_stream, write_stream)
|
||||
)
|
||||
await session.initialize()
|
||||
return session
|
||||
|
||||
if transport == "streamable_http":
|
||||
http_client = await stack.enter_async_context(
|
||||
httpx.AsyncClient(headers=_auth_headers(auth) or None)
|
||||
)
|
||||
read_stream, write_stream, _get_session_id = (
|
||||
await stack.enter_async_context(
|
||||
streamable_http_client(
|
||||
connection.metadata["url"],
|
||||
http_client=http_client,
|
||||
)
|
||||
)
|
||||
)
|
||||
session = await stack.enter_async_context(
|
||||
ClientSession(read_stream, write_stream)
|
||||
)
|
||||
await session.initialize()
|
||||
return session
|
||||
|
||||
raise ValueError(f"unsupported MCP transport {transport!r}")
|
||||
@@ -0,0 +1,88 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from inspect import isawaitable
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from typing import Any, cast
|
||||
|
||||
from ..models import AuthRecord, ConnectionConfig
|
||||
from ..sdk import ToolCallResult
|
||||
from .session import PersistentMcpSession
|
||||
|
||||
SessionFactory = Callable[
|
||||
[ConnectionConfig, AuthRecord | None],
|
||||
PersistentMcpSession | Awaitable[PersistentMcpSession],
|
||||
]
|
||||
|
||||
|
||||
def connection_runtime_fingerprint(
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None = None,
|
||||
) -> str:
|
||||
"""Return the connection identity that decides MCP runtime reuse.
|
||||
|
||||
This is intentionally transport/auth level, not catalog level. Tool list
|
||||
refreshes should not restart a browser-like MCP session, but changing the
|
||||
command, URL, account, or auth payload must create a fresh session.
|
||||
"""
|
||||
|
||||
return json.dumps(
|
||||
{
|
||||
"connection": asdict(connection),
|
||||
"auth": asdict(auth) if auth is not None else None,
|
||||
},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
default=str,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class McpRuntimePool:
|
||||
"""Cache one persistent MCP runtime per unchanged connection fingerprint."""
|
||||
|
||||
session_factory: SessionFactory
|
||||
_sessions: dict[str, tuple[str, PersistentMcpSession]] = field(default_factory=dict)
|
||||
|
||||
async def get_session(
|
||||
self,
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None,
|
||||
) -> PersistentMcpSession:
|
||||
fingerprint = connection_runtime_fingerprint(connection, auth)
|
||||
current = self._sessions.get(connection.id)
|
||||
if current is not None and current[0] == fingerprint:
|
||||
return current[1]
|
||||
if current is not None:
|
||||
await current[1].close()
|
||||
|
||||
created = self.session_factory(connection, auth)
|
||||
if isawaitable(created):
|
||||
session = await created
|
||||
else:
|
||||
session = cast(PersistentMcpSession, created)
|
||||
self._sessions[connection.id] = (fingerprint, session)
|
||||
return session
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None,
|
||||
tool_name: str,
|
||||
payload: dict[str, Any],
|
||||
) -> ToolCallResult:
|
||||
session = await self.get_session(connection, auth)
|
||||
return await session.call_tool(tool_name, payload)
|
||||
|
||||
async def close_connection(self, connection_id: str) -> None:
|
||||
current = self._sessions.pop(connection_id, None)
|
||||
if current is not None:
|
||||
await current[1].close()
|
||||
|
||||
async def close_all(self) -> None:
|
||||
"""Close all live runtimes; useful for server shutdown and tests."""
|
||||
sessions = list(self._sessions.values())
|
||||
self._sessions.clear()
|
||||
for _fingerprint, session in sessions:
|
||||
await session.close()
|
||||
@@ -0,0 +1,33 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from ..models import AuthRecord, ConnectionConfig
|
||||
from ..sdk import ToolCallResult
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PersistentMcpSession:
|
||||
"""Long-lived MCP execution handle for one configured connection."""
|
||||
|
||||
connection: ConnectionConfig
|
||||
auth: AuthRecord | None
|
||||
client: Any
|
||||
close_callback: Callable[[], Awaitable[None]] | None = None
|
||||
|
||||
async def call_tool(self, tool_name: str, payload: dict[str, Any]) -> ToolCallResult:
|
||||
return await self.client.call_tool(tool_name, payload)
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close this runtime without assuming a specific SDK client shape."""
|
||||
if self.close_callback is not None:
|
||||
await self.close_callback()
|
||||
return
|
||||
close = getattr(self.client, "close", None)
|
||||
if close is None:
|
||||
return
|
||||
result = close()
|
||||
if hasattr(result, "__await__"):
|
||||
await result
|
||||
@@ -8,7 +8,7 @@ from typing import Any
|
||||
|
||||
from mcp import types as mcp_types
|
||||
|
||||
from wf_mcp.broker.config import load_broker_config
|
||||
from wf_mcp.broker.config import build_service_from_config, load_broker_config
|
||||
from wf_mcp.models import BrokerConfig, ConnectionConfig
|
||||
from wf_mcp.server import create_server_client
|
||||
|
||||
@@ -21,6 +21,29 @@ def _structured(result: Any) -> dict[str, Any]:
|
||||
return content
|
||||
|
||||
|
||||
def test_config_built_service_uses_persistent_tool_executor() -> None:
|
||||
config = BrokerConfig(
|
||||
store_root=local_temp_root() / "runtime_config_store",
|
||||
connections=[
|
||||
ConnectionConfig(
|
||||
id="fixture.personal",
|
||||
server="fixture",
|
||||
account="personal",
|
||||
metadata={
|
||||
"transport": "stdio",
|
||||
"command": sys.executable,
|
||||
"args": [fixture_server_path()],
|
||||
},
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
service = build_service_from_config(config)
|
||||
|
||||
assert service.adapters["fixture"].__class__.__name__ == "McpSdkAdapter"
|
||||
assert service.tool_executor is not None
|
||||
|
||||
|
||||
async def _assert_safe_tool_maps(
|
||||
client: Any,
|
||||
*,
|
||||
|
||||
@@ -2,12 +2,15 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import shutil
|
||||
from typing import Any, cast
|
||||
|
||||
from wf_artifacts import FileDraftWorkspaceStore
|
||||
from wf_authoring import NodeSpec
|
||||
from wf_core import END, NodeUse, RunStatus
|
||||
from wf_authoring import NodeSpec, build_async_registry
|
||||
from wf_core import END, NodeUse, RunStatus, RuntimeContext
|
||||
from wf_mcp.broker import WfMcpService
|
||||
from wf_mcp.models import AuthRecord, ConnectionConfig, RawWorkflowPlan
|
||||
from wf_mcp.runtime import ToolExecutor
|
||||
from wf_mcp.sdk import ToolCallResult
|
||||
from wf_mcp.shared.errors import error_payload
|
||||
from wf_mcp.storage import FileStore
|
||||
from wf_platform import (
|
||||
@@ -1022,6 +1025,45 @@ def test_service_can_call_upstream_tool_directly() -> None:
|
||||
assert "tool_call_completed" in event_kinds
|
||||
|
||||
|
||||
def test_generated_specs_use_injected_tool_executor() -> None:
|
||||
class RecordingExecutor:
|
||||
def __init__(self) -> None:
|
||||
self.payloads: list[dict[str, Any]] = []
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None,
|
||||
tool_name: str,
|
||||
payload: dict[str, Any],
|
||||
) -> ToolCallResult:
|
||||
self.payloads.append(payload)
|
||||
return ToolCallResult(outcome="ok", output={"echoed": payload["text"]})
|
||||
|
||||
executor = RecordingExecutor()
|
||||
service = WfMcpService(
|
||||
store=FileStore(local_temp_root() / "injected_executor_store"),
|
||||
tool_executor=cast(ToolExecutor, executor),
|
||||
)
|
||||
service.register_connection(
|
||||
ConnectionConfig(id="demo.personal", server="demo", account="personal")
|
||||
)
|
||||
service.register_adapter("demo", FakeAdapter())
|
||||
|
||||
asyncio.run(service.refresh_connection_catalog("demo.personal"))
|
||||
spec = service._get_qualified_spec("demo.personal.echo_tool")
|
||||
handler = build_async_registry(spec)[spec.name]
|
||||
|
||||
async def run_node() -> dict[str, Any]:
|
||||
return await handler({"text": "hello"}, RuntimeContext(current_node_id="echo"))
|
||||
|
||||
result = asyncio.run(run_node())
|
||||
|
||||
assert result["outcome"] == "ok"
|
||||
assert result["output"]["echoed"] == "hello"
|
||||
assert executor.payloads == [{"text": "hello"}]
|
||||
|
||||
|
||||
def test_service_records_catalog_refresh_failures() -> None:
|
||||
service = WfMcpService(store=FileStore(local_temp_root() / "refresh_fail_store"))
|
||||
service.register_connection(
|
||||
|
||||
@@ -8,6 +8,7 @@ from wf_authoring import build_async_registry
|
||||
from wf_core import RuntimeContext
|
||||
from wf_mcp.capabilities import DiscoveredTool
|
||||
from wf_mcp.models import AuthRecord, ConnectionConfig
|
||||
from wf_mcp.runtime import McpRuntimePool, PersistentMcpSession
|
||||
from wf_mcp.sdk import ToolCallResult
|
||||
from wf_mcp.workflow import wrap_discovered_tool
|
||||
|
||||
@@ -37,6 +38,27 @@ class FakeStatefulExecutor:
|
||||
raise KeyError(tool_name)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class FakeStatefulClient:
|
||||
"""Session-client fake with the same call shape as MCP SDK ClientSession."""
|
||||
|
||||
page_open: bool = False
|
||||
closed: bool = False
|
||||
calls: list[tuple[str, dict[str, Any]]] = field(default_factory=list)
|
||||
|
||||
async def call_tool(self, tool_name: str, payload: dict[str, Any]) -> ToolCallResult:
|
||||
self.calls.append((tool_name, payload))
|
||||
if tool_name == "browser_navigate":
|
||||
self.page_open = True
|
||||
return ToolCallResult(outcome="ok", output={"content": "opened"})
|
||||
if tool_name == "browser_snapshot" and self.page_open:
|
||||
return ToolCallResult(outcome="ok", output={"content": "snapshot"})
|
||||
return ToolCallResult(outcome="error", output={"message": "No open page"})
|
||||
|
||||
async def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
def _tool(name: str) -> DiscoveredTool:
|
||||
return DiscoveredTool(
|
||||
name=name,
|
||||
@@ -86,3 +108,70 @@ def test_generated_workflow_specs_share_injected_tool_executor() -> None:
|
||||
assert result["outcome"] == "ok"
|
||||
assert result["output"]["content"] == "snapshot"
|
||||
assert executor.calls == [("browser_navigate", {}), ("browser_snapshot", {})]
|
||||
|
||||
|
||||
def test_runtime_pool_reuses_stateful_session_for_same_connection() -> None:
|
||||
connection = ConnectionConfig(
|
||||
id="playwright.default",
|
||||
server="playwright",
|
||||
account="default",
|
||||
metadata={
|
||||
"transport": "stdio",
|
||||
"command": "pnpx",
|
||||
"args": ["@playwright/mcp"],
|
||||
},
|
||||
)
|
||||
created_clients: list[FakeStatefulClient] = []
|
||||
|
||||
async def factory(
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None,
|
||||
) -> PersistentMcpSession:
|
||||
client = FakeStatefulClient()
|
||||
created_clients.append(client)
|
||||
return PersistentMcpSession(connection=connection, auth=auth, client=client)
|
||||
|
||||
async def run_calls() -> ToolCallResult:
|
||||
pool = McpRuntimePool(factory)
|
||||
await pool.call_tool(connection, None, "browser_navigate", {})
|
||||
return await pool.call_tool(connection, None, "browser_snapshot", {})
|
||||
|
||||
result = asyncio.run(run_calls())
|
||||
|
||||
assert result.outcome == "ok"
|
||||
assert result.output["content"] == "snapshot"
|
||||
assert len(created_clients) == 1
|
||||
|
||||
|
||||
def test_runtime_pool_replaces_session_when_fingerprint_changes() -> None:
|
||||
original = ConnectionConfig(
|
||||
id="playwright.default",
|
||||
server="playwright",
|
||||
account="default",
|
||||
metadata={"transport": "stdio", "command": "pnpx", "args": ["old"]},
|
||||
)
|
||||
changed = ConnectionConfig(
|
||||
id="playwright.default",
|
||||
server="playwright",
|
||||
account="default",
|
||||
metadata={"transport": "stdio", "command": "pnpx", "args": ["new"]},
|
||||
)
|
||||
created_clients: list[FakeStatefulClient] = []
|
||||
|
||||
def factory(
|
||||
connection: ConnectionConfig,
|
||||
auth: AuthRecord | None,
|
||||
) -> PersistentMcpSession:
|
||||
client = FakeStatefulClient()
|
||||
created_clients.append(client)
|
||||
return PersistentMcpSession(connection=connection, auth=auth, client=client)
|
||||
|
||||
async def run_calls() -> None:
|
||||
pool = McpRuntimePool(factory)
|
||||
await pool.call_tool(original, None, "browser_navigate", {})
|
||||
await pool.call_tool(changed, None, "browser_snapshot", {})
|
||||
|
||||
asyncio.run(run_calls())
|
||||
|
||||
assert len(created_clients) == 2
|
||||
assert created_clients[0].closed is True
|
||||
|
||||
Reference in New Issue
Block a user