mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 15:19:10 +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
|
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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -48,9 +48,9 @@ def __getattr__(name: str):
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"FastMCP",
|
|
||||||
"Context",
|
|
||||||
"client",
|
|
||||||
"Client",
|
"Client",
|
||||||
|
"Context",
|
||||||
|
"FastMCP",
|
||||||
|
"client",
|
||||||
"settings",
|
"settings",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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[
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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"]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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",
|
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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"]):
|
||||||
|
|
|
||||||
|
|
@ -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",
|
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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}")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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",
|
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -3,4 +3,4 @@ from .context import Context
|
||||||
from . import dependencies
|
from . import dependencies
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["FastMCP", "Context"]
|
__all__ = ["Context", "FastMCP"]
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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):
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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.")
|
||||||
|
|
|
||||||
|
|
@ -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__)
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,7 @@ from .middleware import (
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"CallNext",
|
||||||
"Middleware",
|
"Middleware",
|
||||||
"MiddlewareContext",
|
"MiddlewareContext",
|
||||||
"CallNext",
|
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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]:
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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__)
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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):
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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"]
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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("", "", "")
|
||||||
|
|
|
||||||
|
|
@ -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}")
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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 (
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue