fix(structured_output): preserve anyOf for nullable enums and respect model_json_schema overrides (#4610)

This commit is contained in:
liramon2
2026-10-01 12:38:20 -04:00
committed by GitHub
parent 3b11e8255f
commit 2fd75ddf09
2 changed files with 200 additions and 2 deletions
@@ -123,6 +123,15 @@ def _process_property(
# For Optional fields, we mark as nullable but copy all properties from the non-null option
result = non_null_type.copy() if isinstance(non_null_type, dict) else {}
# Preserve anyOf when enum/const is present — a null type with non-null enum/const values is invalid.
if "enum" in result or "const" in result:
result = {"anyOf": [non_null_type, {"type": "null"}]}
# Carry over top-level metadata from the original property
for key, value in prop.items():
if key != "anyOf":
result[key] = value
return result
# For type, ensure it includes "null"
if "type" in result and isinstance(result["type"], str):
result["type"] = [result["type"], "null"]
@@ -159,7 +168,7 @@ def _process_property(
if key not in ["$ref", "anyOf"]:
if isinstance(value, dict):
result[key] = _process_nested_dict(value, defs)
elif key == "type" and not is_required and not is_nullable:
elif key == "type" and not is_required and not is_nullable and "enum" not in prop and "const" not in prop:
# For non-required fields, ensure type is a list with "null"
if isinstance(value, str):
result[key] = [value, "null"]
@@ -326,6 +335,10 @@ def _expand_nested_properties(schema: dict[str, Any], model: type[BaseModel]) ->
# If this is a BaseModel field, expand its properties with full details
if isinstance(field_type, type) and issubclass(field_type, BaseModel):
# Skip properties already expanded inline (e.g. via a model_json_schema override).
if "properties" in prop_info:
continue
# Get the nested model's schema with all its properties
nested_model_schema = field_type.model_json_schema()
@@ -1,4 +1,5 @@
from typing import Literal, Optional
from enum import Enum
from typing import Any, Literal, Optional
import pytest
from pydantic import BaseModel, Field
@@ -423,3 +424,187 @@ def test_convert_pydantic_with_refs():
"name": "Person",
}
assert tool_spec == expected_spec
def test_convert_pydantic_with_nullable_enum():
"""Test that nullable Literal/enum fields preserve anyOf instead of creating invalid type/enum combo."""
class NullableEnum(BaseModel):
sentiment: Literal["positive", "negative"] | None
flag: Literal["x"] | None = None
tool_spec = convert_pydantic_to_tool_spec(NullableEnum)
expected_spec = {
"name": "NullableEnum",
"description": "NullableEnum structured output tool",
"inputSchema": {
"json": {
"type": "object",
"properties": {
"sentiment": {
"anyOf": [
{"enum": ["positive", "negative"], "type": "string"},
{"type": "null"},
],
"title": "Sentiment",
},
"flag": {
"anyOf": [
{"const": "x", "type": "string"},
{"type": "null"},
],
"default": None,
"title": "Flag",
},
},
"title": "NullableEnum",
"required": ["sentiment"],
}
},
}
assert tool_spec == expected_spec
# Verify we can construct a valid ToolSpec
tool_spec_obj = ToolSpec(**tool_spec)
assert tool_spec_obj is not None
def test_convert_pydantic_with_model_json_schema_override():
"""Test that a model_json_schema override returning a dereferenced schema is respected."""
filters_prop = {
"filters": {
"type": "array",
"items": {
"type": "object",
"properties": {"name": {"type": "string"}},
"required": ["name"],
},
}
}
class Filter(BaseModel):
name: str
class FilterGroup(BaseModel):
filters: tuple[Filter, ...]
class Scope(BaseModel):
group: FilterGroup
@classmethod
def model_json_schema(cls, *args: Any, **kwargs: Any) -> dict[str, Any]:
return {
"type": "object",
"properties": {
"group": {
"type": "object",
"properties": filters_prop,
"required": ["filters"],
}
},
"required": ["group"],
}
# Should not raise ValueError: Missing reference: Filter
tool_spec = convert_pydantic_to_tool_spec(Scope)
expected_spec = {
"name": "Scope",
"description": "Scope structured output tool",
"inputSchema": {
"json": {
"type": "object",
"properties": {
"group": {
"type": "object",
"description": "",
"properties": filters_prop,
"required": ["filters"],
}
},
"required": ["group"],
}
},
}
assert tool_spec == expected_spec
# Verify we can construct a valid ToolSpec
tool_spec_obj = ToolSpec(**tool_spec)
assert tool_spec_obj is not None
def test_convert_pydantic_with_nullable_str_enum():
"""Test that nullable str Enum fields (resolved via $ref) preserve anyOf."""
class Color(str, Enum):
RED = "red"
GREEN = "green"
BLUE = "blue"
class Palette(BaseModel):
favorite: Color | None = None
tool_spec = convert_pydantic_to_tool_spec(Palette)
expected_spec = {
"name": "Palette",
"description": "Palette structured output tool",
"inputSchema": {
"json": {
"type": "object",
"properties": {
"favorite": {
"anyOf": [
{"enum": ["red", "green", "blue"], "title": "Color", "type": "string"},
{"type": "null"},
],
"default": None,
}
},
"title": "Palette",
}
},
}
assert tool_spec == expected_spec
def test_convert_pydantic_with_non_nullable_literal_default():
"""Test that a non-nullable Literal with a default does not get null injected into its type."""
class Config(BaseModel):
mode: Literal["fast", "slow"] = "fast"
flag: Literal["on"] = "on"
tool_spec = convert_pydantic_to_tool_spec(Config)
expected_spec = {
"name": "Config",
"description": "Config structured output tool",
"inputSchema": {
"json": {
"type": "object",
"properties": {
"mode": {
"default": "fast",
"enum": ["fast", "slow"],
"title": "Mode",
"type": "string",
},
"flag": {
"const": "on",
"default": "on",
"title": "Flag",
"type": "string",
},
},
"title": "Config",
}
},
}
assert tool_spec == expected_spec