From f5e46f621a7be979e5818218cc7a13b7c170006f Mon Sep 17 00:00:00 2001 From: Chesars Date: Sat, 14 Feb 2026 21:50:12 -0300 Subject: [PATCH 1/3] feat: support per-request enable_json_schema_validation for thread safety Allow passing enable_json_schema_validation as a parameter to completion() and acompletion() instead of only relying on the global litellm.enable_json_schema_validation flag. The per-request value takes priority when provided; otherwise falls back to the global (backward compatible). This makes JSON schema validation safe for concurrent usage in FastAPI and other multi-threaded environments. --- litellm/main.py | 5 + litellm/types/utils.py | 1 + litellm/utils.py | 13 +- .../test_json_schema_validation.py | 139 ++++++++++++++++++ 4 files changed, 157 insertions(+), 1 deletion(-) create mode 100644 tests/litellm/litellm_core_utils/test_json_schema_validation.py diff --git a/litellm/main.py b/litellm/main.py index bca023e65ec..362e5b8b263 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -416,6 +416,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]: """ @@ -560,6 +562,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( @@ -1045,6 +1048,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]: """ diff --git a/litellm/types/utils.py b/litellm/types/utils.py index e1f780ffcc3..72705b82fc2 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2939,6 +2939,7 @@ all_litellm_params = ( "shared_session", "search_tool_name", "order", + "enable_json_schema_validation", ] + list(StandardCallbackDynamicParams.__annotations__.keys()) + list(CustomPricingLiteLLMParams.model_fields.keys()) diff --git a/litellm/utils.py b/litellm/utils.py index 6fdd2d88bca..0af844772b9 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1318,7 +1318,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 diff --git a/tests/litellm/litellm_core_utils/test_json_schema_validation.py b/tests/litellm/litellm_core_utils/test_json_schema_validation.py new file mode 100644 index 00000000000..859a1238a61 --- /dev/null +++ b/tests/litellm/litellm_core_utils/test_json_schema_validation.py @@ -0,0 +1,139 @@ +""" +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"], + }, + }, +} + +# Response that does NOT match the schema (wrong field names) +INVALID_RESPONSE = _make_response({"name": "test", "age": 25}) + +# Response that matches the schema +VALID_RESPONSE = _make_response({"title": "Inception", "rating": 9}) + + +@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( + INVALID_RESPONSE, + "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( + INVALID_RESPONSE, + "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( + INVALID_RESPONSE, + "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( + INVALID_RESPONSE, + "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( + VALID_RESPONSE, + "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 From 1fe2e92d3272bcf1c7d21f401efd8b8a32bbd475 Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 3 Mar 2026 14:34:32 -0300 Subject: [PATCH 2/3] fix(main): forward enable_json_schema_validation to acompletion_with_mcp The parameter was declared in completion() signature but not passed to acompletion_with_mcp, causing per-request JSON schema validation to silently fall back to the global default when MCP tools are present. --- litellm/main.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/main.py b/litellm/main.py index 362e5b8b263..75353c1b070 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1170,6 +1170,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) From 909e3ce6c9c0bcc3bbf0ef320766fd3e5d712726 Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 3 Mar 2026 14:54:11 -0300 Subject: [PATCH 3/3] test: create fresh ModelResponse per test to avoid shared mutable state --- .../test_json_schema_validation.py | 17 +++++++---------- 1 file changed, 7 insertions(+), 10 deletions(-) diff --git a/tests/litellm/litellm_core_utils/test_json_schema_validation.py b/tests/litellm/litellm_core_utils/test_json_schema_validation.py index 859a1238a61..f798db6fb43 100644 --- a/tests/litellm/litellm_core_utils/test_json_schema_validation.py +++ b/tests/litellm/litellm_core_utils/test_json_schema_validation.py @@ -46,11 +46,8 @@ STRICT_SCHEMA = { }, } -# Response that does NOT match the schema (wrong field names) -INVALID_RESPONSE = _make_response({"name": "test", "age": 25}) - -# Response that matches the schema -VALID_RESPONSE = _make_response({"title": "Inception", "rating": 9}) +INVALID_CONTENT = {"name": "test", "age": 25} # Does NOT match the schema +VALID_CONTENT = {"title": "Inception", "rating": 9} # Matches the schema @pytest.fixture(autouse=True) @@ -70,7 +67,7 @@ class TestPerRequestJsonSchemaValidation: litellm.enable_json_schema_validation = False # Should NOT raise even though response doesn't match schema post_call_processing( - INVALID_RESPONSE, + _make_response(INVALID_CONTENT), "test-model", {"response_format": STRICT_SCHEMA}, _mock_completion, @@ -82,7 +79,7 @@ class TestPerRequestJsonSchemaValidation: litellm.enable_json_schema_validation = False with pytest.raises(litellm.JSONSchemaValidationError): post_call_processing( - INVALID_RESPONSE, + _make_response(INVALID_CONTENT), "test-model", { "response_format": STRICT_SCHEMA, @@ -97,7 +94,7 @@ class TestPerRequestJsonSchemaValidation: litellm.enable_json_schema_validation = True # Should NOT raise because per-request says False post_call_processing( - INVALID_RESPONSE, + _make_response(INVALID_CONTENT), "test-model", { "response_format": STRICT_SCHEMA, @@ -112,7 +109,7 @@ class TestPerRequestJsonSchemaValidation: litellm.enable_json_schema_validation = True with pytest.raises(litellm.JSONSchemaValidationError): post_call_processing( - INVALID_RESPONSE, + _make_response(INVALID_CONTENT), "test-model", {"response_format": STRICT_SCHEMA}, _mock_completion, @@ -122,7 +119,7 @@ class TestPerRequestJsonSchemaValidation: def test_valid_response_passes_with_per_request_on(self): """Per-request ON + valid response -> no error raised.""" post_call_processing( - VALID_RESPONSE, + _make_response(VALID_CONTENT), "test-model", { "response_format": STRICT_SCHEMA,