mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: shared health check serialization (#21119)
* Add BaseModel serialization to safe_json_dumps * Use safe_dumps in shared health check caching * remove redundant seen.add
This commit is contained in:
parent
5ddba48df9
commit
7a164ba2cc
3 changed files with 42 additions and 2 deletions
|
|
@ -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)
|
||||
return json.dumps(safe_data, default=str)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue