Merge pull request #576 from jlowin/headers

Use lowercase name for headers when comparing
This commit is contained in:
Jeremiah Lowin 2025-05-23 16:16:46 -04:00 committed by GitHub
commit f0fec7f857
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 41 additions and 34 deletions

View file

@ -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:

View file

@ -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

View file

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