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..af5d8d5c2 100644 --- a/src/fastmcp/server/middleware/error_handling.py +++ b/src/fastmcp/server/middleware/error_handling.py @@ -38,7 +38,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. diff --git a/src/fastmcp/settings.py b/src/fastmcp/settings.py index 393b1ff2f..0b76f2a48 100644 --- a/src/fastmcp/settings.py +++ b/src/fastmcp/settings.py @@ -258,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" 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