mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-28 10:18:08 +02:00
Add starlette request to context
This commit is contained in:
parent
9d422e08aa
commit
750c6bf005
2 changed files with 45 additions and 7 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue