From b8e918a86248d56d091d8b8025673f9402d577d2 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 23 May 2025 15:39:45 -0400 Subject: [PATCH] Ensure headers are passed through proxy servers --- src/fastmcp/client/transports.py | 48 ++++++++++++++++++++++++++------ tests/client/test_openapi.py | 36 ++++++++++++++++++++++-- 2 files changed, 74 insertions(+), 10 deletions(-) diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index f82b238d0..ff2c5ded3 100644 --- a/src/fastmcp/client/transports.py +++ b/src/fastmcp/client/transports.py @@ -25,6 +25,7 @@ from pydantic import AnyUrl from typing_extensions import Unpack from fastmcp.server import FastMCP as FastMCPServer +from fastmcp.server.dependencies import get_http_request from fastmcp.server.server import FastMCP from fastmcp.utilities.logging import get_logger from fastmcp.utilities.mcp_config import MCPConfig, infer_transport_type_from_url @@ -34,6 +35,11 @@ if TYPE_CHECKING: logger = get_logger(__name__) +EXCLUDE_HEADERS = { + "content-type", + "content-length", +} + class SessionKwargs(TypedDict, total=False): """Keyword arguments for the MCP ClientSession constructor.""" @@ -132,7 +138,21 @@ class SSETransport(ClientTransport): async def connect_session( self, **session_kwargs: Unpack[SessionKwargs] ) -> AsyncIterator[ClientSession]: - client_kwargs = {} + client_kwargs: dict[str, Any] = { + "headers": self.headers, + } + + # load headers from an active HTTP request, if available. This will only be true + # if the client is used in a FastMCP Proxy, in which case the MCP client headers + # need to be forwarded to the remote server. + try: + active_request = get_http_request() + for name, value in active_request.headers.items(): + if name not in self.headers and name not in EXCLUDE_HEADERS: + client_kwargs["headers"][name] = str(value) + except RuntimeError: + client_kwargs["headers"] = self.headers + # sse_read_timeout has a default value set, so we can't pass None without overriding it # instead we simply leave the kwarg out if it's not provided if self.sse_read_timeout is not None: @@ -143,9 +163,7 @@ class SSETransport(ClientTransport): ) client_kwargs["timeout"] = read_timeout_seconds.total_seconds() - async with sse_client( - self.url, headers=self.headers, **client_kwargs - ) as transport: + async with sse_client(self.url, **client_kwargs) as transport: read_stream, write_stream = transport async with ClientSession( read_stream, write_stream, **session_kwargs @@ -180,7 +198,23 @@ class StreamableHttpTransport(ClientTransport): async def connect_session( self, **session_kwargs: Unpack[SessionKwargs] ) -> AsyncIterator[ClientSession]: - client_kwargs = {} + client_kwargs: dict[str, Any] = { + "headers": self.headers, + } + + # load headers from an active HTTP request, if available. This will only be true + # if the client is used in a FastMCP Proxy, in which case the MCP client headers + # need to be forwarded to the remote server. + try: + active_request = get_http_request() + for name, value in active_request.headers.items(): + if name not in self.headers and name not in EXCLUDE_HEADERS: + client_kwargs["headers"][name] = str(value) + + except RuntimeError: + client_kwargs["headers"] = self.headers + print(client_kwargs) + # sse_read_timeout has a default value set, so we can't pass None without overriding it # instead we simply leave the kwarg out if it's not provided if self.sse_read_timeout is not None: @@ -188,9 +222,7 @@ class StreamableHttpTransport(ClientTransport): if session_kwargs.get("read_timeout_seconds", None) is not None: client_kwargs["timeout"] = session_kwargs.get("read_timeout_seconds") - async with streamablehttp_client( - self.url, headers=self.headers, **client_kwargs - ) as transport: + async with streamablehttp_client(self.url, **client_kwargs) as transport: read_stream, write_stream, _ = transport async with ClientSession( read_stream, write_stream, **session_kwargs diff --git a/tests/client/test_openapi.py b/tests/client/test_openapi.py index b2639e614..5ee33ce7e 100644 --- a/tests/client/test_openapi.py +++ b/tests/client/test_openapi.py @@ -71,16 +71,40 @@ class TestClientHeaders: sys.exit(1) sys.exit(0) - @pytest.fixture(autouse=True, scope="class") + def run_proxy_server(self, host: str, port: int, remote_url: str) -> None: + try: + client = Client(transport=StreamableHttpTransport(remote_url)) + app = FastMCP.as_proxy(client).http_app(transport="streamable-http") + server = uvicorn.Server( + config=uvicorn.Config( + app=app, + host=host, + port=port, + log_level="error", + lifespan="on", + ) + ) + server.run() + except Exception as e: + print(f"Server error: {e}") + sys.exit(1) + sys.exit(0) + + @pytest.fixture(scope="class") def shttp_server(self) -> Generator[str, None, None]: with run_server_in_process(self.run_shttp_server) as url: yield f"{url}/mcp" - @pytest.fixture(autouse=True, scope="class") + @pytest.fixture(scope="class") def sse_server(self) -> Generator[str, None, None]: with run_server_in_process(self.run_sse_server) as url: yield f"{url}/sse" + @pytest.fixture(scope="class") + def proxy_server(self, shttp_server: str) -> Generator[str, None, None]: + with run_server_in_process(self.run_proxy_server, shttp_server + "/mcp") as url: + yield f"{url}/mcp" + async def test_client_headers_sse_resource(self, sse_server: str): async with Client( transport=SSETransport(sse_server, headers={"X-TEST": "test-123"}) @@ -155,3 +179,11 @@ class TestClientHeaders: assert isinstance(result[0], TextResourceContents) headers = json.loads(result[0].text) assert headers["x-server"] == "test-abc" + + async def test_client_headers_proxy(self, proxy_server: str): + async with Client(transport=StreamableHttpTransport(proxy_server)) as client: + await client.ping() + result = await client.read_resource("resource://get_headers_headers_get") + assert isinstance(result[0], TextResourceContents) + headers = json.loads(result[0].text) + assert headers["x-server"] == "test-abc"