refactor: share mcp session opener

This commit is contained in:
lda
2026-06-07 15:43:01 +07:00 Verified
parent ac2163bf30
commit 008cc6d139
8 changed files with 356 additions and 90 deletions
+7 -4
View File
@@ -225,10 +225,13 @@ implementation state.
`McpSourceConnection` seam in `wf_sources_mcp`, not by moving `McpSourceConnection` seam in `wf_sources_mcp`, not by moving
`runtime/factory.py` as-is. The active plan was `runtime/factory.py` as-is. The active plan was
[2026-06-07 MCP source connection seam](./historical/superpowers/plans/2026-06-07-mcp-source-connection-seam.md). [2026-06-07 MCP source connection seam](./historical/superpowers/plans/2026-06-07-mcp-source-connection-seam.md).
- Planned next: share one MCP session opener between the one-shot SDK - Completed: shared MCP session opener exists in `wf_sources_mcp.client`.
adapter and persistent runtime before moving runtime files. The active One-shot adapter (`McpSdkAdapter`) and persistent runtime
plan is (`PersistentSessionFactory`) both use it. Runtime files remain in
[2026-06-07 MCP client session opener](./superpowers/plans/2026-06-07-mcp-client-session-opener.md). `wf_mcp` for compatibility. Next slice can move `PersistentSessionFactory`,
`PersistentMcpSession`, and `McpRuntimePool` to
`wf_sources_mcp.runtime`. The completed plan was
[2026-06-07 MCP client session opener](./historical/superpowers/plans/2026-06-07-mcp-client-session-opener.md).
- Auth/source secrets boundary: keep registry desired state separate from - Auth/source secrets boundary: keep registry desired state separate from
upstream credentials, and surface missing auth as validation diagnostics. upstream credentials, and surface missing auth as validation diagnostics.
The contract is now specified in The contract is now specified in
@@ -93,7 +93,13 @@ First slices should move leaf modules only and leave `wf_mcp` re-export shims:
3. Complete: upstream MCP catalog/discovery DTOs moved to `wf_sources_mcp.catalog`, with `wf_mcp.capabilities` and `wf_mcp.catalog.models` retained as shims. 3. Complete: upstream MCP catalog/discovery DTOs moved to `wf_sources_mcp.catalog`, with `wf_mcp.capabilities` and `wf_mcp.catalog.models` retained as shims.
4. Complete: upstream SDK protocol/result types moved to `wf_sources_mcp.sdk`, with `wf_mcp.sdk` and `wf_mcp.runtime.protocols` retained as shims. 4. Complete: upstream SDK protocol/result types moved to `wf_sources_mcp.sdk`, with `wf_mcp.sdk` and `wf_mcp.runtime.protocols` retained as shims.
5. Complete: MCP SDK conversion helpers moved to `wf_sources_mcp.sdk.converters`, with `wf_mcp.sdk.converters` retained as a shim. 5. Complete: MCP SDK conversion helpers moved to `wf_sources_mcp.sdk.converters`, with `wf_mcp.sdk.converters` retained as a shim.
6. Upstream transport/discovery/session services. 6. Complete: shared MCP session opener in `wf_sources_mcp.client`. One-shot
adapter (`McpSdkAdapter`) and persistent runtime
(`PersistentSessionFactory`) both use `open_mcp_session`. Runtime files
remain in `wf_mcp` for compatibility; next slice moves
`PersistentSessionFactory`, `PersistentMcpSession`, and `McpRuntimePool`
to `wf_sources_mcp.runtime`.
7. Upstream transport/discovery/session services.
Each slice should add import-direction tests so the new source-provider package Each slice should add import-direction tests so the new source-provider package
does not depend on `wf_mcp.workflow_surface`, `wf_mcp.admin_surface`, does not depend on `wf_mcp.workflow_surface`, `wf_mcp.admin_surface`,
+8 -46
View File
@@ -4,13 +4,12 @@ import asyncio
from contextlib import AsyncExitStack from contextlib import AsyncExitStack
from dataclasses import dataclass, field from dataclasses import dataclass, field
import httpx
from mcp.client.session import ClientSession 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 mcp.types import CallToolResult
from wf_sources_mcp.auth import AuthRecord, mcp_auth_env, mcp_auth_headers from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.client import open_mcp_session
from wf_sources_mcp.connections import mcp_source_connection_from_connection_config
from ..models import ConnectionConfig from ..models import ConnectionConfig
from .session import PersistentMcpSession from .session import PersistentMcpSession
@@ -46,48 +45,11 @@ class PersistentSessionFactory:
connection: ConnectionConfig, connection: ConnectionConfig,
auth: AuthRecord | None, auth: AuthRecord | None,
) -> ClientSession: ) -> ClientSession:
transport = connection.metadata.get("transport", "stdio") source_connection = mcp_source_connection_from_connection_config(connection)
if transport == "stdio": session = await stack.enter_async_context(
env = connection.metadata.get("env") open_mcp_session(source_connection, auth)
auth_env = mcp_auth_env(auth) )
if auth_env: return session
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=mcp_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}")
@dataclass(slots=True) @dataclass(slots=True)
+4 -39
View File
@@ -3,11 +3,7 @@ from __future__ import annotations
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import Any from typing import Any
import httpx
from mcp import ClientResult 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 ( from mcp.types import (
ClientNotification, ClientNotification,
ClientRequest, ClientRequest,
@@ -17,8 +13,9 @@ from mcp.types import (
) )
from pydantic import AnyUrl from pydantic import AnyUrl
from wf_sources_mcp.auth import AuthRecord, mcp_auth_env, mcp_auth_headers from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool from wf_sources_mcp.catalog import DiscoveredPrompt, DiscoveredResource, DiscoveredTool
from wf_sources_mcp.client import open_mcp_session
from wf_sources_mcp.connections import McpSourceConnection from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.sdk import BackendAdapter, ToolCallResult from wf_sources_mcp.sdk import BackendAdapter, ToolCallResult
from wf_sources_mcp.sdk.converters import ( from wf_sources_mcp.sdk.converters import (
@@ -27,7 +24,6 @@ from wf_sources_mcp.sdk.converters import (
tool_result_to_call_result, tool_result_to_call_result,
tool_to_discovered, tool_to_discovered,
) )
from wf_sources_mcp.transports import HttpSourceTransport, StdioSourceTransport
class McpSdkAdapter(BackendAdapter): class McpSdkAdapter(BackendAdapter):
@@ -37,39 +33,8 @@ class McpSdkAdapter(BackendAdapter):
connection: McpSourceConnection, connection: McpSourceConnection,
auth: AuthRecord | None, auth: AuthRecord | None,
): ):
transport = connection.transport async with open_mcp_session(connection, auth) as session:
if transport is None: yield session
raise ValueError(f"connection {connection.id!r} requires metadata.transport")
if isinstance(transport, StdioSourceTransport):
auth_env = mcp_auth_env(auth)
env = dict(transport.env)
if auth_env:
env = {**env, **auth_env}
params = StdioServerParameters(
command=transport.command,
args=list(transport.args),
env=env,
)
async with stdio_client(params) as (read_stream, write_stream):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
yield session
return
if isinstance(transport, HttpSourceTransport):
headers = mcp_auth_headers(auth)
http_client = httpx.AsyncClient(headers=headers or None)
async with http_client:
async with streamable_http_client(
str(transport.url),
http_client=http_client,
) as (read_stream, write_stream, _get_session_id):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
yield session
return
raise ValueError(f"unsupported MCP transport {type(transport).__name__}")
async def list_tools( async def list_tools(
self, self,
+5
View File
@@ -0,0 +1,5 @@
from __future__ import annotations
from .transport import open_mcp_session
__all__ = ["open_mcp_session"]
+68
View File
@@ -0,0 +1,68 @@
"""Shared MCP session opener for one-shot and persistent runtimes."""
from __future__ import annotations
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
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 wf_sources_mcp.auth import AuthRecord, mcp_auth_env, mcp_auth_headers
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.transports import HttpSourceTransport, StdioSourceTransport
@asynccontextmanager
async def open_mcp_session(
connection: McpSourceConnection,
auth: AuthRecord | None,
) -> AsyncIterator[ClientSession]:
"""Open and initialize an MCP client session for the given connection.
For stdio transports, merges transport env with auth env (auth wins on
duplicate keys) and passes command, args, env, and cwd to
StdioServerParameters.
For HTTP transports, creates an httpx.AsyncClient with auth headers and
enters streamable_http_client.
Yields an initialized ClientSession. Caller owns the session lifetime.
"""
transport = connection.transport
if transport is None:
raise ValueError(f"connection {connection.id!r} requires metadata.transport")
if isinstance(transport, StdioSourceTransport):
auth_env = mcp_auth_env(auth)
env = dict(transport.env)
if auth_env:
env = {**env, **auth_env}
params = StdioServerParameters(
command=transport.command,
args=list(transport.args),
env=env,
cwd=transport.cwd,
)
async with stdio_client(params) as (read_stream, write_stream):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
yield session
return
if isinstance(transport, HttpSourceTransport):
headers = mcp_auth_headers(auth)
http_client = httpx.AsyncClient(headers=headers or None)
async with http_client:
async with streamable_http_client(
str(transport.url),
http_client=http_client,
) as (read_stream, write_stream, _get_session_id):
async with ClientSession(read_stream, write_stream) as session:
await session.initialize()
yield session
return
raise ValueError(f"unsupported MCP transport {transport.kind!r}")
@@ -0,0 +1,257 @@
from __future__ import annotations
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from dataclasses import dataclass
from typing import Any
import pytest
from wf_sources_mcp.auth import AuthRecord
from wf_sources_mcp.client.transport import open_mcp_session
from wf_sources_mcp.connections import McpSourceConnection
from wf_sources_mcp.transports import HttpSourceTransport, StdioSourceTransport
@dataclass
class _FakeSession:
initialized: bool = False
calls: list[str] | None = None
async def initialize(self) -> None:
self.initialized = True
@asynccontextmanager
async def _fake_stdio_client(params: Any) -> AsyncIterator[tuple[Any, Any]]:
yield "read", "write"
@asynccontextmanager
async def _fake_streamable_http_client(
url: str, *, http_client: Any = None
) -> AsyncIterator[tuple[Any, Any, Any]]:
yield "read", "write", lambda: None
@asynccontextmanager
async def _fake_client_session(
read: Any, write: Any
) -> AsyncIterator[_FakeSession]:
yield _FakeSession()
def _stdio_connection(
*,
command: str = "uvx",
args: tuple[str, ...] = ("server",),
env: dict[str, str] | None = None,
cwd: str | None = None,
) -> McpSourceConnection:
return McpSourceConnection(
id="test.server",
provider="test",
account="server",
transport=StdioSourceTransport(
command=command,
args=args,
env=env or {},
cwd=cwd,
),
)
def _http_connection(
url: str = "http://127.0.0.1:8000/mcp",
headers: dict[str, str] | None = None,
) -> McpSourceConnection:
return McpSourceConnection(
id="test.server",
provider="test",
account="server",
transport=HttpSourceTransport(
url=url,
headers=headers or {},
),
)
@pytest.fixture(autouse=True)
def _patch_mcp(monkeypatch: pytest.MonkeyPatch) -> None:
import wf_sources_mcp.client.transport as mod
monkeypatch.setattr(mod, "stdio_client", _fake_stdio_client)
monkeypatch.setattr(mod, "streamable_http_client", _fake_streamable_http_client)
monkeypatch.setattr(mod, "ClientSession", _fake_client_session)
@pytest.mark.asyncio
async def test_stdio_session_initializes_before_yielding() -> None:
connection = _stdio_connection()
async with open_mcp_session(connection, None) as session:
assert session.initialized is True
@pytest.mark.asyncio
async def test_stdio_env_merges_transport_and_auth_wins_on_duplicate() -> None:
connection = _stdio_connection(env={"A": "transport", "B": "transport_only"})
auth = AuthRecord(
connection_id="test.server",
scheme="env",
payload={"env": {"A": "auth_wins", "C": "auth_only"}},
)
import wf_sources_mcp.client.transport as mod
captured_params: list[Any] = []
@asynccontextmanager
async def _capturing_stdio_client(
params: Any,
) -> AsyncIterator[tuple[Any, Any]]:
captured_params.append(params)
yield "read", "write"
mod.stdio_client = _capturing_stdio_client # type: ignore[assignment]
async with open_mcp_session(connection, auth) as session:
assert session.initialized is True
params = captured_params[0]
assert params.env["A"] == "auth_wins"
assert params.env["B"] == "transport_only"
assert params.env["C"] == "auth_only"
@pytest.mark.asyncio
async def test_stdio_cwd_propagated_to_server_parameters() -> None:
connection = _stdio_connection(cwd="/workspace")
import wf_sources_mcp.client.transport as mod
captured_params: list[Any] = []
@asynccontextmanager
async def _capturing_stdio_client(
params: Any,
) -> AsyncIterator[tuple[Any, Any]]:
captured_params.append(params)
yield "read", "write"
mod.stdio_client = _capturing_stdio_client # type: ignore[assignment]
async with open_mcp_session(connection, None) as session:
assert session.initialized is True
assert captured_params[0].cwd == "/workspace"
@pytest.mark.asyncio
async def test_http_session_initializes_before_yielding() -> None:
connection = _http_connection()
async with open_mcp_session(connection, None) as session:
assert session.initialized is True
@pytest.mark.asyncio
async def test_http_auth_headers_passed_to_client(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection = _http_connection()
auth = AuthRecord(
connection_id="test.server",
scheme="bearer",
payload={"token": "secret123"},
)
import httpx as _httpx
captured_clients: list[_httpx.AsyncClient] = []
_original_client = _httpx.AsyncClient
class _CapturingClient(_original_client): # type: ignore[type-arg]
def __init__(self, **kwargs: Any) -> None:
super().__init__(**kwargs)
captured_clients.append(self)
import wf_sources_mcp.client.transport as mod
class _PatchedHttpx:
AsyncClient = _CapturingClient
monkeypatch.setattr(mod, "httpx", _PatchedHttpx())
async with open_mcp_session(connection, auth) as session:
assert session.initialized is True
assert len(captured_clients) == 1
assert captured_clients[0].headers["Authorization"] == "Bearer secret123"
@pytest.mark.asyncio
async def test_unsupported_transport_raises_value_error() -> None:
connection = McpSourceConnection(
id="test.server",
provider="test",
account="server",
)
with pytest.raises(ValueError, match="requires metadata.transport"):
async with open_mcp_session(connection, None):
pass
@pytest.mark.asyncio
async def test_stdio_no_auth_uses_transport_env_only() -> None:
connection = _stdio_connection(env={"TOKEN": "abc"})
import wf_sources_mcp.client.transport as mod
captured_params: list[Any] = []
@asynccontextmanager
async def _capturing_stdio_client(
params: Any,
) -> AsyncIterator[tuple[Any, Any]]:
captured_params.append(params)
yield "read", "write"
mod.stdio_client = _capturing_stdio_client # type: ignore[assignment]
async with open_mcp_session(connection, None):
pass
assert captured_params[0].env == {"TOKEN": "abc"}
@pytest.mark.asyncio
async def test_http_no_auth_creates_client_with_no_auth_headers(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection = _http_connection()
import httpx as _httpx
captured_clients: list[_httpx.AsyncClient] = []
_original_client = _httpx.AsyncClient
class _CapturingClient(_original_client): # type: ignore[type-arg]
def __init__(self, **kwargs: Any) -> None:
super().__init__(**kwargs)
captured_clients.append(self)
import wf_sources_mcp.client.transport as mod
monkeypatch.setattr(
mod,
"httpx",
type("_PatchedHttpx", (), {"AsyncClient": _CapturingClient})(),
)
async with open_mcp_session(connection, None):
pass
assert len(captured_clients) == 1
assert "Authorization" not in captured_clients[0].headers