mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): inline $ref request body schemas in OpenAPI tool generation
FastAPI/Pydantic specs emit request bodies as $ref, so the generated MCP inputSchema exposed an empty body object and models had to guess field names. Inline component schema refs (including nested, array, and combinator refs) and carry the body-level required list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
5f2986a1f3
commit
b387ae2737
4 changed files with 214 additions and 11 deletions
|
|
@ -1919,7 +1919,7 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
# Build input schema using imported function
|
||||
input_schema = build_input_schema(resolved_operation)
|
||||
input_schema = build_input_schema(resolved_operation, components)
|
||||
|
||||
# Create tool function with headers using imported function
|
||||
tool_func = create_tool_function(path, method, resolved_operation, base_url, headers=headers)
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import os
|
|||
import re
|
||||
from collections.abc import Mapping, Sequence
|
||||
from pathlib import PurePosixPath
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, TypeAlias, TypedDict
|
||||
from urllib.parse import quote
|
||||
|
||||
|
|
@ -50,12 +51,8 @@ from litellm.proxy._experimental.mcp_server.tool_registry import (
|
|||
_OpenAPIParameter: TypeAlias = Mapping[str, Any]
|
||||
|
||||
|
||||
class _OpenAPIJSONSchema(TypedDict, total=False):
|
||||
properties: Mapping[str, object]
|
||||
|
||||
|
||||
class _OpenAPIMediaType(TypedDict, total=False):
|
||||
schema: _OpenAPIJSONSchema
|
||||
schema: Mapping[str, Any]
|
||||
|
||||
|
||||
class _OpenAPIRequestBody(TypedDict, total=False):
|
||||
|
|
@ -80,6 +77,7 @@ class _OpenAPIPathItem(TypedDict, total=False):
|
|||
|
||||
class _OpenAPIComponents(TypedDict, total=False):
|
||||
parameters: Mapping[str, _OpenAPIParameter]
|
||||
schemas: Mapping[str, Mapping[str, Any]]
|
||||
|
||||
|
||||
# Store the base URL and headers globally
|
||||
|
|
@ -266,6 +264,46 @@ def resolve_operation_params(
|
|||
return result
|
||||
|
||||
|
||||
_SCHEMA_REF_PREFIX: Final = "#/components/schemas/"
|
||||
_EMPTY_SCHEMA: Final[Mapping[str, Any]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _inline_schema_refs(
|
||||
schema: Mapping[str, Any],
|
||||
component_schemas: Mapping[str, Mapping[str, Any]],
|
||||
seen: frozenset[str] = frozenset(),
|
||||
) -> Mapping[str, Any]:
|
||||
"""Inline ``#/components/schemas/...`` refs into *schema*.
|
||||
|
||||
MCP hands the client an ``inputSchema`` on its own, without the OpenAPI
|
||||
``components`` section, so a surviving ``$ref`` is unresolvable there and the
|
||||
client sees a field with no definition. A ref missing from *component_schemas*,
|
||||
or one recursing into a schema already being inlined, is dropped: inlining it
|
||||
would either invent a definition or never terminate.
|
||||
"""
|
||||
ref: Final[str] = schema.get("$ref", "")
|
||||
name: Final = ref[len(_SCHEMA_REF_PREFIX) :] if ref.startswith(_SCHEMA_REF_PREFIX) else ""
|
||||
target: Final = component_schemas.get(name) if name and name not in seen else None
|
||||
if target is not None:
|
||||
return _inline_schema_refs(target, component_schemas, seen | frozenset((name,)))
|
||||
# mutable-ok: the result is a JSON Schema payload handed to MCP clients, so it stays plain JSON types
|
||||
return {
|
||||
key: _inlined_schema_value(value, component_schemas, seen) for key, value in schema.items() if key != "$ref"
|
||||
}
|
||||
|
||||
|
||||
def _inlined_schema_value(
|
||||
value: object,
|
||||
component_schemas: Mapping[str, Mapping[str, Any]],
|
||||
seen: frozenset[str],
|
||||
) -> object:
|
||||
if isinstance(value, Mapping):
|
||||
return _inline_schema_refs(value, component_schemas, seen)
|
||||
if isinstance(value, Sequence) and not isinstance(value, (str, bytes)):
|
||||
return tuple(_inlined_schema_value(item, component_schemas, seen) for item in value)
|
||||
return value
|
||||
|
||||
|
||||
def extract_parameters(operation: Mapping[str, Any]) -> tuple[Sequence[str], Sequence[str], Sequence[str]]:
|
||||
"""Extract parameter names from OpenAPI operation."""
|
||||
path_params: Final = []
|
||||
|
|
@ -292,10 +330,13 @@ def extract_parameters(operation: Mapping[str, Any]) -> tuple[Sequence[str], Seq
|
|||
return path_params, query_params, body_params
|
||||
|
||||
|
||||
def build_input_schema(operation: Mapping[str, Any]) -> dict[str, Any]:
|
||||
def build_input_schema(operation: Mapping[str, Any], components: _OpenAPIComponents | None = None) -> dict[str, Any]:
|
||||
"""Build MCP input schema from OpenAPI operation."""
|
||||
properties: Final = {}
|
||||
required: Final = []
|
||||
component_schemas: Final[Mapping[str, Mapping[str, Any]]] = (components or _EMPTY_SCHEMA).get(
|
||||
"schemas", _EMPTY_SCHEMA
|
||||
)
|
||||
|
||||
# Process parameters
|
||||
if "parameters" in operation:
|
||||
|
|
@ -321,11 +362,14 @@ def build_input_schema(operation: Mapping[str, Any]) -> dict[str, Any]:
|
|||
|
||||
# Try to get JSON schema
|
||||
if "application/json" in content:
|
||||
schema: Final[_OpenAPIJSONSchema] = content["application/json"].get("schema", {})
|
||||
schema: Final[Mapping[str, Any]] = _inline_schema_refs(
|
||||
content["application/json"].get("schema", _EMPTY_SCHEMA), component_schemas
|
||||
)
|
||||
properties["body"] = {
|
||||
"type": "object",
|
||||
"description": request_body.get("description", "Request body"),
|
||||
"properties": schema.get("properties", {}),
|
||||
"properties": schema.get("properties", _EMPTY_SCHEMA),
|
||||
"required": tuple(schema.get("required", ())),
|
||||
}
|
||||
if request_body.get("required", False):
|
||||
required.append("body")
|
||||
|
|
@ -493,6 +537,7 @@ def create_tool_function(
|
|||
def register_tools_from_openapi(spec: Mapping[str, Any], base_url: str) -> None:
|
||||
"""Register MCP tools from OpenAPI specification."""
|
||||
paths: Final[Mapping[str, Mapping[str, Any]]] = spec.get("paths", {})
|
||||
components: Final[_OpenAPIComponents] = spec.get("components", {})
|
||||
used_names: Final = set()
|
||||
|
||||
for path, path_item in paths.items():
|
||||
|
|
@ -526,7 +571,7 @@ def register_tools_from_openapi(spec: Mapping[str, Any], base_url: str) -> None:
|
|||
description = operation.get("summary", operation.get("description", f"{method.upper()} {path}"))
|
||||
|
||||
# Build input schema
|
||||
input_schema = build_input_schema(operation)
|
||||
input_schema = build_input_schema(operation, components)
|
||||
|
||||
# Create tool function
|
||||
tool_func = create_tool_function(path, method, operation, base_url)
|
||||
|
|
|
|||
|
|
@ -1292,7 +1292,7 @@ if MCP_AVAILABLE:
|
|||
used_names.add(op_id)
|
||||
summary = operation.get("summary", "")
|
||||
description = operation.get("description", summary)
|
||||
input_schema = build_input_schema(resolved_op)
|
||||
input_schema = build_input_schema(resolved_op, components)
|
||||
tools.append(
|
||||
{
|
||||
"name": op_id,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ This test suite ensures that:
|
|||
5. Path parameters are properly URL encoded
|
||||
"""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
|
|
@ -411,6 +412,163 @@ class TestBuildInputSchema:
|
|||
# Required should include original names
|
||||
assert "repository-id" in schema["required"]
|
||||
|
||||
def test_request_body_ref_is_dereferenced(self):
|
||||
"""A FastAPI/Pydantic-style $ref request body must expose its fields.
|
||||
|
||||
Without inlining, the body reaches the MCP client as an empty object, so
|
||||
models guess field names and the upstream API answers 422.
|
||||
"""
|
||||
operation = {
|
||||
"operationId": "tool_kubectl_get_post",
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {"$ref": "#/components/schemas/kubectl_get_form_model"}
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
components = {
|
||||
"schemas": {
|
||||
"kubectl_get_form_model": {
|
||||
"type": "object",
|
||||
"required": ["resourceType"],
|
||||
"properties": {
|
||||
"resourceType": {
|
||||
"type": "string",
|
||||
"description": "Type of resource to get",
|
||||
},
|
||||
"name": {"type": "string"},
|
||||
"namespace": {"type": "string", "default": "default"},
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
schema = build_input_schema(operation, components)
|
||||
|
||||
body = schema["properties"]["body"]
|
||||
assert set(body["properties"]) == {"resourceType", "name", "namespace"}
|
||||
assert body["properties"]["resourceType"]["description"] == "Type of resource to get"
|
||||
assert body["properties"]["namespace"]["default"] == "default"
|
||||
assert body["required"] == ("resourceType",)
|
||||
assert schema["required"] == ["body"]
|
||||
|
||||
def test_request_body_nested_refs_are_dereferenced(self):
|
||||
"""Refs nested inside body fields must be inlined too.
|
||||
|
||||
MCP ships inputSchema without the spec's components section, so a
|
||||
surviving $ref is unresolvable on the client side.
|
||||
"""
|
||||
operation = {
|
||||
"requestBody": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {"$ref": "#/components/schemas/Outer"}
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
components = {
|
||||
"schemas": {
|
||||
"Outer": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"inner": {"$ref": "#/components/schemas/Inner"},
|
||||
"inners": {
|
||||
"type": "array",
|
||||
"items": {"$ref": "#/components/schemas/Inner"},
|
||||
},
|
||||
"maybe_inner": {
|
||||
"anyOf": [
|
||||
{"$ref": "#/components/schemas/Inner"},
|
||||
{"type": "null"},
|
||||
]
|
||||
},
|
||||
},
|
||||
},
|
||||
"Inner": {
|
||||
"type": "object",
|
||||
"properties": {"label": {"type": "string"}},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
body = build_input_schema(operation, components)["properties"]["body"]
|
||||
|
||||
inner_props = {"label": {"type": "string"}}
|
||||
assert body["properties"]["inner"]["properties"] == inner_props
|
||||
assert body["properties"]["inners"]["items"]["properties"] == inner_props
|
||||
assert body["properties"]["maybe_inner"]["anyOf"][0]["properties"] == inner_props
|
||||
assert "$ref" not in json.dumps(body)
|
||||
|
||||
def test_recursive_request_body_ref_terminates(self):
|
||||
"""A self-referencing schema must not recurse forever."""
|
||||
operation = {
|
||||
"requestBody": {
|
||||
"content": {
|
||||
"application/json": {"schema": {"$ref": "#/components/schemas/Node"}}
|
||||
},
|
||||
},
|
||||
}
|
||||
components = {
|
||||
"schemas": {
|
||||
"Node": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"value": {"type": "string"},
|
||||
"child": {"$ref": "#/components/schemas/Node"},
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
body = build_input_schema(operation, components)["properties"]["body"]
|
||||
|
||||
assert "value" in body["properties"]
|
||||
assert body["properties"]["child"] == {}
|
||||
|
||||
def test_unresolvable_request_body_ref_does_not_raise(self):
|
||||
operation = {
|
||||
"requestBody": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {"$ref": "#/components/schemas/Missing"}
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
body = build_input_schema(operation, {"schemas": {}})["properties"]["body"]
|
||||
|
||||
assert body["properties"] == {}
|
||||
|
||||
def test_inline_request_body_schema_still_works(self):
|
||||
operation = {
|
||||
"requestBody": {
|
||||
"required": True,
|
||||
"description": "The payload",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"type": "object",
|
||||
"required": ["name"],
|
||||
"properties": {"name": {"type": "string"}},
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
schema = build_input_schema(operation)
|
||||
|
||||
body = schema["properties"]["body"]
|
||||
assert body["description"] == "The payload"
|
||||
assert body["properties"] == {"name": {"type": "string"}}
|
||||
assert body["required"] == ("name",)
|
||||
assert schema["required"] == ["body"]
|
||||
|
||||
|
||||
class TestExtractParameters:
|
||||
"""Test parameter extraction from OpenAPI operations."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue