Compare commits

...

1 commit

Author SHA1 Message Date
Jake Kaplan
2b3520f8d0 Make proxy output validation errors payload-safe
🤖 Generated with Codex
2026-07-15 20:34:06 -04:00
4 changed files with 410 additions and 9 deletions

View file

@ -7,6 +7,8 @@ import weakref
from typing import TYPE_CHECKING, Any, Literal, cast, overload from typing import TYPE_CHECKING, Any, Literal, cast, overload
import mcp_types import mcp_types
from jsonschema.exceptions import SchemaError as JSONSchemaError
from jsonschema.exceptions import ValidationError as JSONSchemaValidationError
from mcp.client.caching import CacheMode from mcp.client.caching import CacheMode
from opentelemetry.trace import Status, StatusCode from opentelemetry.trace import Status, StatusCode
from pydantic import RootModel from pydantic import RootModel
@ -18,7 +20,11 @@ if TYPE_CHECKING:
from fastmcp.client.progress import ProgressHandler from fastmcp.client.progress import ProgressHandler
from fastmcp.client.tasks import ToolTask from fastmcp.client.tasks import ToolTask
from fastmcp.client.telemetry import client_span from fastmcp.client.telemetry import client_span
from fastmcp.exceptions import ToolError from fastmcp.exceptions import (
InvalidToolOutputSchemaError,
ToolError,
ToolOutputValidationError,
)
from fastmcp.telemetry import inject_trace_context from fastmcp.telemetry import inject_trace_context
from fastmcp.utilities.json_schema_type import json_schema_to_type from fastmcp.utilities.json_schema_type import json_schema_to_type
from fastmcp.utilities.logging import get_logger from fastmcp.utilities.logging import get_logger
@ -33,6 +39,72 @@ AUTO_PAGINATION_MAX_PAGES = 250
ToolTaskResponseUnion = RootModel[mcp_types.CreateTaskResult | mcp_types.CallToolResult] ToolTaskResponseUnion = RootModel[mcp_types.CreateTaskResult | mcp_types.CallToolResult]
def _translate_tool_output_error(
tool_name: str, error: RuntimeError
) -> ToolOutputValidationError | InvalidToolOutputSchemaError | None:
"""Translate SDK output-schema failures without retaining returned data."""
current = error.__cause__ or error.__context__
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
if isinstance(current, JSONSchemaError):
return InvalidToolOutputSchemaError(tool_name)
if isinstance(current, JSONSchemaValidationError):
path = tuple(current.absolute_path)
rule = current.validator if isinstance(current.validator, str) else None
expected_types: tuple[str, ...] = ()
if rule == "type":
value = current.validator_value
if isinstance(value, str):
candidates = (value,)
elif isinstance(value, list) and all(
isinstance(item, str) for item in value
):
candidates = tuple(value)
else:
candidates = ()
json_types = {
"array",
"boolean",
"integer",
"null",
"number",
"object",
"string",
}
if all(candidate in json_types for candidate in candidates):
expected_types = candidates
instance = current.instance
if instance is None:
received_type = "null"
elif isinstance(instance, bool):
received_type = "boolean"
elif isinstance(instance, str):
received_type = "string"
elif isinstance(instance, dict):
received_type = "object"
elif isinstance(instance, list):
received_type = "array"
elif isinstance(instance, int):
received_type = "integer"
elif isinstance(instance, float):
received_type = "number"
else:
received_type = None
return ToolOutputValidationError(
tool_name=tool_name,
path=path,
rule=rule,
expected_types=expected_types,
received_type=received_type,
)
current = current.__cause__ or current.__context__
return None
class ClientToolsMixin: class ClientToolsMixin:
"""Mixin providing tool-related methods for Client.""" """Mixin providing tool-related methods for Client."""
@ -210,10 +282,25 @@ class ClientToolsMixin:
allow_input_required=True, allow_input_required=True,
) )
first = await self._await_with_session_monitoring(_retry(None, None)) result: mcp_types.CallToolResult | None = None
result = await self._await_with_session_monitoring( output_error: (
self._drive_input_required(first, _retry) ToolOutputValidationError | InvalidToolOutputSchemaError | None
) ) = None
try:
first = await self._await_with_session_monitoring(_retry(None, None))
result = await self._await_with_session_monitoring(
self._drive_input_required(first, _retry)
)
except RuntimeError as error:
output_error = _translate_tool_output_error(name, error)
if output_error is None:
raise
# Raise outside the SDK exception handler so the payload-bearing
# RuntimeError is not attached to the safe exception's context.
if output_error is not None:
raise output_error
result = cast(mcp_types.CallToolResult, result)
# Reflect tool-level errors on the span so callers see ERROR # Reflect tool-level errors on the span so callers see ERROR
# status even though the MCP protocol call itself succeeded. # status even though the MCP protocol call itself succeeded.

View file

@ -72,6 +72,44 @@ class ClientError(Exception):
"""Error in client operations.""" """Error in client operations."""
class ToolOutputValidationError(ClientError):
"""An upstream tool returned data that violates its output schema."""
def __init__(
self,
tool_name: str,
path: tuple[str | int, ...],
rule: str | None,
expected_types: tuple[str, ...],
received_type: str | None,
) -> None:
self.tool_name = tool_name
self.path = path
self.rule = rule
self.expected_types = expected_types
self.received_type = received_type
location = ".".join(str(part) for part in path) or "$"
message = (
f"Tool {tool_name!r} returned data that does not match its declared "
f"output schema at {location!r}"
)
if expected_types and received_type is not None:
expected = " or ".join(expected_types)
message += f": expected {expected}, received {received_type}"
elif rule is not None:
message += f": validation rule {rule!r} failed"
super().__init__(f"{message}.")
class InvalidToolOutputSchemaError(ClientError):
"""An upstream tool advertised an invalid output schema."""
def __init__(self, tool_name: str) -> None:
self.tool_name = tool_name
super().__init__(f"Tool {tool_name!r} advertised an invalid output schema.")
class NotFoundError(Exception): class NotFoundError(Exception):
"""Object not found.""" """Object not found."""

View file

@ -34,7 +34,12 @@ from fastmcp.client.roots import RootsList, create_roots_callback
from fastmcp.client.sampling import create_sampling_callback from fastmcp.client.sampling import create_sampling_callback
from fastmcp.client.telemetry import client_span from fastmcp.client.telemetry import client_span
from fastmcp.client.transports import ClientTransportT from fastmcp.client.transports import ClientTransportT
from fastmcp.exceptions import ResourceError from fastmcp.exceptions import (
InvalidToolOutputSchemaError,
ResourceError,
ToolError,
ToolOutputValidationError,
)
from fastmcp.mcp_config import MCPConfig from fastmcp.mcp_config import MCPConfig
from fastmcp.prompts import Message, Prompt, PromptResult from fastmcp.prompts import Message, Prompt, PromptResult
from fastmcp.prompts.base import PromptArgument from fastmcp.prompts.base import PromptArgument
@ -209,9 +214,15 @@ class ProxyTool(Tool):
dict(req_ctx.meta) if req_ctx is not None and req_ctx.meta else None dict(req_ctx.meta) if req_ctx is not None and req_ctx.meta else None
) )
result = await client.call_tool_mcp( try:
name=backend_name, arguments=arguments, meta=meta result = await client.call_tool_mcp(
) name=backend_name, arguments=arguments, meta=meta
)
except (
ToolOutputValidationError,
InvalidToolOutputSchemaError,
) as error:
raise ToolError(str(error)) from None
# Pass an upstream error result through faithfully rather than # Pass an upstream error result through faithfully rather than
# collapsing it into a raised ToolError — this preserves the # collapsing it into a raised ToolError — this preserves the
# backend's content (including non-text and structured content), # backend's content (including non-text and structured content),

View file

@ -0,0 +1,265 @@
from __future__ import annotations
import datetime
from typing import Any
from unittest.mock import AsyncMock, patch
import httpx
import mcp_types
import pytest
from mcp.server import Server as LowLevelServer
from mcp.server.mcpserver import MCPServer
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from pydantic import TypeAdapter
from pydantic import ValidationError as PydanticValidationError
from fastmcp import Client
from fastmcp.client.progress import ProgressHandler
from fastmcp.exceptions import (
InvalidToolOutputSchemaError,
ToolError,
ToolOutputValidationError,
)
from fastmcp.server import create_proxy
from fastmcp.server.providers.proxy import FastMCPProxy
OUTPUT_SCHEMA = {
"type": "object",
"properties": {"result": {"type": "string"}},
"required": ["result"],
}
SENTINEL = "returned-payload-sentinel-9eabda"
def make_upstream(
result: mcp_types.CallToolResult,
*,
output_schema: dict[str, Any] = OUTPUT_SCHEMA,
) -> MCPServer:
async def list_tools(_context: Any, _params: Any) -> mcp_types.ListToolsResult:
return mcp_types.ListToolsResult(
tools=[
mcp_types.Tool(
name="get_headers",
input_schema={"type": "object"},
output_schema=output_schema,
)
]
)
async def call_tool(
_context: Any, _params: mcp_types.CallToolRequestParams
) -> mcp_types.CallToolResult:
return result
server = MCPServer("upstream")
server._lowlevel_server = LowLevelServer(
"upstream", on_list_tools=list_tools, on_call_tool=call_tool
)
return server
def text_of(result: mcp_types.CallToolResult) -> str:
assert result.content
assert isinstance(result.content[0], mcp_types.TextContent)
return result.content[0].text
@pytest.mark.parametrize(
("returned", "received_type"),
[
({"secret": SENTINEL}, "object"),
([SENTINEL], "array"),
],
)
async def test_client_safely_translates_output_validation_errors(
returned: object, received_type: str
):
upstream = make_upstream(
mcp_types.CallToolResult(content=[], structured_content={"result": returned})
)
async with Client(upstream) as client:
tools = await client.list_tools()
assert tools[0].output_schema == OUTPUT_SCHEMA
with pytest.raises(ToolOutputValidationError) as exc_info:
await client.call_tool_mcp("get_headers", {})
error = exc_info.value
assert error.tool_name == "get_headers"
assert error.path == ("result",)
assert error.rule == "type"
assert error.expected_types == ("string",)
assert error.received_type == received_type
assert SENTINEL not in str(error)
assert error.__context__ is None
assert error.__cause__ is None
async def test_proxy_returns_payload_safe_output_validation_error(
trace_exporter: InMemorySpanExporter,
):
upstream = make_upstream(
mcp_types.CallToolResult(
content=[],
structured_content={"result": {"authorization": SENTINEL}},
)
)
async with Client(upstream) as upstream_client:
await upstream_client.list_tools()
with pytest.raises(ToolOutputValidationError) as exc_info:
await upstream_client.call_tool_mcp("get_headers", {})
assert SENTINEL not in str(exc_info.value)
proxy = create_proxy(upstream)
async with Client(proxy) as client:
tools = await client.list_tools()
assert tools[0].output_schema == OUTPUT_SCHEMA
result = await client.call_tool_mcp("get_headers", {})
assert result.is_error is True
assert text_of(result) == (
"Tool 'get_headers' returned data that does not match its declared output "
"schema at 'result': expected string, received object."
)
assert SENTINEL not in result.model_dump_json()
for span in trace_exporter.get_finished_spans():
assert SENTINEL not in str(span.status.description)
assert SENTINEL not in repr(span.attributes)
for event in span.events:
assert SENTINEL not in repr(event.attributes)
if event.attributes is not None:
assert SENTINEL not in str(
event.attributes.get("exception.stacktrace", "")
)
async def test_proxy_reports_invalid_output_schema_separately():
invalid_schema = {
"type": "object",
"properties": {"result": {"type": "not-a-json-type"}},
}
upstream = make_upstream(
mcp_types.CallToolResult(content=[], structured_content={"result": SENTINEL}),
output_schema=invalid_schema,
)
async with Client(upstream) as upstream_client:
await upstream_client.list_tools()
with pytest.raises(InvalidToolOutputSchemaError) as exc_info:
await upstream_client.call_tool_mcp("get_headers", {})
assert str(exc_info.value) == (
"Tool 'get_headers' advertised an invalid output schema."
)
assert SENTINEL not in str(exc_info.value)
proxy = create_proxy(upstream)
async with Client(proxy) as client:
tools = await client.list_tools()
assert tools[0].output_schema == invalid_schema
result = await client.call_tool_mcp("get_headers", {})
assert result.is_error is True
assert text_of(result) == (
"Tool 'get_headers' advertised an invalid output schema."
)
assert SENTINEL not in result.model_dump_json()
def make_pydantic_error() -> PydanticValidationError:
try:
TypeAdapter(int).validate_python("not an integer")
except PydanticValidationError as error:
return error
raise AssertionError("Expected Pydantic validation to fail")
@pytest.mark.parametrize(
"failure",
[
RuntimeError("unrelated runtime failure"),
ToolError("upstream-owned tool failure"),
make_pydantic_error(),
httpx.ConnectError("backend transport failure"),
],
)
async def test_client_does_not_reclassify_unrelated_failures(failure: Exception):
upstream = make_upstream(
mcp_types.CallToolResult(content=[], structured_content={"result": "valid"})
)
async with Client(upstream) as client:
with patch.object(
client.session,
"call_tool",
new_callable=AsyncMock,
side_effect=failure,
):
with pytest.raises(type(failure)) as exc_info:
await client.call_tool_mcp("get_headers", {})
assert exc_info.value is failure
class RuntimeFailureClient(Client):
async def call_tool_mcp(
self,
name: str,
arguments: dict[str, Any],
progress_handler: ProgressHandler | None = None,
timeout: datetime.timedelta | float | int | None = None,
meta: dict[str, Any] | None = None,
) -> mcp_types.CallToolResult:
raise RuntimeError("unrelated backend details")
@pytest.mark.parametrize("mask_error_details", [False, True])
async def test_proxy_preserves_masking_for_unrelated_runtime_errors(
mask_error_details: bool,
):
upstream = make_upstream(
mcp_types.CallToolResult(content=[], structured_content={"result": "valid"})
)
proxy = FastMCPProxy(
client_factory=lambda: RuntimeFailureClient(upstream),
mask_error_details=mask_error_details,
)
async with Client(proxy) as client:
result = await client.call_tool_mcp("get_headers", {})
assert result.is_error is True
if mask_error_details:
assert text_of(result) == "Error calling tool 'get_headers'"
else:
assert text_of(result) == (
"Error calling tool 'get_headers': unrelated backend details"
)
async def test_proxy_preserves_valid_results_and_upstream_tool_errors():
valid_result = mcp_types.CallToolResult(
content=[mcp_types.TextContent(type="text", text="valid")],
structured_content={"result": "valid"},
)
valid_proxy = create_proxy(make_upstream(valid_result))
async with Client(valid_proxy) as client:
result = await client.call_tool_mcp("get_headers", {})
assert result == valid_result
upstream_error = mcp_types.CallToolResult(
content=[mcp_types.TextContent(type="text", text="upstream rejected call")],
structured_content={"result": {"detail": "owned by upstream"}},
is_error=True,
)
error_proxy = create_proxy(make_upstream(upstream_error))
async with Client(error_proxy) as client:
result = await client.call_tool_mcp("get_headers", {})
assert result == upstream_error