Feat: Configurable LoggingMiddleware payload serialization

This commit is contained in:
William Easton 2025-08-27 10:31:08 -05:00 committed by GitHub
commit 5b433f57cf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 361 additions and 76 deletions

View file

@ -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=<non-serializable>")
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"] = "<non-serializable>"
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:

View file

@ -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