Detach forwarded proxy metadata

🤖 Generated with OpenAI Codex
This commit is contained in:
Jake Kaplan 2026-08-06 11:16:55 -04:00
commit 3f4e9aac27
2 changed files with 68 additions and 4 deletions

View file

@ -12,6 +12,7 @@ import inspect
import time
import warnings
from collections.abc import Awaitable, Callable, Sequence
from copy import deepcopy
from dataclasses import dataclass, replace
from typing import TYPE_CHECKING, Any, Literal, cast
@ -1119,10 +1120,13 @@ class _UpstreamServerMetadata:
result: mcp_types.InitializeResult | mcp_types.DiscoverResult,
server_info: mcp_types.Implementation | None,
) -> _UpstreamServerMetadata:
"""Detach forwarded values from the backend session's adopted result."""
return cls(
instructions=result.instructions,
server_info=server_info,
meta=dict(result.meta or {}),
server_info=(
server_info.model_copy(deep=True) if server_info is not None else None
),
meta=deepcopy(result.meta or {}),
)
@classmethod

View file

@ -38,6 +38,9 @@ UPSTREAM_INFO = mcp_types.Implementation(
class UpstreamMetadataMiddleware(Middleware):
"""Advertise metadata that differs from the gateway's own claims."""
def __init__(self, server_info: mcp_types.Implementation = UPSTREAM_INFO) -> None:
self.server_info = server_info
def _updates(self, result: mcp_types.Result) -> dict[str, Any]:
meta = {
**(result.meta or {}),
@ -53,13 +56,13 @@ class UpstreamMetadataMiddleware(Middleware):
}
if isinstance(result, mcp_types.InitializeResult):
updates.update(
server_info=UPSTREAM_INFO,
server_info=self.server_info,
capabilities=mcp_types.ServerCapabilities(
experimental={"upstream": {"claimed": True}}
),
)
else:
meta[mcp_types.SERVER_INFO_META_KEY] = UPSTREAM_INFO.model_dump(
meta[mcp_types.SERVER_INFO_META_KEY] = self.server_info.model_dump(
by_alias=True, mode="json", exclude_none=True
)
updates.update(
@ -326,6 +329,63 @@ async def test_connected_pinned_client_probes_without_adopting_metadata():
assert backend_client.instructions is None
@pytest.mark.parametrize("mode", ["legacy", "auto"])
async def test_forwarded_metadata_does_not_alias_connected_backend(mode: str):
backend_info = mcp_types.Implementation(name="shared-backend", version="1.0")
upstream = FastMCP(
"upstream",
middleware=[UpstreamMetadataMiddleware(backend_info)],
)
class MutateForwardedMetadata(Middleware):
def _mutate(self, result: ResultT) -> ResultT:
assert result.meta is not None
nested = result.meta["com.example/upstream"]
assert isinstance(nested, dict)
nested["enabled"] = False
if isinstance(result, mcp_types.InitializeResult):
result.server_info.name = "frontend mutation"
else:
server_info = result.meta[mcp_types.SERVER_INFO_META_KEY]
assert isinstance(server_info, dict)
server_info["name"] = "frontend mutation"
return result
async def on_initialize(self, context, call_next):
result = await call_next(context)
assert result is not None
return self._mutate(result)
async def on_discover(self, context, call_next):
result = await call_next(context)
if not isinstance(result, mcp_types.DiscoverResult):
return result
return self._mutate(result)
async with Client(upstream, mode=mode) as backend_client:
provider = ProxyProvider(lambda: backend_client)
gateway = FastMCP(
"gateway",
providers=[provider],
middleware=[
MutateForwardedMetadata(),
ProxyMetadataMiddleware(provider, identity="upstream"),
],
)
async with Client(gateway, mode=mode):
pass
backend_result = (
backend_client.session.initialize_result
or backend_client.session.discover_result
)
assert backend_result is not None
assert backend_result.meta is not None
assert backend_result.meta["com.example/upstream"] == {"enabled": True}
assert backend_client.server_info == backend_info
async def test_disconnected_pinned_client_is_not_cloned():
class UnclonableProxyClient(ProxyClient):
def new(self) -> ProxyClient: