Pass client headers through to OpenAPI client

This commit is contained in:
Jeremiah Lowin 2025-05-23 12:38:24 -04:00
commit 213abc4244
6 changed files with 296 additions and 53 deletions

View file

@ -17,6 +17,7 @@ from pydantic.networks import AnyUrl
from fastmcp.exceptions import ToolError
from fastmcp.resources import Resource, ResourceTemplate
from fastmcp.server.dependencies import get_http_request
from fastmcp.server.server import FastMCP
from fastmcp.tools.tool import Tool, _convert_to_content
from fastmcp.utilities import openapi
@ -365,6 +366,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)
for p in self._route.parameters:
if (
p.location == "header"
@ -519,10 +534,24 @@ class OpenAPIResource(Resource):
if value is not None and value != "":
query_params[param.name] = value
# 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
response = await self._client.request(
method=self._route.method,
url=path,
params=query_params,
headers=headers,
timeout=self._timeout,
)
@ -733,6 +762,7 @@ class FastMCPOpenAPI(FastMCP):
) -> str:
"""Generate a default name from the route path."""
# First check for OpenAPI operationId which takes precedence
if route.operation_id:
return route.operation_id
@ -848,7 +878,7 @@ class FastMCPOpenAPI(FastMCP):
# Get a unique resource name
resource_name = self._get_unique_name(name, "resources")
resource_uri = f"resource://openapi/{resource_name}"
resource_uri = f"resource://{resource_name}"
base_description = (
route.description or route.summary or f"Represents {route.path}"
)
@ -896,7 +926,7 @@ class FastMCPOpenAPI(FastMCP):
path_params = [p.name for p in route.parameters if p.location == "path"]
path_params.sort() # Sort for consistent URIs
uri_template_str = f"resource://openapi/{template_name}"
uri_template_str = f"resource://{template_name}"
if path_params:
uri_template_str += "/" + "/".join(f"{{{p}}}" for p in path_params)

View file

@ -0,0 +1,157 @@
import json
import sys
from collections.abc import Generator
import pytest
import uvicorn
from fastapi import FastAPI, Request
from mcp.types import TextContent, TextResourceContents
from fastmcp import Client, FastMCP
from fastmcp.client.transports import SSETransport, StreamableHttpTransport
from fastmcp.utilities.tests import run_server_in_process
def fastmcp_server_for_headers() -> FastMCP:
app = FastAPI()
@app.get("/headers")
def get_headers(request: Request):
return request.headers
@app.get("/headers/{header_name}")
def get_header_by_name(header_name: str, request: Request):
return request.headers[header_name]
@app.post("/headers")
def post_headers(request: Request):
return request.headers
mcp = FastMCP.from_fastapi(
app, httpx_client_kwargs={"headers": {"X-SERVER": "test-abc"}}
)
return mcp
class TestClientHeaders:
def run_shttp_server(self, host: str, port: int) -> None:
try:
app = fastmcp_server_for_headers().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)
def run_sse_server(self, host: str, port: int) -> None:
try:
app = fastmcp_server_for_headers().http_app(transport="sse")
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(autouse=True, 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")
def sse_server(self) -> Generator[str, None, None]:
with run_server_in_process(self.run_sse_server) as url:
yield f"{url}/sse"
async def test_client_headers_sse_resource(self, sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
) as client:
result = await client.read_resource("resource://get_headers_headers_get")
assert isinstance(result[0], TextResourceContents)
headers = json.loads(result[0].text)
assert headers["x-test"] == "test-123"
async def test_client_headers_shttp_resource(self, shttp_server: str):
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-TEST": "test-123"}
)
) as client:
result = await client.read_resource("resource://get_headers_headers_get")
assert isinstance(result[0], TextResourceContents)
headers = json.loads(result[0].text)
assert headers["x-test"] == "test-123"
async def test_client_headers_sse_resource_template(self, sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
) as client:
result = await client.read_resource(
"resource://get_header_by_name_headers__header_name__get/x-test"
)
assert isinstance(result[0], TextResourceContents)
header = json.loads(result[0].text)
assert header == "test-123"
async def test_client_headers_shttp_resource_template(self, shttp_server: str):
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-TEST": "test-123"}
)
) as client:
result = await client.read_resource(
"resource://get_header_by_name_headers__header_name__get/x-test"
)
assert isinstance(result[0], TextResourceContents)
header = json.loads(result[0].text)
assert header == "test-123"
async def test_client_headers_sse_tool(self, sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
) as client:
result = await client.call_tool("post_headers_headers_post")
assert isinstance(result[0], TextContent)
headers = json.loads(result[0].text)
assert headers["x-test"] == "test-123"
async def test_client_headers_shttp_tool(self, shttp_server: str):
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-TEST": "test-123"}
)
) as client:
result = await client.call_tool("post_headers_headers_post")
assert isinstance(result[0], TextContent)
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 with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-SERVER": "test-client"}
)
) as client:
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"

View file

@ -109,5 +109,5 @@ class TestRouteTypeIgnoreDeprecation:
resource_uris = [str(r.uri) for r in resources.values()]
# Analytics should be excluded
assert "resource://openapi/get_items" in resource_uris
assert "resource://openapi/get_analytics" not in resource_uris
assert "resource://get_items" in resource_uris
assert "resource://get_analytics" not in resource_uris

View file

@ -7,7 +7,7 @@ import uvicorn
from mcp.types import TextContent, TextResourceContents
from fastmcp.client import Client
from fastmcp.client.transports import StreamableHttpTransport
from fastmcp.client.transports import SSETransport, StreamableHttpTransport
from fastmcp.server.dependencies import get_http_request
from fastmcp.server.server import FastMCP
from fastmcp.utilities.tests import run_server_in_process
@ -41,9 +41,28 @@ def fastmcp_server():
return server
def run_server(host: str, port: int) -> None:
def run_shttp_server(host: str, port: int) -> None:
try:
app = fastmcp_server().http_app()
app = fastmcp_server().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)
def run_sse_server(host: str, port: int) -> None:
try:
app = fastmcp_server().http_app(transport="sse")
server = uvicorn.Server(
config=uvicorn.Config(
app=app,
@ -61,15 +80,23 @@ def run_server(host: str, port: int) -> None:
@pytest.fixture(autouse=True, scope="module")
def sse_server() -> Generator[str, None, None]:
with run_server_in_process(run_server) as url:
def shttp_server() -> Generator[str, None, None]:
with run_server_in_process(run_shttp_server) as url:
yield f"{url}/mcp"
async def test_http_headers_resource(sse_server: str):
@pytest.fixture(autouse=True, scope="module")
def sse_server() -> Generator[str, None, None]:
with run_server_in_process(run_sse_server) as url:
yield f"{url}/sse"
async def test_http_headers_resource_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
raw_result = await client.read_resource("request://headers")
assert isinstance(raw_result[0], TextResourceContents)
@ -78,10 +105,24 @@ async def test_http_headers_resource(sse_server: str):
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_tool(sse_server: str):
async def test_http_headers_resource_sse(sse_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
raw_result = await client.read_resource("request://headers")
assert isinstance(raw_result[0], TextResourceContents)
json_result = json.loads(raw_result[0].text)
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_tool_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
result = await client.call_tool("get_headers_tool")
assert isinstance(result[0], TextContent)
@ -90,10 +131,35 @@ async def test_http_headers_tool(sse_server: str):
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_prompt(sse_server: str):
async def test_http_headers_tool_sse(sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
result = await client.call_tool("get_headers_tool")
assert isinstance(result[0], TextContent)
json_result = json.loads(result[0].text)
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_prompt_shttp(shttp_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=StreamableHttpTransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
transport=StreamableHttpTransport(
shttp_server, headers={"X-DEMO-HEADER": "ABC"}
)
) as client:
result = await client.get_prompt("get_headers_prompt")
assert isinstance(result.messages[0].content, TextContent)
json_result = json.loads(result.messages[0].content.text)
assert "x-demo-header" in json_result
assert json_result["x-demo-header"] == "ABC"
async def test_http_headers_prompt_sse(sse_server: str):
"""Test getting HTTP headers from the server."""
async with Client(
transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
) as client:
result = await client.get_prompt("get_headers_prompt")
assert isinstance(result.messages[0].content, TextContent)

View file

@ -249,7 +249,7 @@ class TestTools:
# Check that the user was created via MCP
async with Client(fastmcp_openapi_server) as client:
user_response = await client.read_resource(
"resource://openapi/get_user_users__user_id__get/4"
"resource://get_user_users__user_id__get/4"
)
assert isinstance(user_response[0], TextResourceContents)
response_text = user_response[0].text
@ -283,7 +283,7 @@ class TestTools:
# Check that the user was updated via MCP
async with Client(fastmcp_openapi_server) as client:
user_response = await client.read_resource(
"resource://openapi/get_user_users__user_id__get/1"
"resource://get_user_users__user_id__get/1"
)
assert isinstance(user_response[0], TextResourceContents)
response_text = user_response[0].text
@ -325,7 +325,7 @@ class TestResources:
async with Client(fastmcp_openapi_server) as client:
resources = await client.list_resources()
assert len(resources) == 4
assert resources[0].uri == AnyUrl("resource://openapi/get_users_users_get")
assert resources[0].uri == AnyUrl("resource://get_users_users_get")
assert resources[0].name == "get_users_users_get"
async def test_get_resource(
@ -343,7 +343,7 @@ class TestResources:
)
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
"resource://openapi/get_users_users_get"
"resource://get_users_users_get"
)
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
@ -360,7 +360,7 @@ class TestResources:
"""Test reading a resource that returns bytes."""
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
"resource://openapi/ping_bytes_ping_bytes_get"
"resource://ping_bytes_ping_bytes_get"
)
assert isinstance(resource_response[0], BlobResourceContents)
assert base64.b64decode(resource_response[0].blob) == b"pong"
@ -372,9 +372,7 @@ class TestResources:
):
"""Test reading a resource that returns a string."""
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
"resource://openapi/ping_ping_get"
)
resource_response = await client.read_resource("resource://ping_ping_get")
assert isinstance(resource_response[0], TextResourceContents)
assert resource_response[0].text == "pong"
@ -392,7 +390,7 @@ class TestResourceTemplates:
assert resource_templates[0].name == "get_user_users__user_id__get"
assert (
resource_templates[0].uriTemplate
== r"resource://openapi/get_user_users__user_id__get/{user_id}"
== r"resource://get_user_users__user_id__get/{user_id}"
)
assert (
resource_templates[1].name
@ -400,7 +398,7 @@ class TestResourceTemplates:
)
assert (
resource_templates[1].uriTemplate
== r"resource://openapi/get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
== r"resource://get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
)
async def test_get_resource_template(
@ -415,7 +413,7 @@ class TestResourceTemplates:
user_id = 2
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
f"resource://openapi/get_user_users__user_id__get/{user_id}"
f"resource://get_user_users__user_id__get/{user_id}"
)
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
@ -438,7 +436,7 @@ class TestResourceTemplates:
is_active = True
async with Client(fastmcp_openapi_server) as client:
resource_response = await client.read_resource(
f"resource://openapi/get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
f"resource://get_user_active_state_users__user_id___is_active__get/{is_active}/{user_id}"
)
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
@ -555,7 +553,7 @@ class TestTagTransfer:
# Manually create a resource from template
params = {"user_id": 1}
resource = await get_user_template.create_resource(
"resource://openapi/get_user_users__user_id__get/1", params
"resource://get_user_users__user_id__get/1", params
)
# Verify tags are preserved from template to resource
@ -672,7 +670,7 @@ class TestOpenAPI30Compatibility:
async with Client(openapi_30_server) as client:
resources = await client.list_resources()
assert len(resources) == 1
assert resources[0].uri == AnyUrl("resource://openapi/listProducts")
assert resources[0].uri == AnyUrl("resource://listProducts")
async def test_resource_template_discovery(self, openapi_30_server):
"""Test that resource templates are correctly discovered from an OpenAPI 3.0 spec."""
@ -680,7 +678,7 @@ class TestOpenAPI30Compatibility:
templates = await client.list_resource_templates()
assert len(templates) == 1
assert templates[0].name == "getProduct"
assert templates[0].uriTemplate == r"resource://openapi/getProduct/{product_id}"
assert templates[0].uriTemplate == r"resource://getProduct/{product_id}"
async def test_tool_discovery(self, openapi_30_server):
"""Test that tools are correctly discovered from an OpenAPI 3.0 spec."""
@ -694,9 +692,7 @@ class TestOpenAPI30Compatibility:
async def test_resource_access(self, openapi_30_server):
"""Test reading a resource from an OpenAPI 3.0 server."""
async with Client(openapi_30_server) as client:
resource_response = await client.read_resource(
"resource://openapi/listProducts"
)
resource_response = await client.read_resource("resource://listProducts")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)
@ -707,9 +703,7 @@ class TestOpenAPI30Compatibility:
async def test_resource_template_access(self, openapi_30_server):
"""Test reading a resource from template from an OpenAPI 3.0 server."""
async with Client(openapi_30_server) as client:
resource_response = await client.read_resource(
"resource://openapi/getProduct/p1"
)
resource_response = await client.read_resource("resource://getProduct/p1")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)
@ -852,7 +846,7 @@ class TestOpenAPI31Compatibility:
async with Client(openapi_31_server) as client:
resources = await client.list_resources()
assert len(resources) == 1
assert resources[0].uri == AnyUrl("resource://openapi/listOrders")
assert resources[0].uri == AnyUrl("resource://listOrders")
async def test_resource_template_discovery(self, openapi_31_server):
"""Test that resource templates are correctly discovered from an OpenAPI 3.1 spec."""
@ -860,7 +854,7 @@ class TestOpenAPI31Compatibility:
templates = await client.list_resource_templates()
assert len(templates) == 1
assert templates[0].name == "getOrder"
assert templates[0].uriTemplate == r"resource://openapi/getOrder/{order_id}"
assert templates[0].uriTemplate == r"resource://getOrder/{order_id}"
async def test_tool_discovery(self, openapi_31_server):
"""Test that tools are correctly discovered from an OpenAPI 3.1 spec."""
@ -874,9 +868,7 @@ class TestOpenAPI31Compatibility:
async def test_resource_access(self, openapi_31_server):
"""Test reading a resource from an OpenAPI 3.1 server."""
async with Client(openapi_31_server) as client:
resource_response = await client.read_resource(
"resource://openapi/listOrders"
)
resource_response = await client.read_resource("resource://listOrders")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)
@ -887,9 +879,7 @@ class TestOpenAPI31Compatibility:
async def test_resource_template_access(self, openapi_31_server):
"""Test reading a resource from template from an OpenAPI 3.1 server."""
async with Client(openapi_31_server) as client:
resource_response = await client.read_resource(
"resource://openapi/getOrder/o1"
)
resource_response = await client.read_resource("resource://getOrder/o1")
assert isinstance(resource_response[0], TextResourceContents)
response_text = resource_response[0].text
content = json.loads(response_text)

View file

@ -9,7 +9,7 @@ from fastmcp.utilities.openapi import parse_openapi_to_http_routes
@pytest.fixture
def fastapi_server() -> FastAPI:
def fastapi_app() -> FastAPI:
"""Fixture that returns a FastAPI app for live OpenAPI schema testing."""
from enum import Enum
@ -228,9 +228,9 @@ def fastapi_server() -> FastAPI:
@pytest.fixture
def fastapi_openapi_schema(fastapi_server) -> dict[str, Any]:
def fastapi_openapi_schema(fastapi_app) -> dict[str, Any]:
"""Fixture that returns the OpenAPI schema from a live FastAPI server."""
return fastapi_server.openapi()
return fastapi_app.openapi()
@pytest.fixture
@ -472,11 +472,11 @@ def test_tag_consistency_across_related_endpoints(route_map):
)
def test_tag_order_preservation(fastapi_server):
def test_tag_order_preservation(fastapi_app):
"""Test that tag order is preserved in the parsed routes."""
# Add a new endpoint with specifically ordered tags
@fastapi_server.get(
@fastapi_app.get(
"/test-tag-order",
tags=["first", "second", "third"],
operation_id="test_tag_order",
@ -485,7 +485,7 @@ def test_tag_order_preservation(fastapi_server):
return {"result": "testing tag order"}
# Get the updated schema and parse routes
routes = parse_openapi_to_http_routes(fastapi_server.openapi())
routes = parse_openapi_to_http_routes(fastapi_app.openapi())
# Find our test route
test_route = next((r for r in routes if r.path == "/test-tag-order"), None)
@ -497,11 +497,11 @@ def test_tag_order_preservation(fastapi_server):
)
def test_duplicate_tags_handling(fastapi_server):
def test_duplicate_tags_handling(fastapi_app):
"""Test handling of duplicate tags in the OpenAPI schema."""
# Add an endpoint with duplicate tags
@fastapi_server.get(
@fastapi_app.get(
"/test-duplicate-tags",
tags=["duplicate", "items", "duplicate"],
operation_id="test_duplicate_tags",
@ -510,7 +510,7 @@ def test_duplicate_tags_handling(fastapi_server):
return {"result": "testing duplicate tags"}
# Get the updated schema and parse routes
routes = parse_openapi_to_http_routes(fastapi_server.openapi())
routes = parse_openapi_to_http_routes(fastapi_app.openapi())
# Find our test route
test_route = next((r for r in routes if r.path == "/test-duplicate-tags"), None)