Merge pull request #476 from jlowin/strict-typing

strict typing for `server.py`
This commit is contained in:
nate nowack 2025-05-15 20:12:50 -05:00 committed by GitHub
commit d27d331c9d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 36 additions and 17 deletions

View file

@ -96,6 +96,7 @@ reportMissingTypeStubs = false
useLibraryCodeForTypes = true
venvPath = "."
venv = ".venv"
strict = ["src/fastmcp/server/server.py"]
[tool.ruff.lint]
extend-select = ["I", "UP"]

View file

@ -10,9 +10,15 @@ 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.streamable_http_manager import StreamableHTTPSessionManager
from starlette.applications import Starlette
from starlette.middleware import Middleware
@ -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]
]:
@ -226,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)
@ -237,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)
@ -250,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)
@ -263,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)
@ -741,7 +743,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,
middleware: list[Middleware] | None = None,
) -> None:
"""Run the server using HTTP transport.
@ -778,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."""
@ -900,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(
@ -1127,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,