mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(streaming): normalize provider logprobs to prevent MockValSer serialization crash
This commit is contained in:
parent
69a491e168
commit
e8356c18d0
2 changed files with 124 additions and 1 deletions
|
|
@ -39,6 +39,7 @@ from litellm.types.utils import (
|
||||||
)
|
)
|
||||||
from litellm.types.utils import GenericStreamingChunk as GChunk
|
from litellm.types.utils import GenericStreamingChunk as GChunk
|
||||||
from litellm.types.utils import (
|
from litellm.types.utils import (
|
||||||
|
ChoiceLogprobs,
|
||||||
LlmProviders,
|
LlmProviders,
|
||||||
ModelResponse,
|
ModelResponse,
|
||||||
ModelResponseStream,
|
ModelResponseStream,
|
||||||
|
|
@ -1041,6 +1042,17 @@ class CustomStreamWrapper:
|
||||||
del model_response.choices[0].delta.reasoning_content
|
del model_response.choices[0].delta.reasoning_content
|
||||||
return
|
return
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_logprobs(logprobs: Any) -> Any:
|
||||||
|
"""Convert raw provider/SDK logprobs into litellm's own ``ChoiceLogprobs`` so the chunk stays serializable."""
|
||||||
|
if logprobs is None or isinstance(logprobs, ChoiceLogprobs):
|
||||||
|
return logprobs
|
||||||
|
if isinstance(logprobs, BaseModel):
|
||||||
|
return ChoiceLogprobs(**logprobs.model_dump())
|
||||||
|
if isinstance(logprobs, dict):
|
||||||
|
return ChoiceLogprobs(**logprobs)
|
||||||
|
return logprobs
|
||||||
|
|
||||||
def _dispatch_provider_chunk(
|
def _dispatch_provider_chunk(
|
||||||
self,
|
self,
|
||||||
chunk: Any,
|
chunk: Any,
|
||||||
|
|
@ -1345,7 +1357,7 @@ class CustomStreamWrapper:
|
||||||
model_response.system_fingerprint = response_obj["original_chunk"].system_fingerprint
|
model_response.system_fingerprint = response_obj["original_chunk"].system_fingerprint
|
||||||
self.system_fingerprint = response_obj["original_chunk"].system_fingerprint
|
self.system_fingerprint = response_obj["original_chunk"].system_fingerprint
|
||||||
if response_obj["logprobs"] is not None:
|
if response_obj["logprobs"] is not None:
|
||||||
model_response.choices[0].logprobs = response_obj["logprobs"]
|
model_response.choices[0].logprobs = self._normalize_logprobs(response_obj["logprobs"])
|
||||||
|
|
||||||
if response_obj["usage"] is not None:
|
if response_obj["usage"] is not None:
|
||||||
if isinstance(response_obj["usage"], dict):
|
if isinstance(response_obj["usage"], dict):
|
||||||
|
|
|
||||||
|
|
@ -22,6 +22,7 @@ from litellm.litellm_core_utils.streaming_handler import (
|
||||||
_ProviderChunkParsed,
|
_ProviderChunkParsed,
|
||||||
)
|
)
|
||||||
from litellm.types.utils import (
|
from litellm.types.utils import (
|
||||||
|
ChoiceLogprobs,
|
||||||
CompletionTokensDetailsWrapper,
|
CompletionTokensDetailsWrapper,
|
||||||
Delta,
|
Delta,
|
||||||
ModelResponseStream,
|
ModelResponseStream,
|
||||||
|
|
@ -3059,3 +3060,113 @@ async def test_stream_chunk_builder_raise_and_usage_recovery_failure_does_not_cr
|
||||||
chunks = [c async for c in response]
|
chunks = [c async for c in response]
|
||||||
|
|
||||||
assert len(chunks) > 0
|
assert len(chunks) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_logprobs_converts_raw_sdk_object():
|
||||||
|
"""
|
||||||
|
Regression test for https://github.com/BerriAI/litellm/issues/33456
|
||||||
|
|
||||||
|
openai SDK models are built with defer_build=True, so an instance parsed via
|
||||||
|
model_validate() (how the SDK parses every streaming chunk) keeps a MockValSer
|
||||||
|
placeholder serializer. If such a raw object is stored on a litellm chunk, a
|
||||||
|
later model_dump() blows up with
|
||||||
|
"'MockValSer' object cannot be converted to 'SchemaSerializer'".
|
||||||
|
|
||||||
|
_normalize_logprobs must convert it to litellm's own ChoiceLogprobs.
|
||||||
|
"""
|
||||||
|
from openai.types.chat.chat_completion_chunk import (
|
||||||
|
ChoiceLogprobs as OpenAIChoiceLogprobs,
|
||||||
|
)
|
||||||
|
|
||||||
|
raw = OpenAIChoiceLogprobs.model_validate(
|
||||||
|
{
|
||||||
|
"content": [
|
||||||
|
{"token": "hi", "bytes": [104, 105], "logprob": -0.1, "top_logprobs": []}
|
||||||
|
],
|
||||||
|
"refusal": None,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
normalized = CustomStreamWrapper._normalize_logprobs(raw)
|
||||||
|
|
||||||
|
assert isinstance(normalized, ChoiceLogprobs)
|
||||||
|
assert type(normalized).__module__.startswith("litellm")
|
||||||
|
assert normalized.content is not None and normalized.content[0].token == "hi"
|
||||||
|
|
||||||
|
assert isinstance(
|
||||||
|
CustomStreamWrapper._normalize_logprobs({"content": None, "refusal": None}),
|
||||||
|
ChoiceLogprobs,
|
||||||
|
)
|
||||||
|
assert CustomStreamWrapper._normalize_logprobs(normalized) is normalized
|
||||||
|
assert CustomStreamWrapper._normalize_logprobs(None) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_streaming_logprobs_do_not_store_raw_sdk_object(logging_obj: Logging):
|
||||||
|
"""
|
||||||
|
Regression test for https://github.com/BerriAI/litellm/issues/33456
|
||||||
|
|
||||||
|
A provider (e.g. SAP AI Core) that emits a final chunk carrying logprobs +
|
||||||
|
usage with an empty delta used to leave the raw openai SDK ChoiceLogprobs on
|
||||||
|
the litellm chunk. The subsequent model_dump() in __anext__ then crashed with
|
||||||
|
a MockValSer TypeError, which the router surfaced as MidStreamFallbackError and
|
||||||
|
truncated the stream mid-way.
|
||||||
|
|
||||||
|
The forwarded chunk's logprobs must be litellm's own ChoiceLogprobs and every
|
||||||
|
chunk must stay serializable via model_dump()/model_dump_json().
|
||||||
|
"""
|
||||||
|
from litellm.types.llms.openai import OpenAIChatCompletionChunk
|
||||||
|
|
||||||
|
def _chunk(delta, finish=None, usage=None, logprobs=None):
|
||||||
|
payload = {
|
||||||
|
"id": "chatcmpl-33456",
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": 1,
|
||||||
|
"model": "gpt-4o",
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": delta,
|
||||||
|
"finish_reason": finish,
|
||||||
|
"logprobs": logprobs,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
if usage is not None:
|
||||||
|
payload["usage"] = usage
|
||||||
|
return OpenAIChatCompletionChunk.model_validate(payload)
|
||||||
|
|
||||||
|
logprobs_payload = {
|
||||||
|
"content": [
|
||||||
|
{"token": "hi", "bytes": [104, 105], "logprob": -0.1, "top_logprobs": []}
|
||||||
|
],
|
||||||
|
"refusal": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
async def _stream():
|
||||||
|
yield _chunk({"role": "assistant", "content": "hi"})
|
||||||
|
yield _chunk(
|
||||||
|
{},
|
||||||
|
finish="stop",
|
||||||
|
logprobs=logprobs_payload,
|
||||||
|
usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||||
|
)
|
||||||
|
|
||||||
|
response = CustomStreamWrapper(
|
||||||
|
completion_stream=_stream(),
|
||||||
|
model="gpt-4o",
|
||||||
|
custom_llm_provider="openai",
|
||||||
|
logging_obj=logging_obj,
|
||||||
|
)
|
||||||
|
|
||||||
|
logprobs_chunks = []
|
||||||
|
async for chunk in response:
|
||||||
|
chunk.model_dump()
|
||||||
|
chunk.model_dump_json()
|
||||||
|
if chunk.choices and chunk.choices[0].logprobs is not None:
|
||||||
|
logprobs_chunks.append(chunk.choices[0].logprobs)
|
||||||
|
|
||||||
|
assert logprobs_chunks, "expected a forwarded chunk carrying logprobs"
|
||||||
|
for lp in logprobs_chunks:
|
||||||
|
assert isinstance(lp, ChoiceLogprobs)
|
||||||
|
assert "openai" not in type(lp).__module__
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue