mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 15:19:10 +02:00
Compare commits
1 commit
main
...
fix/proxy-
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2b3520f8d0 |
4 changed files with 410 additions and 9 deletions
|
|
@ -7,6 +7,8 @@ import weakref
|
|||
from typing import TYPE_CHECKING, Any, Literal, cast, overload
|
||||
|
||||
import mcp_types
|
||||
from jsonschema.exceptions import SchemaError as JSONSchemaError
|
||||
from jsonschema.exceptions import ValidationError as JSONSchemaValidationError
|
||||
from mcp.client.caching import CacheMode
|
||||
from opentelemetry.trace import Status, StatusCode
|
||||
from pydantic import RootModel
|
||||
|
|
@ -18,7 +20,11 @@ if TYPE_CHECKING:
|
|||
from fastmcp.client.progress import ProgressHandler
|
||||
from fastmcp.client.tasks import ToolTask
|
||||
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.utilities.json_schema_type import json_schema_to_type
|
||||
from fastmcp.utilities.logging import get_logger
|
||||
|
|
@ -33,6 +39,72 @@ AUTO_PAGINATION_MAX_PAGES = 250
|
|||
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:
|
||||
"""Mixin providing tool-related methods for Client."""
|
||||
|
||||
|
|
@ -210,10 +282,25 @@ class ClientToolsMixin:
|
|||
allow_input_required=True,
|
||||
)
|
||||
|
||||
first = await self._await_with_session_monitoring(_retry(None, None))
|
||||
result = await self._await_with_session_monitoring(
|
||||
self._drive_input_required(first, _retry)
|
||||
)
|
||||
result: mcp_types.CallToolResult | None = None
|
||||
output_error: (
|
||||
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
|
||||
# status even though the MCP protocol call itself succeeded.
|
||||
|
|
|
|||
|
|
@ -72,6 +72,44 @@ class ClientError(Exception):
|
|||
"""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):
|
||||
"""Object not found."""
|
||||
|
||||
|
|
|
|||
|
|
@ -34,7 +34,12 @@ from fastmcp.client.roots import RootsList, create_roots_callback
|
|||
from fastmcp.client.sampling import create_sampling_callback
|
||||
from fastmcp.client.telemetry import client_span
|
||||
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.prompts import Message, Prompt, PromptResult
|
||||
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
|
||||
)
|
||||
|
||||
result = await client.call_tool_mcp(
|
||||
name=backend_name, arguments=arguments, meta=meta
|
||||
)
|
||||
try:
|
||||
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
|
||||
# collapsing it into a raised ToolError — this preserves the
|
||||
# backend's content (including non-text and structured content),
|
||||
|
|
|
|||
265
tests/server/providers/proxy/test_output_schema_errors.py
Normal file
265
tests/server/providers/proxy/test_output_schema_errors.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue