fix(lint): resolve the basedpyright/ruff-strict budget gate failures in both tag rate limit hooks

Fixes the specific type errors the gate was blaming this branch for,
rather than leaving them as a flagged pre-existing gap:

- Both async_log_success_event/async_log_failure_event overrides now carry
  full parameter type annotations (Mapping[str, object]/object), matching
  CustomLogger's own contract and clearing 16 ANN001/reportMissingParameterType
  violations. Reading standard_logging_object off the now-typed kwargs needs
  the same reportAssignmentType suppression shadow_eval_logger.py already
  uses for the identical pattern; litellm_params_for_metadata now narrows via
  isinstance(..., Mapping) instead of a bare `or` fallback, which also
  resolves several downstream reportArgumentType errors on
  _extract_key_hash/_extract_key_alias/_extract_team_id.
- Passed the already-list-typed group_operations directly instead of
  wrapping it in a redundant tuple(), which was the actual reportArgumentType
  mismatch against async_increment_tokens_with_ttl_preservation's
  list[RedisPipelineIncrementOperation] parameter.
- Added a small _model_name helper so sorting/grouping deployments by
  model_name compares an actual str instead of an untyped object, clearing
  the "no overloads for sorted match" reportCallIssue.

Net effect confirmed via a before/after basedpyright diff scoped to both
files: every affected rule's count went down or stayed flat, none went up.
Both scripts/type_check_gate.py and scripts/ruff_strict_gate.py now report
OK against origin/litellm_internal_staging; make lint-budget-update ran
afterward but had nothing to ratchet since these are new files with no
prior baseline to have improved on.
This commit is contained in:
Deepanshu 2026-08-27 17:46:29 -04:00
parent fee2c41359
commit 0f9b6f0833
2 changed files with 34 additions and 15 deletions

View file

@ -711,7 +711,9 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
"""
await self._release_pending_for_call_id(request_data)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
async def async_log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
# Always release regardless of which hook raised: this hook's own
# rejection never reserves a slot, so pending_concurrency_keys is
# already empty in that case and the check below no-ops; a rejection
@ -719,7 +721,9 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
# land after this hook already reserved its own slot.
await self._release_pending_for_call_id(kwargs)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
stash: Final = _stash_for_call(_call_id_from_kwargs(kwargs))
if stash is not None and stash.pending_concurrency_keys:
release_task: Final = asyncio.create_task(self._release_pending_for_call_id(kwargs))
@ -730,7 +734,9 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
if config is None:
return
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object")
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get( # pyright: ignore[reportAssignmentType] # untyped callback kwargs, same as shadow_eval_logger.py's identical read
"standard_logging_object"
)
if standard_logging_object is None:
return
@ -738,7 +744,8 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
# request kwargs admission sees: metadata/litellm_metadata are never
# top-level here, only nested under kwargs["litellm_params"] (see
# Logging.update_environment_variables).
litellm_params_for_metadata: Final = kwargs.get("litellm_params") or kwargs
litellm_params_raw: Final = kwargs.get("litellm_params")
litellm_params_for_metadata: Final = litellm_params_raw if isinstance(litellm_params_raw, Mapping) else kwargs
metadata_variable_name: Final = _resolve_success_event_metadata_variable_name(litellm_params_for_metadata)
key_hash: Final = _extract_key_hash(litellm_params_for_metadata, metadata_variable_name)
key_alias: Final = _extract_key_alias(litellm_params_for_metadata, metadata_variable_name)
@ -803,7 +810,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
partition = await self._partition_for(partition_key) # not Final: rebound each loop iteration
accounting_task = asyncio.create_task( # not Final: rebound each loop iteration
partition.v3.async_increment_tokens_with_ttl_preservation(
pipeline_operations=tuple(group_operations), parent_otel_span=None
pipeline_operations=group_operations, parent_otel_span=None
)
)
_BACKGROUND_TASKS.add(accounting_task) # mutable-ok: see _BACKGROUND_TASKS's own docstring

View file

@ -580,6 +580,13 @@ class _LimitsIndex:
return tuple(deduped.values())
def _model_name(deployment: Mapping[str, object]) -> str:
name = deployment["model_name"]
if not isinstance(name, str):
raise TypeError(f"deployment model_name must be a string, got {type(name)!r}")
return name
def _team_alias_key(deployment: Mapping[str, object]) -> tuple[str, str] | None:
model_info: Final = deployment.get("model_info") or _EMPTY_MAPPING
team_id: Final = model_info.get("team_id")
@ -624,7 +631,7 @@ def _build_limits_index(model_list: Sequence[Mapping[str, object]]) -> _LimitsIn
would have seen them in without the sort, which is what keeps this safe
(that relative order decides first-seen signature order downstream).
"""
sorted_by_model_name: Final = sorted(model_list, key=lambda deployment: deployment["model_name"])
sorted_by_model_name: Final = sorted(model_list, key=_model_name)
by_model_name: Final[Mapping[str, tuple[_ConfiguredLimit, ...]]] = MappingProxyType(
{
model_name: (
@ -641,9 +648,7 @@ def _build_limits_index(model_list: Sequence[Mapping[str, object]]) -> _LimitsIn
if (team_scope := next((key[0] for dep in group if (key := _team_alias_key(dep))), None)) is not None
else configured
)
for model_name, deployment_group in groupby(
sorted_by_model_name, key=lambda deployment: deployment["model_name"]
)
for model_name, deployment_group in groupby(sorted_by_model_name, key=_model_name)
for group in (tuple(deployment_group),)
if (configured := tuple(limit for unit in _LIMIT_UNITS for limit in _build_group_limits(group, unit)))
}
@ -1818,7 +1823,9 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
)
await self._release_keys(release_keys)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
async def async_log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
# No special-case skip for this hook's own tag_rate_limit_exceeded
# rejection: a hop whose own admission rejects never reaches the
# point where a concurrency reservation is queued (see
@ -1832,7 +1839,9 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
if release_keys:
await self._release_keys(release_keys)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
release_keys: Final = await self._pop_pending_concurrency_keys(kwargs)
if release_keys:
release_task: Final = asyncio.create_task(self._release_keys(release_keys))
@ -1842,7 +1851,9 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
if self.llm_router is None:
return
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object")
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get( # pyright: ignore[reportAssignmentType] # untyped callback kwargs, same as shadow_eval_logger.py's identical read
"standard_logging_object"
)
if standard_logging_object is None:
return
@ -1854,7 +1865,8 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
# request kwargs admission sees: metadata/litellm_metadata are never
# top-level here, only nested under kwargs["litellm_params"] (see
# Logging.update_environment_variables).
litellm_params_for_metadata: Final = kwargs.get("litellm_params") or kwargs
litellm_params_raw: Final = kwargs.get("litellm_params")
litellm_params_for_metadata: Final = litellm_params_raw if isinstance(litellm_params_raw, Mapping) else kwargs
metadata_variable_name: Final = _resolve_success_event_metadata_variable_name(litellm_params_for_metadata)
team_id: Final = _extract_team_id(litellm_params_for_metadata, metadata_variable_name)
key_hash: Final = _extract_key_hash(litellm_params_for_metadata, metadata_variable_name)
@ -1932,12 +1944,12 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
operations = operations_by_partition.setdefault(partition_key, []) # mutable-ok: see above
operations.append(operation) # mutable-ok: see comment above
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs if isinstance(kwargs, dict) else None)
for partition_key, group_operations in operations_by_partition.items():
partition = await self._partition_for(partition_key) # not Final: rebound each loop iteration
accounting_task = asyncio.create_task( # not Final: rebound each loop iteration
partition.v3.async_increment_tokens_with_ttl_preservation(
pipeline_operations=tuple(group_operations),
pipeline_operations=group_operations,
parent_otel_span=parent_otel_span,
)
)