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): """