diff --git a/litellm/litellm_core_utils/safe_json_dumps.py b/litellm/litellm_core_utils/safe_json_dumps.py index 8b50e41a795..051aa2f27a5 100644 --- a/litellm/litellm_core_utils/safe_json_dumps.py +++ b/litellm/litellm_core_utils/safe_json_dumps.py @@ -1,6 +1,8 @@ import json from typing import Any, Union +from pydantic import BaseModel + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH @@ -41,6 +43,11 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str: result = sorted([_serialize(item, seen, depth + 1) for item in obj]) seen.remove(id(obj)) return result + elif isinstance(obj, BaseModel): + dumped = obj.model_dump() + result = _serialize(dumped, seen, depth + 1) + seen.remove(id(obj)) + return result else: # Fall back to string conversion for non-serializable objects. try: @@ -49,4 +56,4 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str: return "Unserializable Object" safe_data = _serialize(data, set(), 0) - return json.dumps(safe_data, default=str) \ No newline at end of file + return json.dumps(safe_data, default=str) diff --git a/litellm/proxy/health_check_utils/shared_health_check_manager.py b/litellm/proxy/health_check_utils/shared_health_check_manager.py index cfd03a4d178..d0c99d84e94 100644 --- a/litellm/proxy/health_check_utils/shared_health_check_manager.py +++ b/litellm/proxy/health_check_utils/shared_health_check_manager.py @@ -4,6 +4,7 @@ import time from typing import Any, Dict, List, Optional, Tuple from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.caching.redis_cache import RedisCache from litellm.constants import ( DEFAULT_SHARED_HEALTH_CHECK_TTL, @@ -177,7 +178,7 @@ class SharedHealthCheckManager: cache_key = self.get_health_check_cache_key() await self.redis_cache.async_set_cache( cache_key, - json.dumps(cache_data), + safe_dumps(cache_data), ttl=self.health_check_ttl, ) diff --git a/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py b/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py index 56937e62666..7e48e3b88b2 100644 --- a/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py +++ b/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py @@ -138,3 +138,35 @@ def test_non_standard_dict_keys_complex(): traceback.print_exc() raise e + + +def test_pydantic_base_model(): + from pydantic import BaseModel + + class InnerModel(BaseModel): + value: int + label: str + + class OuterModel(BaseModel): + name: str + inner: InnerModel + tags: list + + outer = OuterModel(name="test", inner=InnerModel(value=42, label="hello"), tags=["a", "b"]) + + # Test a pydantic model at the top level + result = json.loads(safe_dumps(outer)) + assert result["name"] == "test" + assert result["inner"] == {"value": 42, "label": "hello"} + assert result["tags"] == ["a", "b"] + + # Test pydantic models nested inside dicts and lists + data = { + "healthy_endpoints": [outer, InnerModel(value=1, label="one")], + "count": 2, + } + result = json.loads(safe_dumps(data)) + assert result["count"] == 2 + assert len(result["healthy_endpoints"]) == 2 + assert result["healthy_endpoints"][0]["name"] == "test" + assert result["healthy_endpoints"][1] == {"value": 1, "label": "one"}