From f72e6464e8fa2a56bd8849234f231124ee48750b Mon Sep 17 00:00:00 2001 From: shin-bot-litellm Date: Sat, 31 Jan 2026 08:06:30 +0000 Subject: [PATCH] litellm_fix: refactor base_process_llm_request to fix PLR0915 - Extract response metadata extraction into _extract_response_metadata helper - Add ResponseMetadata NamedTuple for type-safe metadata handling - Use metadata object instead of individual variables to reduce statement count - Remove noqa: PLR0915 suppression as the function now passes ruff check This properly addresses the 'too many statements' warning by refactoring instead of suppressing it. --- litellm/proxy/common_request_processing.py | 307 ++++++++------------- 1 file changed, 108 insertions(+), 199 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index d3f01f3bd2c..d0850106b89 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -9,6 +9,7 @@ from typing import ( AsyncGenerator, Callable, Literal, + NamedTuple, Optional, Tuple, Union, @@ -55,20 +56,14 @@ from litellm.types.utils import ModelResponse, ModelResponseStream, Usage async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]: """Parses an event line and returns an error code if present, else None.""" - event_line = ( - event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line - ) + event_line = event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line if event_line.startswith("data: "): json_str = event_line[len("data: ") :].strip() if not json_str or json_str == "[DONE]": # handle empty data or [DONE] message return None try: data = orjson.loads(json_str) - if ( - isinstance(data, dict) - and "error" in data - and isinstance(data["error"], dict) - ): + if isinstance(data, dict) and "error" in data and isinstance(data["error"], dict): error_code_raw = data["error"].get("code") error_code: Optional[int] = None @@ -87,12 +82,8 @@ async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional # Ensure error_code is a valid HTTP status code if error_code is not None and 100 <= error_code <= 599: return error_code - elif ( - error_code_raw is not None - ): # Log if original code was present but not valid - verbose_proxy_logger.warning( - f"Error has invalid or non-convertible code: {error_code_raw}" - ) + elif error_code_raw is not None: # Log if original code was present but not valid + verbose_proxy_logger.warning(f"Error has invalid or non-convertible code: {error_code_raw}") except (orjson.JSONDecodeError, json.JSONDecodeError): # not a known error chunk pass @@ -109,9 +100,7 @@ def _extract_error_from_sse_chunk(event_line: Union[str, bytes]) -> dict: Returns: Error dictionary in OpenAI API format """ - event_line = ( - event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line - ) + event_line = event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line # Default error format default_error = { @@ -162,9 +151,7 @@ async def create_response( if first_chunk_value is not None: try: - error_code_from_chunk = await _parse_event_data_for_error( - first_chunk_value - ) + error_code_from_chunk = await _parse_event_data_for_error(first_chunk_value) if error_code_from_chunk is not None: # First chunk is an error, stream hasn't really started yet # Should return standard JSON error response instead of SSE format @@ -205,9 +192,7 @@ async def create_response( ) except Exception as e: # Unexpected error consuming first chunk. - verbose_proxy_logger.exception( - f"Error consuming first chunk from generator: {e}" - ) + verbose_proxy_logger.exception(f"Error consuming first chunk from generator: {e}") # Fallback to a generic error stream async def error_gen_message() -> AsyncGenerator[str, None]: @@ -325,6 +310,18 @@ def _get_cost_breakdown_from_logging_obj( return original_cost, discount_amount, margin_total_amount, margin_percent +class ResponseMetadata(NamedTuple): + """Metadata extracted from LLM response hidden_params.""" + + hidden_params: dict + model_id: str + cache_key: str + api_base: str + response_cost: str + fastest_response_batch_completion: Optional[bool] + additional_headers: dict + + class ProxyBaseLLMRequestProcessing: def __init__(self, data: dict): self.data = data @@ -356,9 +353,7 @@ class ProxyBaseLLMRequestProcessing: discount_amount, margin_total_amount, margin_percent, - ) = _get_cost_breakdown_from_logging_obj( - litellm_logging_obj=litellm_logging_obj - ) + ) = _get_cost_breakdown_from_logging_obj(litellm_logging_obj=litellm_logging_obj) # Calculate updated spend for header (include current response_cost) current_spend = user_api_key_dict.spend or 0.0 @@ -366,11 +361,7 @@ class ProxyBaseLLMRequestProcessing: if response_cost is not None: try: # Convert response_cost to float if it's a string - cost_value = ( - float(response_cost) - if isinstance(response_cost, str) - else response_cost - ) + cost_value = float(response_cost) if isinstance(response_cost, str) else response_cost if cost_value > 0: updated_spend = current_spend + cost_value except (ValueError, TypeError): @@ -387,40 +378,26 @@ class ProxyBaseLLMRequestProcessing: "x-litellm-version": version, "x-litellm-model-region": model_region, "x-litellm-response-cost": str(response_cost), - "x-litellm-response-cost-original": ( - str(original_cost) if original_cost is not None else None - ), - "x-litellm-response-cost-discount-amount": ( - str(discount_amount) if discount_amount is not None else None - ), + "x-litellm-response-cost-original": (str(original_cost) if original_cost is not None else None), + "x-litellm-response-cost-discount-amount": (str(discount_amount) if discount_amount is not None else None), "x-litellm-response-cost-margin-amount": ( str(margin_total_amount) if margin_total_amount is not None else None ), - "x-litellm-response-cost-margin-percent": ( - str(margin_percent) if margin_percent is not None else None - ), + "x-litellm-response-cost-margin-percent": (str(margin_percent) if margin_percent is not None else None), "x-litellm-key-tpm-limit": str(user_api_key_dict.tpm_limit), "x-litellm-key-rpm-limit": str(user_api_key_dict.rpm_limit), "x-litellm-key-max-budget": str(user_api_key_dict.max_budget), "x-litellm-key-spend": str(updated_spend), - "x-litellm-response-duration-ms": str( - hidden_params.get("_response_ms", None) - ), - "x-litellm-overhead-duration-ms": str( - hidden_params.get("litellm_overhead_time_ms", None) - ), + "x-litellm-response-duration-ms": str(hidden_params.get("_response_ms", None)), + "x-litellm-overhead-duration-ms": str(hidden_params.get("litellm_overhead_time_ms", None)), "x-litellm-fastest_response_batch_completion": ( - str(fastest_response_batch_completion) - if fastest_response_batch_completion is not None - else None + str(fastest_response_batch_completion) if fastest_response_batch_completion is not None else None ), "x-litellm-timeout": str(timeout) if timeout is not None else None, **{k: str(v) for k, v in kwargs.items()}, } if request_data: - remaining_tokens_header = ( - get_remaining_tokens_and_requests_from_request_data(request_data) - ) + remaining_tokens_header = get_remaining_tokens_and_requests_from_request_data(request_data) headers.update(remaining_tokens_header) logging_caching_headers = get_logging_caching_headers(request_data) @@ -428,11 +405,7 @@ class ProxyBaseLLMRequestProcessing: headers.update(logging_caching_headers) try: - return { - key: str(value) - for key, value in headers.items() - if value not in exclude_values - } + return {key: str(value) for key, value in headers.items() if value not in exclude_values} except Exception as e: verbose_proxy_logger.error(f"Error setting custom headers: {e}") return {} @@ -539,9 +512,7 @@ class ProxyBaseLLMRequestProcessing: self.data[_metadata_variable_name] = {} if not isinstance(self.data[_metadata_variable_name], dict): self.data[_metadata_variable_name] = {} - self.data[_metadata_variable_name][ - "queue_time_seconds" - ] = queue_time_seconds + self.data[_metadata_variable_name]["queue_time_seconds"] = queue_time_seconds self.data["model"] = ( general_settings.get("completion_model", None) # server default @@ -563,10 +534,7 @@ class ProxyBaseLLMRequestProcessing: ### MODEL ALIAS MAPPING ### # check if model name in model alias map # get the actual model name - if ( - isinstance(self.data["model"], str) - and self.data["model"] in litellm.model_alias_map - ): + if isinstance(self.data["model"], str) and self.data["model"] in litellm.model_alias_map: self.data["model"] = litellm.model_alias_map[self.data["model"]] # Check key-specific aliases @@ -578,9 +546,7 @@ class ProxyBaseLLMRequestProcessing: ): self.data["model"] = user_api_key_dict.aliases[self.data["model"]] - self.data["litellm_call_id"] = request.headers.get( - "x-litellm-call-id", str(uuid.uuid4()) - ) + self.data["litellm_call_id"] = request.headers.get("x-litellm-call-id", str(uuid.uuid4())) ### CALL HOOKS ### - modify/reject incoming data before calling the model ## LOGGING OBJECT ## - initialize logging object for logging success/failure events for call @@ -595,7 +561,9 @@ class ProxyBaseLLMRequestProcessing: self.data["litellm_logging_obj"] = logging_obj self.data = await proxy_logging_obj.pre_call_hook( # type: ignore - user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type # type: ignore + user_api_key_dict=user_api_key_dict, + data=self.data, + call_type=route_type, # type: ignore ) # Apply hierarchical router_settings (Key > Team > Global) @@ -623,7 +591,29 @@ class ProxyBaseLLMRequestProcessing: return self.data, logging_obj - async def base_process_llm_request( # noqa: PLR0915 + @staticmethod + def _extract_response_metadata(response: Any, data: dict) -> ResponseMetadata: + """Extract metadata from LLM response hidden_params.""" + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) or "" + + # Fallback: extract model_id from litellm_metadata if not in hidden_params + if not model_id: + litellm_metadata = data.get("litellm_metadata", {}) or {} + model_info = litellm_metadata.get("model_info", {}) or {} + model_id = model_info.get("id", "") or "" + + return ResponseMetadata( + hidden_params=hidden_params, + model_id=model_id, + cache_key=hidden_params.get("cache_key", None) or "", + api_base=hidden_params.get("api_base", None) or "", + response_cost=hidden_params.get("response_cost", None) or "", + fastest_response_batch_completion=hidden_params.get("fastest_response_batch_completion", None), + additional_headers=hidden_params.get("additional_headers", {}) or {}, + ) + + async def base_process_llm_request( self, request: Request, fastapi_response: Response, @@ -698,9 +688,7 @@ class ProxyBaseLLMRequestProcessing: ) if verbose_proxy_logger.isEnabledFor(logging.DEBUG): verbose_proxy_logger.debug( - "Request received by LiteLLM:\n{}".format( - json.dumps(self.data, indent=4, default=str) - ), + "Request received by LiteLLM:\n{}".format(json.dumps(self.data, indent=4, default=str)), ) self.data, logging_obj = await self.common_processing_pre_call_logic( @@ -748,36 +736,18 @@ class ProxyBaseLLMRequestProcessing: tasks.append(llm_call) # wait for call to end - llm_responses = asyncio.gather( - *tasks - ) # run the moderation check in parallel to the actual llm api call + llm_responses = asyncio.gather(*tasks) # run the moderation check in parallel to the actual llm api call responses = await llm_responses response = responses[1] - hidden_params = getattr(response, "_hidden_params", {}) or {} - model_id = hidden_params.get("model_id", None) or "" - - # Fallback: extract model_id from litellm_metadata if not in hidden_params - if not model_id: - litellm_metadata = self.data.get("litellm_metadata", {}) or {} - model_info = litellm_metadata.get("model_info", {}) or {} - model_id = model_info.get("id", "") or "" - - cache_key, api_base, response_cost = ( - hidden_params.get("cache_key", None) or "", - hidden_params.get("api_base", None) or "", - hidden_params.get("response_cost", None) or "", - ) - fastest_response_batch_completion, additional_headers = ( - hidden_params.get("fastest_response_batch_completion", None), - hidden_params.get("additional_headers", {}) or {}, - ) + # Extract response metadata + metadata = self._extract_response_metadata(response, self.data) # Post Call Processing if llm_router is not None: - self.data["deployment"] = llm_router.get_deployment(model_id=model_id) + self.data["deployment"] = llm_router.get_deployment(model_id=metadata.model_id) asyncio.create_task( proxy_logging_obj.update_request_status( litellm_call_id=self.data.get("litellm_call_id", ""), status="success" @@ -785,23 +755,21 @@ class ProxyBaseLLMRequestProcessing: ) if self._is_streaming_request( data=self.data, is_streaming_request=is_streaming_request - ) or self._is_streaming_response( - response - ): # use generate_responses to stream responses + ) or self._is_streaming_response(response): # use generate_responses to stream responses custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( user_api_key_dict=user_api_key_dict, call_id=logging_obj.litellm_call_id, - model_id=model_id, - cache_key=cache_key, - api_base=api_base, + model_id=metadata.model_id, + cache_key=metadata.cache_key, + api_base=metadata.api_base, version=version, - response_cost=response_cost, + response_cost=metadata.response_cost, model_region=getattr(user_api_key_dict, "allowed_model_region", ""), - fastest_response_batch_completion=fastest_response_batch_completion, + fastest_response_batch_completion=metadata.fastest_response_batch_completion, request_data=self.data, - hidden_params=hidden_params, + hidden_params=metadata.hidden_params, litellm_logging_obj=logging_obj, - **additional_headers, + **metadata.additional_headers, ) # Call response headers hook for streaming success @@ -847,13 +815,11 @@ class ProxyBaseLLMRequestProcessing: # This handles cases like websearch_interception agentic loop # which returns a non-streaming dict even for streaming requests if self._is_streaming_response(response): - selected_data_generator = ( - ProxyBaseLLMRequestProcessing.async_sse_data_generator( - response=response, - user_api_key_dict=user_api_key_dict, - request_data=self.data, - proxy_logging_obj=proxy_logging_obj, - ) + selected_data_generator = ProxyBaseLLMRequestProcessing.async_sse_data_generator( + response=response, + user_api_key_dict=user_api_key_dict, + request_data=self.data, + proxy_logging_obj=proxy_logging_obj, ) return await create_response( generator=selected_data_generator, @@ -887,26 +853,25 @@ class ProxyBaseLLMRequestProcessing: log_context=f"litellm_call_id={logging_obj.litellm_call_id}", ) - hidden_params = ( - getattr(response, "_hidden_params", {}) or {} - ) # get any updated response headers - additional_headers = hidden_params.get("additional_headers", {}) or {} + # Get any updated response headers + updated_hidden_params = getattr(response, "_hidden_params", {}) or {} + updated_additional_headers = updated_hidden_params.get("additional_headers", {}) or {} fastapi_response.headers.update( ProxyBaseLLMRequestProcessing.get_custom_headers( user_api_key_dict=user_api_key_dict, call_id=logging_obj.litellm_call_id, - model_id=model_id, - cache_key=cache_key, - api_base=api_base, + model_id=metadata.model_id, + cache_key=metadata.cache_key, + api_base=metadata.api_base, version=version, - response_cost=response_cost, + response_cost=metadata.response_cost, model_region=getattr(user_api_key_dict, "allowed_model_region", ""), - fastest_response_batch_completion=fastest_response_batch_completion, + fastest_response_batch_completion=metadata.fastest_response_batch_completion, request_data=self.data, - hidden_params=hidden_params, + hidden_params=updated_hidden_params, litellm_logging_obj=logging_obj, - **additional_headers, + **updated_additional_headers, ) ) @@ -998,9 +963,7 @@ class ProxyBaseLLMRequestProcessing: return False - def _is_streaming_request( - self, data: dict, is_streaming_request: Optional[bool] = False - ) -> bool: + def _is_streaming_request(self, data: dict, is_streaming_request: Optional[bool] = False) -> bool: """ Check if the request is a streaming request. @@ -1043,9 +1006,7 @@ class ProxyBaseLLMRequestProcessing: timeout = getattr( e, "timeout", None ) # returns the timeout set by the wrapper. Used for testing if model-specific timeout are set correctly - _litellm_logging_obj: Optional[LiteLLMLoggingObj] = self.data.get( - "litellm_logging_obj", None - ) + _litellm_logging_obj: Optional[LiteLLMLoggingObj] = self.data.get("litellm_logging_obj", None) # Attempt to get model_id from logging object # @@ -1055,9 +1016,7 @@ class ProxyBaseLLMRequestProcessing: custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( user_api_key_dict=user_api_key_dict, - call_id=( - _litellm_logging_obj.litellm_call_id if _litellm_logging_obj else None - ), + call_id=(_litellm_logging_obj.litellm_call_id if _litellm_logging_obj else None), model_id=model_id, version=version, response_cost=0, @@ -1113,18 +1072,9 @@ class ProxyBaseLLMRequestProcessing: # 1. Direct AttributeError (already handled above) # 2. In underlying exception (__cause__, __context__, original_exception) has_attribute_error = ( - ( - isinstance(e, Exception) - and isinstance(getattr(e, "__cause__", None), AttributeError) - ) - or ( - isinstance(e, Exception) - and isinstance(getattr(e, "__context__", None), AttributeError) - ) - or ( - isinstance(e, Exception) - and isinstance(getattr(e, "original_exception", None), AttributeError) - ) + (isinstance(e, Exception) and isinstance(getattr(e, "__cause__", None), AttributeError)) + or (isinstance(e, Exception) and isinstance(getattr(e, "__context__", None), AttributeError)) + or (isinstance(e, Exception) and isinstance(getattr(e, "original_exception", None), AttributeError)) ) if has_attribute_error: @@ -1181,16 +1131,12 @@ class ProxyBaseLLMRequestProcessing: verbose_proxy_logger.debug("inside generator") try: str_so_far = "" - async for ( - chunk - ) in proxy_logging_obj.async_post_call_streaming_iterator_hook( + async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( user_api_key_dict=user_api_key_dict, response=response, request_data=request_data, ): - verbose_proxy_logger.debug( - "async_data_generator: received streaming chunk - {}".format(chunk) - ) + verbose_proxy_logger.debug("async_data_generator: received streaming chunk - {}".format(chunk)) ### CALL HOOKS ### - modify outgoing data chunk = await proxy_logging_obj.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, @@ -1205,19 +1151,13 @@ class ProxyBaseLLMRequestProcessing: # Inject cost into Anthropic-style SSE usage for /v1/messages for any provider model_name = request_data.get("model", "") - chunk = ( - ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection( - chunk, model_name - ) - ) + chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(chunk, model_name) # Format chunk using helper function yield ProxyBaseLLMRequestProcessing.return_sse_chunk(chunk) except Exception as e: verbose_proxy_logger.exception( - "litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format( - str(e) - ) + "litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(str(e)) ) # Allow callbacks to transform the error response transformed_exception = await proxy_logging_obj.post_call_failure_hook( @@ -1264,42 +1204,24 @@ class ProxyBaseLLMRequestProcessing: try: if isinstance(chunk, dict): - maybe_modified = ( - ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict( - chunk, model_name - ) - ) + maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(chunk, model_name) if maybe_modified is not None: return maybe_modified elif isinstance(chunk, (bytes, bytearray)): # Decode to str, inject, and rebuild as bytes try: s = chunk.decode("utf-8", errors="ignore") - maybe_mod = ( - ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str( - s, model_name - ) - ) + maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(s, model_name) if maybe_mod is not None: - return ( - maybe_mod + ("" if maybe_mod.endswith("\n\n") else "\n\n") - ).encode("utf-8") + return (maybe_mod + ("" if maybe_mod.endswith("\n\n") else "\n\n")).encode("utf-8") except Exception: pass elif isinstance(chunk, str): # Try to parse SSE frame and inject cost into the data line - maybe_mod = ( - ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str( - chunk, model_name - ) - ) + maybe_mod = ProxyBaseLLMRequestProcessing._inject_cost_into_sse_frame_str(chunk, model_name) if maybe_mod is not None: # Ensure trailing frame separator - return ( - maybe_mod - if maybe_mod.endswith("\n\n") - else (maybe_mod + "\n\n") - ) + return maybe_mod if maybe_mod.endswith("\n\n") else (maybe_mod + "\n\n") except Exception: # Never break streaming on optional cost injection pass @@ -1307,9 +1229,7 @@ class ProxyBaseLLMRequestProcessing: return chunk @staticmethod - def _inject_cost_into_sse_frame_str( - frame_str: str, model_name: str - ) -> Optional[str]: + def _inject_cost_into_sse_frame_str(frame_str: str, model_name: str) -> Optional[str]: """ Inject cost information into an SSE frame string by modifying the JSON in the 'data:' line. @@ -1329,11 +1249,7 @@ class ProxyBaseLLMRequestProcessing: json_part = stripped_ln.split("data:", 1)[1].strip() if json_part and json_part != "[DONE]": obj = json.loads(json_part) - maybe_modified = ( - ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict( - obj, model_name - ) - ) + maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict(obj, model_name) if maybe_modified is not None: # Replace just this line with updated JSON using safe_dumps lines[idx] = f"data: {safe_dumps(maybe_modified)}" @@ -1359,8 +1275,7 @@ class ProxyBaseLLMRequestProcessing: prompt_tokens = int(_usage.get("input_tokens", 0) or 0) completion_tokens = int(_usage.get("output_tokens", 0) or 0) total_tokens = int( - _usage.get("total_tokens", prompt_tokens + completion_tokens) - or (prompt_tokens + completion_tokens) + _usage.get("total_tokens", prompt_tokens + completion_tokens) or (prompt_tokens + completion_tokens) ) # Extract additional usage fields @@ -1384,15 +1299,11 @@ class ProxyBaseLLMRequestProcessing: # Handle web_search_requests by wrapping in ServerToolUse if web_search_requests is not None: - usage_kwargs["server_tool_use"] = ServerToolUse( - web_search_requests=web_search_requests - ) + usage_kwargs["server_tool_use"] = ServerToolUse(web_search_requests=web_search_requests) # Add cache-related fields to **params (handled by Usage.__init__) if cache_creation_input_tokens is not None: - usage_kwargs[ - "cache_creation_input_tokens" - ] = cache_creation_input_tokens + usage_kwargs["cache_creation_input_tokens"] = cache_creation_input_tokens if cache_read_input_tokens is not None: usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens @@ -1411,9 +1322,7 @@ class ProxyBaseLLMRequestProcessing: return obj return None - def maybe_get_model_id( - self, _logging_obj: Optional[LiteLLMLoggingObj] - ) -> Optional[str]: + def maybe_get_model_id(self, _logging_obj: Optional[LiteLLMLoggingObj]) -> Optional[str]: """ Get model_id from logging object or request metadata.