mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
d384f7c320
4 changed files with 155 additions and 1 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
136
tests/litellm/litellm_core_utils/test_json_schema_validation.py
Normal file
136
tests/litellm/litellm_core_utils/test_json_schema_validation.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue