mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge fce092a885 into 632b69b5c8
This commit is contained in:
commit
af567e3bd4
3 changed files with 96 additions and 21 deletions
|
|
@ -54,7 +54,8 @@ FUNCTION_CALL_ATTRIBUTE: Final = "function_call"
|
|||
|
||||
_SYNC_ITER_EXHAUSTED: Final = object()
|
||||
|
||||
_GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__)
|
||||
_GCHUNK_FIELDS: Final[frozenset[str]] = frozenset(GChunk.__annotations__)
|
||||
_GCHUNK_REQUIRED_FIELDS: Final[frozenset[str]] = frozenset(GChunk.__required_keys__)
|
||||
_USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value})
|
||||
|
||||
|
||||
|
|
@ -1109,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
|
||||
|
|
@ -1132,10 +1132,22 @@ 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)
|
||||
or (is_registered_custom_provider and _GCHUNK_REQUIRED_FIELDS <= chunk.keys())
|
||||
)
|
||||
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", ""))
|
||||
|
|
@ -2467,7 +2479,7 @@ def generic_chunk_has_all_required_fields(chunk: dict) -> bool:
|
|||
:param chunk: The dictionary to check.
|
||||
:return: True if all required fields are present, False otherwise.
|
||||
"""
|
||||
return all(key in _GCHUNK_FIELDS for key in chunk)
|
||||
return _GCHUNK_REQUIRED_FIELDS <= chunk.keys() <= _GCHUNK_FIELDS
|
||||
|
||||
|
||||
def convert_generic_chunk_to_model_response_stream(
|
||||
|
|
|
|||
|
|
@ -2511,6 +2511,74 @@ 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 test_custom_provider_complete_generic_chunk_with_extra_fields_is_preserved(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
custom_llm_provider = "my-custom-llm"
|
||||
monkeypatch.setattr(litellm, "_custom_providers", [custom_llm_provider])
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=iter(
|
||||
[
|
||||
{
|
||||
"text": "hello",
|
||||
"is_finished": True,
|
||||
"finish_reason": "stop",
|
||||
"usage": None,
|
||||
"custom_metadata": {"trace_id": "trace-1"},
|
||||
}
|
||||
]
|
||||
),
|
||||
model="custom-model",
|
||||
logging_obj=MagicMock(),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
chunks = list(wrapper)
|
||||
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "hello"
|
||||
assert chunks[-1].choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
def _run_dispatch(wrapper: CustomStreamWrapper, chunk):
|
||||
model_response = wrapper.model_response_creator()
|
||||
completion_obj = {"content": ""}
|
||||
|
|
|
|||
|
|
@ -126,25 +126,20 @@ def test_gchunk_fields_is_frozenset():
|
|||
assert _GCHUNK_FIELDS == frozenset(GChunk.__annotations__)
|
||||
|
||||
|
||||
def test_generic_chunk_has_all_required_fields_uses_module_constant(monkeypatch):
|
||||
"""generic_chunk_has_all_required_fields must use _GCHUNK_FIELDS, not __annotations__.
|
||||
|
||||
The check semantics: every key in `chunk` must be a known GChunk field.
|
||||
This identifies GChunk-shaped dicts (all keys are valid GChunk fields).
|
||||
"""
|
||||
def test_generic_chunk_has_all_required_fields_rejects_incomplete_chunks():
|
||||
valid_chunk = _make_generic_chunk("hello")
|
||||
assert generic_chunk_has_all_required_fields(valid_chunk) is True
|
||||
|
||||
# A dict with an extra unknown key should return False — the unknown key
|
||||
# is not a GChunk field, so the chunk is not a pure GChunk.
|
||||
extra_key_chunk = dict(valid_chunk)
|
||||
extra_key_chunk["unknown_extra_key"] = "value"
|
||||
assert generic_chunk_has_all_required_fields(extra_key_chunk) is False
|
||||
|
||||
# A dict with only known GChunk fields but fewer keys still passes because
|
||||
# all its keys are valid (subset of GChunk fields).
|
||||
partial_chunk = {"text": "hi", "is_finished": False}
|
||||
assert generic_chunk_has_all_required_fields(partial_chunk) is True
|
||||
for required_field in GChunk.__required_keys__:
|
||||
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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue