diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index ff2c5ded3..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", } @@ -148,7 +148,10 @@ class SSETransport(ClientTransport): 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: + name = name.lower() + 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 @@ -208,7 +211,10 @@ class StreamableHttpTransport(ClientTransport): 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: + name = name.lower() + 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: 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): """