Merge pull request #21233 from Chesars/feat/per-request-json-schema-validation

feat: support per-request enable_json_schema_validation for thread safety
This commit is contained in:
Cesar Garcia 2026-03-03 15:29:54 -03:00 • committed by GitHub
commit d384f7c320
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 155 additions and 1 deletions

View file

@ -418,6 +418,8 @@ async def acompletion( # noqa: PLR0915
web_search_options: Optional[OpenAIWebSearchOptions] = None,
# Session management
shared_session: Optional["ClientSession"] = None,
# Per-request JSON schema validation (overrides litellm.enable_json_schema_validation)
enable_json_schema_validation: Optional[bool] = None,
**kwargs,
) -> Union[ModelResponse, CustomStreamWrapper]:
"""
@ -562,6 +564,7 @@ async def acompletion( # noqa: PLR0915
"thinking": thinking,
"web_search_options": web_search_options,
"shared_session": shared_session,
"enable_json_schema_validation": enable_json_schema_validation,
}
if custom_llm_provider is None:
_, custom_llm_provider, _, _ = get_llm_provider(
@ -1047,6 +1050,8 @@ def completion( # type: ignore # noqa: PLR0915
thinking: Optional[AnthropicThinkingParam] = None,
# Session management
shared_session: Optional["ClientSession"] = None,
# Per-request JSON schema validation (overrides litellm.enable_json_schema_validation)
enable_json_schema_validation: Optional[bool] = None,
**kwargs,
) -> Union[ModelResponse, CustomStreamWrapper]:
"""
@ -1167,6 +1172,7 @@ def completion( # type: ignore # noqa: PLR0915
thinking=thinking,
web_search_options=web_search_options,
shared_session=shared_session,
enable_json_schema_validation=enable_json_schema_validation,
**kwargs,
)
api_base = kwargs.get("api_base", None)

View file

@ -3026,6 +3026,7 @@ all_litellm_params = (
"shared_session",
"search_tool_name",
"order",
"enable_json_schema_validation",
]
+ list(StandardCallbackDynamicParams.__annotations__.keys())
+ list(CustomPricingLiteLLMParams.model_fields.keys())

View file

@ -1319,7 +1319,18 @@ def post_call_processing(
### POST-CALL RULES ###
rules_obj.post_call_rules(input=model_response, model=model)
### JSON SCHEMA VALIDATION ###
if litellm.enable_json_schema_validation is True:
# Per-request flag takes priority over global flag
_per_request_validation = (
optional_params.get("enable_json_schema_validation")
if optional_params is not None
else None
)
_enable_json_schema_validation = (
_per_request_validation
if _per_request_validation is not None
else litellm.enable_json_schema_validation
)
if _enable_json_schema_validation is True:
try:
if (
optional_params is not None

View file

@ -0,0 +1,136 @@
"""
Tests for per-request enable_json_schema_validation parameter.
Ensures the per-request flag overrides the global litellm.enable_json_schema_validation,
making JSON schema validation thread-safe for concurrent usage.
Related issue: https://github.com/BerriAI/litellm/issues/XXXX
"""
import json
import pytest
import litellm
from litellm.types.utils import ModelResponse
from litellm.utils import Rules, post_call_processing
def _make_response(content: dict) -> ModelResponse:
"""Create a ModelResponse with the given content as JSON string."""
response = ModelResponse()
response.choices[0].message.content = json.dumps(content)
return response
def _mock_completion():
"""Mock function with __name__ == 'completion' for post_call_processing."""
pass
_mock_completion.__name__ = "completion"
# Schema that requires 'title' (string) and 'rating' (integer)
STRICT_SCHEMA = {
"type": "json_schema",
"json_schema": {
"name": "MovieReview",
"schema": {
"type": "object",
"properties": {
"title": {"type": "string"},
"rating": {"type": "integer"},
},
"required": ["title", "rating"],
},
},
}
INVALID_CONTENT = {"name": "test", "age": 25} # Does NOT match the schema
VALID_CONTENT = {"title": "Inception", "rating": 9} # Matches the schema
@pytest.fixture(autouse=True)
def _reset_global_flag():
"""Reset the global flag before and after each test."""
original = litellm.enable_json_schema_validation
litellm.enable_json_schema_validation = False
yield
litellm.enable_json_schema_validation = original
class TestPerRequestJsonSchemaValidation:
"""Test that per-request enable_json_schema_validation overrides the global flag."""
def test_global_off_no_per_request_skips_validation(self):
"""Global OFF + no per-request flag -> no validation (default behavior)."""
litellm.enable_json_schema_validation = False
# Should NOT raise even though response doesn't match schema
post_call_processing(
_make_response(INVALID_CONTENT),
"test-model",
{"response_format": STRICT_SCHEMA},
_mock_completion,
Rules(),
)
def test_per_request_on_overrides_global_off(self):
"""Global OFF + per-request ON -> validation runs and catches invalid response."""
litellm.enable_json_schema_validation = False
with pytest.raises(litellm.JSONSchemaValidationError):
post_call_processing(
_make_response(INVALID_CONTENT),
"test-model",
{
"response_format": STRICT_SCHEMA,
"enable_json_schema_validation": True,
},
_mock_completion,
Rules(),
)
def test_per_request_off_overrides_global_on(self):
"""Global ON + per-request OFF -> validation skipped (per-request wins)."""
litellm.enable_json_schema_validation = True
# Should NOT raise because per-request says False
post_call_processing(
_make_response(INVALID_CONTENT),
"test-model",
{
"response_format": STRICT_SCHEMA,
"enable_json_schema_validation": False,
},
_mock_completion,
Rules(),
)
def test_global_on_no_per_request_validates(self):
"""Global ON + no per-request flag -> validation runs (backward compatible)."""
litellm.enable_json_schema_validation = True
with pytest.raises(litellm.JSONSchemaValidationError):
post_call_processing(
_make_response(INVALID_CONTENT),
"test-model",
{"response_format": STRICT_SCHEMA},
_mock_completion,
Rules(),
)
def test_valid_response_passes_with_per_request_on(self):
"""Per-request ON + valid response -> no error raised."""
post_call_processing(
_make_response(VALID_CONTENT),
"test-model",
{
"response_format": STRICT_SCHEMA,
"enable_json_schema_validation": True,
},
_mock_completion,
Rules(),
)
def test_per_request_flag_is_in_all_litellm_params(self):
"""Ensure the param is registered so it doesn't leak to provider APIs."""
from litellm.types.utils import all_litellm_params
assert "enable_json_schema_validation" in all_litellm_params