Add "High Value" Ruff Rules (#2255)

* Safe Fixes from ruff

* Fix remaining issues

* lint/check

* Fix mysterious ty check errors

* small cleanup

* pr fixes
This commit is contained in:
William Easton 2025-10-26 09:20:31 -05:00 committed by GitHub
commit 9d4c378e1b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
65 changed files with 340 additions and 333 deletions

View file

@ -7,14 +7,14 @@ from ._read import fetch_notifications, fetch_timeline, search_for_posts
from ._social import follow_user_by_handle, like_post_by_uri, repost_by_uri from ._social import follow_user_by_handle, like_post_by_uri, repost_by_uri
__all__ = [ __all__ = [
"get_client",
"get_profile_info",
"create_post", "create_post",
"create_thread", "create_thread",
"fetch_timeline",
"search_for_posts",
"fetch_notifications", "fetch_notifications",
"fetch_timeline",
"follow_user_by_handle", "follow_user_by_handle",
"get_client",
"get_profile_info",
"like_post_by_uri", "like_post_by_uri",
"repost_by_uri", "repost_by_uri",
"search_for_posts",
] ]

View file

@ -1,3 +1,9 @@
# /// script
# dependencies = ["aiohttp", "fastmcp"]
# ///
# uv pip install aiohttp fastmcp
import aiohttp import aiohttp
from fastmcp.server import FastMCP from fastmcp.server import FastMCP

View file

@ -19,7 +19,7 @@ from typing import Annotated, Any, Self
import asyncpg import asyncpg
import numpy as np import numpy as np
from openai import AsyncOpenAI from openai import AsyncOpenAI
from pgvector.asyncpg import register_vector # Import register_vector from pgvector.asyncpg import register_vector
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from pydantic_ai import Agent from pydantic_ai import Agent
@ -149,7 +149,9 @@ class MemoryNode(BaseModel):
) )
self.importance += other.importance self.importance += other.importance
self.access_count += other.access_count self.access_count += other.access_count
self.embedding = [(a + b) / 2 for a, b in zip(self.embedding, other.embedding)] self.embedding = [
(a + b) / 2 for a, b in zip(self.embedding, other.embedding, strict=True)
]
self.summary = await do_ai( self.summary = await do_ai(
self.content, "Summarize the following text concisely.", str, deps self.content, "Summarize the following text concisely.", str, deps
) )
@ -281,9 +283,9 @@ async def display_memory_tree(deps: Deps) -> str:
@mcp.tool @mcp.tool
async def remember( async def remember(
contents: list[str] = Field( contents: Annotated[
description="List of observations or memories to store" list[str], Field(description="List of observations or memories to store")
), ],
): ):
deps = Deps(openai=AsyncOpenAI(), pool=await get_db_pool()) deps = Deps(openai=AsyncOpenAI(), pool=await get_db_pool())
try: try:

View file

@ -137,12 +137,34 @@ unknown-argument = "ignore" # 61 errors
call-non-callable = "ignore" # 7 errors call-non-callable = "ignore" # 7 errors
[tool.ruff.lint] [tool.ruff.lint]
extend-select = ["I", "UP"] fixable = ["ALL"]
ignore = [
"COM812",
"PLR0913", # Too many arguments, MCP Servers have a lot of arguments, OKAY?!
"SIM102", # Dont require combining if statements
]
extend-select = [
"B", # flake8-bugbear: Catches actual bugs like mutable default arguments
"C4", # flake8-comprehensions: More efficient/readable comprehensions
"I", # flake8-builtins: Catches builtins that are not explicitly imported
"PIE", # flake8-pie: More idiomatic Python code
"RUF", # Ruff-specific: Modern best practices unique to Ruff
"SIM", # flake8-simplify: Simplifies verbose code patterns
"UP" # flake8-unused-imports: Catches unused imports
]
[tool.ruff.lint.per-file-ignores] [tool.ruff.lint.per-file-ignores]
"__init__.py" = ["F401", "I001", "RUF013"] "__init__.py" = ["F401", "I001", "RUF013"]
# allow imports not at the top of the file # allow imports not at the top of the file
"src/fastmcp/__init__.py" = ["E402"] "src/fastmcp/__init__.py" = ["E402"]
"!src/**.py" = [ # Only enforce extended ruff rules for code in src/
"B", # flake8-bugbear
"C4", # flake8-comprehensions
"PIE", # flake8-pie
"RUF", # Ruff-specific
"SIM", # flake8-simplify
]
[tool.codespell] [tool.codespell]
ignore-words-list = "asend,shttp,te" ignore-words-list = "asend,shttp,te"

View file

@ -48,9 +48,9 @@ def __getattr__(name: str):
__all__ = [ __all__ = [
"FastMCP",
"Context",
"client",
"Client", "Client",
"Context",
"FastMCP",
"client",
"settings", "settings",
] ]

View file

@ -78,7 +78,7 @@ def with_argv(args: list[str] | None):
original = sys.argv[:] original = sys.argv[:]
try: try:
# Preserve the script name (sys.argv[0]) and replace the rest # Preserve the script name (sys.argv[0]) and replace the rest
sys.argv = [sys.argv[0]] + args sys.argv = [sys.argv[0], *args]
yield yield
finally: finally:
sys.argv = original sys.argv = original
@ -277,7 +277,7 @@ async def dev(
# Run the MCP Inspector command # Run the MCP Inspector command
process = subprocess.run( process = subprocess.run(
[npx_cmd, inspector_cmd] + uv_cmd, [npx_cmd, inspector_cmd, *uv_cmd],
check=True, check=True,
env=env, env=env,
) )

View file

@ -15,18 +15,18 @@ from .transports import (
from .auth import OAuth, BearerAuth from .auth import OAuth, BearerAuth
__all__ = [ __all__ = [
"BearerAuth",
"Client", "Client",
"ClientTransport", "ClientTransport",
"WSTransport", "FastMCPTransport",
"NodeStdioTransport",
"NpxStdioTransport",
"OAuth",
"PythonStdioTransport",
"SSETransport", "SSETransport",
"StdioTransport", "StdioTransport",
"PythonStdioTransport",
"NodeStdioTransport",
"UvxStdioTransport",
"UvStdioTransport",
"NpxStdioTransport",
"FastMCPTransport",
"StreamableHttpTransport", "StreamableHttpTransport",
"OAuth", "UvStdioTransport",
"BearerAuth", "UvxStdioTransport",
"WSTransport",
] ]

View file

@ -36,8 +36,6 @@ logger = get_logger(__name__)
class ClientNotFoundError(Exception): class ClientNotFoundError(Exception):
"""Raised when OAuth client credentials are not found on the server.""" """Raised when OAuth client credentials are not found on the server."""
pass
async def check_if_auth_required( async def check_if_auth_required(
mcp_url: str, httpx_kwargs: dict[str, Any] | None = None mcp_url: str, httpx_kwargs: dict[str, Any] | None = None
@ -58,7 +56,7 @@ async def check_if_auth_required(
return True return True
# Check for WWW-Authenticate header # Check for WWW-Authenticate header
if "WWW-Authenticate" in response.headers: if "WWW-Authenticate" in response.headers: # noqa: SIM103
return True return True
# If we get a successful response, auth may not be required # If we get a successful response, auth may not be required
@ -195,7 +193,8 @@ class OAuth(OAuthClientProvider):
warn( warn(
message="Using in-memory token storage is not recommended for production use -- " message="Using in-memory token storage is not recommended for production use -- "
+ "tokens will be lost on server restart." + "tokens will be lost on server restart.",
stacklevel=2,
) )
self.token_storage_adapter: TokenStorageAdapter = TokenStorageAdapter( self.token_storage_adapter: TokenStorageAdapter = TokenStorageAdapter(
@ -272,8 +271,10 @@ class OAuth(OAuthClientProvider):
if result.error: if result.error:
raise result.error raise result.error
return result.code, result.state # type: ignore return result.code, result.state # type: ignore
except TimeoutError: except TimeoutError as e:
raise TimeoutError(f"OAuth callback timed out after {TIMEOUT} seconds") raise TimeoutError(
f"OAuth callback timed out after {TIMEOUT} seconds"
) from e
finally: finally:
server.should_exit = True server.should_exit = True
await anyio.sleep(0.1) # Allow server to shut down gracefully await anyio.sleep(0.1) # Allow server to shut down gracefully

View file

@ -61,15 +61,15 @@ from .transports import (
__all__ = [ __all__ = [
"Client", "Client",
"SessionKwargs", "ClientSamplingHandler",
"RootsHandler", "ElicitationHandler",
"RootsList",
"LogHandler", "LogHandler",
"MessageHandler", "MessageHandler",
"ClientSamplingHandler",
"SamplingHandler",
"ElicitationHandler",
"ProgressHandler", "ProgressHandler",
"RootsHandler",
"RootsList",
"SamplingHandler",
"SessionKwargs",
] ]
logger = get_logger(__name__) logger = get_logger(__name__)
@ -362,10 +362,10 @@ class Client(Generic[ClientTransportT]):
await self._session_state.session.initialize() await self._session_state.session.initialize()
) )
yield yield
except anyio.ClosedResourceError: except anyio.ClosedResourceError as e:
raise RuntimeError("Server session was closed unexpectedly") raise RuntimeError("Server session was closed unexpectedly") from e
except TimeoutError: except TimeoutError as e:
raise RuntimeError("Failed to initialize server session") raise RuntimeError("Failed to initialize server session") from e
finally: finally:
self._session_state.session = None self._session_state.session = None
self._session_state.initialize_result = None self._session_state.initialize_result = None

View file

@ -11,7 +11,7 @@ from mcp.types import SamplingMessage
from fastmcp.server.sampling.handler import ServerSamplingHandler from fastmcp.server.sampling.handler import ServerSamplingHandler
__all__ = ["SamplingMessage", "SamplingParams", "SamplingHandler"] __all__ = ["SamplingHandler", "SamplingMessage", "SamplingParams"]
ClientSamplingHandler: TypeAlias = Callable[ ClientSamplingHandler: TypeAlias = Callable[

View file

@ -46,16 +46,16 @@ ClientTransportT = TypeVar("ClientTransportT", bound="ClientTransport")
__all__ = [ __all__ = [
"ClientTransport", "ClientTransport",
"SSETransport",
"StreamableHttpTransport",
"StdioTransport",
"PythonStdioTransport",
"FastMCPStdioTransport", "FastMCPStdioTransport",
"NodeStdioTransport",
"UvxStdioTransport",
"UvStdioTransport",
"NpxStdioTransport",
"FastMCPTransport", "FastMCPTransport",
"NodeStdioTransport",
"NpxStdioTransport",
"PythonStdioTransport",
"SSETransport",
"StdioTransport",
"StreamableHttpTransport",
"UvStdioTransport",
"UvxStdioTransport",
"infer_transport", "infer_transport",
] ]
@ -109,9 +109,8 @@ class ClientTransport(abc.ABC):
# Basic representation for subclasses # Basic representation for subclasses
return f"<{self.__class__.__name__}>" return f"<{self.__class__.__name__}>"
async def close(self): async def close(self): # noqa: B027
"""Close the transport.""" """Close the transport."""
pass
def _set_auth(self, auth: httpx.Auth | Literal["oauth"] | str | None): def _set_auth(self, auth: httpx.Auth | Literal["oauth"] | str | None):
if auth is not None: if auth is not None:
@ -141,10 +140,10 @@ class WSTransport(ClientTransport):
) -> AsyncIterator[ClientSession]: ) -> AsyncIterator[ClientSession]:
try: try:
from mcp.client.websocket import websocket_client from mcp.client.websocket import websocket_client
except ImportError: except ImportError as e:
raise ImportError( raise ImportError(
"The websocket transport is not available. Please install fastmcp[websockets] or install the websockets package manually." "The websocket transport is not available. Please install fastmcp[websockets] or install the websockets package manually."
) ) from e
async with websocket_client(self.url) as transport: async with websocket_client(self.url) as transport:
read_stream, write_stream = transport read_stream, write_stream = transport
@ -207,7 +206,7 @@ class SSETransport(ClientTransport):
# instead we simply leave the kwarg out if it's not provided # instead we simply leave the kwarg out if it's not provided
if self.sse_read_timeout is not None: if self.sse_read_timeout is not None:
client_kwargs["sse_read_timeout"] = self.sse_read_timeout.total_seconds() client_kwargs["sse_read_timeout"] = self.sse_read_timeout.total_seconds()
if session_kwargs.get("read_timeout_seconds", None) is not None: if session_kwargs.get("read_timeout_seconds") is not None:
read_timeout_seconds = cast( read_timeout_seconds = cast(
datetime.timedelta, session_kwargs.get("read_timeout_seconds") datetime.timedelta, session_kwargs.get("read_timeout_seconds")
) )
@ -277,7 +276,7 @@ class StreamableHttpTransport(ClientTransport):
# instead we simply leave the kwarg out if it's not provided # instead we simply leave the kwarg out if it's not provided
if self.sse_read_timeout is not None: if self.sse_read_timeout is not None:
client_kwargs["sse_read_timeout"] = self.sse_read_timeout client_kwargs["sse_read_timeout"] = self.sse_read_timeout
if session_kwargs.get("read_timeout_seconds", None) is not None: if session_kwargs.get("read_timeout_seconds") is not None:
client_kwargs["timeout"] = session_kwargs.get("read_timeout_seconds") client_kwargs["timeout"] = session_kwargs.get("read_timeout_seconds")
if self.httpx_client_factory is not None: if self.httpx_client_factory is not None:
@ -451,8 +450,7 @@ async def _stdio_transport_connect_task(
if log_file is None: if log_file is None:
log_file_handle = sys.stderr log_file_handle = sys.stderr
elif isinstance(log_file, Path): elif isinstance(log_file, Path):
log_file_handle = open(log_file, "a") log_file_handle = stack.enter_context(log_file.open("a"))
stack.callback(log_file_handle.close)
else: else:
# Must be TextIO - use it directly # Must be TextIO - use it directly
log_file_handle = log_file log_file_handle = log_file
@ -852,26 +850,28 @@ class FastMCPTransport(ClientTransport):
server_read, server_write = server_streams server_read, server_write = server_streams
# Create a cancel scope for the server task # Create a cancel scope for the server task
async with anyio.create_task_group() as tg: async with (
async with _enter_server_lifespan(server=self.server): anyio.create_task_group() as tg,
tg.start_soon( _enter_server_lifespan(server=self.server),
lambda: self.server._mcp_server.run( ):
server_read, tg.start_soon(
server_write, lambda: self.server._mcp_server.run(
self.server._mcp_server.create_initialization_options(), server_read,
raise_exceptions=self.raise_exceptions, server_write,
) self.server._mcp_server.create_initialization_options(),
raise_exceptions=self.raise_exceptions,
) )
)
try: try:
async with ClientSession( async with ClientSession(
read_stream=client_read, read_stream=client_read,
write_stream=client_write, write_stream=client_write,
**session_kwargs, **session_kwargs,
) as client_session: ) as client_session:
yield client_session yield client_session
finally: finally:
tg.cancel_scope.cancel() tg.cancel_scope.cancel()
def __repr__(self) -> str: def __repr__(self) -> str:
return f"<FastMCPTransport(server='{self.server.name}')>" return f"<FastMCPTransport(server='{self.server.name}')>"
@ -952,7 +952,7 @@ class MCPConfigTransport(ClientTransport):
# if there's exactly one server, create a client for that server # if there's exactly one server, create a client for that server
elif len(self.config.mcpServers) == 1: elif len(self.config.mcpServers) == 1:
self.transport = list(self.config.mcpServers.values())[0].to_transport() self.transport = next(iter(self.config.mcpServers.values())).to_transport()
self._underlying_transports.append(self.transport) self._underlying_transports.append(self.transport)
# otherwise create a composite client # otherwise create a composite client

View file

@ -1,4 +1,4 @@
from .component_manager import set_up_component_manager from .component_manager import set_up_component_manager
from .component_service import ComponentService from .component_service import ComponentService
__all__ = ["set_up_component_manager", "ComponentService"] __all__ = ["ComponentService", "set_up_component_manager"]

View file

@ -97,11 +97,11 @@ def make_endpoint(action, component, config):
return JSONResponse( return JSONResponse(
{"message": f"{action.capitalize()}d {component}: {name}"} {"message": f"{action.capitalize()}d {component}: {name}"}
) )
except NotFoundError: except NotFoundError as e:
raise StarletteHTTPException( raise StarletteHTTPException(
status_code=404, status_code=404,
detail=f"Unknown {component}: {name}", detail=f"Unknown {component}: {name}",
) ) from e
return endpoint return endpoint

View file

@ -2,7 +2,7 @@ from .mcp_mixin import MCPMixin, mcp_tool, mcp_resource, mcp_prompt
__all__ = [ __all__ = [
"MCPMixin", "MCPMixin",
"mcp_tool",
"mcp_resource",
"mcp_prompt", "mcp_prompt",
"mcp_resource",
"mcp_tool",
] ]

View file

@ -21,10 +21,10 @@ try:
ChatCompletionUserMessageParam, ChatCompletionUserMessageParam,
) )
from openai.types.shared.chat_model import ChatModel from openai.types.shared.chat_model import ChatModel
except ImportError: except ImportError as e:
raise ImportError( raise ImportError(
"The `openai` package is not installed. Please install `fastmcp[openai]` or add `openai` to your dependencies manually." "The `openai` package is not installed. Please install `fastmcp[openai]` or add `openai` to your dependencies manually."
) ) from e
from typing_extensions import override from typing_extensions import override

View file

@ -22,17 +22,14 @@ from .components import (
# Export public symbols - maintaining backward compatibility # Export public symbols - maintaining backward compatibility
__all__ = [ __all__ = [
# Server
"FastMCPOpenAPI",
# Routing
"MCPType",
"RouteMap",
"RouteMapFn",
"ComponentFn",
"DEFAULT_ROUTE_MAPPINGS", "DEFAULT_ROUTE_MAPPINGS",
"_determine_route_type", "ComponentFn",
# Components "FastMCPOpenAPI",
"OpenAPITool", "MCPType",
"OpenAPIResource", "OpenAPIResource",
"OpenAPIResourceTemplate", "OpenAPIResourceTemplate",
"OpenAPITool",
"RouteMap",
"RouteMapFn",
"_determine_route_type",
] ]

View file

@ -146,11 +146,11 @@ class OpenAPITool(Tool):
if e.response.text: if e.response.text:
error_message += f" - {e.response.text}" error_message += f" - {e.response.text}"
raise ValueError(error_message) raise ValueError(error_message) from e
except httpx.RequestError as e: except httpx.RequestError as e:
# Handle request errors (connection, timeout, etc.) # Handle request errors (connection, timeout, etc.)
raise ValueError(f"Request error: {str(e)}") raise ValueError(f"Request error: {e!s}") from e
class OpenAPIResource(Resource): class OpenAPIResource(Resource):
@ -165,9 +165,11 @@ class OpenAPIResource(Resource):
name: str, name: str,
description: str, description: str,
mime_type: str = "application/json", mime_type: str = "application/json",
tags: set[str] = set(), tags: set[str] | None = None,
timeout: float | None = None, timeout: float | None = None,
): ):
if tags is None:
tags = set()
super().__init__( super().__init__(
uri=AnyUrl(uri), # Convert string to AnyUrl uri=AnyUrl(uri), # Convert string to AnyUrl
name=name, name=name,
@ -276,11 +278,11 @@ class OpenAPIResource(Resource):
if e.response.text: if e.response.text:
error_message += f" - {e.response.text}" error_message += f" - {e.response.text}"
raise ValueError(error_message) raise ValueError(error_message) from e
except httpx.RequestError as e: except httpx.RequestError as e:
# Handle request errors (connection, timeout, etc.) # Handle request errors (connection, timeout, etc.)
raise ValueError(f"Request error: {str(e)}") raise ValueError(f"Request error: {e!s}") from e
class OpenAPIResourceTemplate(ResourceTemplate): class OpenAPIResourceTemplate(ResourceTemplate):
@ -295,9 +297,11 @@ class OpenAPIResourceTemplate(ResourceTemplate):
name: str, name: str,
description: str, description: str,
parameters: dict[str, Any], parameters: dict[str, Any],
tags: set[str] = set(), tags: set[str] | None = None,
timeout: float | None = None, timeout: float | None = None,
): ):
if tags is None:
tags = set()
super().__init__( super().__init__(
uri_template=uri_template, uri_template=uri_template,
name=name, name=name,
@ -342,7 +346,7 @@ class OpenAPIResourceTemplate(ResourceTemplate):
# Export public symbols # Export public symbols
__all__ = [ __all__ = [
"OpenAPITool",
"OpenAPIResource", "OpenAPIResource",
"OpenAPIResourceTemplate", "OpenAPIResourceTemplate",
"OpenAPITool",
] ]

View file

@ -121,10 +121,10 @@ def _determine_route_type(
# Export public symbols # Export public symbols
__all__ = [ __all__ = [
"DEFAULT_ROUTE_MAPPINGS",
"ComponentFn",
"MCPType", "MCPType",
"RouteMap", "RouteMap",
"RouteMapFn", "RouteMapFn",
"ComponentFn",
"DEFAULT_ROUTE_MAPPINGS",
"_determine_route_type", "_determine_route_type",
] ]

View file

@ -40,29 +40,24 @@ from .json_schema_converter import (
# Export public symbols - maintaining backward compatibility # Export public symbols - maintaining backward compatibility
__all__ = [ __all__ = [
# Models
"HTTPRoute", "HTTPRoute",
"HttpMethod",
"JsonSchema",
"ParameterInfo", "ParameterInfo",
"ParameterLocation",
"RequestBodyInfo", "RequestBodyInfo",
"ResponseInfo", "ResponseInfo",
"HttpMethod", "_combine_schemas",
"ParameterLocation", "_make_optional_parameter_nullable",
"JsonSchema", "clean_schema_for_display",
# Parser "convert_openapi_schema_to_json_schema",
"parse_openapi_to_http_routes", "convert_schema_definitions",
# Formatters "extract_output_schema_from_responses",
"format_array_parameter", "format_array_parameter",
"format_deep_object_parameter", "format_deep_object_parameter",
"format_description_with_responses", "format_description_with_responses",
"format_json_for_description", "format_json_for_description",
"format_simple_description", "format_simple_description",
"generate_example_from_schema", "generate_example_from_schema",
# Schemas "parse_openapi_to_http_routes",
"_combine_schemas",
"extract_output_schema_from_responses",
"clean_schema_for_display",
"_make_optional_parameter_nullable",
# JSON Schema Converter
"convert_openapi_schema_to_json_schema",
"convert_schema_definitions",
] ]

View file

@ -63,7 +63,7 @@ class RequestDirector:
# Step 4: Handle request body # Step 4: Handle request body
if body is not None: if body is not None:
if isinstance(body, dict) or isinstance(body, list): if isinstance(body, dict | list):
request_data["json"] = body request_data["json"] = body
else: else:
request_data["content"] = body request_data["content"] = body

View file

@ -164,10 +164,10 @@ def _convert_nullable_field(schema: dict[str, Any]) -> dict[str, Any]:
if isinstance(current_type, str): if isinstance(current_type, str):
result["type"] = [current_type, "null"] result["type"] = [current_type, "null"]
elif isinstance(current_type, list) and "null" not in current_type: elif isinstance(current_type, list) and "null" not in current_type:
result["type"] = current_type + ["null"] result["type"] = [*current_type, "null"]
elif "oneOf" in result: elif "oneOf" in result:
# Convert oneOf to anyOf with null # Convert oneOf to anyOf with null
result["anyOf"] = result.pop("oneOf") + [{"type": "null"}] result["anyOf"] = [*result.pop("oneOf"), {"type": "null"}]
elif "anyOf" in result: elif "anyOf" in result:
# Add null to anyOf if not present # Add null to anyOf if not present
if not any(item.get("type") == "null" for item in result["anyOf"]): if not any(item.get("type") == "null" for item in result["anyOf"]):

View file

@ -79,10 +79,10 @@ class HTTPRoute(FastMCPBaseModel):
# Export public symbols # Export public symbols
__all__ = [ __all__ = [
"HTTPRoute", "HTTPRoute",
"HttpMethod",
"JsonSchema",
"ParameterInfo", "ParameterInfo",
"ParameterLocation",
"RequestBodyInfo", "RequestBodyInfo",
"ResponseInfo", "ResponseInfo",
"HttpMethod",
"ParameterLocation",
"JsonSchema",
] ]

View file

@ -178,7 +178,7 @@ class OpenAPIParser(
else: else:
# Special handling for components # Special handling for components
if part == "components" and hasattr(target, "components"): if part == "components" and hasattr(target, "components"):
target = getattr(target, "components") target = target.components
elif hasattr(target, part): # Fallback check elif hasattr(target, part): # Fallback check
target = getattr(target, part, None) target = getattr(target, part, None)
else: else:
@ -554,9 +554,7 @@ class OpenAPIParser(
if "$ref" in obj and isinstance(obj["$ref"], str): if "$ref" in obj and isinstance(obj["$ref"], str):
ref = obj["$ref"] ref = obj["$ref"]
# Handle both converted and unconverted refs # Handle both converted and unconverted refs
if ref.startswith("#/$defs/"): if ref.startswith(("#/$defs/", "#/components/schemas/")):
schema_name = ref.split("/")[-1]
elif ref.startswith("#/components/schemas/"):
schema_name = ref.split("/")[-1] schema_name = ref.split("/")[-1]
else: else:
return return
@ -815,6 +813,6 @@ class OpenAPIParser(
# Export public symbols # Export public symbols
__all__ = [ __all__ = [
"parse_openapi_to_http_routes",
"OpenAPIParser", "OpenAPIParser",
"parse_openapi_to_http_routes",
] ]

View file

@ -585,9 +585,9 @@ def extract_output_schema_from_responses(
# Export public symbols # Export public symbols
__all__ = [ __all__ = [
"clean_schema_for_display",
"_combine_schemas", "_combine_schemas",
"_combine_schemas_and_map_params", "_combine_schemas_and_map_params",
"extract_output_schema_from_responses",
"_make_optional_parameter_nullable", "_make_optional_parameter_nullable",
"clean_schema_for_display",
"extract_output_schema_from_responses",
] ]

View file

@ -288,9 +288,8 @@ class MCPConfig(BaseModel):
@classmethod @classmethod
def from_file(cls, file_path: Path) -> Self: def from_file(cls, file_path: Path) -> Self:
"""Load configuration from JSON file.""" """Load configuration from JSON file."""
if file_path.exists(): if file_path.exists() and (content := file_path.read_text().strip()):
if content := file_path.read_text().strip(): return cls.model_validate_json(content)
return cls.model_validate_json(content)
raise ValueError(f"No MCP servers defined in the config: {file_path}") raise ValueError(f"No MCP servers defined in the config: {file_path}")

View file

@ -2,8 +2,8 @@ from .prompt import Prompt, PromptMessage, Message
from .prompt_manager import PromptManager from .prompt_manager import PromptManager
__all__ = [ __all__ = [
"Message",
"Prompt", "Prompt",
"PromptManager", "PromptManager",
"PromptMessage", "PromptMessage",
"Message",
] ]

View file

@ -207,10 +207,7 @@ class FunctionPrompt(Prompt):
# Auto-detect context parameter if not provided # Auto-detect context parameter if not provided
context_kwarg = find_kwarg_by_type(fn, kwarg_type=Context) context_kwarg = find_kwarg_by_type(fn, kwarg_type=Context)
if context_kwarg: prune_params = [context_kwarg] if context_kwarg else None
prune_params = [context_kwarg]
else:
prune_params = None
parameters = compress_schema(parameters, prune_params=prune_params) parameters = compress_schema(parameters, prune_params=prune_params)
@ -290,10 +287,7 @@ class FunctionPrompt(Prompt):
if ( if (
param.annotation == inspect.Parameter.empty param.annotation == inspect.Parameter.empty
or param.annotation is str or param.annotation is str
): ) or not isinstance(param_value, str):
converted_kwargs[param_name] = param_value
# If argument is not a string, pass as-is (already properly typed)
elif not isinstance(param_value, str):
converted_kwargs[param_name] = param_value converted_kwargs[param_name] = param_value
else: else:
# Try to convert string argument using type adapter # Try to convert string argument using type adapter
@ -314,7 +308,7 @@ class FunctionPrompt(Prompt):
raise PromptError( raise PromptError(
f"Could not convert argument '{param_name}' with value '{param_value}' " f"Could not convert argument '{param_name}' with value '{param_value}' "
f"to expected type {param.annotation}. Error: {e}" f"to expected type {param.annotation}. Error: {e}"
) ) from e
else: else:
# Parameter not in function signature, pass as-is # Parameter not in function signature, pass as-is
converted_kwargs[param_name] = param_value converted_kwargs[param_name] = param_value
@ -376,10 +370,12 @@ class FunctionPrompt(Prompt):
content=TextContent(type="text", text=content), content=TextContent(type="text", text=content),
) )
) )
except Exception: except Exception as e:
raise PromptError("Could not convert prompt result to message.") raise PromptError(
"Could not convert prompt result to message."
) from e
return messages return messages
except Exception: except Exception as e:
logger.exception(f"Error rendering prompt {self.name}") logger.exception(f"Error rendering prompt {self.name}")
raise PromptError(f"Error rendering prompt {self.name}.") raise PromptError(f"Error rendering prompt {self.name}.") from e

View file

@ -10,13 +10,13 @@ from .types import (
from .resource_manager import ResourceManager from .resource_manager import ResourceManager
__all__ = [ __all__ = [
"Resource",
"TextResource",
"BinaryResource", "BinaryResource",
"FunctionResource",
"FileResource",
"HttpResource",
"DirectoryResource", "DirectoryResource",
"ResourceTemplate", "FileResource",
"FunctionResource",
"HttpResource",
"Resource",
"ResourceManager", "ResourceManager",
"ResourceTemplate",
"TextResource",
] ]

View file

@ -217,9 +217,7 @@ class FunctionResource(Resource):
if isinstance(result, Resource): if isinstance(result, Resource):
return await result.read() return await result.read()
elif isinstance(result, bytes): elif isinstance(result, bytes | str):
return result
elif isinstance(result, str):
return result return result
else: else:
return pydantic_core.to_json(result, fallback=str).decode() return pydantic_core.to_json(result, fallback=str).decode()

View file

@ -235,7 +235,7 @@ class ResourceManager:
# Then check templates (local and mounted) only if not found in concrete resources # Then check templates (local and mounted) only if not found in concrete resources
templates = await self.get_resource_templates() templates = await self.get_resource_templates()
for template_key in templates.keys(): for template_key in templates:
if match_uri_template(uri_str, template_key): if match_uri_template(uri_str, template_key):
return True return True

View file

@ -3,4 +3,4 @@ from .context import Context
from . import dependencies from . import dependencies
__all__ = ["FastMCP", "Context"] __all__ = ["Context", "FastMCP"]

View file

@ -10,14 +10,14 @@ from .oauth_proxy import OAuthProxy
__all__ = [ __all__ = [
"AuthProvider",
"OAuthProvider",
"TokenVerifier",
"JWTVerifier",
"StaticTokenVerifier",
"RemoteAuthProvider",
"AccessToken", "AccessToken",
"AuthProvider",
"JWTVerifier",
"OAuthProvider",
"OAuthProxy", "OAuthProxy",
"RemoteAuthProvider",
"StaticTokenVerifier",
"TokenVerifier",
] ]

View file

@ -23,7 +23,7 @@ from mcp.server.auth.settings import (
ClientRegistrationOptions, ClientRegistrationOptions,
RevocationOptions, RevocationOptions,
) )
from pydantic import AnyHttpUrl from pydantic import AnyHttpUrl, Field
from starlette.middleware import Middleware from starlette.middleware import Middleware
from starlette.middleware.authentication import AuthenticationMiddleware from starlette.middleware.authentication import AuthenticationMiddleware
from starlette.routing import Route from starlette.routing import Route
@ -32,7 +32,7 @@ from starlette.routing import Route
class AccessToken(_SDKAccessToken): class AccessToken(_SDKAccessToken):
"""AccessToken that includes all JWT claims.""" """AccessToken that includes all JWT claims."""
claims: dict[str, Any] = {} claims: dict[str, Any] = Field(default_factory=dict)
class AuthProvider(TokenVerifierProtocol): class AuthProvider(TokenVerifierProtocol):

View file

@ -123,10 +123,10 @@ class OIDCConfiguration(BaseModel):
try: try:
AnyHttpUrl(value) AnyHttpUrl(value)
except Exception: except Exception as e:
message = f"Invalid URL for configuration metadata: {attr}" message = f"Invalid URL for configuration metadata: {attr}"
logger.error(message) logger.error(message)
raise ValueError(message) raise ValueError(message) from e
enforce("issuer", True) enforce("issuer", True)
enforce("authorization_endpoint", True) enforce("authorization_endpoint", True)

View file

@ -113,12 +113,12 @@ class AzureProvider(OAuthProxy):
client_id: str | NotSetT = NotSet, client_id: str | NotSetT = NotSet,
client_secret: str | NotSetT = NotSet, client_secret: str | NotSetT = NotSet,
tenant_id: str | NotSetT = NotSet, tenant_id: str | NotSetT = NotSet,
identifier_uri: str | None | NotSetT = NotSet, identifier_uri: str | NotSetT | None = NotSet,
base_url: str | NotSetT = NotSet, base_url: str | NotSetT = NotSet,
issuer_url: str | NotSetT = NotSet, issuer_url: str | NotSetT = NotSet,
redirect_path: str | NotSetT = NotSet, redirect_path: str | NotSetT = NotSet,
required_scopes: list[str] | None | NotSetT = NotSet, required_scopes: list[str] | NotSetT | None = NotSet,
additional_authorize_scopes: list[str] | None | NotSetT = NotSet, additional_authorize_scopes: list[str] | NotSetT | None = NotSet,
allowed_client_redirect_uris: list[str] | NotSetT = NotSet, allowed_client_redirect_uris: list[str] | NotSetT = NotSet,
client_storage: AsyncKeyValue | None = None, client_storage: AsyncKeyValue | None = None,
jwt_signing_key: str | bytes | NotSetT = NotSet, jwt_signing_key: str | bytes | NotSetT = NotSet,

View file

@ -11,7 +11,7 @@ from fastmcp.server.auth.providers.jwt import JWKData, JWKSData, RSAKeyPair
from fastmcp.server.auth.providers.jwt import JWTVerifier as BearerAuthProvider from fastmcp.server.auth.providers.jwt import JWTVerifier as BearerAuthProvider
# Re-export for backwards compatibility # Re-export for backwards compatibility
__all__ = ["BearerAuthProvider", "RSAKeyPair", "JWKData", "JWKSData"] __all__ = ["BearerAuthProvider", "JWKData", "JWKSData", "RSAKeyPair"]
# Deprecated in 2.11 # Deprecated in 2.11
if fastmcp.settings.deprecation_warnings: if fastmcp.settings.deprecation_warnings:

View file

@ -96,10 +96,10 @@ class InMemoryOAuthProvider(OAuthProvider):
# or if params.redirect_uri is None and client has a default. # or if params.redirect_uri is None and client has a default.
# However, the AuthorizationHandler handles the primary validation. # However, the AuthorizationHandler handles the primary validation.
pass # Let's assume AuthorizationHandler did its job. pass # Let's assume AuthorizationHandler did its job.
except Exception: # Replace with specific validation error if client.validate_redirect_uri existed except Exception as e: # Replace with specific validation error if client.validate_redirect_uri existed
raise AuthorizeError( raise AuthorizeError(
error="invalid_request", error_description="Invalid redirect_uri." error="invalid_request", error_description="Invalid redirect_uri."
) ) from e
auth_code_value = f"test_auth_code_{secrets.token_hex(16)}" auth_code_value = f"test_auth_code_{secrets.token_hex(16)}"
expires_at = time.time() + DEFAULT_AUTH_CODE_EXPIRY_SECONDS expires_at = time.time() + DEFAULT_AUTH_CODE_EXPIRY_SECONDS

View file

@ -97,8 +97,8 @@ class IntrospectionTokenVerifier(TokenVerifier):
client_id: str | NotSetT = NotSet, client_id: str | NotSetT = NotSet,
client_secret: str | NotSetT = NotSet, client_secret: str | NotSetT = NotSet,
timeout_seconds: int | NotSetT = NotSet, timeout_seconds: int | NotSetT = NotSet,
required_scopes: list[str] | None | NotSetT = NotSet, required_scopes: list[str] | NotSetT | None = NotSet,
base_url: AnyHttpUrl | str | None | NotSetT = NotSet, base_url: AnyHttpUrl | str | NotSetT | None = NotSet,
): ):
""" """
Initialize the introspection token verifier. Initialize the introspection token verifier.

View file

@ -184,13 +184,13 @@ class JWTVerifier(TokenVerifier):
def __init__( def __init__(
self, self,
*, *,
public_key: str | None | NotSetT = NotSet, public_key: str | NotSetT | None = NotSet,
jwks_uri: str | None | NotSetT = NotSet, jwks_uri: str | NotSetT | None = NotSet,
issuer: str | None | NotSetT = NotSet, issuer: str | NotSetT | None = NotSet,
audience: str | list[str] | None | NotSetT = NotSet, audience: str | list[str] | NotSetT | None = NotSet,
algorithm: str | None | NotSetT = NotSet, algorithm: str | NotSetT | None = NotSet,
required_scopes: list[str] | None | NotSetT = NotSet, required_scopes: list[str] | NotSetT | None = NotSet,
base_url: AnyHttpUrl | str | None | NotSetT = NotSet, base_url: AnyHttpUrl | str | NotSetT | None = NotSet,
): ):
""" """
Initialize the JWT token verifier. Initialize the JWT token verifier.
@ -283,7 +283,7 @@ class JWTVerifier(TokenVerifier):
return await self._get_jwks_key(kid) return await self._get_jwks_key(kid)
except Exception as e: except Exception as e:
raise ValueError(f"Failed to extract key ID from token: {e}") raise ValueError(f"Failed to extract key ID from token: {e}") from e
async def _get_jwks_key(self, kid: str | None) -> str: async def _get_jwks_key(self, kid: str | None) -> str:
"""Fetch key from JWKS with simple caching.""" """Fetch key from JWKS with simple caching."""
@ -342,10 +342,10 @@ class JWTVerifier(TokenVerifier):
raise ValueError("No keys found in JWKS") raise ValueError("No keys found in JWKS")
except httpx.HTTPError as e: except httpx.HTTPError as e:
raise ValueError(f"Failed to fetch JWKS: {e}") raise ValueError(f"Failed to fetch JWKS: {e}") from e
except Exception as e: except Exception as e:
self.logger.debug(f"JWKS fetch failed: {e}") self.logger.debug(f"JWKS fetch failed: {e}")
raise ValueError(f"Failed to fetch JWKS: {e}") raise ValueError(f"Failed to fetch JWKS: {e}") from e
def _extract_scopes(self, claims: dict[str, Any]) -> list[str]: def _extract_scopes(self, claims: dict[str, Any]) -> list[str]:
""" """
@ -400,14 +400,13 @@ class JWTVerifier(TokenVerifier):
# Validate issuer - note we use issuer instead of issuer_url here because # Validate issuer - note we use issuer instead of issuer_url here because
# issuer is optional, allowing users to make this check optional # issuer is optional, allowing users to make this check optional
if self.issuer: if self.issuer and claims.get("iss") != self.issuer:
if claims.get("iss") != self.issuer: self.logger.debug(
self.logger.debug( "Token validation failed: issuer mismatch for client %s",
"Token validation failed: issuer mismatch for client %s", client_id,
client_id, )
) self.logger.info("Bearer token rejected for client %s", client_id)
self.logger.info("Bearer token rejected for client %s", client_id) return None
return None
# Validate audience if configured # Validate audience if configured
if self.audience: if self.audience:

View file

@ -83,7 +83,7 @@ class SupabaseProvider(RemoteAuthProvider):
*, *,
project_url: AnyHttpUrl | str | NotSetT = NotSet, project_url: AnyHttpUrl | str | NotSetT = NotSet,
base_url: AnyHttpUrl | str | NotSetT = NotSet, base_url: AnyHttpUrl | str | NotSetT = NotSet,
required_scopes: list[str] | None | NotSetT = NotSet, required_scopes: list[str] | NotSetT | None = NotSet,
token_verifier: TokenVerifier | None = None, token_verifier: TokenVerifier | None = None,
): ):
"""Initialize Supabase metadata provider. """Initialize Supabase metadata provider.

View file

@ -169,7 +169,7 @@ class WorkOSProvider(OAuthProxy):
base_url: AnyHttpUrl | str | NotSetT = NotSet, base_url: AnyHttpUrl | str | NotSetT = NotSet,
issuer_url: AnyHttpUrl | str | NotSetT = NotSet, issuer_url: AnyHttpUrl | str | NotSetT = NotSet,
redirect_path: str | NotSetT = NotSet, redirect_path: str | NotSetT = NotSet,
required_scopes: list[str] | None | NotSetT = NotSet, required_scopes: list[str] | NotSetT | None = NotSet,
timeout_seconds: int | NotSetT = NotSet, timeout_seconds: int | NotSetT = NotSet,
allowed_client_redirect_uris: list[str] | NotSetT = NotSet, allowed_client_redirect_uris: list[str] | NotSetT = NotSet,
client_storage: AsyncKeyValue | None = None, client_storage: AsyncKeyValue | None = None,
@ -338,7 +338,7 @@ class AuthKitProvider(RemoteAuthProvider):
*, *,
authkit_domain: AnyHttpUrl | str | NotSetT = NotSet, authkit_domain: AnyHttpUrl | str | NotSetT = NotSet,
base_url: AnyHttpUrl | str | NotSetT = NotSet, base_url: AnyHttpUrl | str | NotSetT = NotSet,
required_scopes: list[str] | None | NotSetT = NotSet, required_scopes: list[str] | NotSetT | None = NotSet,
token_verifier: TokenVerifier | None = None, token_verifier: TokenVerifier | None = None,
): ):
"""Initialize AuthKit metadata provider. """Initialize AuthKit metadata provider.

View file

@ -188,8 +188,8 @@ class Context:
""" """
try: try:
return request_ctx.get() return request_ctx.get()
except LookupError: except LookupError as e:
raise ValueError("Context is not available outside of a request") raise ValueError("Context is not available outside of a request") from e
async def report_progress( async def report_progress(
self, progress: float, total: float | None = None, message: str | None = None self, progress: float, total: float | None = None, message: str | None = None
@ -342,7 +342,7 @@ class Context:
session_id = str(uuid4()) session_id = str(uuid4())
# Save the session id to the session attributes # Save the session id to the session attributes
setattr(session, "_fastmcp_id", session_id) session._fastmcp_id = session_id
return session_id return session_id
@property @property
@ -595,13 +595,11 @@ class Context:
choice_literal = Literal[tuple(response_type)] # type: ignore choice_literal = Literal[tuple(response_type)] # type: ignore
response_type = ScalarElicitationType[choice_literal] # type: ignore response_type = ScalarElicitationType[choice_literal] # type: ignore
# if the user provided a primitive scalar, wrap it in an object schema # if the user provided a primitive scalar, wrap it in an object schema
elif response_type in {bool, int, float, str}: elif (
response_type = ScalarElicitationType[response_type] # type: ignore response_type in {bool, int, float, str}
# if the user provided a Literal type, wrap it in an object schema or get_origin(response_type) is Literal
elif get_origin(response_type) is Literal: or (isinstance(response_type, type) and issubclass(response_type, Enum))
response_type = ScalarElicitationType[response_type] # type: ignore ):
# if the user provided an Enum type, wrap it in an object schema
elif isinstance(response_type, type) and issubclass(response_type, Enum):
response_type = ScalarElicitationType[response_type] # type: ignore response_type = ScalarElicitationType[response_type] # type: ignore
response_type = cast(type[T], response_type) response_type = cast(type[T], response_type)

View file

@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import contextlib
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from mcp.server.auth.middleware.auth_context import ( from mcp.server.auth.middleware.auth_context import (
@ -16,11 +17,11 @@ if TYPE_CHECKING:
from fastmcp.server.context import Context from fastmcp.server.context import Context
__all__ = [ __all__ = [
"get_context",
"get_http_request",
"get_http_headers",
"get_access_token",
"AccessToken", "AccessToken",
"get_access_token",
"get_context",
"get_http_headers",
"get_http_request",
] ]
@ -43,10 +44,8 @@ def get_http_request() -> Request:
from mcp.server.lowlevel.server import request_ctx from mcp.server.lowlevel.server import request_ctx
request = None request = None
try: with contextlib.suppress(LookupError):
request = request_ctx.get().request request = request_ctx.get().request
except LookupError:
pass
if request is None: if request is None:
raise RuntimeError("No active HTTP request found.") raise RuntimeError("No active HTTP request found.")

View file

@ -20,8 +20,8 @@ __all__ = [
"AcceptedElicitation", "AcceptedElicitation",
"CancelledElicitation", "CancelledElicitation",
"DeclinedElicitation", "DeclinedElicitation",
"get_elicitation_schema",
"ScalarElicitationType", "ScalarElicitationType",
"get_elicitation_schema",
] ]
logger = get_logger(__name__) logger = get_logger(__name__)

View file

@ -342,9 +342,8 @@ def create_streamable_http_app(
# Create a lifespan manager to start and stop the session manager # Create a lifespan manager to start and stop the session manager
@asynccontextmanager @asynccontextmanager
async def lifespan(app: Starlette) -> AsyncGenerator[None, None]: async def lifespan(app: Starlette) -> AsyncGenerator[None, None]:
async with server._lifespan_manager(): async with server._lifespan_manager(), session_manager.run():
async with session_manager.run(): yield
yield
# Create and return the app with lifespan # Create and return the app with lifespan
app = create_base_app( app = create_base_app(

View file

@ -5,7 +5,7 @@ from .middleware import (
) )
__all__ = [ __all__ = [
"CallNext",
"Middleware", "Middleware",
"MiddlewareContext", "MiddlewareContext",
"CallNext",
] ]

View file

@ -46,7 +46,7 @@ class CachableReadResourceContents(BaseModel):
@classmethod @classmethod
def get_sizes(cls, values: Sequence[Self]) -> int: def get_sizes(cls, values: Sequence[Self]) -> int:
return sum([item.get_size() for item in values]) return sum(item.get_size() for item in values)
@classmethod @classmethod
def wrap(cls, values: Sequence[ReadResourceContents]) -> list[Self]: def wrap(cls, values: Sequence[ReadResourceContents]) -> list[Self]:

View file

@ -64,7 +64,7 @@ class ErrorHandlingMiddleware(Middleware):
error_key = f"{error_type}:{method}" error_key = f"{error_type}:{method}"
self.error_counts[error_key] = self.error_counts.get(error_key, 0) + 1 self.error_counts[error_key] = self.error_counts.get(error_key, 0) + 1
base_message = f"Error in {method}: {error_type}: {str(error)}" base_message = f"Error in {method}: {error_type}: {error!s}"
if self.include_traceback: if self.include_traceback:
self.logger.error(f"{base_message}\n{traceback.format_exc()}") self.logger.error(f"{base_message}\n{traceback.format_exc()}")
@ -91,24 +91,24 @@ class ErrorHandlingMiddleware(Middleware):
if error_type in (ValueError, TypeError): if error_type in (ValueError, TypeError):
return McpError( return McpError(
ErrorData(code=-32602, message=f"Invalid params: {str(error)}") ErrorData(code=-32602, message=f"Invalid params: {error!s}")
) )
elif error_type in (FileNotFoundError, KeyError, NotFoundError): elif error_type in (FileNotFoundError, KeyError, NotFoundError):
return McpError( return McpError(
ErrorData(code=-32001, message=f"Resource not found: {str(error)}") ErrorData(code=-32001, message=f"Resource not found: {error!s}")
) )
elif error_type is PermissionError: elif error_type is PermissionError:
return McpError( return McpError(
ErrorData(code=-32000, message=f"Permission denied: {str(error)}") ErrorData(code=-32000, message=f"Permission denied: {error!s}")
) )
# asyncio.TimeoutError is a subclass of TimeoutError in Python 3.10, alias in 3.11+ # asyncio.TimeoutError is a subclass of TimeoutError in Python 3.10, alias in 3.11+
elif error_type in (TimeoutError, asyncio.TimeoutError): elif error_type in (TimeoutError, asyncio.TimeoutError):
return McpError( return McpError(
ErrorData(code=-32000, message=f"Request timeout: {str(error)}") ErrorData(code=-32000, message=f"Request timeout: {error!s}")
) )
else: else:
return McpError( return McpError(
ErrorData(code=-32603, message=f"Internal error: {str(error)}") ErrorData(code=-32603, message=f"Internal error: {error!s}")
) )
async def on_message(self, context: MiddlewareContext, call_next: CallNext) -> Any: async def on_message(self, context: MiddlewareContext, call_next: CallNext) -> Any:
@ -120,7 +120,7 @@ class ErrorHandlingMiddleware(Middleware):
# Transform and re-raise # Transform and re-raise
transformed_error = self._transform_error(error) transformed_error = self._transform_error(error)
raise transformed_error raise transformed_error from error
def get_error_stats(self) -> dict[str, int]: def get_error_stats(self) -> dict[str, int]:
"""Get error statistics for monitoring.""" """Get error statistics for monitoring."""
@ -200,7 +200,7 @@ class RetryMiddleware(Middleware):
delay = self._calculate_delay(attempt) delay = self._calculate_delay(attempt)
self.logger.warning( self.logger.warning(
f"Request {context.method} failed (attempt {attempt + 1}/{self.max_retries + 1}): " f"Request {context.method} failed (attempt {attempt + 1}/{self.max_retries + 1}): "
f"{type(error).__name__}: {str(error)}. Retrying in {delay:.1f}s..." f"{type(error).__name__}: {error!s}. Retrying in {delay:.1f}s..."
) )
await anyio.sleep(delay) await anyio.sleep(delay)

View file

@ -27,9 +27,9 @@ if TYPE_CHECKING:
from fastmcp.server.context import Context from fastmcp.server.context import Context
__all__ = [ __all__ = [
"CallNext",
"Middleware", "Middleware",
"MiddlewareContext", "MiddlewareContext",
"CallNext",
] ]
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)

View file

@ -513,11 +513,11 @@ class OpenAPITool(Tool):
if e.response.text: if e.response.text:
error_message += f" - {e.response.text}" error_message += f" - {e.response.text}"
raise ValueError(error_message) raise ValueError(error_message) from e
except httpx.RequestError as e: except httpx.RequestError as e:
# Handle request errors (connection, timeout, etc.) # Handle request errors (connection, timeout, etc.)
raise ValueError(f"Request error: {str(e)}") raise ValueError(f"Request error: {e!s}") from e
class OpenAPIResource(Resource): class OpenAPIResource(Resource):
@ -531,9 +531,11 @@ class OpenAPIResource(Resource):
name: str, name: str,
description: str, description: str,
mime_type: str = "application/json", mime_type: str = "application/json",
tags: set[str] = set(), tags: set[str] | None = None,
timeout: float | None = None, timeout: float | None = None,
): ):
if tags is None:
tags = set()
super().__init__( super().__init__(
uri=AnyUrl(uri), # Convert string to AnyUrl uri=AnyUrl(uri), # Convert string to AnyUrl
name=name, name=name,
@ -632,11 +634,11 @@ class OpenAPIResource(Resource):
if e.response.text: if e.response.text:
error_message += f" - {e.response.text}" error_message += f" - {e.response.text}"
raise ValueError(error_message) raise ValueError(error_message) from e
except httpx.RequestError as e: except httpx.RequestError as e:
# Handle request errors (connection, timeout, etc.) # Handle request errors (connection, timeout, etc.)
raise ValueError(f"Request error: {str(e)}") raise ValueError(f"Request error: {e!s}") from e
class OpenAPIResourceTemplate(ResourceTemplate): class OpenAPIResourceTemplate(ResourceTemplate):
@ -650,9 +652,11 @@ class OpenAPIResourceTemplate(ResourceTemplate):
name: str, name: str,
description: str, description: str,
parameters: dict[str, Any], parameters: dict[str, Any],
tags: set[str] = set(), tags: set[str] | None = None,
timeout: float | None = None, timeout: float | None = None,
): ):
if tags is None:
tags = set()
super().__init__( super().__init__(
uri_template=uri_template, uri_template=uri_template,
name=name, name=name,

View file

@ -198,7 +198,9 @@ class ProxyResourceManager(ResourceManager, ProxyManagerMixin):
elif isinstance(result[0], BlobResourceContents): elif isinstance(result[0], BlobResourceContents):
return result[0].blob return result[0].blob
else: else:
raise ResourceError(f"Unsupported content type: {type(result[0])}") raise ResourceError(
f"Unsupported content type: {type(result[0])}"
) from None
class ProxyPromptManager(PromptManager, ProxyManagerMixin): class ProxyPromptManager(PromptManager, ProxyManagerMixin):
@ -558,7 +560,7 @@ class ProxyClient(Client[ClientTransportT]):
kwargs["log_handler"] = ProxyClient.default_log_handler kwargs["log_handler"] = ProxyClient.default_log_handler
if "progress_handler" not in kwargs: if "progress_handler" not in kwargs:
kwargs["progress_handler"] = ProxyClient.default_progress_handler kwargs["progress_handler"] = ProxyClient.default_progress_handler
super().__init__(**kwargs | dict(transport=transport)) super().__init__(**kwargs | {"transport": transport})
@classmethod @classmethod
async def default_sampling_handler( async def default_sampling_handler(
@ -572,7 +574,7 @@ class ProxyClient(Client[ClientTransportT]):
""" """
ctx = get_context() ctx = get_context()
content = await ctx.sample( content = await ctx.sample(
[msg for msg in messages], list(messages),
system_prompt=params.systemPrompt, system_prompt=params.systemPrompt,
temperature=params.temperature, temperature=params.temperature,
max_tokens=params.maxTokens, max_tokens=params.maxTokens,
@ -649,7 +651,6 @@ class StatefulProxyClient(ProxyClient[ClientTransportT]):
The stateful proxy client will be forced disconnected when the session is exited. The stateful proxy client will be forced disconnected when the session is exited.
So we do nothing here. So we do nothing here.
""" """
pass
async def clear(self): async def clear(self):
""" """

View file

@ -15,7 +15,11 @@ from collections.abc import (
Mapping, Mapping,
Sequence, Sequence,
) )
from contextlib import AbstractAsyncContextManager, AsyncExitStack, asynccontextmanager from contextlib import (
AbstractAsyncContextManager,
AsyncExitStack,
asynccontextmanager,
)
from dataclasses import dataclass from dataclasses import dataclass
from functools import partial from functools import partial
from pathlib import Path from pathlib import Path
@ -150,7 +154,7 @@ class FastMCP(Generic[LifespanResultT]):
version: str | None = None, version: str | None = None,
website_url: str | None = None, website_url: str | None = None,
icons: list[mcp.types.Icon] | None = None, icons: list[mcp.types.Icon] | None = None,
auth: AuthProvider | None | NotSetT = NotSet, auth: AuthProvider | NotSetT | None = NotSet,
middleware: Sequence[Middleware] | None = None, middleware: Sequence[Middleware] | None = None,
lifespan: LifespanCallable | None = None, lifespan: LifespanCallable | None = None,
dependencies: list[str] | None = None, dependencies: list[str] | None = None,
@ -1062,10 +1066,10 @@ class FastMCP(Generic[LifespanResultT]):
try: try:
result = await self._call_tool_middleware(key, arguments) result = await self._call_tool_middleware(key, arguments)
return result.to_mcp_result() return result.to_mcp_result()
except DisabledError: except DisabledError as e:
raise NotFoundError(f"Unknown tool: {key}") raise NotFoundError(f"Unknown tool: {key}") from e
except NotFoundError: except NotFoundError as e:
raise NotFoundError(f"Unknown tool: {key}") raise NotFoundError(f"Unknown tool: {key}") from e
async def _call_tool_middleware( async def _call_tool_middleware(
self, self,
@ -1142,12 +1146,12 @@ class FastMCP(Generic[LifespanResultT]):
return list[ReadResourceContents]( return list[ReadResourceContents](
await self._read_resource_middleware(uri) await self._read_resource_middleware(uri)
) )
except DisabledError: except DisabledError as e:
# convert to NotFoundError to avoid leaking resource presence # convert to NotFoundError to avoid leaking resource presence
raise NotFoundError(f"Unknown resource: {str(uri)!r}") raise NotFoundError(f"Unknown resource: {str(uri)!r}") from e
except NotFoundError: except NotFoundError as e:
# standardize NotFound message # standardize NotFound message
raise NotFoundError(f"Unknown resource: {str(uri)!r}") raise NotFoundError(f"Unknown resource: {str(uri)!r}") from e
async def _read_resource_middleware( async def _read_resource_middleware(
self, self,
@ -1158,10 +1162,7 @@ class FastMCP(Generic[LifespanResultT]):
""" """
# Convert string URI to AnyUrl if needed # Convert string URI to AnyUrl if needed
if isinstance(uri, str): uri_param = AnyUrl(uri) if isinstance(uri, str) else uri
uri_param = AnyUrl(uri)
else:
uri_param = uri
mw_context = MiddlewareContext( mw_context = MiddlewareContext(
message=mcp.types.ReadResourceRequestParams(uri=uri_param), message=mcp.types.ReadResourceRequestParams(uri=uri_param),
@ -1241,12 +1242,12 @@ class FastMCP(Generic[LifespanResultT]):
async with fastmcp.server.context.Context(fastmcp=self): async with fastmcp.server.context.Context(fastmcp=self):
try: try:
return await self._get_prompt_middleware(name, arguments) return await self._get_prompt_middleware(name, arguments)
except DisabledError: except DisabledError as e:
# convert to NotFoundError to avoid leaking prompt presence # convert to NotFoundError to avoid leaking prompt presence
raise NotFoundError(f"Unknown prompt: {name}") raise NotFoundError(f"Unknown prompt: {name}") from e
except NotFoundError: except NotFoundError as e:
# standardize NotFound message # standardize NotFound message
raise NotFoundError(f"Unknown prompt: {name}") raise NotFoundError(f"Unknown prompt: {name}") from e
async def _get_prompt_middleware( async def _get_prompt_middleware(
self, name: str, arguments: dict[str, Any] | None = None self, name: str, arguments: dict[str, Any] | None = None
@ -1369,7 +1370,7 @@ class FastMCP(Generic[LifespanResultT]):
description: str | None = None, description: str | None = None,
icons: list[mcp.types.Icon] | None = None, icons: list[mcp.types.Icon] | None = None,
tags: set[str] | None = None, tags: set[str] | None = None,
output_schema: dict[str, Any] | None | NotSetT = NotSet, output_schema: dict[str, Any] | NotSetT | None = NotSet,
annotations: ToolAnnotations | dict[str, Any] | None = None, annotations: ToolAnnotations | dict[str, Any] | None = None,
exclude_args: list[str] | None = None, exclude_args: list[str] | None = None,
meta: dict[str, Any] | None = None, meta: dict[str, Any] | None = None,
@ -1386,7 +1387,7 @@ class FastMCP(Generic[LifespanResultT]):
description: str | None = None, description: str | None = None,
icons: list[mcp.types.Icon] | None = None, icons: list[mcp.types.Icon] | None = None,
tags: set[str] | None = None, tags: set[str] | None = None,
output_schema: dict[str, Any] | None | NotSetT = NotSet, output_schema: dict[str, Any] | NotSetT | None = NotSet,
annotations: ToolAnnotations | dict[str, Any] | None = None, annotations: ToolAnnotations | dict[str, Any] | None = None,
exclude_args: list[str] | None = None, exclude_args: list[str] | None = None,
meta: dict[str, Any] | None = None, meta: dict[str, Any] | None = None,
@ -1402,7 +1403,7 @@ class FastMCP(Generic[LifespanResultT]):
description: str | None = None, description: str | None = None,
icons: list[mcp.types.Icon] | None = None, icons: list[mcp.types.Icon] | None = None,
tags: set[str] | None = None, tags: set[str] | None = None,
output_schema: dict[str, Any] | None | NotSetT = NotSet, output_schema: dict[str, Any] | NotSetT | None = NotSet,
annotations: ToolAnnotations | dict[str, Any] | None = None, annotations: ToolAnnotations | dict[str, Any] | None = None,
exclude_args: list[str] | None = None, exclude_args: list[str] | None = None,
meta: dict[str, Any] | None = None, meta: dict[str, Any] | None = None,
@ -2029,14 +2030,14 @@ class FastMCP(Generic[LifespanResultT]):
port=port, port=port,
path=server_path, path=server_path,
) )
_uvicorn_config_from_user = uvicorn_config or {} uvicorn_config_from_user = uvicorn_config or {}
config_kwargs: dict[str, Any] = { config_kwargs: dict[str, Any] = {
"timeout_graceful_shutdown": 0, "timeout_graceful_shutdown": 0,
"lifespan": "on", "lifespan": "on",
"ws": "websockets-sansio", "ws": "websockets-sansio",
} }
config_kwargs.update(_uvicorn_config_from_user) config_kwargs.update(uvicorn_config_from_user)
if "log_config" not in config_kwargs and "log_level" not in config_kwargs: if "log_config" not in config_kwargs and "log_level" not in config_kwargs:
config_kwargs["log_level"] = default_log_level_to_use config_kwargs["log_level"] = default_log_level_to_use
@ -2605,8 +2606,8 @@ class FastMCP(Generic[LifespanResultT]):
# - Connected clients: reuse existing session for all requests # - Connected clients: reuse existing session for all requests
# - Disconnected clients: create fresh sessions per request for isolation # - Disconnected clients: create fresh sessions per request for isolation
if client.is_connected(): if client.is_connected():
_proxy_logger = get_logger(__name__) proxy_logger = get_logger(__name__)
_proxy_logger.info( proxy_logger.info(
"Proxy detected connected client - reusing existing session for all requests. " "Proxy detected connected client - reusing existing session for all requests. "
"This may cause context mixing in concurrent scenarios." "This may cause context mixing in concurrent scenarios."
) )
@ -2678,10 +2679,7 @@ class FastMCP(Generic[LifespanResultT]):
return False return False
if self.include_tags is not None: if self.include_tags is not None:
if any(itag in component.tags for itag in self.include_tags): return bool(any(itag in component.tags for itag in self.include_tags))
return True
else:
return False
return True return True

View file

@ -2,4 +2,4 @@ from .tool import Tool, FunctionTool
from .tool_manager import ToolManager from .tool_manager import ToolManager
from .tool_transform import forward, forward_raw from .tool_transform import forward, forward_raw
__all__ = ["Tool", "ToolManager", "FunctionTool", "forward", "forward_raw"] __all__ = ["FunctionTool", "Tool", "ToolManager", "forward", "forward_raw"]

View file

@ -173,7 +173,7 @@ class Tool(FastMCPComponent):
tags: set[str] | None = None, tags: set[str] | None = None,
annotations: ToolAnnotations | None = None, annotations: ToolAnnotations | None = None,
exclude_args: list[str] | None = None, exclude_args: list[str] | None = None,
output_schema: dict[str, Any] | None | NotSetT | Literal[False] = NotSet, output_schema: dict[str, Any] | Literal[False] | NotSetT | None = NotSet,
serializer: ToolResultSerializerType | None = None, serializer: ToolResultSerializerType | None = None,
meta: dict[str, Any] | None = None, meta: dict[str, Any] | None = None,
enabled: bool | None = None, enabled: bool | None = None,
@ -212,13 +212,13 @@ class Tool(FastMCPComponent):
tool: Tool, tool: Tool,
*, *,
name: str | None = None, name: str | None = None,
title: str | None | NotSetT = NotSet, title: str | NotSetT | None = NotSet,
description: str | None | NotSetT = NotSet, description: str | NotSetT | None = NotSet,
tags: set[str] | None = None, tags: set[str] | None = None,
annotations: ToolAnnotations | None | NotSetT = NotSet, annotations: ToolAnnotations | NotSetT | None = NotSet,
output_schema: dict[str, Any] | None | NotSetT | Literal[False] = NotSet, output_schema: dict[str, Any] | Literal[False] | NotSetT | None = NotSet,
serializer: ToolResultSerializerType | None = None, serializer: ToolResultSerializerType | None = None,
meta: dict[str, Any] | None | NotSetT = NotSet, meta: dict[str, Any] | NotSetT | None = NotSet,
transform_args: dict[str, ArgTransform] | None = None, transform_args: dict[str, ArgTransform] | None = None,
enabled: bool | None = None, enabled: bool | None = None,
transform_fn: Callable[..., Any] | None = None, transform_fn: Callable[..., Any] | None = None,
@ -255,7 +255,7 @@ class FunctionTool(Tool):
tags: set[str] | None = None, tags: set[str] | None = None,
annotations: ToolAnnotations | None = None, annotations: ToolAnnotations | None = None,
exclude_args: list[str] | None = None, exclude_args: list[str] | None = None,
output_schema: dict[str, Any] | None | NotSetT | Literal[False] = NotSet, output_schema: dict[str, Any] | Literal[False] | NotSetT | None = NotSet,
serializer: ToolResultSerializerType | None = None, serializer: ToolResultSerializerType | None = None,
meta: dict[str, Any] | None = None, meta: dict[str, Any] | None = None,
enabled: bool | None = None, enabled: bool | None = None,
@ -446,9 +446,8 @@ class ParsedFunction:
# we ensure that no output schema is automatically generated. # we ensure that no output schema is automatically generated.
clean_output_type = replace_type( clean_output_type = replace_type(
output_type, output_type,
{ dict.fromkeys( # type: ignore[arg-type]
t: _UnserializableType (
for t in (
Image, Image,
Audio, Audio,
File, File,
@ -458,8 +457,9 @@ class ParsedFunction:
mcp.types.AudioContent, mcp.types.AudioContent,
mcp.types.ResourceLink, mcp.types.ResourceLink,
mcp.types.EmbeddedResource, mcp.types.EmbeddedResource,
) ),
}, _UnserializableType,
),
) )
try: try:

View file

@ -365,15 +365,15 @@ class TransformedTool(Tool):
cls, cls,
tool: Tool, tool: Tool,
name: str | None = None, name: str | None = None,
title: str | None | NotSetT = NotSet, title: str | NotSetT | None = NotSet,
description: str | None | NotSetT = NotSet, description: str | NotSetT | None = NotSet,
tags: set[str] | None = None, tags: set[str] | None = None,
transform_fn: Callable[..., Any] | None = None, transform_fn: Callable[..., Any] | None = None,
transform_args: dict[str, ArgTransform] | None = None, transform_args: dict[str, ArgTransform] | None = None,
annotations: ToolAnnotations | None | NotSetT = NotSet, annotations: ToolAnnotations | NotSetT | None = NotSet,
output_schema: dict[str, Any] | None | NotSetT | Literal[False] = NotSet, output_schema: dict[str, Any] | Literal[False] | NotSetT | None = NotSet,
serializer: Callable[[Any], str] | None | NotSetT = NotSet, serializer: Callable[[Any], str] | NotSetT | None = NotSet,
meta: dict[str, Any] | None | NotSetT = NotSet, meta: dict[str, Any] | NotSetT | None = NotSet,
enabled: bool | None = None, enabled: bool | None = None,
) -> TransformedTool: ) -> TransformedTool:
"""Create a transformed tool from a parent tool. """Create a transformed tool from a parent tool.

View file

@ -240,12 +240,11 @@ def log_server_banner(
info_table.add_row("📦", "Transport:", display_transport) info_table.add_row("📦", "Transport:", display_transport)
# Show connection info based on transport # Show connection info based on transport
if transport in ("http", "streamable-http", "sse"): if transport in ("http", "streamable-http", "sse") and host and port:
if host and port: server_url = f"http://{host}:{port}"
server_url = f"http://{host}:{port}" if path:
if path: server_url += f"/{path.lstrip('/')}"
server_url += f"/{path.lstrip('/')}" info_table.add_row("🔗", "Server URL:", server_url)
info_table.add_row("🔗", "Server URL:", server_url)
# Add documentation link # Add documentation link
info_table.add_row("", "", "") info_table.add_row("", "", "")

View file

@ -412,7 +412,7 @@ class InspectFormat(str, Enum):
MCP = "mcp" MCP = "mcp"
async def format_fastmcp_info(info: FastMCPInfo) -> bytes: def format_fastmcp_info(info: FastMCPInfo) -> bytes:
"""Format FastMCPInfo as FastMCP-specific JSON. """Format FastMCPInfo as FastMCP-specific JSON.
This includes FastMCP-specific fields like tags, enabled, annotations, etc. This includes FastMCP-specific fields like tags, enabled, annotations, etc.
@ -501,6 +501,6 @@ async def format_info(
# This works for both v1 and v2 servers # This works for both v1 and v2 servers
if info is None: if info is None:
info = await inspect_fastmcp(mcp) info = await inspect_fastmcp(mcp)
return await format_fastmcp_info(info) return format_fastmcp_info(info)
else: else:
raise ValueError(f"Unknown format: {format}") raise ValueError(f"Unknown format: {format}")

View file

@ -61,7 +61,7 @@ from pydantic import (
) )
from typing_extensions import NotRequired, TypedDict from typing_extensions import NotRequired, TypedDict
__all__ = ["json_schema_to_type", "JSONSchema"] __all__ = ["JSONSchema", "json_schema_to_type"]
FORMAT_TYPES: dict[str, Any] = { FORMAT_TYPES: dict[str, Any] = {
@ -368,7 +368,7 @@ def _schema_to_type(
return types[0] return types[0]
else: else:
if has_null: if has_null:
return Union[tuple(types + [type(None)])] # type: ignore # noqa: UP007 return Union[(*types, type(None))] # type: ignore
else: else:
return Union[tuple(types)] # type: ignore # noqa: UP007 return Union[tuple(types)] # type: ignore # noqa: UP007
@ -389,7 +389,7 @@ def _schema_to_type(
if len(types) == 1: if len(types) == 1:
return types[0] | None # type: ignore return types[0] | None # type: ignore
else: else:
return Union[tuple(types + [type(None)])] # type: ignore # noqa: UP007 return Union[(*types, type(None))] # type: ignore
return Union[tuple(types)] # type: ignore # noqa: UP007 return Union[tuple(types)] # type: ignore # noqa: UP007
return _get_from_type_handler(schema, schemas)(schema) return _get_from_type_handler(schema, schemas)(schema)
@ -578,7 +578,7 @@ def _create_dataclass(
return _merge_defaults(data, original_schema) return _merge_defaults(data, original_schema)
return data return data
setattr(cls, "_apply_defaults", _apply_defaults) cls._apply_defaults = _apply_defaults # type: ignore[attr-defined]
# Store completed class # Store completed class
_classes[cache_key] = cls _classes[cache_key] = cls

View file

@ -147,6 +147,18 @@ def temporary_log_level(
yield yield
_level_to_no: dict[
Literal["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"] | None, int | None
] = {
"DEBUG": logging.DEBUG,
"INFO": logging.INFO,
"WARNING": logging.WARNING,
"ERROR": logging.ERROR,
"CRITICAL": logging.CRITICAL,
None: None,
}
class _ClampedLogFilter(logging.Filter): class _ClampedLogFilter(logging.Filter):
min_level: tuple[int, str] | None min_level: tuple[int, str] | None
max_level: tuple[int, str] | None max_level: tuple[int, str] | None
@ -161,29 +173,13 @@ class _ClampedLogFilter(logging.Filter):
self.min_level = None self.min_level = None
self.max_level = None self.max_level = None
if min_level_no := self._level_to_no(level=min_level): if min_level_no := _level_to_no.get(min_level):
self.min_level = (min_level_no, str(min_level)) self.min_level = (min_level_no, str(min_level))
if max_level_no := self._level_to_no(level=max_level): if max_level_no := _level_to_no.get(max_level):
self.max_level = (max_level_no, str(max_level)) self.max_level = (max_level_no, str(max_level))
super().__init__() super().__init__()
def _level_to_no(
self, level: Literal["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"] | None
) -> int | None:
if level == "DEBUG":
return logging.DEBUG
elif level == "INFO":
return logging.INFO
elif level == "WARNING":
return logging.WARNING
elif level == "ERROR":
return logging.ERROR
elif level == "CRITICAL":
return logging.CRITICAL
else:
return None
@override @override
def filter(self, record: logging.LogRecord) -> bool: def filter(self, record: logging.LogRecord) -> bool:
if self.max_level: if self.max_level:

View file

@ -15,11 +15,11 @@ from fastmcp.utilities.mcp_server_config.v1.sources.base import Source
from fastmcp.utilities.mcp_server_config.v1.sources.filesystem import FileSystemSource from fastmcp.utilities.mcp_server_config.v1.sources.filesystem import FileSystemSource
__all__ = [ __all__ = [
"Source",
"Deployment", "Deployment",
"Environment", "Environment",
"UVEnvironment",
"MCPServerConfig",
"FileSystemSource", "FileSystemSource",
"MCPServerConfig",
"Source",
"UVEnvironment",
"generate_schema", "generate_schema",
] ]

View file

@ -19,7 +19,6 @@ class Environment(BaseModel, ABC):
Returns: Returns:
Full command ready for subprocess execution Full command ready for subprocess execution
""" """
pass
async def prepare(self, output_dir: Path | None = None) -> None: async def prepare(self, output_dir: Path | None = None) -> None:
"""Prepare the environment (optional, can be no-op). """Prepare the environment (optional, can be no-op).
@ -27,4 +26,4 @@ class Environment(BaseModel, ABC):
Args: Args:
output_dir: Directory for persistent environment setup output_dir: Directory for persistent environment setup
""" """
pass # Default no-op implementation # Default no-op implementation

View file

@ -17,7 +17,6 @@ class Source(BaseModel, ABC):
need preparation (e.g., local files), this is a no-op. need preparation (e.g., local files), this is a no-op.
""" """
# Default implementation for sources that don't need preparation # Default implementation for sources that don't need preparation
pass
@abstractmethod @abstractmethod
async def load_server(self) -> Any: async def load_server(self) -> Any:

View file

@ -175,16 +175,16 @@ class HTTPRoute(FastMCPBaseModel):
# Export public symbols # Export public symbols
__all__ = [ __all__ = [
"HTTPRoute", "HTTPRoute",
"HttpMethod",
"JsonSchema",
"ParameterInfo", "ParameterInfo",
"ParameterLocation",
"RequestBodyInfo", "RequestBodyInfo",
"ResponseInfo", "ResponseInfo",
"HttpMethod", "_handle_nullable_fields",
"ParameterLocation",
"JsonSchema",
"parse_openapi_to_http_routes",
"extract_output_schema_from_responses", "extract_output_schema_from_responses",
"format_deep_object_parameter", "format_deep_object_parameter",
"_handle_nullable_fields", "parse_openapi_to_http_routes",
] ]
# Type variables for generic parser # Type variables for generic parser
@ -321,7 +321,7 @@ class OpenAPIParser(
else: else:
# Special handling for components # Special handling for components
if part == "components" and hasattr(target, "components"): if part == "components" and hasattr(target, "components"):
target = getattr(target, "components") target = target.components
elif hasattr(target, part): # Fallback check elif hasattr(target, part): # Fallback check
target = getattr(target, part, None) target = getattr(target, part, None)
else: else:
@ -1178,10 +1178,10 @@ def _add_null_to_type(schema: dict[str, Any]) -> None:
elif isinstance(current_type, list): elif isinstance(current_type, list):
# Add null to array if not already present # Add null to array if not already present
if "null" not in current_type: if "null" not in current_type:
schema["type"] = current_type + ["null"] schema["type"] = [*current_type, "null"]
elif "oneOf" in schema: elif "oneOf" in schema:
# Convert oneOf to anyOf with null type # Convert oneOf to anyOf with null type
schema["anyOf"] = schema.pop("oneOf") + [{"type": "null"}] schema["anyOf"] = [*schema.pop("oneOf"), {"type": "null"}]
elif "anyOf" in schema: elif "anyOf" in schema:
# Add null type to anyOf if not already present # Add null type to anyOf if not already present
if not any(item.get("type") == "null" for item in schema["anyOf"]): if not any(item.get("type") == "null" for item in schema["anyOf"]):
@ -1233,7 +1233,7 @@ def _handle_nullable_fields(schema: dict[str, Any] | Any) -> dict[str, Any] | An
# Handle properties nullable fields # Handle properties nullable fields
if has_property_nullable_field and "properties" in result: if has_property_nullable_field and "properties" in result:
for prop_name, prop_schema in result["properties"].items(): for _prop_name, prop_schema in result["properties"].items():
if isinstance(prop_schema, dict) and "nullable" in prop_schema: if isinstance(prop_schema, dict) and "nullable" in prop_schema:
nullable_value = prop_schema.pop("nullable") nullable_value = prop_schema.pop("nullable")
if nullable_value and ( if nullable_value and (

View file

@ -6,7 +6,7 @@ import multiprocessing
import socket import socket
import time import time
from collections.abc import AsyncGenerator, Callable, Generator from collections.abc import AsyncGenerator, Callable, Generator
from contextlib import asynccontextmanager, contextmanager from contextlib import asynccontextmanager, contextmanager, suppress
from typing import TYPE_CHECKING, Any, Literal from typing import TYPE_CHECKING, Any, Literal
from urllib.parse import parse_qs, urlparse from urllib.parse import parse_qs, urlparse
@ -216,10 +216,8 @@ async def run_server_async(
finally: finally:
# Cleanup: cancel the task # Cleanup: cancel the task
server_task.cancel() server_task.cancel()
try: with suppress(asyncio.CancelledError):
await server_task await server_task
except asyncio.CancelledError:
pass
@contextmanager @contextmanager

View file

@ -887,7 +887,7 @@ class TestIconExtraction:
return "icon" return "icon"
info = await inspect_fastmcp(mcp) info = await inspect_fastmcp(mcp)
json_bytes = await format_fastmcp_info(info) json_bytes = format_fastmcp_info(info)
import json import json
@ -915,7 +915,7 @@ class TestIconExtraction:
return "none" return "none"
info = await inspect_fastmcp(mcp) info = await inspect_fastmcp(mcp)
json_bytes = await format_fastmcp_info(info) json_bytes = format_fastmcp_info(info)
import json import json
@ -945,7 +945,7 @@ class TestFormatFunctions:
return {"result": x * 2} return {"result": x * 2}
info = await inspect_fastmcp(mcp) info = await inspect_fastmcp(mcp)
json_bytes = await format_fastmcp_info(info) json_bytes = format_fastmcp_info(info)
# Verify it's valid JSON # Verify it's valid JSON
import json import json
@ -1104,7 +1104,7 @@ class TestFormatFunctions:
assert "result" in info.tools[0].output_schema["properties"] assert "result" in info.tools[0].output_schema["properties"]
# Verify it's included in FastMCP format # Verify it's included in FastMCP format
json_bytes = await format_fastmcp_info(info) json_bytes = format_fastmcp_info(info)
import json import json
data = json.loads(json_bytes) data = json.loads(json_bytes)