From 4cfeb63b958d433038f5069cb6c66de99f2042e8 Mon Sep 17 00:00:00 2001 From: William Easton Date: Mon, 13 Apr 2026 23:11:11 -0500 Subject: [PATCH 1/2] Fix input validation for ports, Optional[bool] form backfill, async error callbacks MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 🤖 Generated with Claude Code Co-Authored-By: Claude Opus 4.6 (1M context) --- src/fastmcp/apps/form.py | 14 +++++- .../server/middleware/error_handling.py | 11 +++-- src/fastmcp/settings.py | 7 +++ tests/apps/test_form.py | 21 +++++++++ .../server/middleware/test_error_handling.py | 45 ++++++++++++++----- tests/test_settings.py | 29 ++++++++++++ 6 files changed, 110 insertions(+), 17 deletions(-) create mode 100644 tests/test_settings.py diff --git a/src/fastmcp/apps/form.py b/src/fastmcp/apps/form.py index e5e6d1151..159e10770 100644 --- a/src/fastmcp/apps/form.py +++ b/src/fastmcp/apps/form.py @@ -25,6 +25,8 @@ Usage:: from __future__ import annotations import json +import types +import typing from collections.abc import Callable from typing import Any @@ -54,6 +56,16 @@ import pydantic from fastmcp.apps.app import FastMCPApp +def _is_bool_type(annotation: Any) -> bool: + """Check if annotation is bool or Optional[bool].""" + if annotation is bool: + return True + origin = typing.get_origin(annotation) + if origin is typing.Union or isinstance(annotation, types.UnionType): + return bool in typing.get_args(annotation) + return False + + def _backfill_boolean_defaults( model: type[pydantic.BaseModel], data: dict[str, Any], @@ -67,7 +79,7 @@ def _backfill_boolean_defaults( for name, field_info in model.model_fields.items(): if name in data: continue - if field_info.annotation is bool: + if _is_bool_type(field_info.annotation): if field_info.default is not pydantic.fields.PydanticUndefined: data[name] = field_info.default else: diff --git a/src/fastmcp/server/middleware/error_handling.py b/src/fastmcp/server/middleware/error_handling.py index 81a13bb84..4bd69e252 100644 --- a/src/fastmcp/server/middleware/error_handling.py +++ b/src/fastmcp/server/middleware/error_handling.py @@ -1,6 +1,7 @@ """Error handling middleware for consistent error responses and tracking.""" import asyncio +import inspect import logging import traceback from collections.abc import Callable @@ -38,7 +39,7 @@ class ErrorHandlingMiddleware(Middleware): self, logger: logging.Logger | None = None, include_traceback: bool = False, - error_callback: Callable[[Exception, MiddlewareContext], None] | None = None, + error_callback: Callable[[Exception, MiddlewareContext], Any] | None = None, transform_errors: bool = True, ): """Initialize error handling middleware. @@ -55,7 +56,7 @@ class ErrorHandlingMiddleware(Middleware): self.transform_errors = transform_errors self.error_counts = {} - def _log_error(self, error: Exception, context: MiddlewareContext) -> None: + async def _log_error(self, error: Exception, context: MiddlewareContext) -> None: """Log error with appropriate detail level.""" error_type = type(error).__name__ method = context.method or "unknown" @@ -74,7 +75,9 @@ class ErrorHandlingMiddleware(Middleware): # Call custom error callback if provided if self.error_callback: try: - self.error_callback(error, context) + result = self.error_callback(error, context) + if inspect.isawaitable(result): + await result except Exception as callback_error: self.logger.error(f"Error in error callback: {callback_error}") @@ -122,7 +125,7 @@ class ErrorHandlingMiddleware(Middleware): try: return await call_next(context) except Exception as error: - self._log_error(error, context) + await self._log_error(error, context) # Transform and re-raise transformed_error = self._transform_error(error, context) diff --git a/src/fastmcp/settings.py b/src/fastmcp/settings.py index 393b1ff2f..dbb2d5b54 100644 --- a/src/fastmcp/settings.py +++ b/src/fastmcp/settings.py @@ -185,6 +185,13 @@ class Settings(BaseSettings): return v.upper() return v + @field_validator("port") + @classmethod + def _validate_port(cls, v: int) -> int: + if not (1 <= v <= 65535): + raise ValueError(f"Port must be between 1 and 65535, got {v}") + return v + docket: DocketSettings = DocketSettings() enable_rich_logging: Annotated[ diff --git a/tests/apps/test_form.py b/tests/apps/test_form.py index 5d3af18af..b42010cd7 100644 --- a/tests/apps/test_form.py +++ b/tests/apps/test_form.py @@ -168,3 +168,24 @@ class TestBackfillBooleanDefaults: data = {"title": "Note"} result = _backfill_boolean_defaults(NoteForm, data) assert "content" not in result + + def test_optional_bool_backfilled(self): + class ModelWithOptionalBool(pydantic.BaseModel): + flag: bool = False + opt: bool | None = False + + data: dict = {} + _backfill_boolean_defaults(ModelWithOptionalBool, data) + assert "flag" in data + assert "opt" in data + assert data["flag"] is False + assert data["opt"] is False + + def test_optional_bool_without_default(self): + class ModelWithRequiredOptionalBool(pydantic.BaseModel): + opt: bool | None + + data: dict = {} + _backfill_boolean_defaults(ModelWithRequiredOptionalBool, data) + assert "opt" in data + assert data["opt"] is False diff --git a/tests/server/middleware/test_error_handling.py b/tests/server/middleware/test_error_handling.py index dfc551fbf..6e70df7d1 100644 --- a/tests/server/middleware/test_error_handling.py +++ b/tests/server/middleware/test_error_handling.py @@ -58,48 +58,48 @@ class TestErrorHandlingMiddleware: assert middleware.error_callback is callback assert middleware.transform_errors is False - def test_log_error_basic(self, mock_context, caplog): + async def test_log_error_basic(self, mock_context, caplog): """Test basic error logging.""" middleware = ErrorHandlingMiddleware() error = ValueError("test error") with caplog.at_level(logging.ERROR): - middleware._log_error(error, mock_context) + await middleware._log_error(error, mock_context) assert "Error in test_method: ValueError: test error" in caplog.text assert "ValueError:test_method" in middleware.error_counts assert middleware.error_counts["ValueError:test_method"] == 1 - def test_log_error_with_traceback(self, mock_context, caplog): + async def test_log_error_with_traceback(self, mock_context, caplog): """Test error logging with traceback.""" middleware = ErrorHandlingMiddleware(include_traceback=True) error = ValueError("test error") with caplog.at_level(logging.ERROR): - middleware._log_error(error, mock_context) + await middleware._log_error(error, mock_context) assert "Error in test_method: ValueError: test error" in caplog.text # The traceback is added to the log message assert "Error in test_method: ValueError: test error" in caplog.text - def test_log_error_with_callback(self, mock_context): + async def test_log_error_with_callback(self, mock_context): """Test error logging with callback.""" callback = MagicMock() middleware = ErrorHandlingMiddleware(error_callback=callback) error = ValueError("test error") - middleware._log_error(error, mock_context) + await middleware._log_error(error, mock_context) callback.assert_called_once_with(error, mock_context) - def test_log_error_callback_exception(self, mock_context, caplog): + async def test_log_error_callback_exception(self, mock_context, caplog): """Test error logging when callback raises exception.""" callback = MagicMock(side_effect=RuntimeError("callback error")) middleware = ErrorHandlingMiddleware(error_callback=callback) error = ValueError("test error") with caplog.at_level(logging.ERROR): - middleware._log_error(error, mock_context) + await middleware._log_error(error, mock_context) assert "Error in error callback: callback error" in caplog.text @@ -236,16 +236,16 @@ class TestErrorHandlingMiddleware: assert "Invalid params: test error" in exc_info.value.error.message assert "Error in test_method: ToolError: test error" in caplog.text - def test_get_error_stats(self, mock_context): + async def test_get_error_stats(self, mock_context): """Test getting error statistics.""" middleware = ErrorHandlingMiddleware() error1 = ValueError("error1") error2 = ValueError("error2") error3 = RuntimeError("error3") - middleware._log_error(error1, mock_context) - middleware._log_error(error2, mock_context) - middleware._log_error(error3, mock_context) + await middleware._log_error(error1, mock_context) + await middleware._log_error(error2, mock_context) + await middleware._log_error(error3, mock_context) stats = middleware.get_error_stats() assert stats["ValueError:test_method"] == 2 @@ -566,6 +566,27 @@ class TestErrorHandlingMiddlewareIntegration: # Error should still exist (may be wrapped by FastMCP) assert exc_info.value is not None + async def test_async_error_callback(self): + """Test that async error callbacks are properly awaited.""" + called = False + + async def callback(error, context): + nonlocal called + called = True + + mcp = FastMCP("test") + mcp.add_middleware(ErrorHandlingMiddleware(error_callback=callback)) + + @mcp.tool + def bad_tool() -> str: + raise ValueError("test") + + async with Client(mcp) as client: + with pytest.raises(Exception): + await client.call_tool("bad_tool") + + assert called, "Async callback was not awaited" + class TestRetryMiddlewareIntegration: """Integration tests for retry middleware with real FastMCP server.""" diff --git a/tests/test_settings.py b/tests/test_settings.py new file mode 100644 index 000000000..ecf2a7b97 --- /dev/null +++ b/tests/test_settings.py @@ -0,0 +1,29 @@ +"""Tests for FastMCP settings validation.""" + +import pytest +from pydantic import ValidationError + +from fastmcp.settings import Settings + + +class TestPortValidation: + def test_invalid_port_negative(self): + with pytest.raises(ValidationError): + Settings(port=-1) + + def test_invalid_port_zero(self): + with pytest.raises(ValidationError): + Settings(port=0) + + def test_invalid_port_too_high(self): + with pytest.raises(ValidationError): + Settings(port=99999) + + def test_valid_port_accepted(self): + assert Settings(port=8080).port == 8080 + + def test_valid_port_min(self): + assert Settings(port=1).port == 1 + + def test_valid_port_max(self): + assert Settings(port=65535).port == 65535 From 212425a13d79afebaa1b208d86eaf26627e72af2 Mon Sep 17 00:00:00 2001 From: strawgate Date: Tue, 14 Apr 2026 00:18:30 -0500 Subject: [PATCH 2/2] refactor: use Pydantic Field for port, revert async error callback --- src/fastmcp/server/middleware/error_handling.py | 9 +++------ src/fastmcp/settings.py | 9 +-------- 2 files changed, 4 insertions(+), 14 deletions(-) diff --git a/src/fastmcp/server/middleware/error_handling.py b/src/fastmcp/server/middleware/error_handling.py index 4bd69e252..af5d8d5c2 100644 --- a/src/fastmcp/server/middleware/error_handling.py +++ b/src/fastmcp/server/middleware/error_handling.py @@ -1,7 +1,6 @@ """Error handling middleware for consistent error responses and tracking.""" import asyncio -import inspect import logging import traceback from collections.abc import Callable @@ -56,7 +55,7 @@ class ErrorHandlingMiddleware(Middleware): self.transform_errors = transform_errors self.error_counts = {} - async def _log_error(self, error: Exception, context: MiddlewareContext) -> None: + def _log_error(self, error: Exception, context: MiddlewareContext) -> None: """Log error with appropriate detail level.""" error_type = type(error).__name__ method = context.method or "unknown" @@ -75,9 +74,7 @@ class ErrorHandlingMiddleware(Middleware): # Call custom error callback if provided if self.error_callback: try: - result = self.error_callback(error, context) - if inspect.isawaitable(result): - await result + self.error_callback(error, context) except Exception as callback_error: self.logger.error(f"Error in error callback: {callback_error}") @@ -125,7 +122,7 @@ class ErrorHandlingMiddleware(Middleware): try: return await call_next(context) except Exception as error: - await self._log_error(error, context) + self._log_error(error, context) # Transform and re-raise transformed_error = self._transform_error(error, context) diff --git a/src/fastmcp/settings.py b/src/fastmcp/settings.py index dbb2d5b54..0b76f2a48 100644 --- a/src/fastmcp/settings.py +++ b/src/fastmcp/settings.py @@ -185,13 +185,6 @@ class Settings(BaseSettings): return v.upper() return v - @field_validator("port") - @classmethod - def _validate_port(cls, v: int) -> int: - if not (1 <= v <= 65535): - raise ValueError(f"Port must be between 1 and 65535, got {v}") - return v - docket: DocketSettings = DocketSettings() enable_rich_logging: Annotated[ @@ -265,7 +258,7 @@ class Settings(BaseSettings): # HTTP settings host: str = "127.0.0.1" - port: int = 8000 + port: int = Field(default=8000, ge=1, le=65535) sse_path: str = "/sse" message_path: str = "/messages/" streamable_http_path: str = "/mcp"