mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-23 14:04:18 +02:00
Merge branch 'main' into oauthclient
This commit is contained in:
commit
d4d6b7da53
8 changed files with 58 additions and 138 deletions
|
|
@ -7,7 +7,7 @@ dependencies = [
|
|||
"python-dotenv>=1.1.0",
|
||||
"exceptiongroup>=1.2.2",
|
||||
"httpx>=0.28.1",
|
||||
"mcp>=1.8.1,<2.0.0",
|
||||
"mcp>=1.9.0,<2.0.0",
|
||||
"openapi-pydantic>=0.5.1",
|
||||
"rich>=13.9.4",
|
||||
"typer>=0.15.2",
|
||||
|
|
@ -97,6 +97,7 @@ reportMissingTypeStubs = false
|
|||
useLibraryCodeForTypes = true
|
||||
venvPath = "."
|
||||
venv = ".venv"
|
||||
strict = ["src/fastmcp/server/server.py"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
extend-select = ["I", "UP"]
|
||||
|
|
|
|||
|
|
@ -263,7 +263,7 @@ def dev(
|
|||
try:
|
||||
# Import server to get dependencies
|
||||
server = _import_server(file, server_object)
|
||||
if hasattr(server, "dependencies"):
|
||||
if hasattr(server, "dependencies") and server.dependencies is not None:
|
||||
with_packages = list(set(with_packages + server.dependencies))
|
||||
|
||||
env_vars = {}
|
||||
|
|
|
|||
|
|
@ -225,9 +225,12 @@ class Client:
|
|||
progress_token: str | int,
|
||||
progress: float,
|
||||
total: float | None = None,
|
||||
message: str | None = None,
|
||||
) -> None:
|
||||
"""Send a progress notification."""
|
||||
await self.session.send_progress_notification(progress_token, progress, total)
|
||||
await self.session.send_progress_notification(
|
||||
progress_token, progress, total, message
|
||||
)
|
||||
|
||||
async def set_logging_level(self, level: mcp.types.LoggingLevel) -> None:
|
||||
"""Send a logging/setLevel request."""
|
||||
|
|
|
|||
|
|
@ -1,104 +0,0 @@
|
|||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
from urllib.parse import quote
|
||||
from uuid import uuid4
|
||||
|
||||
import anyio
|
||||
from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
|
||||
from mcp.server.sse import SseServerTransport as LowLevelSSEServerTransport
|
||||
from mcp.shared.message import SessionMessage
|
||||
from sse_starlette import EventSourceResponse
|
||||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SseServerTransport(LowLevelSSEServerTransport):
|
||||
"""
|
||||
Patched SSE server transport
|
||||
"""
|
||||
|
||||
@asynccontextmanager
|
||||
async def connect_sse(self, scope: Scope, receive: Receive, send: Send):
|
||||
"""
|
||||
See https://github.com/modelcontextprotocol/python-sdk/pull/659/
|
||||
"""
|
||||
if scope["type"] != "http":
|
||||
logger.error("connect_sse received non-HTTP request")
|
||||
raise ValueError("connect_sse can only handle HTTP requests")
|
||||
|
||||
logger.debug("Setting up SSE connection")
|
||||
read_stream: MemoryObjectReceiveStream[SessionMessage | Exception]
|
||||
read_stream_writer: MemoryObjectSendStream[SessionMessage | Exception]
|
||||
|
||||
write_stream: MemoryObjectSendStream[SessionMessage]
|
||||
write_stream_reader: MemoryObjectReceiveStream[SessionMessage]
|
||||
|
||||
read_stream_writer, read_stream = anyio.create_memory_object_stream(0)
|
||||
write_stream, write_stream_reader = anyio.create_memory_object_stream(0)
|
||||
|
||||
session_id = uuid4()
|
||||
self._read_stream_writers[session_id] = read_stream_writer
|
||||
logger.debug(f"Created new session with ID: {session_id}")
|
||||
|
||||
# Determine the full path for the message endpoint to be sent to the client.
|
||||
# scope['root_path'] is the prefix where the current Starlette app
|
||||
# instance is mounted.
|
||||
# e.g., "" if top-level, or "/api_prefix" if mounted under "/api_prefix".
|
||||
root_path = scope.get("root_path", "")
|
||||
|
||||
# self._endpoint is the path *within* this app, e.g., "/messages".
|
||||
# Concatenating them gives the full absolute path from the server root.
|
||||
# e.g., "" + "/messages" -> "/messages"
|
||||
# e.g., "/api_prefix" + "/messages" -> "/api_prefix/messages"
|
||||
full_message_path_for_client = root_path.rstrip("/") + self._endpoint
|
||||
|
||||
# This is the URI (path + query) the client will use to POST messages.
|
||||
client_post_uri_data = (
|
||||
f"{quote(full_message_path_for_client)}?session_id={session_id.hex}"
|
||||
)
|
||||
|
||||
sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[
|
||||
dict[str, Any]
|
||||
](0)
|
||||
|
||||
async def sse_writer():
|
||||
logger.debug("Starting SSE writer")
|
||||
async with sse_stream_writer, write_stream_reader:
|
||||
await sse_stream_writer.send(
|
||||
{"event": "endpoint", "data": client_post_uri_data}
|
||||
)
|
||||
logger.debug(f"Sent endpoint event: {client_post_uri_data}")
|
||||
|
||||
async for session_message in write_stream_reader:
|
||||
logger.debug(f"Sending message via SSE: {session_message}")
|
||||
await sse_stream_writer.send(
|
||||
{
|
||||
"event": "message",
|
||||
"data": session_message.message.model_dump_json(
|
||||
by_alias=True, exclude_none=True
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
async with anyio.create_task_group() as tg:
|
||||
|
||||
async def response_wrapper(scope: Scope, receive: Receive, send: Send):
|
||||
"""
|
||||
The EventSourceResponse returning signals a client close / disconnect.
|
||||
In this case we close our side of the streams to signal the client that
|
||||
the connection has been closed.
|
||||
"""
|
||||
await EventSourceResponse(
|
||||
content=sse_stream_reader, data_sender_callable=sse_writer
|
||||
)(scope, receive, send)
|
||||
await read_stream_writer.aclose()
|
||||
await write_stream_reader.aclose()
|
||||
logging.debug(f"Client session disconnected {session_id}")
|
||||
|
||||
logger.debug("Starting SSE response task")
|
||||
tg.start_soon(response_wrapper, scope, receive, send)
|
||||
|
||||
logger.debug("Yielding read and write streams")
|
||||
yield (read_stream, write_stream)
|
||||
|
|
@ -56,7 +56,7 @@ class Context:
|
|||
ctx.error("Error message")
|
||||
|
||||
# Report progress
|
||||
ctx.report_progress(50, 100)
|
||||
ctx.report_progress(50, 100, "Processing")
|
||||
|
||||
# Access resources
|
||||
data = ctx.read_resource("resource://data")
|
||||
|
|
@ -96,7 +96,7 @@ class Context:
|
|||
return self.fastmcp._mcp_server.request_context
|
||||
|
||||
async def report_progress(
|
||||
self, progress: float, total: float | None = None
|
||||
self, progress: float, total: float | None = None, message: str | None = None
|
||||
) -> None:
|
||||
"""Report progress for the current operation.
|
||||
|
||||
|
|
@ -115,7 +115,10 @@ class Context:
|
|||
return
|
||||
|
||||
await self.request_context.session.send_progress_notification(
|
||||
progress_token=progress_token, progress=progress, total=total
|
||||
progress_token=progress_token,
|
||||
progress=progress,
|
||||
total=total,
|
||||
message=message,
|
||||
)
|
||||
|
||||
async def read_resource(self, uri: str | AnyUrl) -> list[ReadResourceContents]:
|
||||
|
|
|
|||
|
|
@ -10,9 +10,16 @@ from mcp.server.auth.middleware.bearer_auth import (
|
|||
BearerAuthBackend,
|
||||
RequireAuthMiddleware,
|
||||
)
|
||||
from mcp.server.auth.provider import OAuthAuthorizationServerProvider
|
||||
from mcp.server.auth.provider import (
|
||||
AccessTokenT,
|
||||
AuthorizationCodeT,
|
||||
OAuthAuthorizationServerProvider,
|
||||
RefreshTokenT,
|
||||
)
|
||||
from mcp.server.auth.routes import create_auth_routes
|
||||
from mcp.server.auth.settings import AuthSettings
|
||||
from mcp.server.lowlevel.server import LifespanResultT
|
||||
from mcp.server.sse import SseServerTransport
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
|
|
@ -22,7 +29,6 @@ from starlette.responses import Response
|
|||
from starlette.routing import BaseRoute, Mount, Route
|
||||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
from fastmcp.low_level.sse_server_transport import SseServerTransport
|
||||
from fastmcp.utilities.logging import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -30,6 +36,7 @@ if TYPE_CHECKING:
|
|||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
_current_http_request: ContextVar[Request | None] = ContextVar(
|
||||
"http_request",
|
||||
default=None,
|
||||
|
|
@ -62,7 +69,10 @@ class RequestContextMiddleware:
|
|||
|
||||
|
||||
def setup_auth_middleware_and_routes(
|
||||
auth_server_provider: OAuthAuthorizationServerProvider | None,
|
||||
auth_server_provider: OAuthAuthorizationServerProvider[
|
||||
AuthorizationCodeT, RefreshTokenT, AccessTokenT
|
||||
]
|
||||
| None,
|
||||
auth_settings: AuthSettings | None,
|
||||
) -> tuple[list[Middleware], list[BaseRoute], list[str]]:
|
||||
"""Set up authentication middleware and routes if auth is enabled.
|
||||
|
|
@ -136,10 +146,13 @@ def create_base_app(
|
|||
|
||||
|
||||
def create_sse_app(
|
||||
server: FastMCP,
|
||||
server: FastMCP[LifespanResultT],
|
||||
message_path: str,
|
||||
sse_path: str,
|
||||
auth_server_provider: OAuthAuthorizationServerProvider | None = None,
|
||||
auth_server_provider: OAuthAuthorizationServerProvider[
|
||||
AuthorizationCodeT, RefreshTokenT, AccessTokenT
|
||||
]
|
||||
| None = None,
|
||||
auth_settings: AuthSettings | None = None,
|
||||
debug: bool = False,
|
||||
routes: list[BaseRoute] | None = None,
|
||||
|
|
@ -236,10 +249,13 @@ def create_sse_app(
|
|||
|
||||
|
||||
def create_streamable_http_app(
|
||||
server: FastMCP,
|
||||
server: FastMCP[LifespanResultT],
|
||||
streamable_http_path: str,
|
||||
event_store: None = None,
|
||||
auth_server_provider: OAuthAuthorizationServerProvider | None = None,
|
||||
auth_server_provider: OAuthAuthorizationServerProvider[
|
||||
AuthorizationCodeT, RefreshTokenT, AccessTokenT
|
||||
]
|
||||
| None = None,
|
||||
auth_settings: AuthSettings | None = None,
|
||||
json_response: bool = False,
|
||||
stateless_http: bool = False,
|
||||
|
|
|
|||
|
|
@ -66,7 +66,7 @@ DuplicateBehavior = Literal["warn", "error", "replace", "ignore"]
|
|||
|
||||
|
||||
@asynccontextmanager
|
||||
async def default_lifespan(server: FastMCP) -> AsyncIterator[Any]:
|
||||
async def default_lifespan(server: FastMCP[LifespanResultT]) -> AsyncIterator[Any]:
|
||||
"""Default lifespan context manager that does nothing.
|
||||
|
||||
Args:
|
||||
|
|
@ -79,8 +79,10 @@ async def default_lifespan(server: FastMCP) -> AsyncIterator[Any]:
|
|||
|
||||
|
||||
def _lifespan_wrapper(
|
||||
app: FastMCP,
|
||||
lifespan: Callable[[FastMCP], AbstractAsyncContextManager[LifespanResultT]],
|
||||
app: FastMCP[LifespanResultT],
|
||||
lifespan: Callable[
|
||||
[FastMCP[LifespanResultT]], AbstractAsyncContextManager[LifespanResultT]
|
||||
],
|
||||
) -> Callable[
|
||||
[MCPServer[LifespanResultT]], AbstractAsyncContextManager[LifespanResultT]
|
||||
]:
|
||||
|
|
@ -189,15 +191,13 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
"""
|
||||
if transport is None:
|
||||
transport = "stdio"
|
||||
if transport not in ["stdio", "streamable-http", "sse"]:
|
||||
if transport not in {"stdio", "streamable-http", "sse"}:
|
||||
raise ValueError(f"Unknown transport: {transport}")
|
||||
|
||||
if transport == "stdio":
|
||||
await self.run_stdio_async(**transport_kwargs)
|
||||
elif transport == "streamable-http":
|
||||
await self.run_http_async(transport="streamable-http", **transport_kwargs)
|
||||
elif transport == "sse":
|
||||
await self.run_http_async(transport="sse", **transport_kwargs)
|
||||
elif transport in {"streamable-http", "sse"}:
|
||||
await self.run_http_async(transport=transport, **transport_kwargs)
|
||||
else:
|
||||
raise ValueError(f"Unknown transport: {transport}")
|
||||
|
||||
|
|
@ -228,7 +228,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
async def get_tools(self) -> dict[str, Tool]:
|
||||
"""Get all registered tools, indexed by registered key."""
|
||||
if (tools := self._cache.get("tools")) is self._cache.NOT_FOUND:
|
||||
tools = {}
|
||||
tools: dict[str, Tool] = {}
|
||||
for server in self._mounted_servers.values():
|
||||
server_tools = await server.get_tools()
|
||||
tools.update(server_tools)
|
||||
|
|
@ -239,7 +239,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
async def get_resources(self) -> dict[str, Resource]:
|
||||
"""Get all registered resources, indexed by registered key."""
|
||||
if (resources := self._cache.get("resources")) is self._cache.NOT_FOUND:
|
||||
resources = {}
|
||||
resources: dict[str, Resource] = {}
|
||||
for server in self._mounted_servers.values():
|
||||
server_resources = await server.get_resources()
|
||||
resources.update(server_resources)
|
||||
|
|
@ -252,7 +252,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
if (
|
||||
templates := self._cache.get("resource_templates")
|
||||
) is self._cache.NOT_FOUND:
|
||||
templates = {}
|
||||
templates: dict[str, ResourceTemplate] = {}
|
||||
for server in self._mounted_servers.values():
|
||||
server_templates = await server.get_resource_templates()
|
||||
templates.update(server_templates)
|
||||
|
|
@ -265,7 +265,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
List all available prompts.
|
||||
"""
|
||||
if (prompts := self._cache.get("prompts")) is self._cache.NOT_FOUND:
|
||||
prompts = {}
|
||||
prompts: dict[str, Prompt] = {}
|
||||
for server in self._mounted_servers.values():
|
||||
server_prompts = await server.get_prompts()
|
||||
prompts.update(server_prompts)
|
||||
|
|
@ -743,7 +743,8 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
port: int | None = None,
|
||||
log_level: str | None = None,
|
||||
path: str | None = None,
|
||||
uvicorn_config: dict | None = None,
|
||||
uvicorn_config: dict[str, Any] | None = None,
|
||||
middleware: list[Middleware] | None = None,
|
||||
) -> None:
|
||||
"""Run the server using HTTP transport.
|
||||
|
||||
|
|
@ -760,7 +761,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
# lifespan is required for streamable http
|
||||
uvicorn_config["lifespan"] = "on"
|
||||
|
||||
app = self.http_app(path=path, transport=transport)
|
||||
app = self.http_app(path=path, transport=transport, middleware=middleware)
|
||||
|
||||
config = uvicorn.Config(
|
||||
app,
|
||||
|
|
@ -779,7 +780,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
log_level: str | None = None,
|
||||
path: str | None = None,
|
||||
message_path: str | None = None,
|
||||
uvicorn_config: dict | None = None,
|
||||
uvicorn_config: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""Run the server using SSE transport."""
|
||||
|
||||
|
|
@ -901,7 +902,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
port: int | None = None,
|
||||
log_level: str | None = None,
|
||||
path: str | None = None,
|
||||
uvicorn_config: dict | None = None,
|
||||
uvicorn_config: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
# Deprecated since 2.3.2
|
||||
warnings.warn(
|
||||
|
|
@ -1128,7 +1129,7 @@ class MountedServer:
|
|||
def __init__(
|
||||
self,
|
||||
prefix: str,
|
||||
server: FastMCP,
|
||||
server: FastMCP[LifespanResultT],
|
||||
tool_separator: str | None = None,
|
||||
resource_separator: str | None = None,
|
||||
prompt_separator: str | None = None,
|
||||
|
|
|
|||
8
uv.lock
generated
8
uv.lock
generated
|
|
@ -458,7 +458,7 @@ requires-dist = [
|
|||
{ name = "authlib", specifier = ">=1.5.2" },
|
||||
{ name = "exceptiongroup", specifier = ">=1.2.2" },
|
||||
{ name = "httpx", specifier = ">=0.28.1" },
|
||||
{ name = "mcp", specifier = ">=1.8.1,<2.0.0" },
|
||||
{ name = "mcp", specifier = ">=1.9.0,<2.0.0" },
|
||||
{ name = "openapi-pydantic", specifier = ">=0.5.1" },
|
||||
{ name = "python-dotenv", specifier = ">=1.1.0" },
|
||||
{ name = "rich", specifier = ">=13.9.4" },
|
||||
|
|
@ -691,7 +691,7 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "mcp"
|
||||
version = "1.8.1"
|
||||
version = "1.9.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "anyio" },
|
||||
|
|
@ -704,9 +704,9 @@ dependencies = [
|
|||
{ name = "starlette" },
|
||||
{ name = "uvicorn", marker = "sys_platform != 'emscripten'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/7c/13/16b712e8a3be6a736b411df2fc6b4e75eb1d3e99b1cd57a3a1decf17f612/mcp-1.8.1.tar.gz", hash = "sha256:ec0646271d93749f784d2316fb5fe6102fb0d1be788ec70a9e2517e8f2722c0e", size = 265605, upload-time = "2025-05-12T17:33:57.887Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/bc/8d/0f4468582e9e97b0a24604b585c651dfd2144300ecffd1c06a680f5c8861/mcp-1.9.0.tar.gz", hash = "sha256:905d8d208baf7e3e71d70c82803b89112e321581bcd2530f9de0fe4103d28749", size = 281432 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/1c/5d/91cf0d40e40ae9ecf8d4004e0f9611eea86085aa0b5505493e0ff53972da/mcp-1.8.1-py3-none-any.whl", hash = "sha256:948e03783859fa35abe05b9b6c0a1d5519be452fc079dc8d7f682549591c1770", size = 119761, upload-time = "2025-05-12T17:33:56.136Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/a5/d5/22e36c95c83c80eb47c83f231095419cf57cf5cca5416f1c960032074c78/mcp-1.9.0-py3-none-any.whl", hash = "sha256:9dfb89c8c56f742da10a5910a1f64b0d2ac2c3ed2bd572ddb1cfab7f35957178", size = 125082 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue