mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix: preserve serialization fallbacks without false circular references
This commit is contained in:
parent
9dbfb060bd
commit
c9dca39cf9
2 changed files with 54 additions and 5 deletions
|
|
@ -3,6 +3,7 @@ from collections.abc import Callable
|
|||
from typing import Any, Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic_core import PydanticSerializationError
|
||||
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
|
||||
|
|
@ -12,6 +13,16 @@ def strip_null_bytes(value: str) -> str:
|
|||
return value.replace("\x00", "")
|
||||
|
||||
|
||||
def _dump_model(model: BaseModel) -> dict[str, object] | str:
|
||||
try:
|
||||
return model.model_dump()
|
||||
except (PydanticSerializationError, TypeError):
|
||||
try:
|
||||
return model.model_dump(mode="json")
|
||||
except (PydanticSerializationError, TypeError):
|
||||
return str(model)
|
||||
|
||||
|
||||
def safe_dumps(
|
||||
data: Any,
|
||||
max_depth: int = DEFAULT_MAX_RECURSE_DEPTH,
|
||||
|
|
@ -66,7 +77,7 @@ def safe_dumps(
|
|||
seen.remove(id(obj))
|
||||
return result
|
||||
elif isinstance(obj, BaseModel):
|
||||
dumped: Final = obj.model_dump()
|
||||
dumped: Final = _dump_model(obj)
|
||||
result = _serialize(dumped, seen, depth + 1, key)
|
||||
seen.remove(id(obj))
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -2,10 +2,50 @@ import json
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
|
||||
|
||||
|
||||
@pytest.mark.parametrize("json_succeeds", [True, False])
|
||||
def test_repeated_pydantic_model_with_serialization_error(json_succeeds: bool) -> None:
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, SerializationInfo, model_serializer
|
||||
|
||||
class FailingModel(BaseModel):
|
||||
value: int = 42
|
||||
|
||||
@model_serializer
|
||||
def serialize_model(self, info: SerializationInfo) -> dict[str, int]:
|
||||
if json_succeeds and info.mode == "json":
|
||||
return {"value": self.value}
|
||||
raise ValueError("serialization failed")
|
||||
|
||||
model: Final = FailingModel()
|
||||
expected: Final = {"value": 42} if json_succeeds else str(model)
|
||||
assert json.loads(safe_dumps([model, model])) == [expected, expected]
|
||||
|
||||
|
||||
def test_pydantic_string_fallback_preserves_value_transform() -> None:
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, model_serializer
|
||||
|
||||
class FailingModel(BaseModel):
|
||||
@model_serializer
|
||||
def serialize_model(self) -> dict[str, object]:
|
||||
raise ValueError("serialization failed")
|
||||
|
||||
def __str__(self) -> str:
|
||||
return "secret\x00"
|
||||
|
||||
model: Final = FailingModel()
|
||||
assert json.loads(
|
||||
safe_dumps(
|
||||
{"token": model, "other": model}, value_transform=lambda key, value: "redacted" if key == "token" else value
|
||||
)
|
||||
) == {"token": "redacted", "other": "secret"}
|
||||
|
||||
|
||||
def test_primitive_types():
|
||||
# Test basic primitive types
|
||||
assert safe_dumps("test") == '"test"'
|
||||
|
|
@ -205,9 +245,7 @@ def test_pydantic_base_model():
|
|||
inner: InnerModel
|
||||
tags: list
|
||||
|
||||
outer = OuterModel(
|
||||
name="test", inner=InnerModel(value=42, label="hello"), tags=["a", "b"]
|
||||
)
|
||||
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))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue