Forward proxy server metadata across protocol eras (#4776)

* Forward proxy negotiation metadata

🤖 Generated with OpenAI Codex

* Limit forwarded proxy metadata

🤖 Generated with OpenAI Codex

* Tighten negotiation metadata forwarding

🤖 Generated with OpenAI Codex

* Tighten proxy metadata docs

🤖 Generated with OpenAI Codex

* Keep proxy metadata middleware with provider

🤖 Generated with OpenAI Codex

* Simplify proxy negotiation middleware API

🤖 Generated with OpenAI Codex

* Name proxy metadata middleware directly

🤖 Generated with OpenAI Codex

* Preserve discovery middleware contracts

🤖 Generated with OpenAI Codex

* Clarify proxy metadata ownership

🤖 Generated with OpenAI Codex

* Align proxy metadata wording

🤖 Generated with OpenAI Codex

* Call forwarded values server metadata

🤖 Generated with OpenAI Codex

* Harden proxy metadata reads

🤖 Generated with OpenAI Codex

* Expose configured discovery result

🤖 Generated with OpenAI Codex

* Preserve proxy discovery compatibility

🤖 Generated with OpenAI Codex

* Preserve deprecated initialization middleware

🤖 Generated with OpenAI Codex

* Harden proxy metadata boundaries

🤖 Generated with OpenAI Codex

* Restore deprecated middleware location

🤖 Generated with OpenAI Codex

* Simplify proxy metadata client lifecycle

🤖 Generated with OpenAI Codex

* Clarify proxy metadata lifecycle

🤖 Generated with OpenAI Codex

* Preserve proxy factory errors

🤖 Generated with OpenAI Codex

* Detach forwarded proxy metadata

🤖 Generated with OpenAI Codex

* Simplify proxy metadata implementation

🤖 Generated with OpenAI Codex

* Distinguish proxy metadata failures

🤖 Generated with OpenAI Codex

* Narrow proxy metadata validation fallback

🤖 Generated with OpenAI Codex

* Retrigger CI

🤖 Generated with OpenAI Codex
This commit is contained in:
Jake Kaplan 2026-08-06 19:09:05 -04:00 committed by GitHub
commit 9feb1f378b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 1113 additions and 134 deletions

View file

@ -22,7 +22,7 @@ from typing import Any
import pytest
from mcp import ClientSession
from mcp.shared.exceptions import MCPError
from mcp_types import METHOD_NOT_FOUND
from mcp_types import METHOD_NOT_FOUND, DiscoverResult, ServerCapabilities
from mcp_types.version import LATEST_HANDSHAKE_VERSION, LATEST_MODERN_VERSION
from typing_extensions import Unpack
@ -242,6 +242,19 @@ class TestNonConformantModernPeer:
class TestPinnedMode:
def test_prior_discover_is_exposed(self, fastmcp_server):
prior = DiscoverResult(
supported_versions=[LATEST_MODERN_VERSION],
capabilities=ServerCapabilities(),
)
client = Client(
fastmcp_server,
mode=LATEST_MODERN_VERSION,
prior_discover=prior,
)
assert client.prior_discover is prior
async def test_pinned_modern_adopts_without_probe(self, fastmcp_server):
"""Pinning the modern version adopts it directly; a synthesized
DiscoverResult carries no identity, so server_info is absent."""

View file

@ -0,0 +1,108 @@
"""Tests for typed middleware support during modern discovery."""
from typing import Any
import mcp_types
from mcp_types.version import LATEST_MODERN_VERSION
from fastmcp import Client, FastMCP
from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
async def test_on_discover_receives_and_transforms_typed_result():
class DiscoveryMiddleware(Middleware):
def __init__(self) -> None:
self.request: mcp_types.DiscoverRequest | None = None
self.result: mcp_types.DiscoverResult | None = None
async def on_discover(
self,
context: MiddlewareContext[mcp_types.DiscoverRequest],
call_next: CallNext[
mcp_types.DiscoverRequest,
mcp_types.DiscoverResult | dict[str, Any],
],
) -> mcp_types.DiscoverResult | dict[str, Any]:
self.request = context.message
result = await call_next(context)
assert isinstance(result, mcp_types.DiscoverResult)
self.result = result
return result.model_copy(update={"instructions": "discovered"})
middleware = DiscoveryMiddleware()
server = FastMCP("typed-discovery", middleware=[middleware])
async with Client(server, mode="auto") as client:
assert client.instructions == "discovered"
assert isinstance(middleware.request, mcp_types.DiscoverRequest)
assert isinstance(middleware.result, mcp_types.DiscoverResult)
async def test_on_discover_forwards_modified_params():
modified = False
server = FastMCP("modified-discovery")
default_handler = server._mcp_server._handle_discover
async def capture_params(ctx, params):
nonlocal modified
assert params is not None
assert params.meta is not None
modified = params.meta["com.example/modified"] is True
return await default_handler(ctx, params)
server._mcp_server.add_request_handler(
"server/discover", mcp_types.RequestParams, capture_params
)
class ModifyParams(Middleware):
async def on_discover(self, context, call_next):
assert context.message.params is not None
assert context.message.params.meta is not None
context.message.params = mcp_types.RequestParams(
meta={
**context.message.params.meta,
"com.example/modified": True,
}
)
return await call_next(context)
server.add_middleware(ModifyParams())
async with Client(server, mode="auto"):
pass
assert modified
async def test_on_discover_preserves_extension_owned_result():
extension_result = {
"resultType": "com.example/custom",
"payload": {"enabled": True},
}
async def custom_discover(_ctx, _params):
return extension_result
class ObserveExtension(Middleware):
def __init__(self) -> None:
self.result: mcp_types.DiscoverResult | dict[str, Any] | None = None
async def on_discover(self, context, call_next):
self.result = await call_next(context)
return self.result
middleware = ObserveExtension()
server = FastMCP("extension-discovery", middleware=[middleware])
server._mcp_server.add_request_handler(
"server/discover", mcp_types.RequestParams, custom_discover
)
async with Client(server, mode=LATEST_MODERN_VERSION) as client:
result = await client.session.send_discover(LATEST_MODERN_VERSION)
assert isinstance(result, dict)
assert result["resultType"] == "com.example/custom"
assert result["payload"] == {"enabled": True}
assert isinstance(middleware.result, dict)
assert middleware.result["payload"] == {"enabled": True}

View file

@ -159,6 +159,23 @@ class TestUnroutableAndMalformed:
assert ("on_message", "tools/call") in recorder.records
assert ("on_call_tool", "tools/call") not in recorder.records
async def test_malformed_discover_params_observed_by_generic_hooks(self):
server = _adder()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server) as client:
recorder.records.clear()
with pytest.raises(MCPError):
await _raw_request(
client,
"server/discover",
{"_meta": {"progressToken": []}},
)
assert ("on_message", "server/discover") in recorder.records
assert ("on_request", "server/discover") in recorder.records
class TestSingleFire:
async def test_each_hook_fires_once_per_component_call(self):

View file

@ -204,13 +204,7 @@ async def test_create_proxy_with_transport(fastmcp_server):
async def test_proxy_forwards_upstream_instructions():
"""A proxy should surface the upstream server's instructions in the handshake.
`FastMCPProxy` registers a `server/discover` handler that forwards the
upstream's instructions, mirroring what `ProxyInitializeMiddleware.on_initialize`
already does for the legacy handshake, so `client.session.instructions`
(era-neutral) resolves the same way on both protocol eras.
"""
"""The metadata middleware forwards upstream instructions."""
upstream = FastMCP(name="upstream", instructions="USE_THIS_MARKER_123")
proxy = create_proxy(upstream, name="proxy")
@ -274,35 +268,25 @@ async def test_proxy_ping_surfaces_wrong_remote_path():
async with run_server_async(remote, transport="http") as url:
proxy = create_proxy(StreamableHttpTransport(url.removesuffix("/mcp")))
# This asserts the error surfaces from merely *connecting* to the proxy,
# with no operation performed. That only happens on the legacy handshake:
# `ProxyInitializeMiddleware.on_initialize` eagerly probes the backend
# during the front's own `initialize` call. A modern front negotiates
# `server/discover` instead, which never runs that middleware hook, so
# connecting succeeds regardless of backend health and the failure would
# only surface on first real use. Pinned because the subject here is
# that eager, handshake-time probe.
#
# SDK v2 surfaces a wrong remote path as an HTTP "Not Found" rather than
# the v1 "Session terminated" message.
with pytest.raises(MCPError, match="Not Found"):
async with Client(proxy, mode="legacy"):
pass
# Optional metadata lookup is best-effort, so the client can connect. The
# first real proxied operation reports the bad backend path instead.
async with Client(proxy, mode="legacy") as client:
with pytest.raises(MCPError, match="Not Found"):
await client.ping()
async def test_proxy_initialize_forwards_remote_connection_error():
async def test_proxy_initialize_defers_remote_connection_error():
port = find_available_port()
proxy = create_proxy(
StreamableHttpTransport(f"http://127.0.0.1:{port}/mcp"),
provider_error_strategy="raise",
)
# Same reasoning as test_proxy_ping_surfaces_wrong_remote_path above: the
# error surfaces from connecting alone only via the legacy handshake's
# eager backend probe in `ProxyInitializeMiddleware.on_initialize`.
with pytest.raises(MCPError, match="Client failed to connect"):
async with Client(proxy, mode="legacy"):
pass
# The client can connect without optional backend metadata; the first
# component operation reports the unavailable backend.
async with Client(proxy, mode="legacy") as client:
with pytest.raises(MCPError, match="Client failed to connect"):
await client.list_tools()
async def test_proxy_list_tools_surfaces_remote_connection_error():
@ -324,13 +308,10 @@ async def test_proxy_list_tools_surfaces_remote_connection_error():
async def test_proxy_list_tools_client_surfaces_remote_connection_error():
"""With a modern front, connecting succeeds (no eager backend probe — see
test_proxy_ping_surfaces_wrong_remote_path) and the failure only surfaces
once `list_tools()` actually hits the dead backend. `ProxyProvider._list_tools`
now normalizes the raw `httpx2.ConnectError` from the failed backend connect
into the `MCPError("Client failed to connect...")` this test expects, the
same way `ProxyInitializeMiddleware.on_initialize` and `ProxyTool.run`
already did.
"""Connecting succeeds and the first component operation reports the backend.
`ProxyProvider._list_tools` normalizes the raw transport failure into the
`MCPError("Client failed to connect...")` this test expects.
"""
port = find_available_port()
proxy = create_proxy(
@ -1459,13 +1440,7 @@ class TestProxyForwardingAppliesToEveryBackendClient:
class TestProxyModernEraInstructions:
"""Upstream instructions must reach a client on the modern era too.
`ProxyInitializeMiddleware.on_initialize` only fires for the legacy
handshake. A `mode="auto"` client negotiates via `server/discover`, which
the SDK builds from the low-level server's own `instructions`, so without a
discover-side hook the proxy drops its upstream's instructions entirely.
"""
"""Upstream instructions must reach a client on the modern era too."""
async def test_proxy_forwards_upstream_instructions_on_modern_era(self):
upstream = FastMCP(name="upstream", instructions="USE_THIS_MARKER_123")
@ -1491,13 +1466,7 @@ class TestProxyModernEraInstructions:
class TestProxyProviderTransportErrors:
"""A dead backend must surface as an MCPError, not a raw transport error.
`ProxyTool.run` and `ProxyInitializeMiddleware.on_initialize` normalize
connection failures into `MCPError`; the provider's list methods caught
only `MCPError`, so an `httpx2.ConnectError` (or the `RuntimeError` the
client wraps a failed connect in) escaped unwrapped to the caller.
"""
"""A dead backend must surface as an MCPError, not a raw transport error."""
@pytest.fixture
def unreachable_provider(self) -> ProxyProvider:

View file

@ -0,0 +1,637 @@
"""Server metadata forwarding across proxy protocol eras."""
from itertools import product
from typing import Any, Literal, TypeVar
import mcp_types
import pytest
from mcp import MCPError
from mcp_types.version import MODERN_PROTOCOL_VERSIONS
from fastmcp import Client, FastMCP, FastMCPDeprecationWarning
from fastmcp.client.logging import LogMessage
from fastmcp.client.transports import StreamableHttpTransport
from fastmcp.server import create_proxy
from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
from fastmcp.server.providers.proxy import (
FastMCPProxy,
ProxyClient,
ProxyInitializeMiddleware,
ProxyMetadataMiddleware,
ProxyProvider,
StatefulProxyClient,
)
from fastmcp.utilities.http import find_available_port
ResultT = TypeVar("ResultT", bound=mcp_types.Result)
UPSTREAM_INFO = mcp_types.Implementation(
name="upstream",
title="Upstream title",
version="1.2.3",
description="Upstream description",
website_url="https://upstream.example.com",
icons=[mcp_types.Icon(src="https://upstream.example.com/icon.png")],
)
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 {}),
mcp_types.PROTOCOL_VERSION_META_KEY: "upstream-version",
mcp_types.CLIENT_INFO_META_KEY: {"name": "upstream-client"},
mcp_types.CLIENT_CAPABILITIES_META_KEY: {"upstream": True},
"com.example/upstream": {"enabled": True},
"com.example/shared": "upstream",
}
updates: dict[str, Any] = {
"instructions": "upstream instructions",
"meta": meta,
}
if isinstance(result, mcp_types.InitializeResult):
updates.update(
server_info=self.server_info,
capabilities=mcp_types.ServerCapabilities(
experimental={"upstream": {"claimed": True}}
),
)
else:
meta[mcp_types.SERVER_INFO_META_KEY] = self.server_info.model_dump(
by_alias=True, mode="json", exclude_none=True
)
updates.update(
ttl_ms=91_000,
cache_scope="public",
capabilities=mcp_types.ServerCapabilities(
experimental={"upstream": {"claimed": True}}
),
)
return updates
async def on_initialize(
self,
context: MiddlewareContext[mcp_types.InitializeRequest],
call_next: CallNext[
mcp_types.InitializeRequest, mcp_types.InitializeResult | None
],
) -> mcp_types.InitializeResult | None:
result = await call_next(context)
assert result is not None
return result.model_copy(update=self._updates(result))
async def on_discover(
self,
context: MiddlewareContext[mcp_types.DiscoverRequest],
call_next: CallNext[
mcp_types.DiscoverRequest,
mcp_types.DiscoverResult | dict[str, Any],
],
) -> mcp_types.DiscoverResult | dict[str, Any]:
result = await call_next(context)
if not isinstance(result, mcp_types.DiscoverResult):
return result
return result.model_copy(update=self._updates(result))
class FrontendMetadataMiddleware(Middleware):
"""Set frontend values that must win over the upstream on collision."""
def _update(self, result: ResultT) -> ResultT:
return result.model_copy(
update={
"meta": {
**(result.meta or {}),
"com.example/shared": "frontend",
"com.example/frontend": {"enabled": True},
},
}
)
async def on_initialize(
self,
context: MiddlewareContext[mcp_types.InitializeRequest],
call_next: CallNext[
mcp_types.InitializeRequest, mcp_types.InitializeResult | None
],
) -> mcp_types.InitializeResult | None:
result = await call_next(context)
assert result is not None
return self._update(result)
async def on_discover(
self,
context: MiddlewareContext[mcp_types.DiscoverRequest],
call_next: CallNext[
mcp_types.DiscoverRequest,
mcp_types.DiscoverResult | dict[str, Any],
],
) -> mcp_types.DiscoverResult | dict[str, Any]:
result = await call_next(context)
if not isinstance(result, mcp_types.DiscoverResult):
return result
return self._update(result)
def make_upstream() -> FastMCP:
return FastMCP("unmodified-upstream", middleware=[UpstreamMetadataMiddleware()])
def make_gateway(
upstream: FastMCP,
*,
backend_mode: str,
identity: Literal["proxy", "upstream"] = "proxy",
instructions: str | None = None,
frontend_metadata: bool = False,
) -> FastMCP:
provider = ProxyProvider(lambda: ProxyClient(upstream, mode=backend_mode))
metadata = ProxyMetadataMiddleware(provider, identity=identity)
middleware: list[Middleware] = [metadata]
if frontend_metadata:
middleware.append(FrontendMetadataMiddleware())
gateway = FastMCP(
"gateway",
version="9.8.7",
instructions=instructions,
providers=[provider],
middleware=middleware,
cache_ttl=7,
cache_scope="private",
)
return gateway
@pytest.mark.parametrize(
("frontend_mode", "backend_mode"),
list(product(("legacy", "auto"), repeat=2)),
)
async def test_forwards_metadata_across_all_protocol_era_combinations(
frontend_mode: str, backend_mode: str
):
gateway = make_gateway(make_upstream(), backend_mode=backend_mode)
async with Client(gateway, mode=frontend_mode) as client:
result = client.session.initialize_result or client.session.discover_result
assert result is not None
assert client.instructions == "upstream instructions"
assert client.server_info is not None
assert client.server_info.name == "gateway"
assert result.meta is not None
assert result.meta["com.example/upstream"] == {"enabled": True}
for key in (
mcp_types.PROTOCOL_VERSION_META_KEY,
mcp_types.CLIENT_INFO_META_KEY,
mcp_types.CLIENT_CAPABILITIES_META_KEY,
):
assert key not in result.meta
stamped_info = result.meta.get(mcp_types.SERVER_INFO_META_KEY)
assert result.capabilities.experimental is None
if isinstance(result, mcp_types.InitializeResult):
assert result.protocol_version not in MODERN_PROTOCOL_VERSIONS
assert stamped_info is None
else:
assert stamped_info is not None
assert stamped_info["name"] == "gateway"
assert result.supported_versions == list(MODERN_PROTOCOL_VERSIONS)
assert result.ttl_ms == 7_000
assert result.cache_scope == "private"
assert result.result_type == "complete"
@pytest.mark.parametrize("frontend_mode", ["legacy", "auto"])
@pytest.mark.parametrize("identity", ["proxy", "upstream"])
async def test_identity_policy_forwards_full_implementation(
frontend_mode: str, identity: Literal["proxy", "upstream"]
):
gateway = make_gateway(make_upstream(), backend_mode="auto", identity=identity)
async with Client(gateway, mode=frontend_mode) as client:
assert client.server_info is not None
if identity == "proxy":
assert client.server_info.name == "gateway"
assert client.server_info.version == "9.8.7"
else:
assert client.server_info == UPSTREAM_INFO
result = client.session.initialize_result or client.session.discover_result
assert result is not None
if isinstance(result, mcp_types.InitializeResult):
assert mcp_types.SERVER_INFO_META_KEY not in (result.meta or {})
else:
assert result.meta is not None
assert result.meta[mcp_types.SERVER_INFO_META_KEY]["name"] == "upstream"
@pytest.mark.parametrize("frontend_mode", ["legacy", "auto"])
async def test_frontend_values_take_precedence(frontend_mode: str):
gateway = make_gateway(
make_upstream(),
backend_mode="auto",
instructions="frontend instructions",
frontend_metadata=True,
)
async with Client(gateway, mode=frontend_mode) as client:
result = client.session.initialize_result or client.session.discover_result
assert result is not None
assert client.instructions == "frontend instructions"
assert result.meta is not None
assert result.meta["com.example/shared"] == "frontend"
assert result.meta["com.example/frontend"] == {"enabled": True}
assert result.meta["com.example/upstream"] == {"enabled": True}
async def test_forwards_backend_logs_while_reading_metadata():
messages: list[str] = []
class LogOnInitialize(Middleware):
async def on_initialize(
self,
context: MiddlewareContext[mcp_types.InitializeRequest],
call_next: CallNext[
mcp_types.InitializeRequest, mcp_types.InitializeResult | None
],
) -> mcp_types.InitializeResult | None:
result = await call_next(context)
assert context.fastmcp_context is not None
await context.fastmcp_context.log("metadata connection")
return result
async def capture_log(message: LogMessage) -> None:
messages.append(message.data["msg"])
upstream = FastMCP("upstream", middleware=[LogOnInitialize()])
proxy = create_proxy(upstream)
async with Client(proxy, mode="legacy", log_handler=capture_log):
pass
assert messages == ["metadata connection"]
async def test_pinned_client_uses_prior_discover_metadata():
prior_info = mcp_types.Implementation(name="prior", version="1.0")
prior = mcp_types.DiscoverResult(
supported_versions=[MODERN_PROTOCOL_VERSIONS[0]],
capabilities=mcp_types.ServerCapabilities(),
instructions="prior instructions",
meta={
mcp_types.SERVER_INFO_META_KEY: prior_info.model_dump(
by_alias=True, mode="json"
),
"com.example/prior": True,
},
)
provider = ProxyProvider(
lambda: ProxyClient(
make_upstream(),
mode=MODERN_PROTOCOL_VERSIONS[0],
prior_discover=prior,
)
)
gateway = FastMCP(
"gateway",
providers=[provider],
middleware=[ProxyMetadataMiddleware(provider, identity="upstream")],
)
async with Client(gateway, mode="auto") as client:
result = client.session.discover_result
assert result is not None
assert client.instructions == "prior instructions"
assert client.server_info == prior_info
assert result.meta is not None
assert result.meta["com.example/prior"] is True
async def test_connected_pinned_client_probes_without_adopting_metadata():
version = MODERN_PROTOCOL_VERSIONS[0]
upstream = make_upstream()
async with Client(upstream, mode=version) as backend_client:
assert backend_client.instructions is None
proxy = create_proxy(backend_client, identity="upstream")
async with Client(proxy, mode="auto") as client:
result = client.session.discover_result
assert result is not None
assert client.instructions == "upstream instructions"
assert client.server_info == UPSTREAM_INFO
assert result.meta is not None
assert result.meta["com.example/upstream"] == {"enabled": True}
assert backend_client.instructions is None
async def test_invalid_upstream_discovery_metadata_is_ignored(
monkeypatch: pytest.MonkeyPatch,
):
version = MODERN_PROTOCOL_VERSIONS[0]
async def invalid_discover(_version: str) -> dict[str, Any]:
return {
"resultType": "complete",
"supportedVersions": [version],
"capabilities": [],
}
async with ProxyClient(make_upstream(), mode=version) as backend_client:
monkeypatch.setattr(backend_client.session, "send_discover", invalid_discover)
proxy = create_proxy(backend_client)
async with Client(proxy, mode="auto") as client:
assert client.server_info is not None
assert client.server_info.name == proxy.name
assert await client.list_tools() == []
async def test_invalid_backend_client_negotiation_is_not_ignored():
version = MODERN_PROTOCOL_VERSIONS[0]
prior = mcp_types.DiscoverResult(
supported_versions=["2099-01-01"],
capabilities=mcp_types.ServerCapabilities(),
)
provider = ProxyProvider(
lambda: ProxyClient(
make_upstream(),
mode=version,
prior_discover=prior,
)
)
gateway = FastMCP(
"gateway",
providers=[provider],
middleware=[ProxyMetadataMiddleware(provider)],
)
with pytest.raises(MCPError):
async with Client(gateway, mode="auto"):
pass
async def test_unrelated_client_validation_error_is_not_ignored():
class InvalidClient(ProxyClient):
async def __aenter__(self) -> ProxyClient:
mcp_types.Implementation.model_validate({})
return self
provider = ProxyProvider(lambda: InvalidClient(make_upstream()))
gateway = FastMCP(
"gateway",
providers=[provider],
middleware=[ProxyMetadataMiddleware(provider)],
)
with pytest.raises(MCPError):
async with Client(gateway, mode="auto"):
pass
@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:
raise AssertionError("metadata client must not be cloned")
version = MODERN_PROTOCOL_VERSIONS[0]
provider = ProxyProvider(
lambda: UnclonableProxyClient(make_upstream(), mode=version)
)
gateway = FastMCP(
"gateway",
providers=[provider],
middleware=[ProxyMetadataMiddleware(provider, identity="upstream")],
)
async with Client(gateway, mode="auto") as client:
assert client.instructions == "upstream instructions"
assert client.server_info == UPSTREAM_INFO
async def test_stateful_pinned_metadata_uses_registered_client_lifecycle():
created: list[StatefulProxyClient] = []
class TrackingStatefulProxyClient(StatefulProxyClient):
def new(self) -> StatefulProxyClient:
client = super().new()
created.append(client)
return client
version = MODERN_PROTOCOL_VERSIONS[0]
stateful_client = TrackingStatefulProxyClient(make_upstream(), mode=version)
proxy = FastMCPProxy(
name="stateful-proxy",
client_factory=stateful_client.new_stateful,
identity="upstream",
)
async with Client(proxy, mode="auto") as client:
assert client.instructions == "upstream instructions"
assert client.server_info == UPSTREAM_INFO
assert len(created) == 1
assert not created[0].is_connected()
@pytest.mark.parametrize("frontend_mode", ["legacy", "auto"])
@pytest.mark.parametrize("async_factory", [False, True])
@pytest.mark.parametrize("error_kind", ["runtime", "mcp"])
async def test_client_factory_errors_are_not_swallowed(
frontend_mode: str,
async_factory: bool,
error_kind: Literal["runtime", "mcp"],
):
def factory_error() -> Exception:
if error_kind == "mcp":
return MCPError(
code=mcp_types.INTERNAL_ERROR,
message="broken client factory",
)
return RuntimeError("broken client factory")
def broken_factory() -> Client:
raise factory_error()
async def broken_async_factory() -> Client:
raise factory_error()
factory = broken_async_factory if async_factory else broken_factory
provider = ProxyProvider(factory)
gateway = FastMCP(
"gateway",
providers=[provider],
middleware=[ProxyMetadataMiddleware(provider)],
)
with pytest.raises(MCPError):
async with Client(gateway, mode=frontend_mode):
pass
@pytest.mark.parametrize("frontend_mode", ["legacy", "auto"])
async def test_unavailable_backend_does_not_block_connection(frontend_mode: str):
port = find_available_port()
provider = ProxyProvider(
lambda: ProxyClient(
StreamableHttpTransport(f"http://127.0.0.1:{port}/mcp"), mode="auto"
),
cache_ttl=0,
)
gateway = FastMCP(
"available-gateway",
providers=[provider],
middleware=[ProxyMetadataMiddleware(provider)],
)
gateway.provider_error_strategy = "raise"
async with Client(gateway, mode=frontend_mode) as client:
assert client.server_info is not None
assert client.server_info.name == "available-gateway"
with pytest.raises(MCPError, match="Client failed to connect"):
await client.list_tools()
async def test_extension_owned_discovery_result_bypasses_metadata_forwarding():
factory_called = False
def broken_factory() -> Client:
nonlocal factory_called
factory_called = True
raise RuntimeError("metadata should not be read")
async def custom_discover(_ctx, _params):
return {
"resultType": "com.example/custom",
"payload": {"enabled": True},
}
provider = ProxyProvider(broken_factory)
gateway = FastMCP(
"extension-gateway",
middleware=[ProxyMetadataMiddleware(provider)],
)
gateway._mcp_server.add_request_handler(
"server/discover", mcp_types.RequestParams, custom_discover
)
version = MODERN_PROTOCOL_VERSIONS[0]
async with Client(gateway, mode=version) as client:
result = await client.session.send_discover(version)
assert isinstance(result, dict)
assert result["payload"] == {"enabled": True}
assert not factory_called
def test_gateway_construction_does_not_create_backend_client():
calls = 0
def client_factory() -> ProxyClient:
nonlocal calls
calls += 1
return ProxyClient(make_upstream())
provider = ProxyProvider(client_factory)
FastMCP(
"lazy-gateway",
providers=[provider],
middleware=[ProxyMetadataMiddleware(provider)],
)
assert calls == 0
async def test_proxy_initialize_middleware_preserves_legacy_behavior():
upstream = FastMCP("upstream", instructions="legacy instructions")
def client_factory() -> ProxyClient:
return ProxyClient(upstream)
proxy = FastMCPProxy(name="compatibility-proxy", client_factory=client_factory)
with pytest.warns(
FastMCPDeprecationWarning,
match="`ProxyInitializeMiddleware` is deprecated",
):
middleware = ProxyInitializeMiddleware(proxy)
proxy.middleware = [middleware]
async with Client(proxy, mode="legacy") as client:
assert client.instructions == "legacy instructions"
async with Client(proxy, mode="auto") as client:
assert client.instructions is None
assert middleware.proxy is proxy
async def test_fastmcp_proxy_uses_public_metadata_middleware():
proxy = create_proxy(make_upstream(), name="convenience", identity="upstream")
assert any(
isinstance(middleware, ProxyMetadataMiddleware)
for middleware in proxy.middleware
)
async with Client(proxy, mode="auto") as client:
assert client.instructions == "upstream instructions"
assert client.server_info == UPSTREAM_INFO