diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index d2853a625c9..c24b67ac7ec 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -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}) @@ -2467,7 +2468,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( diff --git a/tests/unit/litellm_core_utils/test_streaming_overhead.py b/tests/unit/litellm_core_utils/test_streaming_overhead.py index 8fb0659ab5a..f56c57f6c47 100644 --- a/tests/unit/litellm_core_utils/test_streaming_overhead.py +++ b/tests/unit/litellm_core_utils/test_streaming_overhead.py @@ -126,25 +126,17 @@ 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 # ---------------------------------------------------------------------------