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.
This commit is contained in:
shin-bot-litellm 2026-01-31 08:06:30 +00:00
parent d3cad31441
commit f72e6464e8

View file

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