mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-26 23:44:17 +02:00
Ensure headers are passed through proxy servers
This commit is contained in:
parent
213abc4244
commit
b8e918a862
2 changed files with 74 additions and 10 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue