fastmcp/tests/server/middleware/test_opentelemetry_middleware.py
claude[bot] a9b6e71966 Add trace context propagation to OpenTelemetry middleware
Enable distributed tracing across protocols that don't support HTTP headers (like SSE) by propagating W3C Trace Context through MCP _meta fields.

- Add propagate_context parameter (default: True) to OpenTelemetryMiddleware
- Implement _extract_trace_context() to read traceparent/tracestate from request metadata
- Implement _inject_trace_context() to write trace context to response metadata
- Update all operation handlers to extract parent context and inject into results
- Add comprehensive test coverage for context propagation
- Update documentation with examples and configuration details

This enables trace continuity across MCP calls, allowing clients to link server spans to their traces and propagate context downstream.

Co-authored-by: William Easton <strawgate@users.noreply.github.com>
2025-12-03 23:58:17 +00:00

305 lines
11 KiB
Python

"""Tests for OpenTelemetry middleware."""
import pytest
from fastmcp import Client, FastMCP
class TestOpenTelemetryMiddlewareWithoutOTel:
"""Test OpenTelemetry middleware behavior when opentelemetry is not installed."""
async def test_middleware_no_op_without_opentelemetry(self):
"""Test that middleware works as a no-op when OpenTelemetry is not installed."""
# Import after potentially uninstalling opentelemetry
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.tool()
def test_tool(value: str) -> str:
return f"result: {value}"
# Should work without errors even if OpenTelemetry is not installed
async with Client(mcp) as client:
result = await client.call_tool("test_tool", {"value": "test"})
assert result.content[0].text == "result: test" # type: ignore[attr-defined]
class TestOpenTelemetryMiddlewareConfiguration:
"""Test OpenTelemetry middleware configuration options."""
def test_middleware_can_be_disabled(self):
"""Test that middleware can be explicitly disabled."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
middleware = OpenTelemetryMiddleware(enabled=False)
mcp.add_middleware(middleware)
assert middleware.enabled is False
assert middleware.tracer is None
def test_middleware_respects_include_arguments(self):
"""Test that include_arguments parameter is respected."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
middleware = OpenTelemetryMiddleware(include_arguments=False)
assert middleware.include_arguments is False
middleware = OpenTelemetryMiddleware(include_arguments=True)
assert middleware.include_arguments is True
def test_middleware_respects_max_argument_length(self):
"""Test that max_argument_length parameter is respected."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
middleware = OpenTelemetryMiddleware(max_argument_length=100)
assert middleware.max_argument_length == 100
# Test truncation
long_value = "x" * 200
truncated = middleware._truncate_value(long_value)
assert len(truncated) == 103 # 100 + "..."
assert truncated.endswith("...")
class TestOpenTelemetryMiddlewareOperations:
"""Test that middleware handles different MCP operations."""
async def test_tool_call_without_errors(self):
"""Test that middleware handles tool calls without errors."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.tool()
def test_tool(value: str) -> str:
return f"result: {value}"
async with Client(mcp) as client:
result = await client.call_tool("test_tool", {"value": "test"})
assert result.content[0].text == "result: test" # type: ignore[attr-defined]
async def test_resource_read_without_errors(self):
"""Test that middleware handles resource reads without errors."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.resource("test://resource")
def test_resource() -> str:
return "resource content"
async with Client(mcp) as client:
result = await client.read_resource("test://resource")
assert result[0].text == "resource content" # type: ignore[attr-defined]
async def test_prompt_get_without_errors(self):
"""Test that middleware handles prompt retrieval without errors."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.prompt()
def test_prompt(name: str) -> str:
return f"Hello, {name}!"
async with Client(mcp) as client:
result = await client.get_prompt("test_prompt", {"name": "World"})
assert any("Hello, World!" in str(msg) for msg in result.messages)
async def test_list_tools_without_errors(self):
"""Test that middleware handles list tools without errors."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.tool()
def test_tool() -> str:
return "test"
async with Client(mcp) as client:
result = await client.list_tools()
assert len(result) == 1
assert result[0].name == "test_tool"
async def test_list_resources_without_errors(self):
"""Test that middleware handles list resources without errors."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.resource("test://resource")
def test_resource() -> str:
return "test"
async with Client(mcp) as client:
result = await client.list_resources()
assert len(result) == 1
assert str(result[0].uri) == "test://resource"
async def test_list_prompts_without_errors(self):
"""Test that middleware handles list prompts without errors."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.prompt()
def test_prompt() -> str:
return "test"
async with Client(mcp) as client:
result = await client.list_prompts()
assert len(result) == 1
assert result[0].name == "test_prompt"
class TestOpenTelemetryMiddlewareErrorHandling:
"""Test that middleware properly handles errors."""
async def test_tool_error_propagates(self):
"""Test that errors in tools are properly propagated."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware())
@mcp.tool()
def failing_tool() -> str:
raise ValueError("Test error")
async with Client(mcp) as client:
with pytest.raises(Exception):
await client.call_tool("failing_tool", {})
class TestOpenTelemetryMiddlewareContextPropagation:
"""Test trace context propagation through MCP _meta fields."""
def test_propagate_context_can_be_disabled(self):
"""Test that context propagation can be disabled."""
from fastmcp.server.middleware.opentelemetry import OpenTelemetryMiddleware
middleware = OpenTelemetryMiddleware(propagate_context=False)
assert middleware.propagate_context is False
assert middleware.propagator is None
def test_propagate_context_enabled_by_default(self):
"""Test that context propagation is enabled by default when OTel is available."""
from fastmcp.server.middleware.opentelemetry import (
OPENTELEMETRY_AVAILABLE,
OpenTelemetryMiddleware,
)
middleware = OpenTelemetryMiddleware()
assert middleware.propagate_context is True
if OPENTELEMETRY_AVAILABLE:
assert middleware.propagator is not None
else:
assert middleware.propagator is None
async def test_trace_context_injection_in_tool_result(self):
"""Test that trace context is injected into tool result metadata."""
from fastmcp.server.middleware.opentelemetry import (
OPENTELEMETRY_AVAILABLE,
OpenTelemetryMiddleware,
)
if not OPENTELEMETRY_AVAILABLE:
pytest.skip("OpenTelemetry not available")
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
# Set up a tracer provider
trace.set_tracer_provider(TracerProvider())
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware(propagate_context=True))
@mcp.tool()
def test_tool(value: str) -> str:
return f"result: {value}"
async with Client(mcp) as client:
result = await client.call_tool("test_tool", {"value": "test"})
# Check that result has metadata with trace context
assert result.meta is not None
assert "traceparent" in result.meta
async def test_trace_context_extraction_from_request(self):
"""Test that trace context is extracted from request metadata."""
from fastmcp.server.middleware.opentelemetry import (
OPENTELEMETRY_AVAILABLE,
OpenTelemetryMiddleware,
)
if not OPENTELEMETRY_AVAILABLE:
pytest.skip("OpenTelemetry not available")
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
# Set up a tracer provider
trace_provider = TracerProvider()
trace.set_tracer_provider(trace_provider)
mcp = FastMCP("Test")
middleware = OpenTelemetryMiddleware(propagate_context=True)
mcp.add_middleware(middleware)
# Track whether context was extracted
extracted_context = []
original_extract = middleware._extract_trace_context
def mock_extract(context):
result = original_extract(context)
extracted_context.append(result)
return result
middleware._extract_trace_context = mock_extract # type: ignore[method-assign]
@mcp.tool()
def test_tool(value: str) -> str:
return f"result: {value}"
async with Client(mcp) as client:
# Call tool without trace context - should extract None
await client.call_tool("test_tool", {"value": "test"})
assert len(extracted_context) > 0
async def test_context_propagation_disabled_no_injection(self):
"""Test that no context is injected when propagation is disabled."""
from fastmcp.server.middleware.opentelemetry import (
OPENTELEMETRY_AVAILABLE,
OpenTelemetryMiddleware,
)
if not OPENTELEMETRY_AVAILABLE:
pytest.skip("OpenTelemetry not available")
from opentelemetry import trace
from opentelemetry.sdk.trace import TracerProvider
# Set up a tracer provider
trace.set_tracer_provider(TracerProvider())
mcp = FastMCP("Test")
mcp.add_middleware(OpenTelemetryMiddleware(propagate_context=False))
@mcp.tool()
def test_tool(value: str) -> str:
return f"result: {value}"
async with Client(mcp) as client:
result = await client.call_tool("test_tool", {"value": "test"})
# Check that result has no trace context metadata
assert result.meta is None or "traceparent" not in result.meta