Merge branch 'main' into oauthclient

This commit is contained in:
Jeremiah Lowin 2025-05-15 21:40:06 -04:00
commit d4d6b7da53
8 changed files with 58 additions and 138 deletions

View file

@ -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"]

View file

@ -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 = {}

View file

@ -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."""

View file

@ -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)

View file

@ -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]:

View file

@ -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,

View file

@ -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
View file

@ -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]]