mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
chore(typing): clear basedpyright Any errors in streaming_handler
Replace Any typing seams in litellm_core_utils/streaming_handler.py with concrete types: JSON-boundary helpers for the fake-streaming provider chunk parsers, model_validate instead of ** splats for pydantic construction, and a narrower return type for the provider chunk early-return dataclass. reportAny 384 -> 311, reportExplicitAny 23 -> 20 in this file.
This commit is contained in:
parent
3c2264cfac
commit
5bd9e12c85
1 changed files with 58 additions and 26 deletions
|
|
@ -63,7 +63,8 @@ _SYNC_ITER_EXHAUSTED = object()
|
|||
_GCHUNK_FIELDS: frozenset = frozenset(GChunk.__annotations__)
|
||||
|
||||
|
||||
def _next_sync_or_exhausted(it: Any) -> Any:
|
||||
# any-ok: completion_stream's iterator type varies per provider transport (generator, boto3 EventStream, etc.)
|
||||
def _next_sync_or_exhausted(it: Any) -> object:
|
||||
"""
|
||||
Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration.
|
||||
|
||||
|
|
@ -77,6 +78,7 @@ def _next_sync_or_exhausted(it: Any) -> Any:
|
|||
return _SYNC_ITER_EXHAUSTED
|
||||
|
||||
|
||||
# any-ok: completion_stream's type varies per provider transport (generator, boto3 EventStream, etc.)
|
||||
def is_async_iterable(obj: Any) -> bool:
|
||||
"""
|
||||
Check if an object is an async iterable (can be used with 'async for').
|
||||
|
|
@ -98,6 +100,28 @@ def print_verbose(print_statement):
|
|||
pass
|
||||
|
||||
|
||||
def _json_loads_object(raw: Union[str, bytes]) -> object:
|
||||
return json.loads(raw) # any-ok: json.loads is the untyped-JSON boundary; callers narrow via _require_*
|
||||
|
||||
|
||||
def _require_dict(value: object) -> dict[str, object]:
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
raise ValueError(f"Expected a JSON object, got: {value!r}")
|
||||
|
||||
|
||||
def _require_list(value: object) -> list[object]:
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
raise ValueError(f"Expected a JSON array, got: {value!r}")
|
||||
|
||||
|
||||
def _require_str(value: object) -> str:
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
raise ValueError(f"Expected a string, got: {value!r}")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ProviderChunkParsed:
|
||||
response_obj: dict[str, Any]
|
||||
|
|
@ -105,7 +129,7 @@ class _ProviderChunkParsed:
|
|||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ProviderChunkEarlyReturn:
|
||||
value: Any
|
||||
value: Optional[ModelResponseStream]
|
||||
|
||||
|
||||
_ProviderChunkResult = Union[_ProviderChunkParsed, _ProviderChunkEarlyReturn]
|
||||
|
|
@ -131,6 +155,9 @@ class CustomStreamWrapper:
|
|||
self.sent_last_chunk = False
|
||||
self._stream_created_time: float = time.time()
|
||||
|
||||
# any-ok: model_call_details is Dict[str, Any]; ** splat (not model_validate) is required
|
||||
# here since tests construct logging_obj as a bare MagicMock(), whose **-unpacking yields
|
||||
# no keys but whose validate_python/model_validate would reject as a non-mapping.
|
||||
litellm_params: GenericLiteLLMParams = GenericLiteLLMParams(
|
||||
**self.logging_obj.model_call_details.get("litellm_params", {})
|
||||
)
|
||||
|
|
@ -388,9 +415,11 @@ class CustomStreamWrapper:
|
|||
|
||||
def handle_ai21_chunk(self, chunk): # fake streaming
|
||||
chunk = chunk.decode("utf-8")
|
||||
data_json = json.loads(chunk)
|
||||
data_json = _json_loads_object(chunk)
|
||||
try:
|
||||
text = data_json["completions"][0]["data"]["text"]
|
||||
completions = _require_list(_require_dict(data_json)["completions"])
|
||||
data = _require_dict(_require_dict(completions[0])["data"])
|
||||
text = data["text"]
|
||||
is_finished = True
|
||||
finish_reason = "stop"
|
||||
return {
|
||||
|
|
@ -403,9 +432,9 @@ class CustomStreamWrapper:
|
|||
|
||||
def handle_maritalk_chunk(self, chunk): # fake streaming
|
||||
chunk = chunk.decode("utf-8")
|
||||
data_json = json.loads(chunk)
|
||||
data_json = _json_loads_object(chunk)
|
||||
try:
|
||||
text = data_json["answer"]
|
||||
text = _require_dict(data_json)["answer"]
|
||||
is_finished = True
|
||||
finish_reason = "stop"
|
||||
return {
|
||||
|
|
@ -424,8 +453,8 @@ class CustomStreamWrapper:
|
|||
if self.model and "dolphin" in self.model:
|
||||
chunk = self.process_chunk(chunk=chunk)
|
||||
else:
|
||||
data_json = json.loads(chunk)
|
||||
chunk = data_json["generated_text"]
|
||||
data_json = _require_dict(_json_loads_object(chunk))
|
||||
chunk = _require_str(data_json["generated_text"])
|
||||
text = chunk
|
||||
if "[DONE]" in text:
|
||||
text = text.replace("[DONE]", "")
|
||||
|
|
@ -441,9 +470,10 @@ class CustomStreamWrapper:
|
|||
|
||||
def handle_aleph_alpha_chunk(self, chunk):
|
||||
chunk = chunk.decode("utf-8")
|
||||
data_json = json.loads(chunk)
|
||||
data_json = _json_loads_object(chunk)
|
||||
try:
|
||||
text = data_json["completions"][0]["completion"]
|
||||
completions = _require_list(_require_dict(data_json)["completions"])
|
||||
text = _require_dict(completions[0])["completion"]
|
||||
is_finished = True
|
||||
finish_reason = "stop"
|
||||
return {
|
||||
|
|
@ -469,14 +499,16 @@ class CustomStreamWrapper:
|
|||
"finish_reason": finish_reason,
|
||||
}
|
||||
elif chunk.startswith("data:"):
|
||||
data_json = json.loads(chunk[5:]) # chunk.startswith("data:"):
|
||||
data_json = _require_dict(_json_loads_object(chunk[5:])) # chunk.startswith("data:"):
|
||||
try:
|
||||
if len(data_json["choices"]) > 0:
|
||||
delta = data_json["choices"][0]["delta"]
|
||||
text = "" if delta is None else delta.get("content", "")
|
||||
if data_json["choices"][0].get("finish_reason", None):
|
||||
choices = _require_list(data_json["choices"])
|
||||
if len(choices) > 0:
|
||||
choice = _require_dict(choices[0])
|
||||
delta = choice["delta"]
|
||||
text = "" if delta is None else _require_dict(delta).get("content", "")
|
||||
if choice.get("finish_reason", None):
|
||||
is_finished = True
|
||||
finish_reason = data_json["choices"][0]["finish_reason"]
|
||||
finish_reason = choice["finish_reason"]
|
||||
print_verbose(f"text: {text}; is_finished: {is_finished}; finish_reason: {finish_reason}")
|
||||
return {
|
||||
"text": text,
|
||||
|
|
@ -665,19 +697,19 @@ class CustomStreamWrapper:
|
|||
except Exception as e:
|
||||
raise e
|
||||
|
||||
def model_response_creator(self, chunk: Optional[dict] = None, hidden_params: Optional[dict] = None):
|
||||
def model_response_creator(self, chunk: Optional[dict[str, object]] = None, hidden_params: Optional[dict] = None):
|
||||
_model = self._cached_model_name
|
||||
_logging_obj_llm_provider = self._cached_logging_llm_provider
|
||||
|
||||
if chunk is None:
|
||||
args: Dict[str, Any] = {"model": _model}
|
||||
args: dict[str, object] = {"model": _model}
|
||||
else:
|
||||
chunk.pop("model", None)
|
||||
args = {"model": _model}
|
||||
if chunk:
|
||||
args.update({k: v for k, v in chunk.items() if k != "stream"})
|
||||
|
||||
model_response = ModelResponseStream(**args)
|
||||
model_response = ModelResponseStream.model_validate(args)
|
||||
if self.response_id is not None:
|
||||
model_response.id = self.response_id
|
||||
if self.system_fingerprint is not None:
|
||||
|
|
@ -817,7 +849,7 @@ class CustomStreamWrapper:
|
|||
_initial_delta = model_response.choices[0].delta.model_dump()
|
||||
|
||||
_initial_delta.pop("role", None)
|
||||
model_response.choices[0].delta = Delta(**_initial_delta)
|
||||
model_response.choices[0].delta = Delta.model_validate(_initial_delta)
|
||||
return model_response
|
||||
|
||||
def _has_special_delta_content(self, model_response: ModelResponseStream) -> bool:
|
||||
|
|
@ -911,11 +943,11 @@ class CustomStreamWrapper:
|
|||
for choice in original_chunk.choices:
|
||||
try:
|
||||
if isinstance(choice, BaseModel):
|
||||
choice_json = choice.model_dump() # type: ignore
|
||||
choice_json = choice.model_dump()
|
||||
choice_json.pop(
|
||||
"finish_reason", None
|
||||
) # for mistral etc. which return a value in their last chunk (not-openai compatible).
|
||||
choices.append(StreamingChoices(**choice_json))
|
||||
choices.append(StreamingChoices.model_validate(choice_json))
|
||||
except Exception:
|
||||
choices.append(StreamingChoices())
|
||||
setattr(model_response, "choices", choices)
|
||||
|
|
@ -946,7 +978,7 @@ class CustomStreamWrapper:
|
|||
self.sent_first_chunk = True
|
||||
if response_obj.get("provider_specific_fields") is not None:
|
||||
completion_obj["provider_specific_fields"] = response_obj["provider_specific_fields"]
|
||||
model_response.choices[0].delta = Delta(**completion_obj)
|
||||
model_response.choices[0].delta = Delta.model_validate(completion_obj)
|
||||
_index: Optional[int] = completion_obj.get("index")
|
||||
if _index is not None:
|
||||
model_response.choices[0].index = _index
|
||||
|
|
@ -1443,7 +1475,7 @@ class CustomStreamWrapper:
|
|||
):
|
||||
# if function returned but type set to None - mistral's api returns type: None
|
||||
tool["type"] = "function"
|
||||
model_response.choices[0].delta = Delta(**_json_delta)
|
||||
model_response.choices[0].delta = Delta.model_validate(_json_delta)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"litellm.CustomStreamWrapper.chunk_creator(): Exception occured - {}".format(str(e))
|
||||
|
|
@ -1458,7 +1490,7 @@ class CustomStreamWrapper:
|
|||
if original_chunk.choices[0].delta is None
|
||||
else dict(original_chunk.choices[0].delta)
|
||||
)
|
||||
model_response.choices[0].delta = Delta(**delta)
|
||||
model_response.choices[0].delta = Delta.model_validate(delta)
|
||||
except Exception:
|
||||
model_response.choices[0].delta = Delta()
|
||||
else:
|
||||
|
|
@ -1672,7 +1704,7 @@ class CustomStreamWrapper:
|
|||
else:
|
||||
asyncio.run(self.logging_obj.async_success_handler(processed_chunk, None, None, cache_hit))
|
||||
## SYNC LOGGING — only for sync SDK entrypoints; async proxy paths export via async_success_handler
|
||||
litellm_params = self.logging_obj.model_call_details.get("litellm_params", {})
|
||||
litellm_params = _require_dict(self.logging_obj.model_call_details.get("litellm_params", {}))
|
||||
if self.logging_obj._is_sync_litellm_request(litellm_params):
|
||||
self.logging_obj.success_handler(processed_chunk, None, None, cache_hit)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue