Jeremiah Lowin commited on
Commit
3803803
·
1 Parent(s): 82be79f

Add starlette request to context

Browse files
src/fastmcp/server/context.py CHANGED
@@ -13,8 +13,9 @@ from mcp.types import (
13
  SamplingMessage,
14
  TextContent,
15
  )
16
- from pydantic import BaseModel
17
  from pydantic.networks import AnyUrl
 
18
 
19
  from fastmcp.server.server import FastMCP
20
  from fastmcp.utilities.logging import get_logger
@@ -58,17 +59,22 @@ class Context(BaseModel, Generic[ServerSessionT, LifespanContextT]):
58
 
59
  _request_context: RequestContext[ServerSessionT, LifespanContextT] | None
60
  _fastmcp: FastMCP | None
 
 
 
61
 
62
  def __init__(
63
  self,
64
  *,
65
  request_context: RequestContext[ServerSessionT, LifespanContextT] | None = None,
66
  fastmcp: FastMCP | None = None,
 
67
  **kwargs: Any,
68
  ):
69
  super().__init__(**kwargs)
70
  self._request_context = request_context
71
  self._fastmcp = fastmcp
 
72
 
73
  @property
74
  def fastmcp(self) -> FastMCP:
@@ -84,6 +90,13 @@ class Context(BaseModel, Generic[ServerSessionT, LifespanContextT]):
84
  raise ValueError("Context is not available outside of a request")
85
  return self._request_context
86
 
 
 
 
 
 
 
 
87
  async def report_progress(
88
  self, progress: float, total: float | None = None
89
  ) -> None:
 
13
  SamplingMessage,
14
  TextContent,
15
  )
16
+ from pydantic import BaseModel, ConfigDict
17
  from pydantic.networks import AnyUrl
18
+ from starlette.requests import Request
19
 
20
  from fastmcp.server.server import FastMCP
21
  from fastmcp.utilities.logging import get_logger
 
59
 
60
  _request_context: RequestContext[ServerSessionT, LifespanContextT] | None
61
  _fastmcp: FastMCP | None
62
+ _request: Request | None
63
+
64
+ model_config = ConfigDict(arbitrary_types_allowed=True)
65
 
66
  def __init__(
67
  self,
68
  *,
69
  request_context: RequestContext[ServerSessionT, LifespanContextT] | None = None,
70
  fastmcp: FastMCP | None = None,
71
+ request: Request | None = None,
72
  **kwargs: Any,
73
  ):
74
  super().__init__(**kwargs)
75
  self._request_context = request_context
76
  self._fastmcp = fastmcp
77
+ self._request = request
78
 
79
  @property
80
  def fastmcp(self) -> FastMCP:
 
90
  raise ValueError("Context is not available outside of a request")
91
  return self._request_context
92
 
93
+ @property
94
+ def request(self) -> Request:
95
+ """Access to the underlying request."""
96
+ if self._request is None:
97
+ raise ValueError("Context is not available outside of a request")
98
+ return self._request
99
+
100
  async def report_progress(
101
  self, progress: float, total: float | None = None
102
  ) -> None:
src/fastmcp/server/server.py CHANGED
@@ -9,6 +9,7 @@ from contextlib import (
9
  AsyncExitStack,
10
  asynccontextmanager,
11
  )
 
12
  from functools import partial
13
  from typing import TYPE_CHECKING, Any, Generic, Literal
14
 
@@ -61,6 +62,25 @@ logger = get_logger(__name__)
61
  NOT_FOUND = object()
62
 
63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
64
  class MountedServer:
65
  def __init__(
66
  self,
@@ -291,7 +311,11 @@ class FastMCP(Generic[LifespanResultT]):
291
  request_context = None
292
  from fastmcp.server.context import Context
293
 
294
- return Context(request_context=request_context, fastmcp=self)
 
 
 
 
295
 
296
  async def get_tools(self) -> dict[str, Tool]:
297
  """Get all registered tools, indexed by registered key."""
@@ -760,11 +784,12 @@ class FastMCP(Generic[LifespanResultT]):
760
  request.receive,
761
  request._send, # type: ignore[reportPrivateUsage]
762
  ) as streams:
763
- await self._mcp_server.run(
764
- streams[0],
765
- streams[1],
766
- self._mcp_server.create_initialization_options(),
767
- )
 
768
 
769
  return Starlette(
770
  debug=self.settings.debug,
 
9
  AsyncExitStack,
10
  asynccontextmanager,
11
  )
12
+ from contextvars import ContextVar
13
  from functools import partial
14
  from typing import TYPE_CHECKING, Any, Generic, Literal
15
 
 
62
  NOT_FOUND = object()
63
 
64
 
65
+ _current_starlette_request: ContextVar[Request | None] = ContextVar(
66
+ "starlette_request",
67
+ default=None,
68
+ )
69
+
70
+
71
+ @asynccontextmanager
72
+ async def starlette_request_context(request: Request):
73
+ token = _current_starlette_request.set(request)
74
+ try:
75
+ yield
76
+ finally:
77
+ _current_starlette_request.reset(token)
78
+
79
+
80
+ def get_current_starlette_request() -> Request | None:
81
+ return _current_starlette_request.get()
82
+
83
+
84
  class MountedServer:
85
  def __init__(
86
  self,
 
311
  request_context = None
312
  from fastmcp.server.context import Context
313
 
314
+ return Context(
315
+ request_context=request_context,
316
+ fastmcp=self,
317
+ request=get_current_starlette_request(),
318
+ )
319
 
320
  async def get_tools(self) -> dict[str, Tool]:
321
  """Get all registered tools, indexed by registered key."""
 
784
  request.receive,
785
  request._send, # type: ignore[reportPrivateUsage]
786
  ) as streams:
787
+ async with starlette_request_context(request):
788
+ await self._mcp_server.run(
789
+ streams[0],
790
+ streams[1],
791
+ self._mcp_server.create_initialization_options(),
792
+ )
793
 
794
  return Starlette(
795
  debug=self.settings.debug,