mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(types): coerce BaseModel values in provider_specific_fields to dicts
Previously provider_specific_fields could contain raw BaseModel instances (e.g. Message objects from provider SDKs), which Pydantic's serializer couldn't handle cleanly—triggering PydanticSerializationUnexpectedValue warnings during streaming chunk serialization. Now all BaseModel values are recursively converted to plain dicts at assignment time via _coerce_provider_specific_fields, applied in add_provider_specific_fields, Choices.__init__, ModelResponseStream.__init__, and the streaming handler's copy/setattr paths. Removes the band-aid filterwarnings suppression in meta_llama/chat/transformation.py. Fixes #17631
This commit is contained in:
parent
1103a8c620
commit
65e2ef7a50
4 changed files with 152 additions and 13 deletions
|
|
@ -43,6 +43,7 @@ from litellm.types.utils import (
|
|||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
Usage,
|
||||
_coerce_provider_specific_fields,
|
||||
)
|
||||
|
||||
from ..exceptions import OpenAIError
|
||||
|
|
@ -758,8 +759,9 @@ class CustomStreamWrapper:
|
|||
original_chunk, "provider_specific_fields", None
|
||||
)
|
||||
if provider_specific_fields is not None:
|
||||
model_response.provider_specific_fields = provider_specific_fields
|
||||
for k, v in provider_specific_fields.items():
|
||||
coerced = _coerce_provider_specific_fields(provider_specific_fields)
|
||||
model_response.provider_specific_fields = coerced
|
||||
for k, v in coerced.items():
|
||||
setattr(model_response, k, v)
|
||||
return model_response
|
||||
|
||||
|
|
@ -1137,9 +1139,10 @@ class CustomStreamWrapper:
|
|||
"provider_specific_fields" in anthropic_response_obj
|
||||
and anthropic_response_obj["provider_specific_fields"] is not None
|
||||
):
|
||||
for key, value in anthropic_response_obj[
|
||||
"provider_specific_fields"
|
||||
].items():
|
||||
coerced_fields = _coerce_provider_specific_fields(
|
||||
anthropic_response_obj["provider_specific_fields"]
|
||||
)
|
||||
for key, value in coerced_fields.items():
|
||||
setattr(model_response, key, value)
|
||||
|
||||
response_obj = cast(Dict[str, Any], anthropic_response_obj)
|
||||
|
|
|
|||
|
|
@ -6,11 +6,6 @@ Calls done in OpenAI/openai.py as Llama API is openai-compatible.
|
|||
Docs: https://llama.developer.meta.com/docs/features/compatibility/
|
||||
"""
|
||||
|
||||
import warnings
|
||||
|
||||
# Suppress Pydantic serialization warnings for Meta Llama responses
|
||||
warnings.filterwarnings("ignore", message="Pydantic serializer warnings")
|
||||
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1082,12 +1082,35 @@ ChatCompletionMessage(content='This is a test', role='assistant', function_call=
|
|||
"""
|
||||
|
||||
|
||||
def _coerce_provider_specific_value(value: Any) -> Any:
|
||||
if isinstance(value, BaseModel):
|
||||
return value.model_dump() if hasattr(value, "model_dump") else value.dict()
|
||||
if isinstance(value, dict):
|
||||
return {k: _coerce_provider_specific_value(v) for k, v in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [_coerce_provider_specific_value(v) for v in value]
|
||||
return value
|
||||
|
||||
|
||||
def _coerce_provider_specific_fields(
|
||||
provider_specific_fields: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
return {
|
||||
k: _coerce_provider_specific_value(v)
|
||||
for k, v in provider_specific_fields.items()
|
||||
}
|
||||
|
||||
|
||||
def add_provider_specific_fields(
|
||||
object: BaseModel, provider_specific_fields: Optional[Dict[str, Any]]
|
||||
):
|
||||
if not provider_specific_fields: # set if provider_specific_fields is not empty
|
||||
return
|
||||
setattr(object, "provider_specific_fields", provider_specific_fields)
|
||||
setattr(
|
||||
object,
|
||||
"provider_specific_fields",
|
||||
_coerce_provider_specific_fields(provider_specific_fields),
|
||||
)
|
||||
|
||||
|
||||
class Message(SafeAttributeModel, OpenAIObject):
|
||||
|
|
@ -1360,7 +1383,12 @@ class Choices(SafeAttributeModel, OpenAIObject):
|
|||
if enhancements is not None:
|
||||
self.enhancements = enhancements
|
||||
|
||||
self.provider_specific_fields = provider_specific_fields
|
||||
if provider_specific_fields is not None:
|
||||
self.provider_specific_fields = _coerce_provider_specific_fields(
|
||||
provider_specific_fields
|
||||
)
|
||||
else:
|
||||
self.provider_specific_fields = None
|
||||
|
||||
if self.logprobs is None:
|
||||
del self.logprobs
|
||||
|
|
@ -1776,7 +1804,12 @@ class ModelResponseStream(ModelResponseBase):
|
|||
kwargs["id"] = id
|
||||
kwargs["created"] = created
|
||||
kwargs["object"] = "chat.completion.chunk"
|
||||
kwargs["provider_specific_fields"] = provider_specific_fields
|
||||
if provider_specific_fields is not None:
|
||||
kwargs["provider_specific_fields"] = _coerce_provider_specific_fields(
|
||||
provider_specific_fields
|
||||
)
|
||||
else:
|
||||
kwargs["provider_specific_fields"] = None
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.types.utils import (
|
|||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
_coerce_provider_specific_fields,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -126,3 +127,110 @@ def test_streaming_modelresponsestream_no_pydantic_warnings() -> None:
|
|||
assert pydantic_warnings == [], (
|
||||
f"Unexpected Pydantic serialization warnings: {pydantic_warnings}"
|
||||
)
|
||||
|
||||
|
||||
def test_coerce_provider_specific_fields_converts_basemodel_to_dict() -> None:
|
||||
msg = Message(content="test", role="assistant")
|
||||
result = _coerce_provider_specific_fields({"message": msg})
|
||||
assert isinstance(result["message"], dict)
|
||||
assert result["message"]["content"] == "test"
|
||||
assert result["message"]["role"] == "assistant"
|
||||
|
||||
|
||||
def test_coerce_provider_specific_fields_handles_nested_structures() -> None:
|
||||
msg = Message(content="nested", role="user")
|
||||
result = _coerce_provider_specific_fields({
|
||||
"items": [msg, {"plain": "dict"}],
|
||||
"nested_dict": {"inner": msg},
|
||||
"scalar": 42,
|
||||
})
|
||||
assert isinstance(result["items"][0], dict)
|
||||
assert result["items"][0]["content"] == "nested"
|
||||
assert result["items"][1] == {"plain": "dict"}
|
||||
assert isinstance(result["nested_dict"]["inner"], dict)
|
||||
assert result["scalar"] == 42
|
||||
|
||||
|
||||
def test_coerce_provider_specific_fields_passthrough_plain_values() -> None:
|
||||
result = _coerce_provider_specific_fields({
|
||||
"text": "hello",
|
||||
"count": 5,
|
||||
"flag": True,
|
||||
"nested": {"key": "value"},
|
||||
})
|
||||
assert result == {
|
||||
"text": "hello",
|
||||
"count": 5,
|
||||
"flag": True,
|
||||
"nested": {"key": "value"},
|
||||
}
|
||||
|
||||
|
||||
def test_streaming_provider_specific_fields_with_basemodel_no_warnings() -> None:
|
||||
msg = Message(content="provider-data", role="assistant")
|
||||
response = ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content="hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
provider_specific_fields={"message": msg},
|
||||
)
|
||||
|
||||
assert isinstance(response.provider_specific_fields["message"], dict)
|
||||
|
||||
with warnings.catch_warnings(record=True) as captured:
|
||||
warnings.simplefilter("always")
|
||||
_ = response.model_dump()
|
||||
_ = response.model_dump_json()
|
||||
|
||||
pydantic_warnings = [
|
||||
w
|
||||
for w in captured
|
||||
if "PydanticSerializationUnexpectedValue" in str(w.message)
|
||||
or "Pydantic serializer warnings" in str(w.message)
|
||||
]
|
||||
assert pydantic_warnings == [], (
|
||||
f"Unexpected Pydantic serialization warnings: {pydantic_warnings}"
|
||||
)
|
||||
|
||||
|
||||
def test_choices_provider_specific_fields_with_basemodel_no_warnings() -> None:
|
||||
msg = Message(content="extra-data", role="assistant")
|
||||
choice = Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="main", role="assistant"),
|
||||
provider_specific_fields={"extra": msg},
|
||||
)
|
||||
|
||||
assert isinstance(choice.provider_specific_fields["extra"], dict)
|
||||
|
||||
with warnings.catch_warnings(record=True) as captured:
|
||||
warnings.simplefilter("always")
|
||||
response = ModelResponse(model="test", choices=[choice])
|
||||
_ = response.model_dump()
|
||||
_ = response.model_dump_json()
|
||||
|
||||
pydantic_warnings = [
|
||||
w
|
||||
for w in captured
|
||||
if "PydanticSerializationUnexpectedValue" in str(w.message)
|
||||
or "Pydantic serializer warnings" in str(w.message)
|
||||
]
|
||||
assert pydantic_warnings == [], (
|
||||
f"Unexpected Pydantic serialization warnings: {pydantic_warnings}"
|
||||
)
|
||||
|
||||
|
||||
def test_delta_provider_specific_fields_with_basemodel_coerced() -> None:
|
||||
msg = Message(content="delta-data", role="assistant")
|
||||
delta = Delta(
|
||||
content="hi",
|
||||
role="assistant",
|
||||
provider_specific_fields={"msg": msg},
|
||||
)
|
||||
assert isinstance(delta.provider_specific_fields["msg"], dict)
|
||||
assert delta.provider_specific_fields["msg"]["content"] == "delta-data"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue