perf: address GIL contention and hot-path bottlenecks from profiling data

Based on py-spy GIL profiling (38,800 samples, 2000 concurrent users) and
pyinstrument per-request timing, this commit addresses the top performance
bottlenecks identified:

1. _sanitize_request_body_for_spend_logs_payload (2.9% GIL):
   - Remove redundant inner import (constants already imported at top-level)
   - Remove dead-code branch (len check after already confirmed len > max)
   - Pre-compute truncation ratios outside inner function
   - Reorder isinstance checks: str first (most common leaf type)

2. Pydantic repr in logging (2.3% GIL):
   - Guard print_deployment calls behind isEnabledFor(logging.INFO)
   - Replace copy.deepcopy with shallow dict() copy in print_deployment
   - Use %-style lazy formatting instead of f-strings for logger calls
   - Remove kwargs from prometheus debug log message

3. Prometheus label_factory overhead (1.5% + 0.7% GIL):
   - Cache model_dump() on UserAPIKeyLabelValues via get_label_dict()
   - Convert supported_enum_labels to frozenset for O(1) membership tests
   - Called 37 times per success event; caching avoids 36 redundant dumps

4. pre_call_utils header lookup (1.9% GIL):
   - Replace dict comprehension over all headers with early-exit loop
   - Only lowercase and compare the two target header names

5. safe_json_dumps (0.7% GIL):
   - Replace stdlib json.dumps with orjson.dumps for final serialization

6. Hot-path debug logging:
   - Convert f-string debug logs to %-style in litellm_logging.py
   - Simplify prometheus print_verbose call

7. Cost calculator annotation checks:
   - Optimize response_includes_annotation_type to handle both dict
     and object annotation types without repeated __getattr__ calls

Estimated GIL time reduction: ~11-12% under concurrency.

Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
Cursor Agent 2026-03-08 00:42:11 +00:00
parent 7f4cbf4893
commit a7b7a31b26
9 changed files with 100 additions and 91 deletions

View file

@ -887,7 +887,7 @@ class PrometheusLogger(CustomLogger):
from litellm.types.utils import StandardLoggingPayload
verbose_logger.debug(
f"prometheus Logging - Enters success logging function for kwargs {kwargs}"
"prometheus Logging - Enters success logging function"
)
# unpack kwargs
@ -944,7 +944,8 @@ class PrometheusLogger(CustomLogger):
_tags = []
print_verbose(
f"inside track_prometheus_metrics, model {model}, response_cost {response_cost}, tokens_used {tokens_used}, end_user_id {end_user_id}, user_api_key {user_api_key}"
"inside track_prometheus_metrics, model %s, response_cost %s, tokens_used %s"
% (model, response_cost, tokens_used)
)
enum_values = UserAPIKeyLabelValues(
@ -3056,15 +3057,13 @@ def prometheus_label_factory(
Ensures end_user param is not sent to prometheus if it is not supported.
"""
# Extract dictionary from Pydantic object
enum_dict = enum_values.model_dump()
enum_dict = enum_values.get_label_dict()
# Filter supported labels and sanitize values to prevent breaking
# the Prometheus text format (e.g. U+2028 Line Separator in label values)
supported_set = frozenset(supported_enum_labels) if not isinstance(supported_enum_labels, (set, frozenset)) else supported_enum_labels
filtered_labels = {
label: _sanitize_prometheus_label_value(value)
for label, value in enum_dict.items()
if label in supported_enum_labels
if label in supported_set
}
if UserAPIKeyLabelNames.END_USER.value in filtered_labels:
@ -3079,14 +3078,14 @@ def prometheus_label_factory(
for key, value in enum_values.custom_metadata_labels.items():
# check sanitized key
sanitized_key = _sanitize_prometheus_label_name(key)
if sanitized_key in supported_enum_labels:
if sanitized_key in supported_set:
filtered_labels[sanitized_key] = _sanitize_prometheus_label_value(value)
# Add custom tags if configured
if enum_values.tags is not None:
custom_tag_labels = get_custom_labels_from_tags(enum_values.tags)
for key, value in custom_tag_labels.items():
if key in supported_enum_labels:
if key in supported_set:
filtered_labels[key] = _sanitize_prometheus_label_value(value)
for label in supported_enum_labels:

View file

@ -1476,7 +1476,7 @@ class Logging(LiteLLMLoggingBaseClass):
**response_cost_calculator_kwargs
)
verbose_logger.debug(f"response_cost: {response_cost}")
verbose_logger.debug("response_cost: %s", response_cost)
return response_cost
except Exception as e: # error calculating cost
debug_info = StandardLoggingModelCostFailureDebugInformation(

View file

@ -387,11 +387,12 @@ class StandardBuiltInToolCostTracking:
message: Optional[Message] = getattr(choice, "message", None)
if message is None:
continue
if annotations := getattr(message, "annotations", None):
if len(annotations) > 0:
for annotation in annotations:
if annotation.get("type", None) == annotation_type:
return True
annotations = getattr(message, "annotations", None)
if annotations:
for annotation in annotations:
_type = annotation.get("type") if isinstance(annotation, dict) else getattr(annotation, "type", None)
if _type == annotation_type:
return True
return False
@staticmethod

View file

@ -1,6 +1,6 @@
import json
from typing import Any, Union
import orjson
from pydantic import BaseModel
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
@ -13,13 +13,10 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
"""
def _serialize(obj: Any, seen: set, depth: int) -> Any:
# Check for maximum depth.
if depth > max_depth:
return "MaxDepthExceeded"
# Base-case: if it is a primitive, simply return it.
if isinstance(obj, (str, int, float, bool, type(None))):
return obj
# Check for circular reference.
if id(obj) in seen:
return "CircularReference Detected"
seen.add(id(obj))
@ -27,7 +24,7 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
if isinstance(obj, dict):
result = {}
for k, v in obj.items():
if isinstance(k, (str)):
if isinstance(k, str):
result[k] = _serialize(v, seen, depth + 1)
seen.remove(id(obj))
return result
@ -49,11 +46,10 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
seen.remove(id(obj))
return result
else:
# Fall back to string conversion for non-serializable objects.
try:
return str(obj)
except Exception:
return "Unserializable Object"
safe_data = _serialize(data, set(), 0)
return json.dumps(safe_data, default=str)
return orjson.dumps(safe_data, default=str).decode()

View file

@ -102,10 +102,15 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str
"""
if not headers:
return None
normalized = {k.lower(): v for k, v in headers.items() if isinstance(k, str)}
return normalized.get("x-litellm-trace-id") or normalized.get(
"x-litellm-session-id"
)
session_id = None
for k, v in headers.items():
if isinstance(k, str):
k_lower = k.lower()
if k_lower == "x-litellm-trace-id":
return v
elif session_id is None and k_lower == "x-litellm-session-id":
session_id = v
return session_id
def safe_add_api_version_from_query_params(data: dict, request: Request):

View file

@ -632,61 +632,37 @@ def _sanitize_request_body_for_spend_logs_payload(
Recursively sanitize request body to prevent logging large base64 strings or other large values.
Truncates strings longer than MAX_STRING_LENGTH_PROMPT_IN_DB characters and handles nested dictionaries.
"""
from litellm.constants import (
LITELLM_TRUNCATED_PAYLOAD_FIELD,
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE,
)
if visited is None:
visited = set()
if max_string_length_prompt_in_db is None:
max_string_length_prompt_in_db = _get_max_string_length_prompt_in_db()
# Get the object's memory address to track visited objects
obj_id = id(request_body)
if obj_id in visited:
return {}
visited.add(obj_id)
_max_len = max_string_length_prompt_in_db
_start_chars = int(_max_len * 0.35)
_end_chars = min(int(_max_len * 0.65), _max_len - _start_chars)
def _sanitize_value(value: Any) -> Any:
if isinstance(value, dict):
if isinstance(value, str):
if len(value) > _max_len:
skipped_chars = len(value) - _start_chars - _end_chars
return (
f"{value[:_start_chars]}"
f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. "
f"{LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
f"{value[-_end_chars:]}"
)
return value
elif isinstance(value, dict):
return _sanitize_request_body_for_spend_logs_payload(
value, visited, max_string_length_prompt_in_db
value, visited, _max_len
)
elif isinstance(value, list):
return [_sanitize_value(item) for item in value]
elif isinstance(value, str):
if len(value) > max_string_length_prompt_in_db:
# Keep 35% from beginning and 65% from end (end is usually more important)
# This split ensures we keep more context from the end of conversations
start_ratio = 0.35
end_ratio = 0.65
# Calculate character distribution
start_chars = int(max_string_length_prompt_in_db * start_ratio)
end_chars = int(max_string_length_prompt_in_db * end_ratio)
# Ensure we don't exceed the total limit
total_keep = start_chars + end_chars
if total_keep > max_string_length_prompt_in_db:
end_chars = max_string_length_prompt_in_db - start_chars
# If the string length is less than what we want to keep, just truncate normally
if len(value) <= max_string_length_prompt_in_db:
return value
# Calculate how many characters are being skipped
skipped_chars = len(value) - total_keep
# Build the truncated string: beginning + truncation marker + end
truncated_value = (
f"{value[:start_chars]}"
f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. "
f"{LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
f"{value[-end_chars:]}"
)
return truncated_value
return value
return value
return {k: _sanitize_value(v) for k, v in request_body.items()}

View file

@ -1319,19 +1319,24 @@ class Router:
Only returns 2 characters of the api key and masks the rest with * (10 *).
"""
try:
_deployment_copy = copy.deepcopy(deployment)
litellm_params: dict = _deployment_copy["litellm_params"]
litellm_params: dict = deployment.get("litellm_params", {})
if litellm.redact_user_api_key_info:
masker = SensitiveDataMasker(visible_prefix=2, visible_suffix=0)
_deployment_copy["litellm_params"] = masker.mask_dict(litellm_params)
elif "api_key" in litellm_params:
litellm_params["api_key"] = litellm_params["api_key"][:2] + "*" * 10
return _deployment_copy
masked_params = masker.mask_dict(dict(litellm_params))
else:
masked_params = dict(litellm_params)
if "api_key" in masked_params:
api_key = masked_params["api_key"]
masked_params["api_key"] = (
api_key[:2] + "*" * 10 if api_key else api_key
)
return {
"model_name": deployment.get("model_name"),
"litellm_params": masked_params,
}
except Exception as e:
verbose_router_logger.debug(
f"Error occurred while printing deployment - {str(e)}"
"Error occurred while printing deployment - %s", str(e)
)
raise e
@ -8925,9 +8930,13 @@ class Router:
parent_otel_span=parent_otel_span,
)
raise exception
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}"
)
if verbose_router_logger.isEnabledFor(logging.INFO):
verbose_router_logger.info(
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
model,
self.print_deployment(deployment),
model,
)
end_time = time.time()
_duration = end_time - start_time
@ -9077,9 +9086,12 @@ class Router:
)
raise exception
verbose_router_logger.info(
f"async_get_available_deployment_for_pass_through model: {model}, selected deployment: {self.print_deployment(deployment)}"
)
if verbose_router_logger.isEnabledFor(logging.INFO):
verbose_router_logger.info(
"async_get_available_deployment_for_pass_through model: %s, selected deployment: %s",
model,
self.print_deployment(deployment),
)
end_time = time.perf_counter()
_duration = end_time - start_time
@ -9268,9 +9280,13 @@ class Router:
enable_pre_call_checks=self.enable_pre_call_checks,
cooldown_list=_cooldown_list,
)
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment)} for model: {model}"
)
if verbose_router_logger.isEnabledFor(logging.INFO):
verbose_router_logger.info(
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
model,
self.print_deployment(deployment),
model,
)
return deployment
def get_available_deployment_for_pass_through(
@ -9431,9 +9447,12 @@ class Router:
cooldown_list=_cooldown_list,
)
verbose_router_logger.info(
f"get_available_deployment_for_pass_through model: {model}, selected deployment: {self.print_deployment(deployment)}"
)
if verbose_router_logger.isEnabledFor(logging.INFO):
verbose_router_logger.info(
"get_available_deployment_for_pass_through model: %s, selected deployment: %s",
model,
self.print_deployment(deployment),
)
return deployment
def _filter_cooldown_deployments(

View file

@ -5,6 +5,7 @@ If weights are provided, it will return a deployment based on the weights.
"""
import logging
import random
from typing import TYPE_CHECKING, Any, Dict, List, Union
@ -52,9 +53,13 @@ def simple_shuffle(
selected_index = random.choices(range(len(weights)), weights=weights)[0]
verbose_router_logger.debug(f"\n selected index, {selected_index}")
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {llm_router_instance.print_deployment(deployment) or deployment[0]} for model: {model}"
)
if verbose_router_logger.isEnabledFor(logging.INFO):
verbose_router_logger.info(
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
model,
llm_router_instance.print_deployment(deployment) or deployment[0],
model,
)
return deployment or deployment[0]

View file

@ -3,7 +3,7 @@ from dataclasses import dataclass
from enum import Enum
from typing import Any, Dict, List, Literal, Optional, Tuple
from pydantic import BaseModel, Field, field_validator
from pydantic import BaseModel, Field, PrivateAttr, field_validator
from typing_extensions import Annotated
import litellm
@ -722,6 +722,8 @@ class UserAPIKeyLabelValues(BaseModel):
Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value)
] = None
_cached_dump: Optional[Dict[str, Any]] = PrivateAttr(default=None)
@field_validator("stream", mode="before")
@classmethod
def coerce_stream_to_str(cls, v: Any) -> Optional[str]:
@ -729,6 +731,12 @@ class UserAPIKeyLabelValues(BaseModel):
return None
return str(v)
def get_label_dict(self) -> Dict[str, Any]:
"""Return cached model_dump() dict to avoid re-serializing on every prometheus_label_factory call."""
if self._cached_dump is None:
self._cached_dump = self.model_dump()
return self._cached_dump
class PrometheusMetricsConfig(BaseModel):
"""Configuration for filtering Prometheus metrics"""