diff --git a/src/fastmcp/server/middleware/logging.py b/src/fastmcp/server/middleware/logging.py index f770e2faa..3511a3388 100644 --- a/src/fastmcp/server/middleware/logging.py +++ b/src/fastmcp/server/middleware/logging.py @@ -2,11 +2,20 @@ import json import logging +from collections.abc import Callable +from logging import Logger from typing import Any +import pydantic_core + from .middleware import CallNext, Middleware, MiddlewareContext +def default_serializer(data: Any) -> str: + """The default serializer for Payloads in the logging middleware.""" + return pydantic_core.to_json(data, fallback=str).decode() + + class LoggingMiddleware(Middleware): """Middleware that provides comprehensive request and response logging. @@ -33,6 +42,7 @@ class LoggingMiddleware(Middleware): include_payloads: bool = False, max_payload_length: int = 1000, methods: list[str] | None = None, + payload_serializer: Callable[[Any], str] | None = None, ): """Initialize logging middleware. @@ -43,13 +53,14 @@ class LoggingMiddleware(Middleware): max_payload_length: Maximum length of payload to log (prevents huge logs) methods: List of methods to log. If None, logs all methods. """ - self.logger = logger or logging.getLogger("fastmcp.requests") - self.log_level = log_level - self.include_payloads = include_payloads - self.max_payload_length = max_payload_length - self.methods = methods + self.logger: Logger = logger or logging.getLogger("fastmcp.requests") + self.log_level: int = log_level + self.include_payloads: bool = include_payloads + self.max_payload_length: int = max_payload_length + self.methods: list[str] | None = methods + self.payload_serializer: Callable[[Any], str] | None = payload_serializer - def _format_message(self, context: MiddlewareContext) -> str: + def _format_message(self, context: MiddlewareContext[Any]) -> str: """Format a message for logging.""" parts = [ f"source={context.source}", @@ -57,18 +68,29 @@ class LoggingMiddleware(Middleware): f"method={context.method or 'unknown'}", ] - if self.include_payloads and hasattr(context.message, "__dict__"): - try: - payload = json.dumps(context.message.__dict__, default=str) - if len(payload) > self.max_payload_length: - payload = payload[: self.max_payload_length] + "..." - parts.append(f"payload={payload}") - except (TypeError, ValueError): - parts.append("payload=") + if self.include_payloads: + payload: str + if not self.payload_serializer: + payload = default_serializer(context.message) + else: + try: + payload = self.payload_serializer(context.message) + except Exception as e: + self.logger.warning( + f"Failed {e} to serialize payload: {context.type} {context.method} {context.source}." + ) + payload = default_serializer(context.message) + + if len(payload) > self.max_payload_length: + payload = payload[: self.max_payload_length] + "..." + + parts.append(f"payload={payload}") return " ".join(parts) - async def on_message(self, context: MiddlewareContext, call_next: CallNext) -> Any: + async def on_message( + self, context: MiddlewareContext[Any], call_next: CallNext[Any, Any] + ) -> Any: """Log all messages.""" message_info = self._format_message(context) if self.methods and context.method not in self.methods: @@ -111,6 +133,7 @@ class StructuredLoggingMiddleware(Middleware): log_level: int = logging.INFO, include_payloads: bool = False, methods: list[str] | None = None, + payload_serializer: Callable[[Any], str] | None = None, ): """Initialize structured logging middleware. @@ -119,15 +142,18 @@ class StructuredLoggingMiddleware(Middleware): log_level: Log level for messages (default: INFO) include_payloads: Whether to include message payloads in logs methods: List of methods to log. If None, logs all methods. + serializer: Callable that converts objects to a JSON string for the + payload. If not provided, uses FastMCP's default tool serializer. """ - self.logger = logger or logging.getLogger("fastmcp.structured") - self.log_level = log_level - self.include_payloads = include_payloads - self.methods = methods + self.logger: Logger = logger or logging.getLogger("fastmcp.structured") + self.log_level: int = log_level + self.include_payloads: bool = include_payloads + self.methods: list[str] | None = methods + self.payload_serializer: Callable[[Any], str] | None = payload_serializer def _create_log_entry( - self, context: MiddlewareContext, event: str, **extra_fields - ) -> dict: + self, context: MiddlewareContext[Any], event: str, **extra_fields: Any + ) -> dict[str, Any]: """Create a structured log entry.""" entry = { "event": event, @@ -138,15 +164,27 @@ class StructuredLoggingMiddleware(Middleware): **extra_fields, } - if self.include_payloads and hasattr(context.message, "__dict__"): - try: - entry["payload"] = context.message.__dict__ - except (TypeError, ValueError): - entry["payload"] = "" + if self.include_payloads: + payload: str + + if not self.payload_serializer: + payload = default_serializer(context.message) + else: + try: + payload = self.payload_serializer(context.message) + except Exception as e: + self.logger.warning( + f"Failed {str(e)} to serialize payload: {context.type} {context.method} {context.source}." + ) + payload = default_serializer(context.message) + + entry["payload"] = payload return entry - async def on_message(self, context: MiddlewareContext, call_next: CallNext) -> Any: + async def on_message( + self, context: MiddlewareContext[Any], call_next: CallNext[Any, Any] + ) -> Any: """Log structured message information.""" start_entry = self._create_log_entry(context, "request_start") if self.methods and context.method not in self.methods: diff --git a/tests/server/middleware/test_logging.py b/tests/server/middleware/test_logging.py index 9217bb485..d0118f219 100644 --- a/tests/server/middleware/test_logging.py +++ b/tests/server/middleware/test_logging.py @@ -1,30 +1,58 @@ """Tests for logging middleware.""" +import datetime import json import logging +from typing import Any, Literal, TypeVar from unittest.mock import AsyncMock, MagicMock +import mcp import pytest +from inline_snapshot import snapshot +from pydantic import AnyUrl +from fastmcp.resources.template import ResourceTemplate from fastmcp.server.middleware.logging import ( LoggingMiddleware, StructuredLoggingMiddleware, ) -from fastmcp.server.middleware.middleware import MiddlewareContext +from fastmcp.server.middleware.middleware import CallNext, MiddlewareContext +from fastmcp.server.server import FastMCP + +FIXED_DATE = datetime.datetime(2023, 1, 1, tzinfo=datetime.timezone.utc) + +T = TypeVar("T") + + +def new_mock_context( + message: T, + method: str | None = None, + source: Literal["server", "client"] | None = None, + type: Literal["request", "notification"] | None = None, +) -> MiddlewareContext[T]: + """Create a new mock middleware context.""" + context = MagicMock(spec=MiddlewareContext[T]) + context.method = method or "test_method" + context.source = source or "client" + context.type = type or "request" + context.message = message + context.timestamp = FIXED_DATE + return context @pytest.fixture def mock_context(): """Create a mock middleware context.""" - context = MagicMock(spec=MiddlewareContext) - context.method = "test_method" - context.source = "client" - context.type = "request" - context.message = MagicMock() - context.message.__dict__ = {"param": "value"} - context.timestamp = MagicMock() - context.timestamp.isoformat.return_value = "2023-01-01T00:00:00Z" - return context + + return new_mock_context( + message=mcp.types.CallToolRequest( + method="tools/call", + params=mcp.types.CallToolRequestParams( + name="test_method", + arguments={"param": "value"}, + ), + ) + ) @pytest.fixture @@ -58,7 +86,9 @@ class TestLoggingMiddleware: assert middleware.include_payloads is True assert middleware.max_payload_length == 500 - def test_format_message_without_payloads(self, mock_context): + def test_format_message_without_payloads( + self, mock_context: MiddlewareContext[Any] + ): """Test message formatting without payloads.""" middleware = LoggingMiddleware() formatted = middleware._format_message(mock_context) @@ -68,17 +98,16 @@ class TestLoggingMiddleware: assert "method=test_method" in formatted assert "payload=" not in formatted - def test_format_message_with_payloads(self, mock_context): + def test_format_message_with_payloads(self, mock_context: MiddlewareContext[Any]): """Test message formatting with payloads.""" middleware = LoggingMiddleware(include_payloads=True) formatted = middleware._format_message(mock_context) - assert "source=client" in formatted - assert "type=request" in formatted - assert "method=test_method" in formatted - assert 'payload={"param": "value"}' in formatted + assert formatted == snapshot( + 'source=client type=request method=test_method payload={"method":"tools/call","params":{"_meta":null,"name":"test_method","arguments":{"param":"value"}}}' + ) - def test_format_message_long_payload(self, mock_context): + def test_format_message_long_payload(self, mock_context: MiddlewareContext[Any]): """Test message formatting with long payload truncation.""" middleware = LoggingMiddleware(include_payloads=True, max_payload_length=10) formatted = middleware._format_message(mock_context) @@ -86,7 +115,12 @@ class TestLoggingMiddleware: assert "payload=" in formatted assert "..." in formatted - async def test_on_message_success(self, mock_context, mock_call_next, caplog): + async def test_on_message_success( + self, + mock_context: MiddlewareContext[Any], + mock_call_next: CallNext[Any, Any], + caplog: pytest.LogCaptureFixture, + ): """Test logging successful messages.""" middleware = LoggingMiddleware() @@ -98,7 +132,9 @@ class TestLoggingMiddleware: assert "Processing message:" in caplog.text assert "Completed message: test_method" in caplog.text - async def test_on_message_failure(self, mock_context, caplog): + async def test_on_message_failure( + self, mock_context: MiddlewareContext[Any], caplog: pytest.LogCaptureFixture + ): """Test logging failed messages.""" middleware = LoggingMiddleware() mock_call_next = AsyncMock(side_effect=ValueError("test error")) @@ -121,26 +157,40 @@ class TestStructuredLoggingMiddleware: assert middleware.log_level == logging.INFO assert middleware.include_payloads is False - def test_create_log_entry_basic(self, mock_context): + def test_create_log_entry_basic(self, mock_context: MiddlewareContext[Any]): """Test creating basic log entry.""" middleware = StructuredLoggingMiddleware() entry = middleware._create_log_entry(mock_context, "test_event") - assert entry["event"] == "test_event" - assert entry["timestamp"] == "2023-01-01T00:00:00Z" - assert entry["source"] == "client" - assert entry["type"] == "request" - assert entry["method"] == "test_method" - assert "payload" not in entry + assert entry == snapshot( + { + "event": "test_event", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + } + ) - def test_create_log_entry_with_payload(self, mock_context): + def test_create_log_entry_with_payload(self, mock_context: MiddlewareContext[Any]): """Test creating log entry with payload.""" middleware = StructuredLoggingMiddleware(include_payloads=True) entry = middleware._create_log_entry(mock_context, "test_event") - assert entry["payload"] == {"param": "value"} + assert entry == snapshot( + { + "event": "test_event", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "payload": '{"method":"tools/call","params":{"_meta":null,"name":"test_method","arguments":{"param":"value"}}}', + } + ) - def test_create_log_entry_with_extra_fields(self, mock_context): + def test_create_log_entry_with_extra_fields( + self, mock_context: MiddlewareContext[Any] + ): """Test creating log entry with extra fields.""" middleware = StructuredLoggingMiddleware() entry = middleware._create_log_entry( @@ -149,7 +199,12 @@ class TestStructuredLoggingMiddleware: assert entry["extra_field"] == "extra_value" - async def test_on_message_success(self, mock_context, mock_call_next, caplog): + async def test_on_message_success( + self, + mock_context: MiddlewareContext[Any], + mock_call_next: CallNext[Any, Any], + caplog: pytest.LogCaptureFixture, + ): """Test structured logging of successful messages.""" middleware = StructuredLoggingMiddleware() @@ -160,17 +215,33 @@ class TestStructuredLoggingMiddleware: # Check that we have structured JSON logs log_lines = [record.message for record in caplog.records] + assert len(log_lines) == 2 # start and success entries - start_entry = json.loads(log_lines[0]) - assert start_entry["event"] == "request_start" - assert start_entry["method"] == "test_method" + assert json.loads(log_lines[0]) == snapshot( + { + "event": "request_start", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + } + ) - success_entry = json.loads(log_lines[1]) - assert success_entry["event"] == "request_success" - assert success_entry["result_type"] == "str" + assert json.loads(log_lines[1]) == snapshot( + { + "event": "request_success", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "result_type": "str", + } + ) - async def test_on_message_failure(self, mock_context, caplog): + async def test_on_message_failure( + self, mock_context: MiddlewareContext[Any], caplog: pytest.LogCaptureFixture + ): """Test structured logging of failed messages.""" middleware = StructuredLoggingMiddleware() mock_call_next = AsyncMock(side_effect=ValueError("test error")) @@ -186,10 +257,180 @@ class TestStructuredLoggingMiddleware: start_entry = json.loads(log_lines[0]) assert start_entry["event"] == "request_start" - error_entry = json.loads(log_lines[1]) - assert error_entry["event"] == "request_error" - assert error_entry["error_type"] == "ValueError" - assert error_entry["error_message"] == "test error" + assert json.loads(log_lines[1]) == snapshot( + { + "event": "request_error", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "error_type": "ValueError", + "error_message": "test error", + } + ) + + async def test_on_message_with_pydantic_types_in_payload( + self, + mock_call_next: CallNext[Any, Any], + caplog: pytest.LogCaptureFixture, + ): + """Ensure Pydantic AnyUrl in payload serializes correctly when include_payloads=True.""" + + mock_context = new_mock_context( + message=mcp.types.ReadResourceRequest( + method="resources/read", + params=mcp.types.ReadResourceRequestParams( + uri=AnyUrl("test://example/1"), + ), + ) + ) + + middleware = StructuredLoggingMiddleware(include_payloads=True) + + with caplog.at_level(logging.INFO): + result = await middleware.on_message(mock_context, mock_call_next) + + assert result == "test_result" + + log_lines = [record.message for record in caplog.records] + + assert len(log_lines) == 2 + assert json.loads(log_lines[0]) == snapshot( + { + "event": "request_start", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "payload": '{"method":"resources/read","params":{"_meta":null,"uri":"test://example/1"}}', + } + ) + assert json.loads(log_lines[1]) == snapshot( + { + "event": "request_success", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "result_type": "str", + "payload": '{"method":"resources/read","params":{"_meta":null,"uri":"test://example/1"}}', + } + ) + + async def test_on_message_with_resource_template_in_payload( + self, + mock_call_next: CallNext[Any, Any], + caplog: pytest.LogCaptureFixture, + ): + """Ensure ResourceTemplate in payload serializes via pydantic conversion without errors.""" + + mock_context = new_mock_context( + message=ResourceTemplate( + name="tmpl", + uri_template="tmpl://{id}", + parameters={"id": {"type": "string"}}, + ) + ) + + middleware = StructuredLoggingMiddleware(include_payloads=True) + + with caplog.at_level(logging.INFO): + result = await middleware.on_message(mock_context, mock_call_next) + + assert result == "test_result" + + log_lines = [record.message for record in caplog.records] + assert len(log_lines) == 2 + assert json.loads(log_lines[0]) == snapshot( + { + "event": "request_start", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "payload": '{"name":"tmpl","title":null,"description":null,"tags":[],"meta":null,"enabled":true,"uri_template":"tmpl://{id}","mime_type":"text/plain","parameters":{"id":{"type":"string"}},"annotations":null}', + } + ) + + async def test_on_message_with_nonserializable_payload_falls_back_to_str( + self, mock_call_next: CallNext[Any, Any], caplog: pytest.LogCaptureFixture + ): + """Ensure non-JSONable objects fall back to string serialization in payload.""" + + class NonSerializable: + def __str__(self) -> str: + return "NON_SERIALIZABLE" + + mock_context = new_mock_context( + message=mcp.types.CallToolRequest( + method="tools/call", + params=mcp.types.CallToolRequestParams( + name="test_method", + arguments={"obj": NonSerializable()}, + ), + ) + ) + + middleware = StructuredLoggingMiddleware(include_payloads=True) + + with caplog.at_level(logging.INFO): + result = await middleware.on_message(mock_context, mock_call_next) + + assert result == "test_result" + + log_lines = [record.message for record in caplog.records] + assert len(log_lines) >= 2 + assert json.loads(log_lines[0]) == snapshot( + { + "event": "request_start", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "payload": '{"method":"tools/call","params":{"_meta":null,"name":"test_method","arguments":{"obj":"NON_SERIALIZABLE"}}}', + } + ) + + async def test_on_message_with_custom_serializer_applied( + self, mock_call_next: CallNext[Any, Any], caplog: pytest.LogCaptureFixture + ): + """Ensure a custom serializer is used for non-JSONable payloads.""" + + # Provide a serializer that replaces entire payload with a fixed string + def custom_serializer(_: Any) -> str: + return "CUSTOM_PAYLOAD" + + mock_context = new_mock_context( + message=mcp.types.CallToolRequest( + method="tools/call", + params=mcp.types.CallToolRequestParams( + name="test_method", + arguments={"obj": "OBJECT"}, + ), + ) + ) + + middleware = StructuredLoggingMiddleware( + include_payloads=True, payload_serializer=custom_serializer + ) + + with caplog.at_level(logging.INFO): + result = await middleware.on_message(mock_context, mock_call_next) + + assert result == "test_result" + + log_lines = [record.message for record in caplog.records] + assert len(log_lines) >= 2 + assert json.loads(log_lines[0]) == snapshot( + { + "event": "request_start", + "timestamp": "2023-01-01T00:00:00+00:00", + "source": "client", + "type": "request", + "method": "test_method", + "payload": "CUSTOM_PAYLOAD", + } + ) @pytest.fixture @@ -233,7 +474,7 @@ class TestLoggingMiddlewareIntegration: """Integration tests for logging middleware with real FastMCP server.""" async def test_logging_middleware_logs_successful_operations( - self, logging_server, caplog + self, logging_server: FastMCP, caplog: pytest.LogCaptureFixture ): """Test that logging middleware captures successful operations.""" from fastmcp.client import Client @@ -259,7 +500,9 @@ class TestLoggingMiddlewareIntegration: assert processing_count == 2 assert completion_count == 2 - async def test_logging_middleware_logs_failures(self, logging_server, caplog): + async def test_logging_middleware_logs_failures( + self, logging_server: FastMCP, caplog: pytest.LogCaptureFixture + ): """Test that logging middleware captures failed operations.""" from fastmcp.client import Client @@ -279,7 +522,9 @@ class TestLoggingMiddlewareIntegration: assert "Processing message:" in log_text assert "Failed message: tools/call" in log_text - async def test_logging_middleware_with_payloads(self, logging_server, caplog): + async def test_logging_middleware_with_payloads( + self, logging_server: FastMCP, caplog: pytest.LogCaptureFixture + ): """Test logging middleware when configured to include payloads.""" from fastmcp.client import Client @@ -300,7 +545,7 @@ class TestLoggingMiddlewareIntegration: assert "payload=" in log_text async def test_structured_logging_middleware_produces_json( - self, logging_server, caplog + self, logging_server: FastMCP, caplog: pytest.LogCaptureFixture ): """Test that structured logging middleware produces parseable JSON logs.""" import json @@ -334,7 +579,7 @@ class TestLoggingMiddlewareIntegration: assert "method" in log_entry async def test_structured_logging_middleware_handles_errors( - self, logging_server, caplog + self, logging_server: FastMCP, caplog: pytest.LogCaptureFixture ): """Test structured logging of errors with JSON format.""" import json @@ -375,7 +620,7 @@ class TestLoggingMiddlewareIntegration: assert "error_message" in error_entry async def test_logging_middleware_with_different_operations( - self, logging_server, caplog + self, logging_server: FastMCP, caplog: pytest.LogCaptureFixture ): """Test logging middleware with various MCP operations.""" from fastmcp.client import Client @@ -410,7 +655,9 @@ class TestLoggingMiddlewareIntegration: assert processing_count == 4 assert completion_count == 4 - async def test_logging_middleware_custom_configuration(self, logging_server): + async def test_logging_middleware_custom_configuration( + self, logging_server: FastMCP + ): """Test logging middleware with custom logger configuration.""" import io import logging