Add starlette request to context

This commit is contained in:
Jeremiah Lowin 2025-05-01 16:29:06 -04:00
commit 750c6bf005
2 changed files with 45 additions and 7 deletions

View file

@ -13,8 +13,9 @@ from mcp.types import (
SamplingMessage,
TextContent,
)
from pydantic import BaseModel
from pydantic import BaseModel, ConfigDict
from pydantic.networks import AnyUrl
from starlette.requests import Request
from fastmcp.server.server import FastMCP
from fastmcp.utilities.logging import get_logger
@ -58,17 +59,22 @@ class Context(BaseModel, Generic[ServerSessionT, LifespanContextT]):
_request_context: RequestContext[ServerSessionT, LifespanContextT] | None
_fastmcp: FastMCP | None
_request: Request | None
model_config = ConfigDict(arbitrary_types_allowed=True)
def __init__(
self,
*,
request_context: RequestContext[ServerSessionT, LifespanContextT] | None = None,
fastmcp: FastMCP | None = None,
request: Request | None = None,
**kwargs: Any,
):
super().__init__(**kwargs)
self._request_context = request_context
self._fastmcp = fastmcp
self._request = request
@property
def fastmcp(self) -> FastMCP:
@ -84,6 +90,13 @@ class Context(BaseModel, Generic[ServerSessionT, LifespanContextT]):
raise ValueError("Context is not available outside of a request")
return self._request_context
@property
def request(self) -> Request:
"""Access to the underlying request."""
if self._request is None:
raise ValueError("Context is not available outside of a request")
return self._request
async def report_progress(
self, progress: float, total: float | None = None
) -> None:

View file

@ -9,6 +9,7 @@ from contextlib import (
AsyncExitStack,
asynccontextmanager,
)
from contextvars import ContextVar
from functools import partial
from typing import TYPE_CHECKING, Any, Generic, Literal
@ -61,6 +62,25 @@ logger = get_logger(__name__)
NOT_FOUND = object()
_current_starlette_request: ContextVar[Request | None] = ContextVar(
"starlette_request",
default=None,
)
@asynccontextmanager
async def starlette_request_context(request: Request):
token = _current_starlette_request.set(request)
try:
yield
finally:
_current_starlette_request.reset(token)
def get_current_starlette_request() -> Request | None:
return _current_starlette_request.get()
class MountedServer:
def __init__(
self,
@ -291,7 +311,11 @@ class FastMCP(Generic[LifespanResultT]):
request_context = None
from fastmcp.server.context import Context
return Context(request_context=request_context, fastmcp=self)
return Context(
request_context=request_context,
fastmcp=self,
request=get_current_starlette_request(),
)
async def get_tools(self) -> dict[str, Tool]:
"""Get all registered tools, indexed by registered key."""
@ -760,11 +784,12 @@ class FastMCP(Generic[LifespanResultT]):
request.receive,
request._send, # type: ignore[reportPrivateUsage]
) as streams:
await self._mcp_server.run(
streams[0],
streams[1],
self._mcp_server.create_initialization_options(),
)
async with starlette_request_context(request):
await self._mcp_server.run(
streams[0],
streams[1],
self._mcp_server.create_initialization_options(),
)
return Starlette(
debug=self.settings.debug,