From e0ae84c75f9a8841ffa7db669f2cab7f3a74bc2f Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 23 May 2025 16:02:42 -0400 Subject: [PATCH 1/2] Use lowercase name for headers --- src/fastmcp/client/transports.py | 2 ++ src/fastmcp/server/openapi.py | 59 ++++++++++++++++---------------- tests/client/test_openapi.py | 4 +-- 3 files changed, 34 insertions(+), 31 deletions(-) diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index ff2c5ded3..5720d7fd6 100644 --- a/src/fastmcp/client/transports.py +++ b/src/fastmcp/client/transports.py @@ -148,6 +148,7 @@ class SSETransport(ClientTransport): try: active_request = get_http_request() for name, value in active_request.headers.items(): + name = name.lower() if name not in self.headers and name not in EXCLUDE_HEADERS: client_kwargs["headers"][name] = str(value) except RuntimeError: @@ -208,6 +209,7 @@ class StreamableHttpTransport(ClientTransport): try: active_request = get_http_request() for name, value in active_request.headers.items(): + name = name.lower() if name not in self.headers and name not in EXCLUDE_HEADERS: client_kwargs["headers"][name] = str(value) diff --git a/src/fastmcp/server/openapi.py b/src/fastmcp/server/openapi.py index 33e6440e6..b964f34d1 100644 --- a/src/fastmcp/server/openapi.py +++ b/src/fastmcp/server/openapi.py @@ -35,6 +35,26 @@ logger = get_logger(__name__) HttpMethod = Literal["GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS", "HEAD"] + +def _get_mcp_client_headers() -> dict[str, str]: + """ + Extract headers from the current MCP client HTTP request if available. + + These headers will take precedence over OpenAPI-defined headers when both are present. + + Returns: + Dictionary of header name-value pairs (lowercased names), or empty dict if no HTTP request is active. + """ + try: + http_request = get_http_request() + return { + name.lower(): str(value) for name, value in http_request.headers.items() + } + except RuntimeError: + # No active HTTP request (e.g., STDIO transport), return empty dict + return {} + + # Type definitions for the mapping functions RouteMapFn = Callable[[HTTPRoute, "MCPType"], "MCPType | None"] ComponentFn = Callable[ @@ -367,26 +387,20 @@ class OpenAPITool(Tool): # Prepare headers - fix typing by ensuring all values are strings headers = {} - # Try to get headers from the current MCP client HTTP request - try: - http_request = get_http_request() - # Add headers from the MCP client request - for name, value in http_request.headers.items(): - # Don't override headers that are already set on the client - if name not in self._client.headers: - headers[name] = str(value) - except RuntimeError: - # No active HTTP request (e.g., STDIO transport), continue without client headers - pass - - # Add any OpenAPI-defined header parameters (these take precedence over client headers) + # Start with OpenAPI-defined header parameters + openapi_headers = {} for p in self._route.parameters: if ( p.location == "header" and p.name in kwargs and kwargs[p.name] is not None ): - headers[p.name] = str(kwargs[p.name]) + openapi_headers[p.name.lower()] = str(kwargs[p.name]) + headers.update(openapi_headers) + + # Add headers from the current MCP client HTTP request (these take precedence) + mcp_headers = _get_mcp_client_headers() + headers.update(mcp_headers) # Prepare request body json_data = None @@ -536,16 +550,8 @@ class OpenAPIResource(Resource): # Prepare headers from MCP client request if available headers = {} - try: - http_request = get_http_request() - # Add headers from the MCP client request - for name, value in http_request.headers.items(): - # Don't override headers that are already set on the client - if name not in self._client.headers: - headers[name] = str(value) - except RuntimeError: - # No active HTTP request (e.g., STDIO transport), continue without client headers - pass + mcp_headers = _get_mcp_client_headers() + headers.update(mcp_headers) response = await self._client.request( method=self._route.method, @@ -991,8 +997,3 @@ class FastMCPOpenAPI(FastMCP): logger.debug( f"Registered TEMPLATE: {uri_template_str} ({route.method} {route.path}) with tags: {route.tags}" ) - - async def _mcp_call_tool(self, name: str, arguments: dict[str, Any]) -> Any: - """Override the call_tool method to return the raw result without converting to content.""" - result = await self._tool_manager.call_tool(name, arguments) - return result diff --git a/tests/client/test_openapi.py b/tests/client/test_openapi.py index 903582f59..2344a47f1 100644 --- a/tests/client/test_openapi.py +++ b/tests/client/test_openapi.py @@ -169,7 +169,7 @@ class TestClientHeaders: headers = json.loads(result[0].text) assert headers["x-test"] == "test-123" - async def test_client_doesnt_override_server_headers(self, shttp_server: str): + async def test_client_overrides_server_headers(self, shttp_server: str): async with Client( transport=StreamableHttpTransport( shttp_server, headers={"X-SERVER": "test-client"} @@ -178,7 +178,7 @@ class TestClientHeaders: 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" + assert headers["x-server"] == "test-client" async def test_client_headers_proxy(self, proxy_server: str): """ From 9c3e90f01880f332c5ecc1b234eba36dc760b2b7 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 23 May 2025 16:11:56 -0400 Subject: [PATCH 2/2] Update transports.py --- src/fastmcp/client/transports.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index 5720d7fd6..398c080fc 100644 --- a/src/fastmcp/client/transports.py +++ b/src/fastmcp/client/transports.py @@ -35,8 +35,8 @@ if TYPE_CHECKING: logger = get_logger(__name__) +# these headers, when forwarded to the remote server, can cause issues EXCLUDE_HEADERS = { - "content-type", "content-length", } @@ -149,7 +149,9 @@ class SSETransport(ClientTransport): active_request = get_http_request() for name, value in active_request.headers.items(): name = name.lower() - if name not in self.headers and name not in EXCLUDE_HEADERS: + if name not in self.headers and name not in { + h.lower() for h in EXCLUDE_HEADERS + }: client_kwargs["headers"][name] = str(value) except RuntimeError: client_kwargs["headers"] = self.headers @@ -210,7 +212,9 @@ class StreamableHttpTransport(ClientTransport): active_request = get_http_request() for name, value in active_request.headers.items(): name = name.lower() - if name not in self.headers and name not in EXCLUDE_HEADERS: + if name not in self.headers and name not in { + h.lower() for h in EXCLUDE_HEADERS + }: client_kwargs["headers"][name] = str(value) except RuntimeError: