reusing connection: asyncio complications

This commit is contained in:
lda
2026-05-25 19:04:28 +07:00 Verified
parent 1a5efe6cda
commit dc0799b2ff
3 changed files with 160 additions and 14 deletions
+99 -9
View File
@@ -1,12 +1,14 @@
from __future__ import annotations
import asyncio
from contextlib import AsyncExitStack
from dataclasses import dataclass
from dataclasses import dataclass, field
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 mcp.types import CallToolResult
from ..models import AuthRecord, ConnectionConfig
from .session import PersistentMcpSession
@@ -37,17 +39,13 @@ class PersistentSessionFactory:
connection: ConnectionConfig,
auth: AuthRecord | None,
) -> PersistentMcpSession:
stack = AsyncExitStack()
try:
session = await self._create_with_stack(stack, connection, auth)
except BaseException:
await stack.aclose()
raise
owner = _SessionOwner(factory=self, connection=connection, auth=auth)
await owner.start()
return PersistentMcpSession(
connection=connection,
auth=auth,
client=session,
close_callback=stack.aclose,
call_callback=owner.call_tool,
close_callback=owner.close,
)
async def _create_with_stack(
@@ -99,3 +97,95 @@ class PersistentSessionFactory:
return session
raise ValueError(f"unsupported MCP transport {transport!r}")
@dataclass(slots=True)
class _ToolCallRequest:
"""One request submitted to the task that owns the MCP transport."""
tool_name: str
payload: dict[str, object]
result: asyncio.Future[CallToolResult]
@dataclass(slots=True)
class _SessionOwner:
"""Run one MCP client session entirely inside its owning asyncio task.
MCP SDK transports open AnyIO cancel scopes. Entering a transport in one
inbound MCP request and reusing it from another causes
`ClosedResourceError`/cancel-scope ownership failures. This actor keeps
transport creation, calls, and cleanup in one stable task while the public
workflow surface submits requests through a queue.
"""
factory: PersistentSessionFactory
connection: ConnectionConfig
auth: AuthRecord | None
_requests: asyncio.Queue[_ToolCallRequest | None] = field(
default_factory=asyncio.Queue
)
_task: asyncio.Task[None] | None = None
async def start(self) -> None:
"""Start the owner task and wait until its MCP session is initialized."""
ready = asyncio.get_running_loop().create_future()
self._task = asyncio.create_task(
self._run(ready),
name=f"wf-mcp-session:{self.connection.id}",
)
await ready
async def call_tool(
self,
tool_name: str,
payload: dict[str, object],
) -> CallToolResult:
"""Submit a tool call for execution in the transport owner task."""
task = self._task
if task is None:
raise RuntimeError("persistent MCP session is not started")
if task.done():
await task
raise RuntimeError("persistent MCP session stopped unexpectedly")
result = asyncio.get_running_loop().create_future()
await self._requests.put(
_ToolCallRequest(tool_name=tool_name, payload=payload, result=result)
)
return await result
async def close(self) -> None:
"""Ask the owner task to close the MCP transport in its own scope."""
task = self._task
if task is None:
return
if not task.done():
await self._requests.put(None)
await task
self._task = None
async def _run(self, ready: asyncio.Future[None]) -> None:
"""Own the complete MCP transport lifecycle and serialized call loop."""
try:
async with AsyncExitStack() as stack:
session = await self.factory._create_with_stack(
stack, self.connection, self.auth
)
ready.set_result(None)
while True:
request = await self._requests.get()
if request is None:
return
try:
response = await session.call_tool(
request.tool_name, request.payload
)
except Exception as exc:
request.result.set_exception(exc)
else:
request.result.set_result(response)
except BaseException as exc:
if not ready.done():
ready.set_exception(exc)
else:
raise
+15 -5
View File
@@ -5,30 +5,40 @@ from dataclasses import dataclass
from typing import Any
from mcp.client.session import ClientSession
from mcp.types import CallToolResult
from ..models import AuthRecord, ConnectionConfig
from ..sdk import ToolCallResult
from ..sdk.converters import tool_result_to_call_result
RawToolCaller = Callable[[str, dict[str, Any]], Awaitable[CallToolResult]]
@dataclass(slots=True)
class PersistentMcpSession:
"""Long-lived MCP execution handle for one configured connection.
`client` is an initialized MCP SDK `ClientSession`. `call_tool()` returns
this project's normalized `ToolCallResult`, not the SDK result object, so
generated workflow nodes do not need to know MCP wire result shapes.
Production sessions use `call_callback` because MCP transports are entered
inside an AnyIO cancel scope and must be called and closed by that same
owner task. `client` remains available for simple injected/fake sessions in
tests. `call_tool()` always normalizes SDK results for workflow nodes.
"""
connection: ConnectionConfig
auth: AuthRecord | None
client: ClientSession
client: ClientSession | None = None
call_callback: RawToolCaller | None = None
close_callback: Callable[[], Awaitable[None]] | None = None
async def call_tool(
self, tool_name: str, payload: dict[str, Any]
) -> ToolCallResult:
result = await self.client.call_tool(tool_name, payload)
if self.call_callback is not None:
result = await self.call_callback(tool_name, payload)
elif self.client is not None:
result = await self.client.call_tool(tool_name, payload)
else:
raise RuntimeError("persistent MCP session has no tool call transport")
return tool_result_to_call_result(result)
async def close(self) -> None:
+46
View File
@@ -205,6 +205,52 @@ def test_server_exposes_upstream_admin_and_workflow_tools() -> None:
asyncio.run(run_proxy())
def test_server_reuses_real_upstream_session_across_workflow_requests() -> None:
"""Workflow node calls may share one stateful MCP session across requests."""
config = BrokerConfig(
store_root=local_temp_root() / "workflow_persistent_fixture_store",
connections=[
ConnectionConfig(
id="fixture.personal",
server="fixture",
account="personal",
metadata={
"transport": "stdio",
"command": sys.executable,
"args": [fixture_server_path()],
},
)
],
)
async def run_proxy() -> None:
client = create_server_client(config)
async with client:
await client.call_tool(
"wf.admin.refresh_connection_catalog",
{"connection_id": "fixture.personal"},
)
first = await client.call_tool(
"wf.workflow.call_capability",
{
"qualified_name": "fixture.personal.echo_tool",
"payload": {"text": "one"},
},
)
second = await client.call_tool(
"wf.workflow.call_capability",
{
"qualified_name": "fixture.personal.echo_tool",
"payload": {"text": "two"},
},
)
assert _structured(first)["output"]["echoed"] == "one"
assert _structured(second)["output"]["echoed"] == "two"
asyncio.run(run_proxy())
def test_server_can_hide_admin_tools() -> None:
config = BrokerConfig(
store_root=local_temp_root() / "unified_no_admin_store",