Fix OpenAPI allOf reference fields (#4653)

* Fix OpenAPI allOf reference fields

🤖 Generated with Codex

* Add allOf reference crash regression

🤖 Generated with Codex

* Handle OpenAPI component refs in allOf

🤖 Generated with Codex

---------

Co-authored-by: Jeremiah Lowin <153965+jlowin@users.noreply.github.com>
This commit is contained in:
苏紫辰 2026-07-28 03:25:22 +08:00 committed by GitHub
commit c3cbe8b9a3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 159 additions and 8 deletions

View file

@ -221,6 +221,43 @@ def _make_optional_parameter_nullable(schema: dict[str, Any]) -> dict[str, Any]:
return schema
def _allof_members(
schema: dict[str, Any],
schema_defs: dict[str, Any],
resolving: set[str] | None = None,
) -> list[dict[str, Any]]:
"""Expand local schema references while collecting ``allOf`` members."""
resolving = resolving or set()
ref = schema.get("$ref")
if isinstance(ref, str):
for prefix in ("#/$defs/", "#/components/schemas/"):
if ref.startswith(prefix):
name = ref.removeprefix(prefix)
referenced_schema = schema_defs.get(name)
if isinstance(referenced_schema, dict) and name not in resolving:
siblings = {
key: value for key, value in schema.items() if key != "$ref"
}
members = _allof_members(
referenced_schema, schema_defs, resolving | {name}
)
return members + ([siblings] if siblings else [])
break
all_of = schema.get("allOf")
if isinstance(all_of, list):
members = []
for member in all_of:
if isinstance(member, dict):
members.extend(_allof_members(member, schema_defs, resolving))
siblings = {key: value for key, value in schema.items() if key != "allOf"}
return members + ([siblings] if siblings else [])
return [schema]
def _combine_schemas_and_map_params(
route: HTTPRoute,
convert_refs: bool = True,
@ -273,14 +310,13 @@ def _combine_schemas_and_map_params(
merged_props = {}
merged_required = []
for sub_schema in body_schema["allOf"]:
if isinstance(sub_schema, dict):
# Merge properties
if "properties" in sub_schema:
merged_props.update(sub_schema["properties"])
# Merge required fields
if "required" in sub_schema:
merged_required.extend(sub_schema["required"])
for sub_schema in _allof_members(body_schema, route.request_schemas):
# Merge properties
if "properties" in sub_schema:
merged_props.update(sub_schema["properties"])
# Merge required fields
if "required" in sub_schema:
merged_required.extend(sub_schema["required"])
# Update body_schema with merged properties
body_schema["properties"] = merged_props

View file

@ -1,5 +1,6 @@
"""Tests for OpenAPI feature support in OpenAPIProvider."""
import json
from typing import Any
from unittest.mock import AsyncMock, Mock
@ -1412,3 +1413,117 @@ class TestMultipartUpload:
assert "multipart/form-data" in received["content_type"]
assert b"data" in received["body"]
class TestAllOfReferenceRequestBodies:
"""Request bodies keep fields inherited through an allOf reference."""
SPEC = {
"openapi": "3.1.0",
"info": {"title": "Pet API", "version": "1.0.0"},
"servers": [{"url": "https://api.example.com"}],
"paths": {
"/pets": {
"post": {
"operationId": "create_pet",
"requestBody": {
"required": True,
"content": {
"application/json": {
"schema": {"$ref": "#/components/schemas/Cat"}
}
},
},
"responses": {"200": {"description": "Created"}},
}
}
},
"components": {
"schemas": {
"Animal": {
"type": "object",
"properties": {"animalId": {"type": "string"}},
"required": ["animalId"],
},
"Pet": {
"allOf": [
{"$ref": "#/components/schemas/Animal"},
{
"type": "object",
"properties": {"petType": {"type": "string"}},
"required": ["petType"],
},
]
},
"Cat": {
"allOf": [
{"$ref": "#/components/schemas/Pet"},
{
"type": "object",
"properties": {"meowVolume": {"type": "integer"}},
"required": ["meowVolume"],
},
]
},
}
},
}
async def test_allof_reference_fields_reach_tool_schema_and_request_body(self):
received: dict[str, object] = {}
def handler(request):
received["body"] = json.loads(request.content)
return httpx2.Response(200, json={"ok": True})
async with httpx2.AsyncClient(
transport=httpx2.MockTransport(handler),
base_url="https://api.example.com",
) as client:
server = create_openapi_server(self.SPEC, client)
async with Client(server) as mcp_client:
tools = await mcp_client.list_tools()
tool = next(tool for tool in tools if tool.name == "create_pet")
assert tool.input_schema["properties"].keys() >= {
"animalId",
"petType",
"meowVolume",
}
result = await mcp_client.call_tool(
"create_pet",
{"animalId": "a-1", "petType": "cat", "meowVolume": 11},
)
assert result.structured_content == {"ok": True}
assert received["body"] == {
"animalId": "a-1",
"petType": "cat",
"meowVolume": 11,
}
async def test_allof_reference_request_body_does_not_crash(self):
"""Required fields inherited through a reference can be sent together."""
received: dict[str, object] = {}
def handler(request):
received["body"] = json.loads(request.content)
return httpx2.Response(200, json={"ok": True})
async with httpx2.AsyncClient(
transport=httpx2.MockTransport(handler),
base_url="https://api.example.com",
) as client:
server = create_openapi_server(self.SPEC, client)
async with Client(server) as mcp_client:
result = await mcp_client.call_tool(
"create_pet",
{"animalId": "a-1", "petType": "cat", "meowVolume": 11},
)
assert result.structured_content == {"ok": True}
assert received["body"] == {
"animalId": "a-1",
"petType": "cat",
"meowVolume": 11,
}