mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 07:09:11 +02:00
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:
parent
e74918a544
commit
9d4c378e1b
65 changed files with 340 additions and 333 deletions
|
|
@ -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
|
||||
|
||||
__all__ = [
|
||||
"get_client",
|
||||
"get_profile_info",
|
||||
"create_post",
|
||||
"create_thread",
|
||||
"fetch_timeline",
|
||||
"search_for_posts",
|
||||
"fetch_notifications",
|
||||
"fetch_timeline",
|
||||
"follow_user_by_handle",
|
||||
"get_client",
|
||||
"get_profile_info",
|
||||
"like_post_by_uri",
|
||||
"repost_by_uri",
|
||||
"search_for_posts",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,3 +1,9 @@
|
|||
# /// script
|
||||
# dependencies = ["aiohttp", "fastmcp"]
|
||||
# ///
|
||||
|
||||
# uv pip install aiohttp fastmcp
|
||||
|
||||
import aiohttp
|
||||
|
||||
from fastmcp.server import FastMCP
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ from typing import Annotated, Any, Self
|
|||
import asyncpg
|
||||
import numpy as np
|
||||
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_ai import Agent
|
||||
|
||||
|
|
@ -149,7 +149,9 @@ class MemoryNode(BaseModel):
|
|||
)
|
||||
self.importance += other.importance
|
||||
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.content, "Summarize the following text concisely.", str, deps
|
||||
)
|
||||
|
|
@ -281,9 +283,9 @@ async def display_memory_tree(deps: Deps) -> str:
|
|||
|
||||
@mcp.tool
|
||||
async def remember(
|
||||
contents: list[str] = Field(
|
||||
description="List of observations or memories to store"
|
||||
),
|
||||
contents: Annotated[
|
||||
list[str], Field(description="List of observations or memories to store")
|
||||
],
|
||||
):
|
||||
deps = Deps(openai=AsyncOpenAI(), pool=await get_db_pool())
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -137,12 +137,34 @@ unknown-argument = "ignore" # 61 errors
|
|||
call-non-callable = "ignore" # 7 errors
|
||||
|
||||
[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]
|
||||
"__init__.py" = ["F401", "I001", "RUF013"]
|
||||
# allow imports not at the top of the file
|
||||
"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]
|
||||
ignore-words-list = "asend,shttp,te"
|
||||
|
|
|
|||
|
|
@ -48,9 +48,9 @@ def __getattr__(name: str):
|
|||
|
||||
|
||||
__all__ = [
|
||||
"FastMCP",
|
||||
"Context",
|
||||
"client",
|
||||
"Client",
|
||||
"Context",
|
||||
"FastMCP",
|
||||
"client",
|
||||
"settings",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -78,7 +78,7 @@ def with_argv(args: list[str] | None):
|
|||
original = sys.argv[:]
|
||||
try:
|
||||
# Preserve the script name (sys.argv[0]) and replace the rest
|
||||
sys.argv = [sys.argv[0]] + args
|
||||
sys.argv = [sys.argv[0], *args]
|
||||
yield
|
||||
finally:
|
||||
sys.argv = original
|
||||
|
|
@ -277,7 +277,7 @@ async def dev(
|
|||
|
||||
# Run the MCP Inspector command
|
||||
process = subprocess.run(
|
||||
[npx_cmd, inspector_cmd] + uv_cmd,
|
||||
[npx_cmd, inspector_cmd, *uv_cmd],
|
||||
check=True,
|
||||
env=env,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -15,18 +15,18 @@ from .transports import (
|
|||
from .auth import OAuth, BearerAuth
|
||||
|
||||
__all__ = [
|
||||
"BearerAuth",
|
||||
"Client",
|
||||
"ClientTransport",
|
||||
"WSTransport",
|
||||
"FastMCPTransport",
|
||||
"NodeStdioTransport",
|
||||
"NpxStdioTransport",
|
||||
"OAuth",
|
||||
"PythonStdioTransport",
|
||||
"SSETransport",
|
||||
"StdioTransport",
|
||||
"PythonStdioTransport",
|
||||
"NodeStdioTransport",
|
||||
"UvxStdioTransport",
|
||||
"UvStdioTransport",
|
||||
"NpxStdioTransport",
|
||||
"FastMCPTransport",
|
||||
"StreamableHttpTransport",
|
||||
"OAuth",
|
||||
"BearerAuth",
|
||||
"UvStdioTransport",
|
||||
"UvxStdioTransport",
|
||||
"WSTransport",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -36,8 +36,6 @@ logger = get_logger(__name__)
|
|||
class ClientNotFoundError(Exception):
|
||||
"""Raised when OAuth client credentials are not found on the server."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
async def check_if_auth_required(
|
||||
mcp_url: str, httpx_kwargs: dict[str, Any] | None = None
|
||||
|
|
@ -58,7 +56,7 @@ async def check_if_auth_required(
|
|||
return True
|
||||
|
||||
# Check for WWW-Authenticate header
|
||||
if "WWW-Authenticate" in response.headers:
|
||||
if "WWW-Authenticate" in response.headers: # noqa: SIM103
|
||||
return True
|
||||
|
||||
# If we get a successful response, auth may not be required
|
||||
|
|
@ -195,7 +193,8 @@ class OAuth(OAuthClientProvider):
|
|||
|
||||
warn(
|
||||
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(
|
||||
|
|
@ -272,8 +271,10 @@ class OAuth(OAuthClientProvider):
|
|||
if result.error:
|
||||
raise result.error
|
||||
return result.code, result.state # type: ignore
|
||||
except TimeoutError:
|
||||
raise TimeoutError(f"OAuth callback timed out after {TIMEOUT} seconds")
|
||||
except TimeoutError as e:
|
||||
raise TimeoutError(
|
||||
f"OAuth callback timed out after {TIMEOUT} seconds"
|
||||
) from e
|
||||
finally:
|
||||
server.should_exit = True
|
||||
await anyio.sleep(0.1) # Allow server to shut down gracefully
|
||||
|
|
|
|||
|
|
@ -61,15 +61,15 @@ from .transports import (
|
|||
|
||||
__all__ = [
|
||||
"Client",
|
||||
"SessionKwargs",
|
||||
"RootsHandler",
|
||||
"RootsList",
|
||||
"ClientSamplingHandler",
|
||||
"ElicitationHandler",
|
||||
"LogHandler",
|
||||
"MessageHandler",
|
||||
"ClientSamplingHandler",
|
||||
"SamplingHandler",
|
||||
"ElicitationHandler",
|
||||
"ProgressHandler",
|
||||
"RootsHandler",
|
||||
"RootsList",
|
||||
"SamplingHandler",
|
||||
"SessionKwargs",
|
||||
]
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
|
@ -362,10 +362,10 @@ class Client(Generic[ClientTransportT]):
|
|||
await self._session_state.session.initialize()
|
||||
)
|
||||
yield
|
||||
except anyio.ClosedResourceError:
|
||||
raise RuntimeError("Server session was closed unexpectedly")
|
||||
except TimeoutError:
|
||||
raise RuntimeError("Failed to initialize server session")
|
||||
except anyio.ClosedResourceError as e:
|
||||
raise RuntimeError("Server session was closed unexpectedly") from e
|
||||
except TimeoutError as e:
|
||||
raise RuntimeError("Failed to initialize server session") from e
|
||||
finally:
|
||||
self._session_state.session = None
|
||||
self._session_state.initialize_result = None
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from mcp.types import SamplingMessage
|
|||
|
||||
from fastmcp.server.sampling.handler import ServerSamplingHandler
|
||||
|
||||
__all__ = ["SamplingMessage", "SamplingParams", "SamplingHandler"]
|
||||
__all__ = ["SamplingHandler", "SamplingMessage", "SamplingParams"]
|
||||
|
||||
|
||||
ClientSamplingHandler: TypeAlias = Callable[
|
||||
|
|
|
|||
|
|
@ -46,16 +46,16 @@ ClientTransportT = TypeVar("ClientTransportT", bound="ClientTransport")
|
|||
|
||||
__all__ = [
|
||||
"ClientTransport",
|
||||
"SSETransport",
|
||||
"StreamableHttpTransport",
|
||||
"StdioTransport",
|
||||
"PythonStdioTransport",
|
||||
"FastMCPStdioTransport",
|
||||
"NodeStdioTransport",
|
||||
"UvxStdioTransport",
|
||||
"UvStdioTransport",
|
||||
"NpxStdioTransport",
|
||||
"FastMCPTransport",
|
||||
"NodeStdioTransport",
|
||||
"NpxStdioTransport",
|
||||
"PythonStdioTransport",
|
||||
"SSETransport",
|
||||
"StdioTransport",
|
||||
"StreamableHttpTransport",
|
||||
"UvStdioTransport",
|
||||
"UvxStdioTransport",
|
||||
"infer_transport",
|
||||
]
|
||||
|
||||
|
|
@ -109,9 +109,8 @@ class ClientTransport(abc.ABC):
|
|||
# Basic representation for subclasses
|
||||
return f"<{self.__class__.__name__}>"
|
||||
|
||||
async def close(self):
|
||||
async def close(self): # noqa: B027
|
||||
"""Close the transport."""
|
||||
pass
|
||||
|
||||
def _set_auth(self, auth: httpx.Auth | Literal["oauth"] | str | None):
|
||||
if auth is not None:
|
||||
|
|
@ -141,10 +140,10 @@ class WSTransport(ClientTransport):
|
|||
) -> AsyncIterator[ClientSession]:
|
||||
try:
|
||||
from mcp.client.websocket import websocket_client
|
||||
except ImportError:
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"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:
|
||||
read_stream, write_stream = transport
|
||||
|
|
@ -207,7 +206,7 @@ class SSETransport(ClientTransport):
|
|||
# instead we simply leave the kwarg out if it's not provided
|
||||
if self.sse_read_timeout is not None:
|
||||
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(
|
||||
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
|
||||
if self.sse_read_timeout is not None:
|
||||
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")
|
||||
|
||||
if self.httpx_client_factory is not None:
|
||||
|
|
@ -451,8 +450,7 @@ async def _stdio_transport_connect_task(
|
|||
if log_file is None:
|
||||
log_file_handle = sys.stderr
|
||||
elif isinstance(log_file, Path):
|
||||
log_file_handle = open(log_file, "a")
|
||||
stack.callback(log_file_handle.close)
|
||||
log_file_handle = stack.enter_context(log_file.open("a"))
|
||||
else:
|
||||
# Must be TextIO - use it directly
|
||||
log_file_handle = log_file
|
||||
|
|
@ -852,26 +850,28 @@ class FastMCPTransport(ClientTransport):
|
|||
server_read, server_write = server_streams
|
||||
|
||||
# Create a cancel scope for the server task
|
||||
async with anyio.create_task_group() as tg:
|
||||
async with _enter_server_lifespan(server=self.server):
|
||||
tg.start_soon(
|
||||
lambda: self.server._mcp_server.run(
|
||||
server_read,
|
||||
server_write,
|
||||
self.server._mcp_server.create_initialization_options(),
|
||||
raise_exceptions=self.raise_exceptions,
|
||||
)
|
||||
async with (
|
||||
anyio.create_task_group() as tg,
|
||||
_enter_server_lifespan(server=self.server),
|
||||
):
|
||||
tg.start_soon(
|
||||
lambda: self.server._mcp_server.run(
|
||||
server_read,
|
||||
server_write,
|
||||
self.server._mcp_server.create_initialization_options(),
|
||||
raise_exceptions=self.raise_exceptions,
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
async with ClientSession(
|
||||
read_stream=client_read,
|
||||
write_stream=client_write,
|
||||
**session_kwargs,
|
||||
) as client_session:
|
||||
yield client_session
|
||||
finally:
|
||||
tg.cancel_scope.cancel()
|
||||
try:
|
||||
async with ClientSession(
|
||||
read_stream=client_read,
|
||||
write_stream=client_write,
|
||||
**session_kwargs,
|
||||
) as client_session:
|
||||
yield client_session
|
||||
finally:
|
||||
tg.cancel_scope.cancel()
|
||||
|
||||
def __repr__(self) -> str:
|
||||
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
|
||||
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)
|
||||
|
||||
# otherwise create a composite client
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from .component_manager import set_up_component_manager
|
||||
from .component_service import ComponentService
|
||||
|
||||
__all__ = ["set_up_component_manager", "ComponentService"]
|
||||
__all__ = ["ComponentService", "set_up_component_manager"]
|
||||
|
|
|
|||
|
|
@ -97,11 +97,11 @@ def make_endpoint(action, component, config):
|
|||
return JSONResponse(
|
||||
{"message": f"{action.capitalize()}d {component}: {name}"}
|
||||
)
|
||||
except NotFoundError:
|
||||
except NotFoundError as e:
|
||||
raise StarletteHTTPException(
|
||||
status_code=404,
|
||||
detail=f"Unknown {component}: {name}",
|
||||
)
|
||||
) from e
|
||||
|
||||
return endpoint
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ from .mcp_mixin import MCPMixin, mcp_tool, mcp_resource, mcp_prompt
|
|||
|
||||
__all__ = [
|
||||
"MCPMixin",
|
||||
"mcp_tool",
|
||||
"mcp_resource",
|
||||
"mcp_prompt",
|
||||
"mcp_resource",
|
||||
"mcp_tool",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -21,10 +21,10 @@ try:
|
|||
ChatCompletionUserMessageParam,
|
||||
)
|
||||
from openai.types.shared.chat_model import ChatModel
|
||||
except ImportError:
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"The `openai` package is not installed. Please install `fastmcp[openai]` or add `openai` to your dependencies manually."
|
||||
)
|
||||
) from e
|
||||
|
||||
from typing_extensions import override
|
||||
|
||||
|
|
|
|||
|
|
@ -22,17 +22,14 @@ from .components import (
|
|||
|
||||
# Export public symbols - maintaining backward compatibility
|
||||
__all__ = [
|
||||
# Server
|
||||
"FastMCPOpenAPI",
|
||||
# Routing
|
||||
"MCPType",
|
||||
"RouteMap",
|
||||
"RouteMapFn",
|
||||
"ComponentFn",
|
||||
"DEFAULT_ROUTE_MAPPINGS",
|
||||
"_determine_route_type",
|
||||
# Components
|
||||
"OpenAPITool",
|
||||
"ComponentFn",
|
||||
"FastMCPOpenAPI",
|
||||
"MCPType",
|
||||
"OpenAPIResource",
|
||||
"OpenAPIResourceTemplate",
|
||||
"OpenAPITool",
|
||||
"RouteMap",
|
||||
"RouteMapFn",
|
||||
"_determine_route_type",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -146,11 +146,11 @@ class OpenAPITool(Tool):
|
|||
if e.response.text:
|
||||
error_message += f" - {e.response.text}"
|
||||
|
||||
raise ValueError(error_message)
|
||||
raise ValueError(error_message) from e
|
||||
|
||||
except httpx.RequestError as e:
|
||||
# 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):
|
||||
|
|
@ -165,9 +165,11 @@ class OpenAPIResource(Resource):
|
|||
name: str,
|
||||
description: str,
|
||||
mime_type: str = "application/json",
|
||||
tags: set[str] = set(),
|
||||
tags: set[str] | None = None,
|
||||
timeout: float | None = None,
|
||||
):
|
||||
if tags is None:
|
||||
tags = set()
|
||||
super().__init__(
|
||||
uri=AnyUrl(uri), # Convert string to AnyUrl
|
||||
name=name,
|
||||
|
|
@ -276,11 +278,11 @@ class OpenAPIResource(Resource):
|
|||
if e.response.text:
|
||||
error_message += f" - {e.response.text}"
|
||||
|
||||
raise ValueError(error_message)
|
||||
raise ValueError(error_message) from e
|
||||
|
||||
except httpx.RequestError as e:
|
||||
# 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):
|
||||
|
|
@ -295,9 +297,11 @@ class OpenAPIResourceTemplate(ResourceTemplate):
|
|||
name: str,
|
||||
description: str,
|
||||
parameters: dict[str, Any],
|
||||
tags: set[str] = set(),
|
||||
tags: set[str] | None = None,
|
||||
timeout: float | None = None,
|
||||
):
|
||||
if tags is None:
|
||||
tags = set()
|
||||
super().__init__(
|
||||
uri_template=uri_template,
|
||||
name=name,
|
||||
|
|
@ -342,7 +346,7 @@ class OpenAPIResourceTemplate(ResourceTemplate):
|
|||
|
||||
# Export public symbols
|
||||
__all__ = [
|
||||
"OpenAPITool",
|
||||
"OpenAPIResource",
|
||||
"OpenAPIResourceTemplate",
|
||||
"OpenAPITool",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -121,10 +121,10 @@ def _determine_route_type(
|
|||
|
||||
# Export public symbols
|
||||
__all__ = [
|
||||
"DEFAULT_ROUTE_MAPPINGS",
|
||||
"ComponentFn",
|
||||
"MCPType",
|
||||
"RouteMap",
|
||||
"RouteMapFn",
|
||||
"ComponentFn",
|
||||
"DEFAULT_ROUTE_MAPPINGS",
|
||||
"_determine_route_type",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -40,29 +40,24 @@ from .json_schema_converter import (
|
|||
|
||||
# Export public symbols - maintaining backward compatibility
|
||||
__all__ = [
|
||||
# Models
|
||||
"HTTPRoute",
|
||||
"HttpMethod",
|
||||
"JsonSchema",
|
||||
"ParameterInfo",
|
||||
"ParameterLocation",
|
||||
"RequestBodyInfo",
|
||||
"ResponseInfo",
|
||||
"HttpMethod",
|
||||
"ParameterLocation",
|
||||
"JsonSchema",
|
||||
# Parser
|
||||
"parse_openapi_to_http_routes",
|
||||
# Formatters
|
||||
"_combine_schemas",
|
||||
"_make_optional_parameter_nullable",
|
||||
"clean_schema_for_display",
|
||||
"convert_openapi_schema_to_json_schema",
|
||||
"convert_schema_definitions",
|
||||
"extract_output_schema_from_responses",
|
||||
"format_array_parameter",
|
||||
"format_deep_object_parameter",
|
||||
"format_description_with_responses",
|
||||
"format_json_for_description",
|
||||
"format_simple_description",
|
||||
"generate_example_from_schema",
|
||||
# Schemas
|
||||
"_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",
|
||||
"parse_openapi_to_http_routes",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ class RequestDirector:
|
|||
|
||||
# Step 4: Handle request body
|
||||
if body is not None:
|
||||
if isinstance(body, dict) or isinstance(body, list):
|
||||
if isinstance(body, dict | list):
|
||||
request_data["json"] = body
|
||||
else:
|
||||
request_data["content"] = body
|
||||
|
|
|
|||
|
|
@ -164,10 +164,10 @@ def _convert_nullable_field(schema: dict[str, Any]) -> dict[str, Any]:
|
|||
if isinstance(current_type, str):
|
||||
result["type"] = [current_type, "null"]
|
||||
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:
|
||||
# Convert oneOf to anyOf with null
|
||||
result["anyOf"] = result.pop("oneOf") + [{"type": "null"}]
|
||||
result["anyOf"] = [*result.pop("oneOf"), {"type": "null"}]
|
||||
elif "anyOf" in result:
|
||||
# Add null to anyOf if not present
|
||||
if not any(item.get("type") == "null" for item in result["anyOf"]):
|
||||
|
|
|
|||
|
|
@ -79,10 +79,10 @@ class HTTPRoute(FastMCPBaseModel):
|
|||
# Export public symbols
|
||||
__all__ = [
|
||||
"HTTPRoute",
|
||||
"HttpMethod",
|
||||
"JsonSchema",
|
||||
"ParameterInfo",
|
||||
"ParameterLocation",
|
||||
"RequestBodyInfo",
|
||||
"ResponseInfo",
|
||||
"HttpMethod",
|
||||
"ParameterLocation",
|
||||
"JsonSchema",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -178,7 +178,7 @@ class OpenAPIParser(
|
|||
else:
|
||||
# Special handling for components
|
||||
if part == "components" and hasattr(target, "components"):
|
||||
target = getattr(target, "components")
|
||||
target = target.components
|
||||
elif hasattr(target, part): # Fallback check
|
||||
target = getattr(target, part, None)
|
||||
else:
|
||||
|
|
@ -554,9 +554,7 @@ class OpenAPIParser(
|
|||
if "$ref" in obj and isinstance(obj["$ref"], str):
|
||||
ref = obj["$ref"]
|
||||
# Handle both converted and unconverted refs
|
||||
if ref.startswith("#/$defs/"):
|
||||
schema_name = ref.split("/")[-1]
|
||||
elif ref.startswith("#/components/schemas/"):
|
||||
if ref.startswith(("#/$defs/", "#/components/schemas/")):
|
||||
schema_name = ref.split("/")[-1]
|
||||
else:
|
||||
return
|
||||
|
|
@ -815,6 +813,6 @@ class OpenAPIParser(
|
|||
|
||||
# Export public symbols
|
||||
__all__ = [
|
||||
"parse_openapi_to_http_routes",
|
||||
"OpenAPIParser",
|
||||
"parse_openapi_to_http_routes",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -585,9 +585,9 @@ def extract_output_schema_from_responses(
|
|||
|
||||
# Export public symbols
|
||||
__all__ = [
|
||||
"clean_schema_for_display",
|
||||
"_combine_schemas",
|
||||
"_combine_schemas_and_map_params",
|
||||
"extract_output_schema_from_responses",
|
||||
"_make_optional_parameter_nullable",
|
||||
"clean_schema_for_display",
|
||||
"extract_output_schema_from_responses",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -288,9 +288,8 @@ class MCPConfig(BaseModel):
|
|||
@classmethod
|
||||
def from_file(cls, file_path: Path) -> Self:
|
||||
"""Load configuration from JSON file."""
|
||||
if file_path.exists():
|
||||
if content := file_path.read_text().strip():
|
||||
return cls.model_validate_json(content)
|
||||
if file_path.exists() and (content := file_path.read_text().strip()):
|
||||
return cls.model_validate_json(content)
|
||||
|
||||
raise ValueError(f"No MCP servers defined in the config: {file_path}")
|
||||
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ from .prompt import Prompt, PromptMessage, Message
|
|||
from .prompt_manager import PromptManager
|
||||
|
||||
__all__ = [
|
||||
"Message",
|
||||
"Prompt",
|
||||
"PromptManager",
|
||||
"PromptMessage",
|
||||
"Message",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -207,10 +207,7 @@ class FunctionPrompt(Prompt):
|
|||
# Auto-detect context parameter if not provided
|
||||
|
||||
context_kwarg = find_kwarg_by_type(fn, kwarg_type=Context)
|
||||
if context_kwarg:
|
||||
prune_params = [context_kwarg]
|
||||
else:
|
||||
prune_params = None
|
||||
prune_params = [context_kwarg] if context_kwarg else None
|
||||
|
||||
parameters = compress_schema(parameters, prune_params=prune_params)
|
||||
|
||||
|
|
@ -290,10 +287,7 @@ class FunctionPrompt(Prompt):
|
|||
if (
|
||||
param.annotation == inspect.Parameter.empty
|
||||
or param.annotation is 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):
|
||||
) or not isinstance(param_value, str):
|
||||
converted_kwargs[param_name] = param_value
|
||||
else:
|
||||
# Try to convert string argument using type adapter
|
||||
|
|
@ -314,7 +308,7 @@ class FunctionPrompt(Prompt):
|
|||
raise PromptError(
|
||||
f"Could not convert argument '{param_name}' with value '{param_value}' "
|
||||
f"to expected type {param.annotation}. Error: {e}"
|
||||
)
|
||||
) from e
|
||||
else:
|
||||
# Parameter not in function signature, pass as-is
|
||||
converted_kwargs[param_name] = param_value
|
||||
|
|
@ -376,10 +370,12 @@ class FunctionPrompt(Prompt):
|
|||
content=TextContent(type="text", text=content),
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
raise PromptError("Could not convert prompt result to message.")
|
||||
except Exception as e:
|
||||
raise PromptError(
|
||||
"Could not convert prompt result to message."
|
||||
) from e
|
||||
|
||||
return messages
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
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
|
||||
|
|
|
|||
|
|
@ -10,13 +10,13 @@ from .types import (
|
|||
from .resource_manager import ResourceManager
|
||||
|
||||
__all__ = [
|
||||
"Resource",
|
||||
"TextResource",
|
||||
"BinaryResource",
|
||||
"FunctionResource",
|
||||
"FileResource",
|
||||
"HttpResource",
|
||||
"DirectoryResource",
|
||||
"ResourceTemplate",
|
||||
"FileResource",
|
||||
"FunctionResource",
|
||||
"HttpResource",
|
||||
"Resource",
|
||||
"ResourceManager",
|
||||
"ResourceTemplate",
|
||||
"TextResource",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -217,9 +217,7 @@ class FunctionResource(Resource):
|
|||
|
||||
if isinstance(result, Resource):
|
||||
return await result.read()
|
||||
elif isinstance(result, bytes):
|
||||
return result
|
||||
elif isinstance(result, str):
|
||||
elif isinstance(result, bytes | str):
|
||||
return result
|
||||
else:
|
||||
return pydantic_core.to_json(result, fallback=str).decode()
|
||||
|
|
|
|||
|
|
@ -235,7 +235,7 @@ class ResourceManager:
|
|||
|
||||
# Then check templates (local and mounted) only if not found in concrete resources
|
||||
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):
|
||||
return True
|
||||
|
||||
|
|
|
|||
|
|
@ -3,4 +3,4 @@ from .context import Context
|
|||
from . import dependencies
|
||||
|
||||
|
||||
__all__ = ["FastMCP", "Context"]
|
||||
__all__ = ["Context", "FastMCP"]
|
||||
|
|
|
|||
|
|
@ -10,14 +10,14 @@ from .oauth_proxy import OAuthProxy
|
|||
|
||||
|
||||
__all__ = [
|
||||
"AuthProvider",
|
||||
"OAuthProvider",
|
||||
"TokenVerifier",
|
||||
"JWTVerifier",
|
||||
"StaticTokenVerifier",
|
||||
"RemoteAuthProvider",
|
||||
"AccessToken",
|
||||
"AuthProvider",
|
||||
"JWTVerifier",
|
||||
"OAuthProvider",
|
||||
"OAuthProxy",
|
||||
"RemoteAuthProvider",
|
||||
"StaticTokenVerifier",
|
||||
"TokenVerifier",
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ from mcp.server.auth.settings import (
|
|||
ClientRegistrationOptions,
|
||||
RevocationOptions,
|
||||
)
|
||||
from pydantic import AnyHttpUrl
|
||||
from pydantic import AnyHttpUrl, Field
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.middleware.authentication import AuthenticationMiddleware
|
||||
from starlette.routing import Route
|
||||
|
|
@ -32,7 +32,7 @@ from starlette.routing import Route
|
|||
class AccessToken(_SDKAccessToken):
|
||||
"""AccessToken that includes all JWT claims."""
|
||||
|
||||
claims: dict[str, Any] = {}
|
||||
claims: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class AuthProvider(TokenVerifierProtocol):
|
||||
|
|
|
|||
|
|
@ -123,10 +123,10 @@ class OIDCConfiguration(BaseModel):
|
|||
|
||||
try:
|
||||
AnyHttpUrl(value)
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
message = f"Invalid URL for configuration metadata: {attr}"
|
||||
logger.error(message)
|
||||
raise ValueError(message)
|
||||
raise ValueError(message) from e
|
||||
|
||||
enforce("issuer", True)
|
||||
enforce("authorization_endpoint", True)
|
||||
|
|
|
|||
|
|
@ -113,12 +113,12 @@ class AzureProvider(OAuthProxy):
|
|||
client_id: str | NotSetT = NotSet,
|
||||
client_secret: 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,
|
||||
issuer_url: str | NotSetT = NotSet,
|
||||
redirect_path: str | NotSetT = NotSet,
|
||||
required_scopes: list[str] | None | NotSetT = NotSet,
|
||||
additional_authorize_scopes: list[str] | None | NotSetT = NotSet,
|
||||
required_scopes: list[str] | NotSetT | None = NotSet,
|
||||
additional_authorize_scopes: list[str] | NotSetT | None = NotSet,
|
||||
allowed_client_redirect_uris: list[str] | NotSetT = NotSet,
|
||||
client_storage: AsyncKeyValue | None = None,
|
||||
jwt_signing_key: str | bytes | NotSetT = NotSet,
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from fastmcp.server.auth.providers.jwt import JWKData, JWKSData, RSAKeyPair
|
|||
from fastmcp.server.auth.providers.jwt import JWTVerifier as BearerAuthProvider
|
||||
|
||||
# Re-export for backwards compatibility
|
||||
__all__ = ["BearerAuthProvider", "RSAKeyPair", "JWKData", "JWKSData"]
|
||||
__all__ = ["BearerAuthProvider", "JWKData", "JWKSData", "RSAKeyPair"]
|
||||
|
||||
# Deprecated in 2.11
|
||||
if fastmcp.settings.deprecation_warnings:
|
||||
|
|
|
|||
|
|
@ -96,10 +96,10 @@ class InMemoryOAuthProvider(OAuthProvider):
|
|||
# or if params.redirect_uri is None and client has a default.
|
||||
# However, the AuthorizationHandler handles the primary validation.
|
||||
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(
|
||||
error="invalid_request", error_description="Invalid redirect_uri."
|
||||
)
|
||||
) from e
|
||||
|
||||
auth_code_value = f"test_auth_code_{secrets.token_hex(16)}"
|
||||
expires_at = time.time() + DEFAULT_AUTH_CODE_EXPIRY_SECONDS
|
||||
|
|
|
|||
|
|
@ -97,8 +97,8 @@ class IntrospectionTokenVerifier(TokenVerifier):
|
|||
client_id: str | NotSetT = NotSet,
|
||||
client_secret: str | NotSetT = NotSet,
|
||||
timeout_seconds: int | NotSetT = NotSet,
|
||||
required_scopes: list[str] | None | NotSetT = NotSet,
|
||||
base_url: AnyHttpUrl | str | None | NotSetT = NotSet,
|
||||
required_scopes: list[str] | NotSetT | None = NotSet,
|
||||
base_url: AnyHttpUrl | str | NotSetT | None = NotSet,
|
||||
):
|
||||
"""
|
||||
Initialize the introspection token verifier.
|
||||
|
|
|
|||
|
|
@ -184,13 +184,13 @@ class JWTVerifier(TokenVerifier):
|
|||
def __init__(
|
||||
self,
|
||||
*,
|
||||
public_key: str | None | NotSetT = NotSet,
|
||||
jwks_uri: str | None | NotSetT = NotSet,
|
||||
issuer: str | None | NotSetT = NotSet,
|
||||
audience: str | list[str] | None | NotSetT = NotSet,
|
||||
algorithm: str | None | NotSetT = NotSet,
|
||||
required_scopes: list[str] | None | NotSetT = NotSet,
|
||||
base_url: AnyHttpUrl | str | None | NotSetT = NotSet,
|
||||
public_key: str | NotSetT | None = NotSet,
|
||||
jwks_uri: str | NotSetT | None = NotSet,
|
||||
issuer: str | NotSetT | None = NotSet,
|
||||
audience: str | list[str] | NotSetT | None = NotSet,
|
||||
algorithm: str | NotSetT | None = NotSet,
|
||||
required_scopes: list[str] | NotSetT | None = NotSet,
|
||||
base_url: AnyHttpUrl | str | NotSetT | None = NotSet,
|
||||
):
|
||||
"""
|
||||
Initialize the JWT token verifier.
|
||||
|
|
@ -283,7 +283,7 @@ class JWTVerifier(TokenVerifier):
|
|||
return await self._get_jwks_key(kid)
|
||||
|
||||
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:
|
||||
"""Fetch key from JWKS with simple caching."""
|
||||
|
|
@ -342,10 +342,10 @@ class JWTVerifier(TokenVerifier):
|
|||
raise ValueError("No keys found in JWKS")
|
||||
|
||||
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:
|
||||
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]:
|
||||
"""
|
||||
|
|
@ -400,14 +400,13 @@ class JWTVerifier(TokenVerifier):
|
|||
|
||||
# Validate issuer - note we use issuer instead of issuer_url here because
|
||||
# issuer is optional, allowing users to make this check optional
|
||||
if self.issuer:
|
||||
if claims.get("iss") != self.issuer:
|
||||
self.logger.debug(
|
||||
"Token validation failed: issuer mismatch for client %s",
|
||||
client_id,
|
||||
)
|
||||
self.logger.info("Bearer token rejected for client %s", client_id)
|
||||
return None
|
||||
if self.issuer and claims.get("iss") != self.issuer:
|
||||
self.logger.debug(
|
||||
"Token validation failed: issuer mismatch for client %s",
|
||||
client_id,
|
||||
)
|
||||
self.logger.info("Bearer token rejected for client %s", client_id)
|
||||
return None
|
||||
|
||||
# Validate audience if configured
|
||||
if self.audience:
|
||||
|
|
|
|||
|
|
@ -83,7 +83,7 @@ class SupabaseProvider(RemoteAuthProvider):
|
|||
*,
|
||||
project_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,
|
||||
):
|
||||
"""Initialize Supabase metadata provider.
|
||||
|
|
|
|||
|
|
@ -169,7 +169,7 @@ class WorkOSProvider(OAuthProxy):
|
|||
base_url: AnyHttpUrl | str | NotSetT = NotSet,
|
||||
issuer_url: AnyHttpUrl | 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,
|
||||
allowed_client_redirect_uris: list[str] | NotSetT = NotSet,
|
||||
client_storage: AsyncKeyValue | None = None,
|
||||
|
|
@ -338,7 +338,7 @@ class AuthKitProvider(RemoteAuthProvider):
|
|||
*,
|
||||
authkit_domain: 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,
|
||||
):
|
||||
"""Initialize AuthKit metadata provider.
|
||||
|
|
|
|||
|
|
@ -188,8 +188,8 @@ class Context:
|
|||
"""
|
||||
try:
|
||||
return request_ctx.get()
|
||||
except LookupError:
|
||||
raise ValueError("Context is not available outside of a request")
|
||||
except LookupError as e:
|
||||
raise ValueError("Context is not available outside of a request") from e
|
||||
|
||||
async def report_progress(
|
||||
self, progress: float, total: float | None = None, message: str | None = None
|
||||
|
|
@ -342,7 +342,7 @@ class Context:
|
|||
session_id = str(uuid4())
|
||||
|
||||
# Save the session id to the session attributes
|
||||
setattr(session, "_fastmcp_id", session_id)
|
||||
session._fastmcp_id = session_id
|
||||
return session_id
|
||||
|
||||
@property
|
||||
|
|
@ -595,13 +595,11 @@ class Context:
|
|||
choice_literal = Literal[tuple(response_type)] # type: ignore
|
||||
response_type = ScalarElicitationType[choice_literal] # type: ignore
|
||||
# if the user provided a primitive scalar, wrap it in an object schema
|
||||
elif response_type in {bool, int, float, str}:
|
||||
response_type = ScalarElicitationType[response_type] # type: ignore
|
||||
# if the user provided a Literal type, wrap it in an object schema
|
||||
elif get_origin(response_type) is Literal:
|
||||
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):
|
||||
elif (
|
||||
response_type in {bool, int, float, str}
|
||||
or get_origin(response_type) is Literal
|
||||
or (isinstance(response_type, type) and issubclass(response_type, Enum))
|
||||
):
|
||||
response_type = ScalarElicitationType[response_type] # type: ignore
|
||||
|
||||
response_type = cast(type[T], response_type)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from mcp.server.auth.middleware.auth_context import (
|
||||
|
|
@ -16,11 +17,11 @@ if TYPE_CHECKING:
|
|||
from fastmcp.server.context import Context
|
||||
|
||||
__all__ = [
|
||||
"get_context",
|
||||
"get_http_request",
|
||||
"get_http_headers",
|
||||
"get_access_token",
|
||||
"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
|
||||
|
||||
request = None
|
||||
try:
|
||||
with contextlib.suppress(LookupError):
|
||||
request = request_ctx.get().request
|
||||
except LookupError:
|
||||
pass
|
||||
|
||||
if request is None:
|
||||
raise RuntimeError("No active HTTP request found.")
|
||||
|
|
|
|||
|
|
@ -20,8 +20,8 @@ __all__ = [
|
|||
"AcceptedElicitation",
|
||||
"CancelledElicitation",
|
||||
"DeclinedElicitation",
|
||||
"get_elicitation_schema",
|
||||
"ScalarElicitationType",
|
||||
"get_elicitation_schema",
|
||||
]
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
|
|
|||
|
|
@ -342,9 +342,8 @@ def create_streamable_http_app(
|
|||
# Create a lifespan manager to start and stop the session manager
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: Starlette) -> AsyncGenerator[None, None]:
|
||||
async with server._lifespan_manager():
|
||||
async with session_manager.run():
|
||||
yield
|
||||
async with server._lifespan_manager(), session_manager.run():
|
||||
yield
|
||||
|
||||
# Create and return the app with lifespan
|
||||
app = create_base_app(
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from .middleware import (
|
|||
)
|
||||
|
||||
__all__ = [
|
||||
"CallNext",
|
||||
"Middleware",
|
||||
"MiddlewareContext",
|
||||
"CallNext",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -46,7 +46,7 @@ class CachableReadResourceContents(BaseModel):
|
|||
|
||||
@classmethod
|
||||
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
|
||||
def wrap(cls, values: Sequence[ReadResourceContents]) -> list[Self]:
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ class ErrorHandlingMiddleware(Middleware):
|
|||
error_key = f"{error_type}:{method}"
|
||||
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:
|
||||
self.logger.error(f"{base_message}\n{traceback.format_exc()}")
|
||||
|
|
@ -91,24 +91,24 @@ class ErrorHandlingMiddleware(Middleware):
|
|||
|
||||
if error_type in (ValueError, TypeError):
|
||||
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):
|
||||
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:
|
||||
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+
|
||||
elif error_type in (TimeoutError, asyncio.TimeoutError):
|
||||
return McpError(
|
||||
ErrorData(code=-32000, message=f"Request timeout: {str(error)}")
|
||||
ErrorData(code=-32000, message=f"Request timeout: {error!s}")
|
||||
)
|
||||
else:
|
||||
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:
|
||||
|
|
@ -120,7 +120,7 @@ class ErrorHandlingMiddleware(Middleware):
|
|||
|
||||
# Transform and re-raise
|
||||
transformed_error = self._transform_error(error)
|
||||
raise transformed_error
|
||||
raise transformed_error from error
|
||||
|
||||
def get_error_stats(self) -> dict[str, int]:
|
||||
"""Get error statistics for monitoring."""
|
||||
|
|
@ -200,7 +200,7 @@ class RetryMiddleware(Middleware):
|
|||
delay = self._calculate_delay(attempt)
|
||||
self.logger.warning(
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -27,9 +27,9 @@ if TYPE_CHECKING:
|
|||
from fastmcp.server.context import Context
|
||||
|
||||
__all__ = [
|
||||
"CallNext",
|
||||
"Middleware",
|
||||
"MiddlewareContext",
|
||||
"CallNext",
|
||||
]
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
|
|
|||
|
|
@ -513,11 +513,11 @@ class OpenAPITool(Tool):
|
|||
if e.response.text:
|
||||
error_message += f" - {e.response.text}"
|
||||
|
||||
raise ValueError(error_message)
|
||||
raise ValueError(error_message) from e
|
||||
|
||||
except httpx.RequestError as e:
|
||||
# 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):
|
||||
|
|
@ -531,9 +531,11 @@ class OpenAPIResource(Resource):
|
|||
name: str,
|
||||
description: str,
|
||||
mime_type: str = "application/json",
|
||||
tags: set[str] = set(),
|
||||
tags: set[str] | None = None,
|
||||
timeout: float | None = None,
|
||||
):
|
||||
if tags is None:
|
||||
tags = set()
|
||||
super().__init__(
|
||||
uri=AnyUrl(uri), # Convert string to AnyUrl
|
||||
name=name,
|
||||
|
|
@ -632,11 +634,11 @@ class OpenAPIResource(Resource):
|
|||
if e.response.text:
|
||||
error_message += f" - {e.response.text}"
|
||||
|
||||
raise ValueError(error_message)
|
||||
raise ValueError(error_message) from e
|
||||
|
||||
except httpx.RequestError as e:
|
||||
# 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):
|
||||
|
|
@ -650,9 +652,11 @@ class OpenAPIResourceTemplate(ResourceTemplate):
|
|||
name: str,
|
||||
description: str,
|
||||
parameters: dict[str, Any],
|
||||
tags: set[str] = set(),
|
||||
tags: set[str] | None = None,
|
||||
timeout: float | None = None,
|
||||
):
|
||||
if tags is None:
|
||||
tags = set()
|
||||
super().__init__(
|
||||
uri_template=uri_template,
|
||||
name=name,
|
||||
|
|
|
|||
|
|
@ -198,7 +198,9 @@ class ProxyResourceManager(ResourceManager, ProxyManagerMixin):
|
|||
elif isinstance(result[0], BlobResourceContents):
|
||||
return result[0].blob
|
||||
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):
|
||||
|
|
@ -558,7 +560,7 @@ class ProxyClient(Client[ClientTransportT]):
|
|||
kwargs["log_handler"] = ProxyClient.default_log_handler
|
||||
if "progress_handler" not in kwargs:
|
||||
kwargs["progress_handler"] = ProxyClient.default_progress_handler
|
||||
super().__init__(**kwargs | dict(transport=transport))
|
||||
super().__init__(**kwargs | {"transport": transport})
|
||||
|
||||
@classmethod
|
||||
async def default_sampling_handler(
|
||||
|
|
@ -572,7 +574,7 @@ class ProxyClient(Client[ClientTransportT]):
|
|||
"""
|
||||
ctx = get_context()
|
||||
content = await ctx.sample(
|
||||
[msg for msg in messages],
|
||||
list(messages),
|
||||
system_prompt=params.systemPrompt,
|
||||
temperature=params.temperature,
|
||||
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.
|
||||
So we do nothing here.
|
||||
"""
|
||||
pass
|
||||
|
||||
async def clear(self):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -15,7 +15,11 @@ from collections.abc import (
|
|||
Mapping,
|
||||
Sequence,
|
||||
)
|
||||
from contextlib import AbstractAsyncContextManager, AsyncExitStack, asynccontextmanager
|
||||
from contextlib import (
|
||||
AbstractAsyncContextManager,
|
||||
AsyncExitStack,
|
||||
asynccontextmanager,
|
||||
)
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
|
|
@ -150,7 +154,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
version: str | None = None,
|
||||
website_url: str | None = None,
|
||||
icons: list[mcp.types.Icon] | None = None,
|
||||
auth: AuthProvider | None | NotSetT = NotSet,
|
||||
auth: AuthProvider | NotSetT | None = NotSet,
|
||||
middleware: Sequence[Middleware] | None = None,
|
||||
lifespan: LifespanCallable | None = None,
|
||||
dependencies: list[str] | None = None,
|
||||
|
|
@ -1062,10 +1066,10 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
try:
|
||||
result = await self._call_tool_middleware(key, arguments)
|
||||
return result.to_mcp_result()
|
||||
except DisabledError:
|
||||
raise NotFoundError(f"Unknown tool: {key}")
|
||||
except NotFoundError:
|
||||
raise NotFoundError(f"Unknown tool: {key}")
|
||||
except DisabledError as e:
|
||||
raise NotFoundError(f"Unknown tool: {key}") from e
|
||||
except NotFoundError as e:
|
||||
raise NotFoundError(f"Unknown tool: {key}") from e
|
||||
|
||||
async def _call_tool_middleware(
|
||||
self,
|
||||
|
|
@ -1142,12 +1146,12 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
return list[ReadResourceContents](
|
||||
await self._read_resource_middleware(uri)
|
||||
)
|
||||
except DisabledError:
|
||||
except DisabledError as e:
|
||||
# convert to NotFoundError to avoid leaking resource presence
|
||||
raise NotFoundError(f"Unknown resource: {str(uri)!r}")
|
||||
except NotFoundError:
|
||||
raise NotFoundError(f"Unknown resource: {str(uri)!r}") from e
|
||||
except NotFoundError as e:
|
||||
# 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(
|
||||
self,
|
||||
|
|
@ -1158,10 +1162,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
"""
|
||||
|
||||
# Convert string URI to AnyUrl if needed
|
||||
if isinstance(uri, str):
|
||||
uri_param = AnyUrl(uri)
|
||||
else:
|
||||
uri_param = uri
|
||||
uri_param = AnyUrl(uri) if isinstance(uri, str) else uri
|
||||
|
||||
mw_context = MiddlewareContext(
|
||||
message=mcp.types.ReadResourceRequestParams(uri=uri_param),
|
||||
|
|
@ -1241,12 +1242,12 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
async with fastmcp.server.context.Context(fastmcp=self):
|
||||
try:
|
||||
return await self._get_prompt_middleware(name, arguments)
|
||||
except DisabledError:
|
||||
except DisabledError as e:
|
||||
# convert to NotFoundError to avoid leaking prompt presence
|
||||
raise NotFoundError(f"Unknown prompt: {name}")
|
||||
except NotFoundError:
|
||||
raise NotFoundError(f"Unknown prompt: {name}") from e
|
||||
except NotFoundError as e:
|
||||
# standardize NotFound message
|
||||
raise NotFoundError(f"Unknown prompt: {name}")
|
||||
raise NotFoundError(f"Unknown prompt: {name}") from e
|
||||
|
||||
async def _get_prompt_middleware(
|
||||
self, name: str, arguments: dict[str, Any] | None = None
|
||||
|
|
@ -1369,7 +1370,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
description: str | None = None,
|
||||
icons: list[mcp.types.Icon] | 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,
|
||||
exclude_args: list[str] | None = None,
|
||||
meta: dict[str, Any] | None = None,
|
||||
|
|
@ -1386,7 +1387,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
description: str | None = None,
|
||||
icons: list[mcp.types.Icon] | 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,
|
||||
exclude_args: list[str] | None = None,
|
||||
meta: dict[str, Any] | None = None,
|
||||
|
|
@ -1402,7 +1403,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
description: str | None = None,
|
||||
icons: list[mcp.types.Icon] | 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,
|
||||
exclude_args: list[str] | None = None,
|
||||
meta: dict[str, Any] | None = None,
|
||||
|
|
@ -2029,14 +2030,14 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
port=port,
|
||||
path=server_path,
|
||||
)
|
||||
_uvicorn_config_from_user = uvicorn_config or {}
|
||||
uvicorn_config_from_user = uvicorn_config or {}
|
||||
|
||||
config_kwargs: dict[str, Any] = {
|
||||
"timeout_graceful_shutdown": 0,
|
||||
"lifespan": "on",
|
||||
"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:
|
||||
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
|
||||
# - Disconnected clients: create fresh sessions per request for isolation
|
||||
if client.is_connected():
|
||||
_proxy_logger = get_logger(__name__)
|
||||
_proxy_logger.info(
|
||||
proxy_logger = get_logger(__name__)
|
||||
proxy_logger.info(
|
||||
"Proxy detected connected client - reusing existing session for all requests. "
|
||||
"This may cause context mixing in concurrent scenarios."
|
||||
)
|
||||
|
|
@ -2678,10 +2679,7 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
return False
|
||||
|
||||
if self.include_tags is not None:
|
||||
if any(itag in component.tags for itag in self.include_tags):
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
return bool(any(itag in component.tags for itag in self.include_tags))
|
||||
|
||||
return True
|
||||
|
||||
|
|
|
|||
|
|
@ -2,4 +2,4 @@ from .tool import Tool, FunctionTool
|
|||
from .tool_manager import ToolManager
|
||||
from .tool_transform import forward, forward_raw
|
||||
|
||||
__all__ = ["Tool", "ToolManager", "FunctionTool", "forward", "forward_raw"]
|
||||
__all__ = ["FunctionTool", "Tool", "ToolManager", "forward", "forward_raw"]
|
||||
|
|
|
|||
|
|
@ -173,7 +173,7 @@ class Tool(FastMCPComponent):
|
|||
tags: set[str] | None = None,
|
||||
annotations: ToolAnnotations | 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,
|
||||
meta: dict[str, Any] | None = None,
|
||||
enabled: bool | None = None,
|
||||
|
|
@ -212,13 +212,13 @@ class Tool(FastMCPComponent):
|
|||
tool: Tool,
|
||||
*,
|
||||
name: str | None = None,
|
||||
title: str | None | NotSetT = NotSet,
|
||||
description: str | None | NotSetT = NotSet,
|
||||
title: str | NotSetT | None = NotSet,
|
||||
description: str | NotSetT | None = NotSet,
|
||||
tags: set[str] | None = None,
|
||||
annotations: ToolAnnotations | None | NotSetT = NotSet,
|
||||
output_schema: dict[str, Any] | None | NotSetT | Literal[False] = NotSet,
|
||||
annotations: ToolAnnotations | NotSetT | None = NotSet,
|
||||
output_schema: dict[str, Any] | Literal[False] | NotSetT | None = NotSet,
|
||||
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,
|
||||
enabled: bool | None = None,
|
||||
transform_fn: Callable[..., Any] | None = None,
|
||||
|
|
@ -255,7 +255,7 @@ class FunctionTool(Tool):
|
|||
tags: set[str] | None = None,
|
||||
annotations: ToolAnnotations | 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,
|
||||
meta: dict[str, Any] | None = None,
|
||||
enabled: bool | None = None,
|
||||
|
|
@ -446,9 +446,8 @@ class ParsedFunction:
|
|||
# we ensure that no output schema is automatically generated.
|
||||
clean_output_type = replace_type(
|
||||
output_type,
|
||||
{
|
||||
t: _UnserializableType
|
||||
for t in (
|
||||
dict.fromkeys( # type: ignore[arg-type]
|
||||
(
|
||||
Image,
|
||||
Audio,
|
||||
File,
|
||||
|
|
@ -458,8 +457,9 @@ class ParsedFunction:
|
|||
mcp.types.AudioContent,
|
||||
mcp.types.ResourceLink,
|
||||
mcp.types.EmbeddedResource,
|
||||
)
|
||||
},
|
||||
),
|
||||
_UnserializableType,
|
||||
),
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -365,15 +365,15 @@ class TransformedTool(Tool):
|
|||
cls,
|
||||
tool: Tool,
|
||||
name: str | None = None,
|
||||
title: str | None | NotSetT = NotSet,
|
||||
description: str | None | NotSetT = NotSet,
|
||||
title: str | NotSetT | None = NotSet,
|
||||
description: str | NotSetT | None = NotSet,
|
||||
tags: set[str] | None = None,
|
||||
transform_fn: Callable[..., Any] | None = None,
|
||||
transform_args: dict[str, ArgTransform] | None = None,
|
||||
annotations: ToolAnnotations | None | NotSetT = NotSet,
|
||||
output_schema: dict[str, Any] | None | NotSetT | Literal[False] = NotSet,
|
||||
serializer: Callable[[Any], str] | None | NotSetT = NotSet,
|
||||
meta: dict[str, Any] | None | NotSetT = NotSet,
|
||||
annotations: ToolAnnotations | NotSetT | None = NotSet,
|
||||
output_schema: dict[str, Any] | Literal[False] | NotSetT | None = NotSet,
|
||||
serializer: Callable[[Any], str] | NotSetT | None = NotSet,
|
||||
meta: dict[str, Any] | NotSetT | None = NotSet,
|
||||
enabled: bool | None = None,
|
||||
) -> TransformedTool:
|
||||
"""Create a transformed tool from a parent tool.
|
||||
|
|
|
|||
|
|
@ -240,12 +240,11 @@ def log_server_banner(
|
|||
info_table.add_row("📦", "Transport:", display_transport)
|
||||
|
||||
# Show connection info based on transport
|
||||
if transport in ("http", "streamable-http", "sse"):
|
||||
if host and port:
|
||||
server_url = f"http://{host}:{port}"
|
||||
if path:
|
||||
server_url += f"/{path.lstrip('/')}"
|
||||
info_table.add_row("🔗", "Server URL:", server_url)
|
||||
if transport in ("http", "streamable-http", "sse") and host and port:
|
||||
server_url = f"http://{host}:{port}"
|
||||
if path:
|
||||
server_url += f"/{path.lstrip('/')}"
|
||||
info_table.add_row("🔗", "Server URL:", server_url)
|
||||
|
||||
# Add documentation link
|
||||
info_table.add_row("", "", "")
|
||||
|
|
|
|||
|
|
@ -412,7 +412,7 @@ class InspectFormat(str, Enum):
|
|||
MCP = "mcp"
|
||||
|
||||
|
||||
async def format_fastmcp_info(info: FastMCPInfo) -> bytes:
|
||||
def format_fastmcp_info(info: FastMCPInfo) -> bytes:
|
||||
"""Format FastMCPInfo as FastMCP-specific JSON.
|
||||
|
||||
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
|
||||
if info is None:
|
||||
info = await inspect_fastmcp(mcp)
|
||||
return await format_fastmcp_info(info)
|
||||
return format_fastmcp_info(info)
|
||||
else:
|
||||
raise ValueError(f"Unknown format: {format}")
|
||||
|
|
|
|||
|
|
@ -61,7 +61,7 @@ from pydantic import (
|
|||
)
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
__all__ = ["json_schema_to_type", "JSONSchema"]
|
||||
__all__ = ["JSONSchema", "json_schema_to_type"]
|
||||
|
||||
|
||||
FORMAT_TYPES: dict[str, Any] = {
|
||||
|
|
@ -368,7 +368,7 @@ def _schema_to_type(
|
|||
return types[0]
|
||||
else:
|
||||
if has_null:
|
||||
return Union[tuple(types + [type(None)])] # type: ignore # noqa: UP007
|
||||
return Union[(*types, type(None))] # type: ignore
|
||||
else:
|
||||
return Union[tuple(types)] # type: ignore # noqa: UP007
|
||||
|
||||
|
|
@ -389,7 +389,7 @@ def _schema_to_type(
|
|||
if len(types) == 1:
|
||||
return types[0] | None # type: ignore
|
||||
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 _get_from_type_handler(schema, schemas)(schema)
|
||||
|
|
@ -578,7 +578,7 @@ def _create_dataclass(
|
|||
return _merge_defaults(data, original_schema)
|
||||
return data
|
||||
|
||||
setattr(cls, "_apply_defaults", _apply_defaults)
|
||||
cls._apply_defaults = _apply_defaults # type: ignore[attr-defined]
|
||||
|
||||
# Store completed class
|
||||
_classes[cache_key] = cls
|
||||
|
|
|
|||
|
|
@ -147,6 +147,18 @@ def temporary_log_level(
|
|||
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):
|
||||
min_level: tuple[int, str] | None
|
||||
max_level: tuple[int, str] | None
|
||||
|
|
@ -161,29 +173,13 @@ class _ClampedLogFilter(logging.Filter):
|
|||
self.min_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))
|
||||
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))
|
||||
|
||||
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
|
||||
def filter(self, record: logging.LogRecord) -> bool:
|
||||
if self.max_level:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
__all__ = [
|
||||
"Source",
|
||||
"Deployment",
|
||||
"Environment",
|
||||
"UVEnvironment",
|
||||
"MCPServerConfig",
|
||||
"FileSystemSource",
|
||||
"MCPServerConfig",
|
||||
"Source",
|
||||
"UVEnvironment",
|
||||
"generate_schema",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -19,7 +19,6 @@ class Environment(BaseModel, ABC):
|
|||
Returns:
|
||||
Full command ready for subprocess execution
|
||||
"""
|
||||
pass
|
||||
|
||||
async def prepare(self, output_dir: Path | None = None) -> None:
|
||||
"""Prepare the environment (optional, can be no-op).
|
||||
|
|
@ -27,4 +26,4 @@ class Environment(BaseModel, ABC):
|
|||
Args:
|
||||
output_dir: Directory for persistent environment setup
|
||||
"""
|
||||
pass # Default no-op implementation
|
||||
# Default no-op implementation
|
||||
|
|
|
|||
|
|
@ -17,7 +17,6 @@ class Source(BaseModel, ABC):
|
|||
need preparation (e.g., local files), this is a no-op.
|
||||
"""
|
||||
# Default implementation for sources that don't need preparation
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def load_server(self) -> Any:
|
||||
|
|
|
|||
|
|
@ -175,16 +175,16 @@ class HTTPRoute(FastMCPBaseModel):
|
|||
# Export public symbols
|
||||
__all__ = [
|
||||
"HTTPRoute",
|
||||
"HttpMethod",
|
||||
"JsonSchema",
|
||||
"ParameterInfo",
|
||||
"ParameterLocation",
|
||||
"RequestBodyInfo",
|
||||
"ResponseInfo",
|
||||
"HttpMethod",
|
||||
"ParameterLocation",
|
||||
"JsonSchema",
|
||||
"parse_openapi_to_http_routes",
|
||||
"_handle_nullable_fields",
|
||||
"extract_output_schema_from_responses",
|
||||
"format_deep_object_parameter",
|
||||
"_handle_nullable_fields",
|
||||
"parse_openapi_to_http_routes",
|
||||
]
|
||||
|
||||
# Type variables for generic parser
|
||||
|
|
@ -321,7 +321,7 @@ class OpenAPIParser(
|
|||
else:
|
||||
# Special handling for components
|
||||
if part == "components" and hasattr(target, "components"):
|
||||
target = getattr(target, "components")
|
||||
target = target.components
|
||||
elif hasattr(target, part): # Fallback check
|
||||
target = getattr(target, part, None)
|
||||
else:
|
||||
|
|
@ -1178,10 +1178,10 @@ def _add_null_to_type(schema: dict[str, Any]) -> None:
|
|||
elif isinstance(current_type, list):
|
||||
# Add null to array if not already present
|
||||
if "null" not in current_type:
|
||||
schema["type"] = current_type + ["null"]
|
||||
schema["type"] = [*current_type, "null"]
|
||||
elif "oneOf" in schema:
|
||||
# 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:
|
||||
# Add null type to anyOf if not already present
|
||||
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
|
||||
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:
|
||||
nullable_value = prop_schema.pop("nullable")
|
||||
if nullable_value and (
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ import multiprocessing
|
|||
import socket
|
||||
import time
|
||||
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 urllib.parse import parse_qs, urlparse
|
||||
|
||||
|
|
@ -216,10 +216,8 @@ async def run_server_async(
|
|||
finally:
|
||||
# Cleanup: cancel the task
|
||||
server_task.cancel()
|
||||
try:
|
||||
with suppress(asyncio.CancelledError):
|
||||
await server_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
|
||||
@contextmanager
|
||||
|
|
|
|||
|
|
@ -887,7 +887,7 @@ class TestIconExtraction:
|
|||
return "icon"
|
||||
|
||||
info = await inspect_fastmcp(mcp)
|
||||
json_bytes = await format_fastmcp_info(info)
|
||||
json_bytes = format_fastmcp_info(info)
|
||||
|
||||
import json
|
||||
|
||||
|
|
@ -915,7 +915,7 @@ class TestIconExtraction:
|
|||
return "none"
|
||||
|
||||
info = await inspect_fastmcp(mcp)
|
||||
json_bytes = await format_fastmcp_info(info)
|
||||
json_bytes = format_fastmcp_info(info)
|
||||
|
||||
import json
|
||||
|
||||
|
|
@ -945,7 +945,7 @@ class TestFormatFunctions:
|
|||
return {"result": x * 2}
|
||||
|
||||
info = await inspect_fastmcp(mcp)
|
||||
json_bytes = await format_fastmcp_info(info)
|
||||
json_bytes = format_fastmcp_info(info)
|
||||
|
||||
# Verify it's valid JSON
|
||||
import json
|
||||
|
|
@ -1104,7 +1104,7 @@ class TestFormatFunctions:
|
|||
assert "result" in info.tools[0].output_schema["properties"]
|
||||
|
||||
# Verify it's included in FastMCP format
|
||||
json_bytes = await format_fastmcp_info(info)
|
||||
json_bytes = format_fastmcp_info(info)
|
||||
import json
|
||||
|
||||
data = json.loads(json_bytes)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue