mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(proxy): type summary quota helpers without unknown arguments
Signed-off-by: sinksilk <785976238@qq.com>
This commit is contained in:
parent
4d6f12578d
commit
25ca54ca2b
2 changed files with 45 additions and 13 deletions
|
|
@ -1037,6 +1037,20 @@ def _build_summary_messages(
|
|||
return summary_messages
|
||||
|
||||
|
||||
def _count_summary_message_tokens(
|
||||
model: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
) -> int:
|
||||
"""Count summary-call tokens through a fully annotated wrapper.
|
||||
|
||||
``litellm.token_counter``'s own signature is partially unknown (bare
|
||||
``Sequence`` / ``dict`` parameters). Passing that function into
|
||||
``asyncify`` is a new ``reportUnknownArgumentType``. This wrapper's
|
||||
signature is fully known, so the asyncify boundary stays typed.
|
||||
"""
|
||||
return litellm.token_counter(model=model, messages=messages)
|
||||
|
||||
|
||||
async def _estimate_summary_input_tokens(
|
||||
*,
|
||||
summary_model: str,
|
||||
|
|
@ -1044,7 +1058,7 @@ async def _estimate_summary_input_tokens(
|
|||
fallback_tokens: int,
|
||||
) -> int:
|
||||
try:
|
||||
return await asyncify(litellm.token_counter)(
|
||||
return await asyncify(_count_summary_message_tokens)(
|
||||
model=summary_model,
|
||||
messages=summary_messages,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from typing import (
|
|||
Protocol,
|
||||
TypeAlias,
|
||||
TypedDict,
|
||||
cast,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -740,6 +741,19 @@ def _call_id_from_callback_kwargs(kwargs: object) -> str | None:
|
|||
return call_id if isinstance(call_id, str) else None
|
||||
|
||||
|
||||
def _as_str_object_dict(value: object) -> dict[str, object] | None: # mutable-ok: model-group helper requires a dict
|
||||
"""Return a ``dict[str, object]`` view of an untyped callback payload.
|
||||
|
||||
``isinstance(..., dict)`` narrows to ``dict[Unknown, Unknown]``. Keep an
|
||||
``object`` alias from before that narrowing and cast that alias, so the
|
||||
cast argument stays a known type.
|
||||
"""
|
||||
raw: Final[object] = value
|
||||
if not isinstance(value, dict):
|
||||
return None
|
||||
return cast("dict[str, object]", raw) # cast-ok: success-callback payload is an untyped dict
|
||||
|
||||
|
||||
def _parse_output_cap_value(raw_value: object) -> int | None:
|
||||
if isinstance(raw_value, bool) or not isinstance(raw_value, (int, float, str)):
|
||||
return None
|
||||
|
|
@ -4729,7 +4743,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
def _build_unreserved_project_io_token_ops(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
kwargs: object,
|
||||
response_obj: object,
|
||||
) -> tuple[ReservationAwareIncrementOperation, ...]:
|
||||
"""Charge full actual ITPM/OTPM when no pre-call reservation owns this call.
|
||||
|
|
@ -4746,17 +4760,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
get_model_group_from_litellm_kwargs,
|
||||
)
|
||||
|
||||
standard_logging_object: Final = kwargs.get("standard_logging_object")
|
||||
if not isinstance(standard_logging_object, dict):
|
||||
callback_kwargs: Final = _as_str_object_dict(kwargs)
|
||||
if callback_kwargs is None:
|
||||
return ()
|
||||
standard_logging_metadata: Final = standard_logging_object.get("metadata")
|
||||
if not isinstance(standard_logging_metadata, Mapping):
|
||||
logging_map: Final = _as_str_object_dict(callback_kwargs.get("standard_logging_object"))
|
||||
if logging_map is None:
|
||||
return ()
|
||||
standard_logging_metadata: Final = _as_str_object_dict(logging_map.get("metadata"))
|
||||
if standard_logging_metadata is None:
|
||||
return ()
|
||||
|
||||
model_group: Final = get_model_group_from_litellm_kwargs(kwargs) or (
|
||||
standard_logging_object.get("model_group")
|
||||
if isinstance(standard_logging_object.get("model_group"), str)
|
||||
else None
|
||||
logged_group: Final = logging_map.get("model_group")
|
||||
model_group: Final = get_model_group_from_litellm_kwargs(callback_kwargs) or (
|
||||
logged_group if isinstance(logged_group, str) else None
|
||||
)
|
||||
targets: Final = self._collect_project_io_scope_targets(
|
||||
standard_logging_metadata=standard_logging_metadata,
|
||||
|
|
@ -4766,10 +4782,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
return ()
|
||||
|
||||
response_usage: Final = self._resolve_io_token_reconcile_usage(response_obj)
|
||||
combined_usage: Final = self._resolve_io_token_reconcile_usage(kwargs.get("combined_usage_object"))
|
||||
combined_usage_object: Final = callback_kwargs.get("combined_usage_object")
|
||||
combined_usage: Final = self._resolve_io_token_reconcile_usage(combined_usage_object)
|
||||
aggregate_total: Final = self._aggregate_only_total_tokens(
|
||||
self._response_usage(response_obj)
|
||||
) or self._aggregate_only_total_tokens(self._response_usage(kwargs.get("combined_usage_object")))
|
||||
) or self._aggregate_only_total_tokens(self._response_usage(combined_usage_object))
|
||||
if not response_usage[2] and not combined_usage[2] and aggregate_total <= 0:
|
||||
return ()
|
||||
resolved_usage: Final = (
|
||||
|
|
@ -4807,11 +4824,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
subrequest), falls through to the unreserved metadata rebuild so
|
||||
project IO quotas still receive the summary's actual usage.
|
||||
"""
|
||||
callback_kwargs: Final[object] = kwargs
|
||||
if not isinstance(kwargs, dict):
|
||||
return ()
|
||||
stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs))
|
||||
if stash is None:
|
||||
return self._build_unreserved_project_io_token_ops(kwargs, response_obj)
|
||||
return self._build_unreserved_project_io_token_ops(callback_kwargs, response_obj)
|
||||
|
||||
itpm_reserved: Final = stash.itpm_reserved_tokens
|
||||
otpm_reserved: Final = stash.otpm_reserved_tokens
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue