fix(openai): resolve local $refs and nested combinators when flattening tool schemas

This commit is contained in:
mateo-berri 2026-08-29 13:24:19 -07:00
parent 9b8ad46f37
commit 08574d8e07
2 changed files with 147 additions and 17 deletions

View file

@ -1092,6 +1092,7 @@ def sanitize_input_schema_for_anthropic(input_schema: dict) -> "AnthropicInputSc
_TOP_LEVEL_SCHEMA_COMBINATORS: Final = ("allOf", "anyOf", "oneOf")
_OPENAI_REJECTED_TOP_LEVEL_SCHEMA_KEYS: Final = ("enum", "const", "not")
_LOCAL_SCHEMA_REF_PREFIXES: Final = (("#/$defs/", "$defs"), ("#/definitions/", "definitions"))
_EMPTY_SCHEMA: Final[Mapping[str, object]] = MappingProxyType({})
@ -1123,31 +1124,61 @@ def _combinator_required_names(combinator: str, branches: tuple[Mapping[str, obj
return branch_names[0].intersection(*branch_names[1:])
def flatten_top_level_schema_combinators(schema: Mapping[str, object]) -> Mapping[str, object]:
"""Merge top-level ``allOf``/``anyOf``/``oneOf`` branches into an object tool schema.
def _resolve_local_schema_ref(root: Mapping[str, object], ref: str) -> Mapping[str, object] | None:
matched: Final = next(
((prefix, container) for prefix, container in _LOCAL_SCHEMA_REF_PREFIXES if ref.startswith(prefix)),
None,
)
if matched is None:
return None
prefix, container = matched
definitions: Final = root.get(container)
if not isinstance(definitions, dict):
return None
target: Final = definitions.get(ref[len(prefix) :])
return target if isinstance(target, dict) else None
OpenAI's function-calling validator rejects tool ``parameters`` carrying
'oneOf'/'anyOf'/'allOf'/'enum'/'const'/'not' at the top level (nested uses
are accepted), while lenient backends such as the ChatGPT backend Codex
talks to natively accept them, so an MCP tool declaring a top-level union
400s through LiteLLM. Branch properties merge without clobbering (the
top-level schema wins, then earlier branches); a missing ``required``
becomes the intersection of the branch lists for anyOf/oneOf and their
union for allOf. Non-object schemas pass through unchanged and the input
is never mutated.
"""
branch_groups: Final = tuple(
(combinator, _schema_branches(schema, combinator))
def _mergeable_branch(
root: Mapping[str, object], branch: Mapping[str, object], seen_refs: frozenset[str]
) -> Mapping[str, object] | None:
ref: Final = branch.get("$ref")
if isinstance(ref, str):
if ref in seen_refs:
return None
target: Final = _resolve_local_schema_ref(root, ref)
if target is None:
return None
return _mergeable_branch(root, target, seen_refs | frozenset((ref,)))
flattened: Final = _flatten_schema_against_root(branch, root, seen_refs)
if any(combinator in flattened for combinator in _TOP_LEVEL_SCHEMA_COMBINATORS):
return None
return flattened
def _flatten_schema_against_root(
schema: Mapping[str, object], root: Mapping[str, object], seen_refs: frozenset[str]
) -> Mapping[str, object]:
raw_branch_groups: Final = tuple(
(
combinator,
tuple(_mergeable_branch(root, branch, seen_refs) for branch in _schema_branches(schema, combinator)),
)
for combinator in _TOP_LEVEL_SCHEMA_COMBINATORS
if isinstance(schema.get(combinator), list)
)
dropped: Final = (
*(combinator for combinator, _ in branch_groups),
*(combinator for combinator, _ in raw_branch_groups),
*(key for key in _OPENAI_REJECTED_TOP_LEVEL_SCHEMA_KEYS if key in schema),
)
if not dropped:
return schema
if any(branch is None for _, group in raw_branch_groups for branch in group):
return schema
branch_groups: Final = tuple(
(combinator, tuple(branch for branch in group if branch is not None)) for combinator, group in raw_branch_groups
)
branches: Final = tuple(branch for _, group in branch_groups for branch in group)
is_object_schema: Final = schema.get("type") == "object" or (
"type" not in schema and branches != () and all("properties" in branch for branch in branches)
@ -1175,6 +1206,26 @@ def flatten_top_level_schema_combinators(schema: Mapping[str, object]) -> Mappin
}
def flatten_top_level_schema_combinators(schema: Mapping[str, object]) -> Mapping[str, object]:
"""Merge top-level ``allOf``/``anyOf``/``oneOf`` branches into an object tool schema.
OpenAI's function-calling validator rejects tool ``parameters`` carrying
'oneOf'/'anyOf'/'allOf'/'enum'/'const'/'not' at the top level (nested uses
are accepted), while lenient backends such as the ChatGPT backend Codex
talks to natively accept them, so an MCP tool declaring a top-level union
400s through LiteLLM. Branch properties merge without clobbering (the
top-level schema wins, then earlier branches); a missing ``required``
becomes the intersection of the branch lists for anyOf/oneOf and their
union for allOf. Branches that are local ``$ref``s (``#/$defs/...`` or
``#/definitions/...``) are resolved first and branches that are themselves
combinators are flattened recursively; a branch that cannot be fully
merged (an external or cyclic ``$ref``, or a non-object union) leaves the
whole schema untouched so OpenAI's own validation still applies.
Non-object schemas pass through unchanged and the input is never mutated.
"""
return _flatten_schema_against_root(schema, schema, frozenset())
def _get_image_mime_type_from_url(url: str) -> str | None:
"""
Get mime type for common image URLs

View file

@ -1134,8 +1134,14 @@ class TestFlattenTopLevelSchemaCombinators:
schema = {
"anyOf": [
{"properties": {"id": {"type": "string"}, "enabled": {"type": "boolean"}}, "required": ["id", "enabled"]},
{"properties": {"id": {"type": "string"}, "schedule": {"type": "string"}}, "required": ["id", "schedule"]},
{
"properties": {"id": {"type": "string"}, "enabled": {"type": "boolean"}},
"required": ["id", "enabled"],
},
{
"properties": {"id": {"type": "string"}, "schedule": {"type": "string"}},
"required": ["id", "schedule"],
},
]
}
@ -1202,6 +1208,79 @@ class TestFlattenTopLevelSchemaCombinators:
assert "not" not in result
assert result["properties"] == {"id": {"type": "string"}}
def test_resolves_local_ref_branches_from_defs(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
flatten_top_level_schema_combinators,
)
schema = {
"anyOf": [{"$ref": "#/$defs/Enable"}, {"$ref": "#/$defs/Schedule"}],
"$defs": {
"Enable": {
"properties": {"id": {"type": "string"}, "enabled": {"type": "boolean"}},
"required": ["id", "enabled"],
},
"Schedule": {
"properties": {"id": {"type": "string"}, "schedule": {"type": "string"}},
"required": ["id", "schedule"],
},
},
}
result = flatten_top_level_schema_combinators(schema)
assert "anyOf" not in result
assert result["type"] == "object"
assert set(result["properties"]) == {"id", "enabled", "schedule"}
assert result["required"] == ["id"]
assert "$defs" in result
def test_flattens_nested_combinator_branch_from_definitions(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
flatten_top_level_schema_combinators,
)
schema = {
"type": "object",
"oneOf": [
{"$ref": "#/definitions/Toggle"},
{"allOf": [{"properties": {"schedule": {"type": "string"}}, "required": ["schedule"]}]},
],
"definitions": {"Toggle": {"properties": {"enabled": {"type": "boolean"}}, "required": ["enabled"]}},
}
result = flatten_top_level_schema_combinators(schema)
assert "oneOf" not in result
assert set(result["properties"]) == {"enabled", "schedule"}
assert "required" not in result
def test_unresolvable_ref_branch_leaves_schema_untouched(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
flatten_top_level_schema_combinators,
)
schema = {
"type": "object",
"anyOf": [{"$ref": "https://example.com/schemas/automation.json"}],
"properties": {"id": {"type": "string"}},
}
assert flatten_top_level_schema_combinators(schema) is schema
def test_self_referencing_ref_branch_leaves_schema_untouched(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
flatten_top_level_schema_combinators,
)
schema = {
"type": "object",
"anyOf": [{"$ref": "#/$defs/Node"}],
"$defs": {"Node": {"type": "object", "anyOf": [{"$ref": "#/$defs/Node"}]}},
}
assert flatten_top_level_schema_combinators(schema) is schema
def test_non_object_union_passes_through_unchanged(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
flatten_top_level_schema_combinators,