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:
Zakir Jiwani 2026-03-11 14:05:19 +00:00
parent 1103a8c620
commit 65e2ef7a50
4 changed files with 152 additions and 13 deletions

View file

@ -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)

View file

@ -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

View file

@ -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)

View file

@ -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"