fix(proxy): type summary quota helpers without unknown arguments

Signed-off-by: sinksilk <785976238@qq.com>
This commit is contained in:
sinksilk 2026-10-02 00:52:05 +08:00
parent 4d6f12578d
commit 25ca54ca2b
2 changed files with 45 additions and 13 deletions

View file

@ -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,
)

View file

@ -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