mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 07:09:11 +02:00
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:
parent
75fb116e36
commit
9feb1f378b
11 changed files with 1113 additions and 134 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
108
tests/server/middleware/test_discovery_middleware.py
Normal file
108
tests/server/middleware/test_discovery_middleware.py
Normal 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}
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
637
tests/server/providers/proxy/test_server_metadata.py
Normal file
637
tests/server/providers/proxy/test_server_metadata.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue