Ensure headers are passed through proxy servers

This commit is contained in:
Jeremiah Lowin 2025-05-23 15:39:45 -04:00
commit b8e918a862
2 changed files with 74 additions and 10 deletions

View file

@ -25,6 +25,7 @@ from pydantic import AnyUrl
from typing_extensions import Unpack
from fastmcp.server import FastMCP as FastMCPServer
from fastmcp.server.dependencies import get_http_request
from fastmcp.server.server import FastMCP
from fastmcp.utilities.logging import get_logger
from fastmcp.utilities.mcp_config import MCPConfig, infer_transport_type_from_url
@ -34,6 +35,11 @@ if TYPE_CHECKING:
logger = get_logger(__name__)
EXCLUDE_HEADERS = {
"content-type",
"content-length",
}
class SessionKwargs(TypedDict, total=False):
"""Keyword arguments for the MCP ClientSession constructor."""
@ -132,7 +138,21 @@ class SSETransport(ClientTransport):
async def connect_session(
self, **session_kwargs: Unpack[SessionKwargs]
) -> AsyncIterator[ClientSession]:
client_kwargs = {}
client_kwargs: dict[str, Any] = {
"headers": self.headers,
}
# load headers from an active HTTP request, if available. This will only be true
# if the client is used in a FastMCP Proxy, in which case the MCP client headers
# need to be forwarded to the remote server.
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:
client_kwargs["headers"][name] = str(value)
except RuntimeError:
client_kwargs["headers"] = self.headers
# sse_read_timeout has a default value set, so we can't pass None without overriding it
# instead we simply leave the kwarg out if it's not provided
if self.sse_read_timeout is not None:
@ -143,9 +163,7 @@ class SSETransport(ClientTransport):
)
client_kwargs["timeout"] = read_timeout_seconds.total_seconds()
async with sse_client(
self.url, headers=self.headers, **client_kwargs
) as transport:
async with sse_client(self.url, **client_kwargs) as transport:
read_stream, write_stream = transport
async with ClientSession(
read_stream, write_stream, **session_kwargs
@ -180,7 +198,23 @@ class StreamableHttpTransport(ClientTransport):
async def connect_session(
self, **session_kwargs: Unpack[SessionKwargs]
) -> AsyncIterator[ClientSession]:
client_kwargs = {}
client_kwargs: dict[str, Any] = {
"headers": self.headers,
}
# load headers from an active HTTP request, if available. This will only be true
# if the client is used in a FastMCP Proxy, in which case the MCP client headers
# need to be forwarded to the remote server.
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:
client_kwargs["headers"][name] = str(value)
except RuntimeError:
client_kwargs["headers"] = self.headers
print(client_kwargs)
# sse_read_timeout has a default value set, so we can't pass None without overriding it
# instead we simply leave the kwarg out if it's not provided
if self.sse_read_timeout is not None:
@ -188,9 +222,7 @@ class StreamableHttpTransport(ClientTransport):
if session_kwargs.get("read_timeout_seconds", None) is not None:
client_kwargs["timeout"] = session_kwargs.get("read_timeout_seconds")
async with streamablehttp_client(
self.url, headers=self.headers, **client_kwargs
) as transport:
async with streamablehttp_client(self.url, **client_kwargs) as transport:
read_stream, write_stream, _ = transport
async with ClientSession(
read_stream, write_stream, **session_kwargs

View file

@ -71,16 +71,40 @@ class TestClientHeaders:
sys.exit(1)
sys.exit(0)
@pytest.fixture(autouse=True, scope="class")
def run_proxy_server(self, host: str, port: int, remote_url: str) -> None:
try:
client = Client(transport=StreamableHttpTransport(remote_url))
app = FastMCP.as_proxy(client).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)
@pytest.fixture(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")
@pytest.fixture(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"
@pytest.fixture(scope="class")
def proxy_server(self, shttp_server: str) -> Generator[str, None, None]:
with run_server_in_process(self.run_proxy_server, shttp_server + "/mcp") as url:
yield f"{url}/mcp"
async def test_client_headers_sse_resource(self, sse_server: str):
async with Client(
transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
@ -155,3 +179,11 @@ class TestClientHeaders:
assert isinstance(result[0], TextResourceContents)
headers = json.loads(result[0].text)
assert headers["x-server"] == "test-abc"
async def test_client_headers_proxy(self, proxy_server: str):
async with Client(transport=StreamableHttpTransport(proxy_server)) as client:
await client.ping()
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"