From 08574d8e07459e62e5c0bff0ac7a1e53073e3cea Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 29 Aug 2026 13:24:19 -0700 Subject: [PATCH] fix(openai): resolve local $refs and nested combinators when flattening tool schemas --- .../prompt_templates/common_utils.py | 81 ++++++++++++++---- ...ore_utils_prompt_templates_common_utils.py | 83 ++++++++++++++++++- 2 files changed, 147 insertions(+), 17 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 3a8bf8f328c..dcc79a1ccfe 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index ae3e49a1e44..f3c31c60975 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -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,