refactor: extract success pipeline ops to fix PLR0915 in parallel_request_limiter_v3

Move async_log_success_event pipeline construction into
_build_success_event_pipeline_operations so the async hook stays
under Ruff's max-statement limit.

Made-with: Cursor
This commit is contained in:
shivam 2026-04-18 11:38:05 -07:00
parent 46a04e26f1
commit 0d950612f4
No known key found for this signature in database

View file

@ -1543,6 +1543,186 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return "total" # default to total
return specified_rate_limit_type
def _build_success_event_pipeline_operations(
self,
kwargs: Any,
response_obj: Any,
rate_limit_type: Literal["output", "input", "total"],
) -> List["RedisPipelineIncrementOperation"]:
"""Build Redis pipeline increment ops for TPM / parallel-request counters."""
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
)
from litellm.types.caching import RedisPipelineIncrementOperation
# Get metadata from standard_logging_object - this correctly handles both
# 'metadata' and 'litellm_metadata' fields from litellm_params
standard_logging_object = kwargs.get("standard_logging_object") or {}
standard_logging_metadata = standard_logging_object.get("metadata") or {}
# user_api_key_hash is the same as user_api_key (it's the hash)
user_api_key = standard_logging_metadata.get("user_api_key_hash")
user_api_key_user_id = standard_logging_metadata.get("user_api_key_user_id")
user_api_key_team_id = standard_logging_metadata.get("user_api_key_team_id")
user_api_key_organization_id = standard_logging_metadata.get(
"user_api_key_org_id"
)
user_api_key_project_id = standard_logging_metadata.get(
"user_api_key_project_id"
)
user_api_key_end_user_id = kwargs.get(
"user"
) or standard_logging_metadata.get("user_api_key_end_user_id")
model_group = get_model_group_from_litellm_kwargs(kwargs)
# Get total tokens from response
total_tokens = 0
# spot fix for /responses api
if isinstance(response_obj, ModelResponse) or isinstance(
response_obj, BaseLiteLLMOpenAIResponseObject
):
_usage = getattr(response_obj, "usage", None)
total_tokens = self._get_total_tokens_from_usage(
usage=_usage, rate_limit_type=rate_limit_type
)
pipeline_operations: List[RedisPipelineIncrementOperation] = []
# API Key TPM
if user_api_key:
# MAX PARALLEL REQUESTS - only support for API Key, just decrement the counter
counter_key = self.create_rate_limit_keys(
key="api_key",
value=user_api_key,
rate_limit_type="max_parallel_requests",
)
pipeline_operations.append(
RedisPipelineIncrementOperation(
key=counter_key,
increment_value=-1,
ttl=self.window_size,
)
)
pipeline_operations.extend(
self._create_pipeline_operations(
key="api_key",
value=user_api_key,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# User TPM
if user_api_key_user_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="user",
value=user_api_key_user_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Team TPM
if user_api_key_team_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="team",
value=user_api_key_team_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Team Member TPM
if user_api_key_team_id and user_api_key_user_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="team_member",
value=f"{user_api_key_team_id}:{user_api_key_user_id}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# End User TPM
if user_api_key_end_user_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="end_user",
value=user_api_key_end_user_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Model-specific TPM
if model_group and user_api_key:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_key",
value=f"{user_api_key}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
if model_group and user_api_key_team_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_team",
value=f"{user_api_key_team_id}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
if model_group and user_api_key_organization_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_organization",
value=f"{user_api_key_organization_id}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
if model_group and user_api_key_project_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_project",
value=f"{user_api_key_project_id}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Agent TPM
agent_id = standard_logging_metadata.get("agent_id")
if agent_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="agent",
value=agent_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Agent Session TPM
session_id = standard_logging_metadata.get(
"session_id"
) or standard_logging_metadata.get("trace_id")
if session_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="agent_session",
value=f"{agent_id}:{session_id}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
return pipeline_operations
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
Update TPM usage on successful API calls by incrementing counters using pipeline
@ -1550,10 +1730,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
)
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
)
from litellm.types.caching import RedisPipelineIncrementOperation
rate_limit_type = self.get_rate_limit_type()
@ -1565,175 +1741,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"INSIDE parallel request limiter ASYNC SUCCESS LOGGING"
)
# Get metadata from standard_logging_object - this correctly handles both
# 'metadata' and 'litellm_metadata' fields from litellm_params
standard_logging_object = kwargs.get("standard_logging_object") or {}
standard_logging_metadata = standard_logging_object.get("metadata") or {}
# user_api_key_hash is the same as user_api_key (it's the hash)
user_api_key = standard_logging_metadata.get("user_api_key_hash")
user_api_key_user_id = standard_logging_metadata.get("user_api_key_user_id")
user_api_key_team_id = standard_logging_metadata.get("user_api_key_team_id")
user_api_key_organization_id = standard_logging_metadata.get(
"user_api_key_org_id"
pipeline_operations = self._build_success_event_pipeline_operations(
kwargs=kwargs,
response_obj=response_obj,
rate_limit_type=rate_limit_type,
)
user_api_key_project_id = standard_logging_metadata.get(
"user_api_key_project_id"
)
user_api_key_end_user_id = kwargs.get(
"user"
) or standard_logging_metadata.get("user_api_key_end_user_id")
model_group = get_model_group_from_litellm_kwargs(kwargs)
# Get total tokens from response
total_tokens = 0
# spot fix for /responses api
if isinstance(response_obj, ModelResponse) or isinstance(
response_obj, BaseLiteLLMOpenAIResponseObject
):
_usage = getattr(response_obj, "usage", None)
total_tokens = self._get_total_tokens_from_usage(
usage=_usage, rate_limit_type=rate_limit_type
)
# Create pipeline operations for TPM increments
pipeline_operations: List[RedisPipelineIncrementOperation] = []
# API Key TPM
if user_api_key:
# MAX PARALLEL REQUESTS - only support for API Key, just decrement the counter
counter_key = self.create_rate_limit_keys(
key="api_key",
value=user_api_key,
rate_limit_type="max_parallel_requests",
)
pipeline_operations.append(
RedisPipelineIncrementOperation(
key=counter_key,
increment_value=-1,
ttl=self.window_size,
)
)
pipeline_operations.extend(
self._create_pipeline_operations(
key="api_key",
value=user_api_key,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# User TPM
if user_api_key_user_id:
# TPM
pipeline_operations.extend(
self._create_pipeline_operations(
key="user",
value=user_api_key_user_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Team TPM
if user_api_key_team_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="team",
value=user_api_key_team_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Team Member TPM
if user_api_key_team_id and user_api_key_user_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="team_member",
value=f"{user_api_key_team_id}:{user_api_key_user_id}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# End User TPM
if user_api_key_end_user_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="end_user",
value=user_api_key_end_user_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Model-specific TPM
if model_group and user_api_key:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_key",
value=f"{user_api_key}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
if model_group and user_api_key_team_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_team",
value=f"{user_api_key_team_id}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
if model_group and user_api_key_organization_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_organization",
value=f"{user_api_key_organization_id}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
if model_group and user_api_key_project_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="model_per_project",
value=f"{user_api_key_project_id}:{model_group}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Agent TPM
agent_id = standard_logging_metadata.get("agent_id")
if agent_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="agent",
value=agent_id,
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Agent Session TPM
session_id = standard_logging_metadata.get(
"session_id"
) or standard_logging_metadata.get("trace_id")
if session_id:
pipeline_operations.extend(
self._create_pipeline_operations(
key="agent_session",
value=f"{agent_id}:{session_id}",
rate_limit_type="tokens",
total_tokens=total_tokens,
)
)
# Execute all increments in a single pipeline
if pipeline_operations:
await self.async_increment_tokens_with_ttl_preservation(
pipeline_operations=pipeline_operations,