diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 3bdda0f7406..5c2b3dfe0ee 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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,