mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(mcp): keep one upstream MCP session per gateway session across tool calls
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
860bc7811d
commit
71aa41699a
8 changed files with 332 additions and 11 deletions
|
|
@ -672,6 +672,9 @@ class MCPClient:
|
|||
await anyio.lowlevel.checkpoint_if_cancelled()
|
||||
return result
|
||||
|
||||
def open_persistent_session(self) -> "PersistentMCPSession":
|
||||
return PersistentMCPSession(self)
|
||||
|
||||
def update_auth_value(self, mcp_auth_value: str | dict[str, str]) -> None:
|
||||
"""
|
||||
Set the authentication header for the MCP client.
|
||||
|
|
@ -838,10 +841,15 @@ class MCPClient:
|
|||
call_tool_request_params: MCPCallToolRequestParams,
|
||||
host_progress_callback: Callable | None = None,
|
||||
raise_on_error: bool = False,
|
||||
persistent_session: "PersistentMCPSession | None" = None,
|
||||
) -> MCPCallToolResult:
|
||||
"""
|
||||
Call an MCP Tool.
|
||||
|
||||
persistent_session runs the call inside an already-open upstream session (one
|
||||
upstream mcp-session-id shared by every call of a stateful gateway session) instead of
|
||||
the per-call initialize + teardown that run_with_session performs.
|
||||
|
||||
Args:
|
||||
raise_on_error: When True, re-raise the underlying exception instead of returning an
|
||||
``isError=True`` result. The token-exchange (OBO) tool-call path uses this to detect
|
||||
|
|
@ -871,8 +879,9 @@ class MCPClient:
|
|||
progress_callback=on_progress,
|
||||
)
|
||||
|
||||
run: Final = self.run_with_session if persistent_session is None else persistent_session.run
|
||||
try:
|
||||
tool_result: Final = await self.run_with_session(_call_tool_operation, quiet_on_error=raise_on_error)
|
||||
tool_result: Final = await run(_call_tool_operation, quiet_on_error=raise_on_error)
|
||||
verbose_logger.info("MCP client tool call '%s' completed successfully", call_tool_request_params.name)
|
||||
return tool_result
|
||||
except asyncio.CancelledError:
|
||||
|
|
@ -1178,3 +1187,70 @@ class MCPClient:
|
|||
"the MCP server may have crashed, disconnected, or timed out."
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
_PendingOperation: TypeAlias = tuple[Callable[[ClientSession], Awaitable[object]], "asyncio.Future[object]"]
|
||||
|
||||
|
||||
class PersistentMCPSession:
|
||||
"""One upstream MCP session kept open across operations.
|
||||
|
||||
The SDK transport owns anyio cancel scopes that must be entered and exited by the same
|
||||
task, so a dedicated task holds the session open and serves queued operations in order.
|
||||
"""
|
||||
|
||||
def __init__(self, client: MCPClient) -> None:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
self._client: Final = client
|
||||
self._queue: Final[asyncio.Queue[_PendingOperation | None]] = asyncio.Queue()
|
||||
self._ready: Final[asyncio.Future[None]] = loop.create_future()
|
||||
self._task: Final = loop.create_task(self._serve())
|
||||
|
||||
@property
|
||||
def closed(self) -> bool:
|
||||
return self._task.done()
|
||||
|
||||
async def _serve(self) -> None:
|
||||
async def drain(session: ClientSession) -> None:
|
||||
self._ready.set_result(None)
|
||||
while (pending := await self._queue.get()) is not None:
|
||||
operation, future = pending
|
||||
try:
|
||||
future.set_result(await operation(session))
|
||||
except Exception as e:
|
||||
future.set_exception(e)
|
||||
if isinstance(e, (ValueError, httpx2.HTTPError, OSError, MCPError)):
|
||||
return
|
||||
|
||||
try:
|
||||
await self._client.run_with_session(drain, quiet_on_error=True)
|
||||
except BaseException as e:
|
||||
if not self._ready.done():
|
||||
self._ready.set_exception(e)
|
||||
if not isinstance(e, Exception):
|
||||
raise
|
||||
finally:
|
||||
while not self._queue.empty():
|
||||
pending = self._queue.get_nowait()
|
||||
if pending is not None and not pending[1].done():
|
||||
pending[1].set_exception(RuntimeError("upstream MCP session closed"))
|
||||
|
||||
async def run(
|
||||
self,
|
||||
operation: Callable[[ClientSession], Awaitable[TSessionResult]],
|
||||
*,
|
||||
quiet_on_error: bool = False,
|
||||
) -> TSessionResult:
|
||||
del quiet_on_error
|
||||
await self._ready
|
||||
if self.closed:
|
||||
raise RuntimeError("upstream MCP session closed")
|
||||
future: Final[asyncio.Future[object]] = asyncio.get_running_loop().create_future()
|
||||
self._queue.put_nowait((operation, future))
|
||||
return cast(TSessionResult, await future)
|
||||
|
||||
def close(self) -> None:
|
||||
self._queue.put_nowait(None)
|
||||
|
||||
async def wait_closed(self) -> None:
|
||||
await asyncio.gather(self._task, return_exceptions=True)
|
||||
|
|
|
|||
|
|
@ -62,7 +62,13 @@ from litellm.constants import (
|
|||
MCP_TOOL_LISTING_TIMEOUT,
|
||||
)
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme, to_basic_credentials
|
||||
from litellm.experimental_mcp_client.client import (
|
||||
MCPClient,
|
||||
MCPSigV4Auth,
|
||||
PersistentMCPSession,
|
||||
strip_auth_scheme,
|
||||
to_basic_credentials,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
|
||||
)
|
||||
|
|
@ -1861,6 +1867,10 @@ class MCPServerManager:
|
|||
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
|
||||
)
|
||||
self.registry: dict[str, MCPServer] = {}
|
||||
# (gateway mcp-session-id, server_id, upstream auth fingerprint) -> the one upstream
|
||||
# session every tool call of that gateway session reuses, so stateful upstreams keep
|
||||
# their per-session state between calls. Released with the gateway session.
|
||||
self._upstream_sessions: dict[tuple[str, str, str], PersistentMCPSession] = {}
|
||||
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
|
||||
self.config_mcp_servers: dict[str, MCPServer] = {}
|
||||
"""
|
||||
|
|
@ -5776,6 +5786,30 @@ class MCPServerManager:
|
|||
async with semaphore:
|
||||
yield
|
||||
|
||||
async def _upstream_session_for(
|
||||
self,
|
||||
client: MCPClient,
|
||||
mcp_server: MCPServer,
|
||||
raw_headers: Mapping[str, str] | None,
|
||||
) -> PersistentMCPSession | None:
|
||||
gateway_session_id: Final = next(
|
||||
(value for key, value in (raw_headers or {}).items() if key.lower() == "mcp-session-id" and value),
|
||||
None,
|
||||
)
|
||||
if gateway_session_id is None or mcp_server.transport == MCPTransport.stdio:
|
||||
return None
|
||||
key: Final = (gateway_session_id, mcp_server.server_id, await client.discovery_auth_fingerprint())
|
||||
existing: Final = self._upstream_sessions.get(key)
|
||||
if existing is not None and not existing.closed:
|
||||
return existing
|
||||
opened: Final = client.open_persistent_session()
|
||||
self._upstream_sessions[key] = opened
|
||||
return opened
|
||||
|
||||
def release_upstream_sessions(self, gateway_session_id: str) -> None:
|
||||
for key in tuple(key for key in self._upstream_sessions if key[0] == gateway_session_id):
|
||||
self._upstream_sessions.pop(key).close()
|
||||
|
||||
async def _obo_call_tool_with_retry(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -5982,6 +6016,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
persistent_session: Final = await self._upstream_session_for(client, mcp_server, raw_headers)
|
||||
|
||||
call_tool_params: Final = MCPCallToolRequestParams(
|
||||
name=original_tool_name,
|
||||
|
|
@ -6019,7 +6054,9 @@ class MCPServerManager:
|
|||
async def _call_tool_via_client(client, params):
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
if not relays_upstream_auth:
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
return await client.call_tool(
|
||||
params, host_progress_callback=host_progress_callback, persistent_session=persistent_session
|
||||
)
|
||||
# The client-forwarded modes carry the caller's own upstream token, so an upstream
|
||||
# 401 (expired/invalid token) is the caller's to resolve: relay it as
|
||||
# MCPUpstreamAuthError so single-server REST callers turn it into a 401 +
|
||||
|
|
@ -6031,7 +6068,10 @@ class MCPServerManager:
|
|||
# the same isError degradation the default path produces.
|
||||
try:
|
||||
return await client.call_tool(
|
||||
params, host_progress_callback=host_progress_callback, raise_on_error=True
|
||||
params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
raise_on_error=True,
|
||||
persistent_session=persistent_session,
|
||||
)
|
||||
except Exception as e:
|
||||
auth_info: Final = _extract_upstream_auth_failure(e)
|
||||
|
|
|
|||
|
|
@ -628,6 +628,7 @@ if MCP_AVAILABLE:
|
|||
_stateful_session_locks.pop(session_id, None)
|
||||
_stateful_session_active_request_counts.pop(session_id, None)
|
||||
_stateful_session_client_info.pop(session_id, None)
|
||||
operations.global_mcp_server_manager.release_upstream_sessions(session_id)
|
||||
|
||||
# Keep this alias so existing references to session_manager still work
|
||||
session_manager: Final = session_manager_stateless
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
import queue
|
||||
from collections.abc import Iterator
|
||||
from collections.abc import Generator, Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
|
@ -9,10 +9,11 @@ import httpx
|
|||
from integration._support.asgi import asgi_server
|
||||
from integration._support.client import Gateway, Scenario
|
||||
from integration._support.database import read_rows
|
||||
from mcp.server.mcpserver import MCPServer
|
||||
from mcp.server.mcpserver import Context, MCPServer
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from mcp_tests.mcp_e2e_upstream_server import add, multiply
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
|
||||
|
|
@ -65,6 +66,55 @@ def mcp_peer() -> Iterator[McpPeer]:
|
|||
yield McpPeer(url + "/mcp", observed)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def stateful_mcp_peer() -> Generator[McpPeer]:
|
||||
service: Final = MCPServer("integration-stateful")
|
||||
selected: Final[dict[str, str]] = {} # mutable-ok: per-session state the upstream keeps across tool calls
|
||||
|
||||
def upstream_session(ctx: Context) -> str:
|
||||
assert ctx.headers is not None, "stateful upstream requires HTTP request headers"
|
||||
return ctx.headers["mcp-session-id"]
|
||||
|
||||
@service.tool()
|
||||
def select_project(name: str, ctx: Context) -> str:
|
||||
selected[upstream_session(ctx)] = name
|
||||
return f"selected {name}"
|
||||
|
||||
@service.tool()
|
||||
def create_feature(title: str, ctx: Context) -> str:
|
||||
project: Final = selected.get(upstream_session(ctx))
|
||||
if project is None:
|
||||
raise ValueError("no project selected in this session")
|
||||
return f"{project}/{title}"
|
||||
|
||||
app: Final = service.streamable_http_app(
|
||||
stateless_http=False,
|
||||
json_response=True,
|
||||
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=False),
|
||||
)
|
||||
observed: Final[queue.Queue[dict[str, object]]] = queue.Queue()
|
||||
|
||||
async def capture(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] == "http" and scope["method"] == "GET":
|
||||
await Response(status_code=405)(scope, receive, send)
|
||||
return
|
||||
if scope["type"] != "http" or scope["method"] != "POST":
|
||||
await app(scope, receive, send)
|
||||
return
|
||||
body: Final = await Request(scope, receive).body()
|
||||
observed.put({"body": json.loads(body) if body else None, "headers": dict(scope["headers"])})
|
||||
message: Final[Message] = {"type": "http.request", "body": body, "more_body": False}
|
||||
pending: Final = iter((message,))
|
||||
|
||||
async def replay() -> Message:
|
||||
return next(pending, {"type": "http.disconnect"})
|
||||
|
||||
await app(scope, replay, send)
|
||||
|
||||
with asgi_server(capture) as url:
|
||||
yield McpPeer(url + "/mcp", observed)
|
||||
|
||||
|
||||
def register_mcp(scenario: Scenario, peer: McpPeer, alias: str, **fields: object) -> str:
|
||||
response: Final = scenario.gateway.request(
|
||||
"POST", "/v1/mcp/server", {"server_name": alias, "alias": alias, "url": peer.url, "transport": "http", **fields}
|
||||
|
|
|
|||
|
|
@ -250,6 +250,9 @@
|
|||
"tests/integration/mcp/test_mcp_lifecycle.py::test_generated_mcp_edits_preserve_actual_headers_and_tool_results": [
|
||||
"other.mcp.lifecycle.generated_save_reload_preserves_effective_headers"
|
||||
],
|
||||
"tests/integration/mcp/test_mcp_lifecycle.py::test_stateful_upstream_keeps_selection_across_tool_calls_through_gateway": [
|
||||
"mcp.call_tool.session.stateful_upstream_state_survives_between_tool_calls"
|
||||
],
|
||||
"tests/integration/observability/test_callback_delivery.py::test_concurrent_success_and_failure_join_callbacks_and_rows_without_credentials": [
|
||||
"other.observability.callbacks.credentials_stay_out_of_event_bodies",
|
||||
"other.observability.callbacks.concurrent_results_join_complete_events_and_rows"
|
||||
|
|
|
|||
|
|
@ -1,19 +1,24 @@
|
|||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
from contextlib import ExitStack
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from typing import Final, cast
|
||||
|
||||
import httpx
|
||||
import httpx2
|
||||
import pytest
|
||||
import yaml
|
||||
from hypothesis import strategies as st
|
||||
from hypothesis.stateful import RuleBasedStateMachine, invariant, rule, run_state_machine_as_test
|
||||
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests
|
||||
from integration._support.mcp import call_tool, mcp_peer, register_mcp, stateful_mcp_peer, tool_names
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.mcp import call_tool, mcp_peer, register_mcp, tool_names
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
|
||||
@pytest.mark.covers("mcp.call_tool.saved_headers.reach_actual_transport")
|
||||
|
|
@ -304,3 +309,40 @@ def test_same_url_server_grants_scope_discovery_and_direct_or_virtual_execution(
|
|||
assert calls[0]["headers"][b"x-integration-server"] == aliases[server_index].encode()
|
||||
expected_auth: Final = f"Bearer synthetic-{aliases[server_index]}".encode() if authenticated else None
|
||||
assert all(item["headers"].get(b"authorization") == expected_auth for item in observed)
|
||||
|
||||
|
||||
async def _select_then_create(url: str, key: str) -> tuple[CallToolResult, CallToolResult]:
|
||||
async with httpx2.AsyncClient(headers={"x-litellm-api-key": key}, timeout=30, trust_env=False) as http:
|
||||
async with streamable_http_client(url, http_client=http) as (read, write):
|
||||
async with ClientSession(read, write) as session:
|
||||
await session.initialize()
|
||||
names: Final = {tool.name.rsplit("-", 1)[-1]: tool.name for tool in (await session.list_tools()).tools}
|
||||
selected: Final = await session.call_tool(names["select_project"], {"name": "alpha"})
|
||||
created: Final = await session.call_tool(names["create_feature"], {"title": "login"})
|
||||
return selected, created
|
||||
|
||||
|
||||
def _upstream_session_ids(observed: tuple[dict[str, object], ...]) -> set[str | None]:
|
||||
headers: Final = tuple(cast(dict[bytes, bytes], item["headers"]) for item in observed)
|
||||
return {raw.decode() if isinstance(raw := item.get(b"mcp-session-id"), bytes) else None for item in headers}
|
||||
|
||||
|
||||
@pytest.mark.covers("mcp.call_tool.session.stateful_upstream_state_survives_between_tool_calls")
|
||||
def test_stateful_upstream_keeps_selection_across_tool_calls_through_gateway(gateway: Gateway) -> None:
|
||||
with stateful_mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "integration" + uuid.uuid4().hex
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
peer.drain()
|
||||
selected, created = asyncio.run(_select_then_create(f"{gateway.client.base_url}/{alias}/mcp", key))
|
||||
observed: Final = peer.drain()
|
||||
for session_id in _upstream_session_ids(observed):
|
||||
if session_id is not None:
|
||||
httpx.delete(peer.url, headers={"mcp-session-id": session_id}, trust_env=False)
|
||||
call_sessions: Final = _upstream_session_ids(
|
||||
tuple(item for item in observed if cast(dict[str, object], item["body"]).get("method") == "tools/call")
|
||||
)
|
||||
assert selected.is_error is False, selected.model_dump_json()
|
||||
assert created.is_error is False, created.model_dump_json()
|
||||
assert isinstance(created.content[0], TextContent) and created.content[0].text == "alpha/login"
|
||||
assert len(call_sessions) == 1 and None not in call_sessions, call_sessions
|
||||
|
|
|
|||
|
|
@ -15,6 +15,9 @@ import httpx2
|
|||
import pytest
|
||||
from mcp import MCPError
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.server.mcpserver import Context
|
||||
from mcp.server.mcpserver import MCPServer as UpstreamServer
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.types import (
|
||||
CONNECTION_CLOSED,
|
||||
|
|
@ -2896,3 +2899,64 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
closed: Final = await asyncio.wait_for(asyncio.gather(*connections, return_exceptions=True), 2)
|
||||
assert all(result is None or isinstance(result, asyncio.CancelledError) for result in closed), closed
|
||||
await asyncio.wait_for(listener.wait_closed(), 2)
|
||||
|
||||
|
||||
class _StatefulUpstreamClient(MCPClient):
|
||||
"""An MCPClient whose streamable-HTTP transport talks to an in-process stateful MCP server."""
|
||||
|
||||
def __init__(self, app, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self._app = app
|
||||
|
||||
def _create_transport_context(self) -> tuple[_TransportContext, httpx2.AsyncClient]:
|
||||
http_client: Final = self._create_httpx_client_factory(transport=httpx2.ASGITransport(app=self._app))(
|
||||
headers=self._get_auth_headers(), timeout=httpx2.Timeout(self.timeout)
|
||||
)
|
||||
return streamable_http_client(self.server_url, http_client=http_client), http_client
|
||||
|
||||
|
||||
def _stateful_upstream():
|
||||
service: Final = UpstreamServer("stateful")
|
||||
selected: Final[dict[str, str]] = {}
|
||||
|
||||
@service.tool()
|
||||
def select_project(name: str, ctx: Context) -> str:
|
||||
selected[ctx.headers["mcp-session-id"]] = name
|
||||
return f"selected {name}"
|
||||
|
||||
@service.tool()
|
||||
def create_feature(title: str, ctx: Context) -> str:
|
||||
return f"{selected[ctx.headers['mcp-session-id']]}/{title}"
|
||||
|
||||
return service.streamable_http_app(
|
||||
stateless_http=False,
|
||||
json_response=True,
|
||||
transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=False),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_persistent_session_keeps_upstream_state_across_tool_calls():
|
||||
app: Final = _stateful_upstream()
|
||||
async with app.router.lifespan_context(app):
|
||||
client: Final = _StatefulUpstreamClient(app, server_url="http://upstream/mcp", transport_type=MCPTransport.http)
|
||||
per_call: Final = await client.call_tool(CallToolRequestParams(name="select_project", arguments={"name": "a"}))
|
||||
assert per_call.is_error is False
|
||||
fresh: Final = await client.call_tool(CallToolRequestParams(name="create_feature", arguments={"title": "b"}))
|
||||
assert fresh.is_error is True, "a fresh upstream session per call must not see the earlier selection"
|
||||
|
||||
session: Final = client.open_persistent_session()
|
||||
try:
|
||||
selected: Final = await client.call_tool(
|
||||
CallToolRequestParams(name="select_project", arguments={"name": "a"}), persistent_session=session
|
||||
)
|
||||
created: Final = await client.call_tool(
|
||||
CallToolRequestParams(name="create_feature", arguments={"title": "b"}), persistent_session=session
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
assert selected.is_error is False
|
||||
assert created.is_error is False
|
||||
assert created.content[0].text == "a/b"
|
||||
await asyncio.wait_for(session.wait_closed(), 5)
|
||||
assert session.closed
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import os
|
|||
import sys
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Final, Literal, Optional
|
||||
from typing import Any, Dict, Final, Literal, Optional, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -27,7 +27,7 @@ import contextlib
|
|||
|
||||
import httpx
|
||||
import httpx2
|
||||
from mcp import ReadResourceResult, Resource
|
||||
from mcp import ClientSession, ReadResourceResult, Resource
|
||||
from mcp.types import (
|
||||
CallToolResult,
|
||||
GetPromptResult,
|
||||
|
|
@ -39,6 +39,7 @@ from mcp.types import Tool as MCPTool
|
|||
from pydantic import AnyUrl, TypeAdapter
|
||||
|
||||
from litellm.constants import MCP_METADATA_TIMEOUT
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
_deserialize_json_dict,
|
||||
|
|
@ -14516,3 +14517,47 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
|
|||
assert captured["client_ip"] is None
|
||||
finally:
|
||||
auth_context_var.reset(token)
|
||||
|
||||
|
||||
class _OfflineMCPClient(MCPClient):
|
||||
"""An MCPClient whose sessions never touch the network: operations run against a stand-in session."""
|
||||
|
||||
async def run_with_session(self, operation, *, quiet_on_error: bool = False):
|
||||
del quiet_on_error
|
||||
return await operation(cast(ClientSession, object()))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_session_is_shared_per_gateway_session_and_released_with_it():
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="s1", name="stateful", url="http://upstream/mcp", transport=MCPTransport.http)
|
||||
client: Final = _OfflineMCPClient(server_url=server.url, transport_type=MCPTransport.http)
|
||||
try:
|
||||
assert await manager._upstream_session_for(client, server, None) is None
|
||||
assert await manager._upstream_session_for(client, server, {"accept": "application/json"}) is None
|
||||
|
||||
first: Final = await manager._upstream_session_for(client, server, {"Mcp-Session-Id": "gw-1"})
|
||||
second: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"})
|
||||
other: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-2"})
|
||||
assert first is not None and first is second, "every call of one gateway session must share one upstream session"
|
||||
assert other is not None and other is not first, "distinct gateway sessions must not share upstream state"
|
||||
|
||||
manager.release_upstream_sessions("gw-1")
|
||||
await asyncio.wait_for(first.wait_closed(), 5)
|
||||
assert first.closed and not other.closed
|
||||
replacement: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"})
|
||||
assert replacement is not first and not replacement.closed
|
||||
|
||||
other.close()
|
||||
await asyncio.wait_for(other.wait_closed(), 5)
|
||||
reopened: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-2"})
|
||||
assert reopened is not other and not reopened.closed, "a dead upstream session must be replaced, not reused"
|
||||
|
||||
stdio: Final = MCPServer(server_id="s2", name="local", command="cat", transport=MCPTransport.stdio)
|
||||
assert await manager._upstream_session_for(client, stdio, {"mcp-session-id": "gw-1"}) is None
|
||||
finally:
|
||||
for gateway_session_id in ("gw-1", "gw-2"):
|
||||
manager.release_upstream_sessions(gateway_session_id)
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(*(s.wait_closed() for s in manager._upstream_sessions.values()), return_exceptions=True), 5
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue