fix(streaming): handle incomplete chunks safely

This commit is contained in:
Pradeep Ramola 2026-09-30 00:40:43 -04:00
parent 4b6d870472
commit b594c88fe5
3 changed files with 58 additions and 7 deletions

View file

@ -1110,11 +1110,10 @@ class CustomStreamWrapper:
completion_obj: dict[str, Any],
) -> _ProviderChunkResult:
response_obj: dict[str, Any] = {}
if (
isinstance(chunk, ModelResponseStream)
and self.custom_llm_provider is not None
and self.custom_llm_provider in litellm._custom_providers
):
is_registered_custom_provider: Final = (
self.custom_llm_provider is not None and self.custom_llm_provider in litellm._custom_providers
)
if isinstance(chunk, ModelResponseStream) and is_registered_custom_provider:
_has_content: Final = bool(
chunk.choices
and chunk.choices[0].delta is not None
@ -1133,10 +1132,19 @@ class CustomStreamWrapper:
chunk.choices[0].finish_reason = None
return _ProviderChunkEarlyReturn(chunk)
is_generic_chunk: Final = isinstance(chunk, dict) and generic_chunk_has_all_required_fields(chunk=chunk)
if (
isinstance(chunk, dict)
and generic_chunk_has_all_required_fields(chunk=chunk) # check if chunk is a generic streaming chunk
) or (self.custom_llm_provider and self.custom_llm_provider in litellm._custom_providers):
and not is_generic_chunk
and (chunk.keys() <= _GCHUNK_FIELDS or is_registered_custom_provider)
):
partial_chunk_usage: Final = chunk.get("usage")
if isinstance(partial_chunk_usage, dict):
model_response.usage = litellm.Usage(**partial_chunk_usage)
return _ProviderChunkParsed(cast(dict[str, object], chunk))
return _ProviderChunkEarlyReturn(None)
if is_generic_chunk:
if self.received_finish_reason is not None:
_chunk_has_content: Final = isinstance(chunk, dict) and (
bool(chunk.get("text", ""))

View file

@ -2511,6 +2511,46 @@ def test_usage_only_chunk_not_dropped_when_finish_reason_already_set(
assert result.usage is not None
@pytest.mark.parametrize("custom_llm_provider", [None, "my-custom-llm"])
@pytest.mark.parametrize(
"chunk",
[
pytest.param({}, id="empty"),
pytest.param({"is_finished": False}, id="finish-state-only"),
],
)
def test_chunk_creator_skips_incomplete_generic_chunks(
monkeypatch: pytest.MonkeyPatch,
initialized_custom_stream_wrapper: CustomStreamWrapper,
custom_llm_provider: str | None,
chunk: dict[str, object],
):
if custom_llm_provider is not None:
monkeypatch.setattr(litellm, "_custom_providers", [custom_llm_provider])
initialized_custom_stream_wrapper.custom_llm_provider = custom_llm_provider
result = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk)
assert result is None
assert initialized_custom_stream_wrapper.chunks == []
@pytest.mark.parametrize("custom_llm_provider", [None, "my-custom-llm"])
def test_chunk_creator_records_incomplete_usage_chunk(
monkeypatch: pytest.MonkeyPatch,
initialized_custom_stream_wrapper: CustomStreamWrapper,
custom_llm_provider: str | None,
):
if custom_llm_provider is not None:
monkeypatch.setattr(litellm, "_custom_providers", [custom_llm_provider])
initialized_custom_stream_wrapper.custom_llm_provider = custom_llm_provider
result = initialized_custom_stream_wrapper.chunk_creator(chunk={"usage": {"prompt_tokens": 1}})
assert result is None
assert initialized_custom_stream_wrapper.chunks[-1].usage.prompt_tokens == 1
def _run_dispatch(wrapper: CustomStreamWrapper, chunk):
model_response = wrapper.model_response_creator()
completion_obj = {"content": ""}

View file

@ -138,6 +138,9 @@ def test_generic_chunk_has_all_required_fields_rejects_incomplete_chunks():
incomplete_chunk = {key: value for key, value in valid_chunk.items() if key != required_field}
assert generic_chunk_has_all_required_fields(incomplete_chunk) is False
for degenerate_chunk in ({}, {"is_finished": False}, {"usage": {"prompt_tokens": 1}}):
assert generic_chunk_has_all_required_fields(degenerate_chunk) is False
# ---------------------------------------------------------------------------
# 2. Cached model name and provider at init time