diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index c5290b41f7b..1e2f98b7b09 100644 --- a/litellm/llms/base_llm/base_utils.py +++ b/litellm/llms/base_llm/base_utils.py @@ -8,7 +8,7 @@ from abc import ABC, abstractmethod from typing import Any, Final from openai.lib import _parsing, _pydantic -from pydantic import BaseModel +from pydantic import BaseModel, ValidationError from litellm._logging import verbose_logger from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk @@ -168,6 +168,42 @@ def _dict_to_response_format_helper(response_format: dict, ref_template: str | N return response_format +def _is_basemodel_class(response_format: object) -> bool: + if not isinstance(response_format, type): + return False + try: + return issubclass(response_format, BaseModel) + except TypeError: + return False + + +def _pydantic_model_json_schema( + response_format: type[BaseModel] | type, + ref_template: str | None = None, +) -> dict: + model_json_schema = getattr(response_format, "model_json_schema", None) + if callable(model_json_schema) and ref_template is not None: + return model_json_schema(ref_template=ref_template) + if callable(model_json_schema): + return model_json_schema() + schema_fn = getattr(response_format, "schema", None) + if callable(schema_fn): + return schema_fn() + raise TypeError(f"Unsupported response_format type - {response_format}") + + +def _response_format_json_schema( + response_format: type[BaseModel], + ref_template: str | None = None, +) -> dict: + if ref_template is not None: + return _pydantic_model_json_schema(response_format, ref_template=ref_template) + try: + return _pydantic.to_strict_json_schema(response_format) + except (ValidationError, TypeError, ValueError, AttributeError): + return _pydantic_model_json_schema(response_format) + + def type_to_response_format_param( response_format: type[BaseModel] | dict | None, ref_template: str | None = None, @@ -183,17 +219,12 @@ def type_to_response_format_param( if isinstance(response_format, dict): return _dict_to_response_format_helper(response_format, ref_template) - # type checkers don't narrow the negation of a `TypeGuard` as it isn't - # a safe default behaviour but we know that at this point the `response_format` - # can only be a `type` - if not _parsing._completions.is_basemodel_type(response_format): + if not _is_basemodel_class(response_format) and not _parsing._completions.is_basemodel_type( + response_format + ): raise TypeError(f"Unsupported response_format type - {response_format}") - if ref_template is not None: - schema = response_format.model_json_schema(ref_template=ref_template) - else: - schema = _pydantic.to_strict_json_schema(response_format) - + schema: Final = _response_format_json_schema(response_format, ref_template=ref_template) return { "type": "json_schema", "json_schema": { diff --git a/litellm/utils.py b/litellm/utils.py index be722beff90..a50a7be3c9e 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -242,7 +242,6 @@ from litellm import utils as litellm_utils # These are lazy loaded via __getattr__ from litellm.llms.base_llm.base_utils import ( BaseLLMModelInfo, - _pydantic_model_json_schema, type_to_response_format_param, ) diff --git a/tests/test_litellm/litellm_core_utils/test_json_schema_validation.py b/tests/test_litellm/litellm_core_utils/test_json_schema_validation.py index f798db6fb43..0e641813442 100644 --- a/tests/test_litellm/litellm_core_utils/test_json_schema_validation.py +++ b/tests/test_litellm/litellm_core_utils/test_json_schema_validation.py @@ -77,7 +77,7 @@ class TestPerRequestJsonSchemaValidation: 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): + with pytest.raises(litellm.JSONSchemaValidationError) as exc: post_call_processing( _make_response(INVALID_CONTENT), "test-model", @@ -88,6 +88,7 @@ class TestPerRequestJsonSchemaValidation: _mock_completion, Rules(), ) + assert exc.value.raw_response == json.dumps(INVALID_CONTENT) def test_per_request_off_overrides_global_on(self): """Global ON + per-request OFF -> validation skipped (per-request wins).""" @@ -107,7 +108,7 @@ class TestPerRequestJsonSchemaValidation: 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): + with pytest.raises(litellm.JSONSchemaValidationError) as exc: post_call_processing( _make_response(INVALID_CONTENT), "test-model", @@ -115,6 +116,7 @@ class TestPerRequestJsonSchemaValidation: _mock_completion, Rules(), ) + assert exc.value.raw_response == json.dumps(INVALID_CONTENT) def test_valid_response_passes_with_per_request_on(self): """Per-request ON + valid response -> no error raised.""" diff --git a/tests/test_litellm/test_pydantic_validation.py b/tests/test_litellm/test_pydantic_validation.py index 5be0a70288d..1148217ce9a 100644 --- a/tests/test_litellm/test_pydantic_validation.py +++ b/tests/test_litellm/test_pydantic_validation.py @@ -5,7 +5,10 @@ import pytest from pydantic import BaseModel, field_validator import litellm -from litellm.llms.base_llm.base_utils import type_to_response_format_param +from litellm.llms.base_llm.base_utils import ( + _pydantic_model_json_schema, + type_to_response_format_param, +) from litellm.types.utils import LlmProviders, ModelResponse from litellm.utils import ( ProviderConfigManager, @@ -13,7 +16,6 @@ from litellm.utils import ( _apply_response_format_validation, _is_basemodel_class, _is_pydantic_basemodel_type, - _pydantic_model_json_schema, _should_preserve_pydantic_response_format, normalize_completion_response_format, post_call_processing,