Jeremiah Lowin commited on
Commit
80285d4
·
1 Parent(s): ac33644

Ensure headers are passed through proxy servers

Browse files
src/fastmcp/client/transports.py CHANGED
@@ -25,6 +25,7 @@ from pydantic import AnyUrl
25
  from typing_extensions import Unpack
26
 
27
  from fastmcp.server import FastMCP as FastMCPServer
 
28
  from fastmcp.server.server import FastMCP
29
  from fastmcp.utilities.logging import get_logger
30
  from fastmcp.utilities.mcp_config import MCPConfig, infer_transport_type_from_url
@@ -34,6 +35,11 @@ if TYPE_CHECKING:
34
 
35
  logger = get_logger(__name__)
36
 
 
 
 
 
 
37
 
38
  class SessionKwargs(TypedDict, total=False):
39
  """Keyword arguments for the MCP ClientSession constructor."""
@@ -132,7 +138,21 @@ class SSETransport(ClientTransport):
132
  async def connect_session(
133
  self, **session_kwargs: Unpack[SessionKwargs]
134
  ) -> AsyncIterator[ClientSession]:
135
- client_kwargs = {}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
136
  # sse_read_timeout has a default value set, so we can't pass None without overriding it
137
  # instead we simply leave the kwarg out if it's not provided
138
  if self.sse_read_timeout is not None:
@@ -143,9 +163,7 @@ class SSETransport(ClientTransport):
143
  )
144
  client_kwargs["timeout"] = read_timeout_seconds.total_seconds()
145
 
146
- async with sse_client(
147
- self.url, headers=self.headers, **client_kwargs
148
- ) as transport:
149
  read_stream, write_stream = transport
150
  async with ClientSession(
151
  read_stream, write_stream, **session_kwargs
@@ -180,7 +198,23 @@ class StreamableHttpTransport(ClientTransport):
180
  async def connect_session(
181
  self, **session_kwargs: Unpack[SessionKwargs]
182
  ) -> AsyncIterator[ClientSession]:
183
- client_kwargs = {}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
184
  # sse_read_timeout has a default value set, so we can't pass None without overriding it
185
  # instead we simply leave the kwarg out if it's not provided
186
  if self.sse_read_timeout is not None:
@@ -188,9 +222,7 @@ class StreamableHttpTransport(ClientTransport):
188
  if session_kwargs.get("read_timeout_seconds", None) is not None:
189
  client_kwargs["timeout"] = session_kwargs.get("read_timeout_seconds")
190
 
191
- async with streamablehttp_client(
192
- self.url, headers=self.headers, **client_kwargs
193
- ) as transport:
194
  read_stream, write_stream, _ = transport
195
  async with ClientSession(
196
  read_stream, write_stream, **session_kwargs
 
25
  from typing_extensions import Unpack
26
 
27
  from fastmcp.server import FastMCP as FastMCPServer
28
+ from fastmcp.server.dependencies import get_http_request
29
  from fastmcp.server.server import FastMCP
30
  from fastmcp.utilities.logging import get_logger
31
  from fastmcp.utilities.mcp_config import MCPConfig, infer_transport_type_from_url
 
35
 
36
  logger = get_logger(__name__)
37
 
38
+ EXCLUDE_HEADERS = {
39
+ "content-type",
40
+ "content-length",
41
+ }
42
+
43
 
44
  class SessionKwargs(TypedDict, total=False):
45
  """Keyword arguments for the MCP ClientSession constructor."""
 
138
  async def connect_session(
139
  self, **session_kwargs: Unpack[SessionKwargs]
140
  ) -> AsyncIterator[ClientSession]:
141
+ client_kwargs: dict[str, Any] = {
142
+ "headers": self.headers,
143
+ }
144
+
145
+ # load headers from an active HTTP request, if available. This will only be true
146
+ # if the client is used in a FastMCP Proxy, in which case the MCP client headers
147
+ # need to be forwarded to the remote server.
148
+ try:
149
+ active_request = get_http_request()
150
+ for name, value in active_request.headers.items():
151
+ if name not in self.headers and name not in EXCLUDE_HEADERS:
152
+ client_kwargs["headers"][name] = str(value)
153
+ except RuntimeError:
154
+ client_kwargs["headers"] = self.headers
155
+
156
  # sse_read_timeout has a default value set, so we can't pass None without overriding it
157
  # instead we simply leave the kwarg out if it's not provided
158
  if self.sse_read_timeout is not None:
 
163
  )
164
  client_kwargs["timeout"] = read_timeout_seconds.total_seconds()
165
 
166
+ async with sse_client(self.url, **client_kwargs) as transport:
 
 
167
  read_stream, write_stream = transport
168
  async with ClientSession(
169
  read_stream, write_stream, **session_kwargs
 
198
  async def connect_session(
199
  self, **session_kwargs: Unpack[SessionKwargs]
200
  ) -> AsyncIterator[ClientSession]:
201
+ client_kwargs: dict[str, Any] = {
202
+ "headers": self.headers,
203
+ }
204
+
205
+ # load headers from an active HTTP request, if available. This will only be true
206
+ # if the client is used in a FastMCP Proxy, in which case the MCP client headers
207
+ # need to be forwarded to the remote server.
208
+ try:
209
+ active_request = get_http_request()
210
+ for name, value in active_request.headers.items():
211
+ if name not in self.headers and name not in EXCLUDE_HEADERS:
212
+ client_kwargs["headers"][name] = str(value)
213
+
214
+ except RuntimeError:
215
+ client_kwargs["headers"] = self.headers
216
+ print(client_kwargs)
217
+
218
  # sse_read_timeout has a default value set, so we can't pass None without overriding it
219
  # instead we simply leave the kwarg out if it's not provided
220
  if self.sse_read_timeout is not None:
 
222
  if session_kwargs.get("read_timeout_seconds", None) is not None:
223
  client_kwargs["timeout"] = session_kwargs.get("read_timeout_seconds")
224
 
225
+ async with streamablehttp_client(self.url, **client_kwargs) as transport:
 
 
226
  read_stream, write_stream, _ = transport
227
  async with ClientSession(
228
  read_stream, write_stream, **session_kwargs
tests/client/test_openapi.py CHANGED
@@ -71,16 +71,40 @@ class TestClientHeaders:
71
  sys.exit(1)
72
  sys.exit(0)
73
 
74
- @pytest.fixture(autouse=True, scope="class")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75
  def shttp_server(self) -> Generator[str, None, None]:
76
  with run_server_in_process(self.run_shttp_server) as url:
77
  yield f"{url}/mcp"
78
 
79
- @pytest.fixture(autouse=True, scope="class")
80
  def sse_server(self) -> Generator[str, None, None]:
81
  with run_server_in_process(self.run_sse_server) as url:
82
  yield f"{url}/sse"
83
 
 
 
 
 
 
84
  async def test_client_headers_sse_resource(self, sse_server: str):
85
  async with Client(
86
  transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
@@ -155,3 +179,11 @@ class TestClientHeaders:
155
  assert isinstance(result[0], TextResourceContents)
156
  headers = json.loads(result[0].text)
157
  assert headers["x-server"] == "test-abc"
 
 
 
 
 
 
 
 
 
71
  sys.exit(1)
72
  sys.exit(0)
73
 
74
+ def run_proxy_server(self, host: str, port: int, remote_url: str) -> None:
75
+ try:
76
+ client = Client(transport=StreamableHttpTransport(remote_url))
77
+ app = FastMCP.as_proxy(client).http_app(transport="streamable-http")
78
+ server = uvicorn.Server(
79
+ config=uvicorn.Config(
80
+ app=app,
81
+ host=host,
82
+ port=port,
83
+ log_level="error",
84
+ lifespan="on",
85
+ )
86
+ )
87
+ server.run()
88
+ except Exception as e:
89
+ print(f"Server error: {e}")
90
+ sys.exit(1)
91
+ sys.exit(0)
92
+
93
+ @pytest.fixture(scope="class")
94
  def shttp_server(self) -> Generator[str, None, None]:
95
  with run_server_in_process(self.run_shttp_server) as url:
96
  yield f"{url}/mcp"
97
 
98
+ @pytest.fixture(scope="class")
99
  def sse_server(self) -> Generator[str, None, None]:
100
  with run_server_in_process(self.run_sse_server) as url:
101
  yield f"{url}/sse"
102
 
103
+ @pytest.fixture(scope="class")
104
+ def proxy_server(self, shttp_server: str) -> Generator[str, None, None]:
105
+ with run_server_in_process(self.run_proxy_server, shttp_server + "/mcp") as url:
106
+ yield f"{url}/mcp"
107
+
108
  async def test_client_headers_sse_resource(self, sse_server: str):
109
  async with Client(
110
  transport=SSETransport(sse_server, headers={"X-TEST": "test-123"})
 
179
  assert isinstance(result[0], TextResourceContents)
180
  headers = json.loads(result[0].text)
181
  assert headers["x-server"] == "test-abc"
182
+
183
+ async def test_client_headers_proxy(self, proxy_server: str):
184
+ async with Client(transport=StreamableHttpTransport(proxy_server)) as client:
185
+ await client.ping()
186
+ result = await client.read_resource("resource://get_headers_headers_get")
187
+ assert isinstance(result[0], TextResourceContents)
188
+ headers = json.loads(result[0].text)
189
+ assert headers["x-server"] == "test-abc"