feat(prometheus): bucket latency by input sequence length

This commit is contained in:
Yucheng He 2026-09-06 02:25:43 -07:00
parent 02522a5441
commit 793306e7a8
9 changed files with 343 additions and 28 deletions

View file

@ -471,6 +471,7 @@ prometheus_metrics_config: Optional[List] = None
prometheus_exclude_metrics: Optional[List[str]] = None
prometheus_exclude_labels: Optional[List[str]] = None
prometheus_emit_stream_label: bool = False
prometheus_emit_input_sequence_length_label: bool = False
prometheus_deployment_and_latency_caller_identity: Literal[
"api_key_alias",
"user_email",

View file

@ -1754,6 +1754,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
"anthropic_prompt_caching_ttl",
"max_ui_session_budget",
"budget_rollover",
"prometheus_emit_input_sequence_length_label",
"mcp_tool_search",
]
SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]

View file

@ -1443,6 +1443,7 @@ class PrometheusLogger(CustomLogger):
user_agent=standard_logging_payload["metadata"].get("user_agent"),
stream=(str(standard_logging_payload.get("stream")) if litellm.prometheus_emit_stream_label else None),
service_tier=get_service_tier_from_standard_logging_payload(standard_logging_payload),
input_sequence_length=get_input_sequence_length_bucket(standard_logging_payload.get("prompt_tokens")),
)
if user_api_key is not None and isinstance(user_api_key, str) and user_api_key.startswith("sk-"):

View file

@ -1189,10 +1189,16 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
general_settings=general_settings
)
ProxyStartupEvent._initialize_startup_logging(
store_model_in_db = ( # rebind-ok: startup publishes the combined env, YAML and DB setting before callback construction
await ProxyStartupEvent.resolve_store_model_in_db(prisma_client=prisma_client, configured=store_model_in_db)
)
await ProxyStartupEvent._initialize_startup_logging(
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
redis_usage_cache=transaction_buffer_redis_cache,
prisma_client=prisma_client,
should_load_db_litellm_settings=store_model_in_db,
proxy_config_obj=proxy_config,
)
## V2 OTEL: publish the chosen V2 logger's TracerProvider as the OTel global.
@ -7244,9 +7250,9 @@ class ProxyConfig:
await self._init_hashicorp_vault_config_override(prisma_client=prisma_client)
await self._init_cyberark_config_override(prisma_client=prisma_client)
await self._apply_safe_litellm_settings_overrides_from_db(prisma_client=prisma_client)
await self.apply_safe_litellm_settings_overrides_from_db(prisma_client=prisma_client)
async def _apply_safe_litellm_settings_overrides_from_db(self, prisma_client: PrismaClient) -> None:
async def apply_safe_litellm_settings_overrides_from_db(self, prisma_client: PrismaClient) -> None:
config_record: Final = await get_config_param(prisma_client, "litellm_settings")
if config_record is None or config_record.param_value is None:
return
@ -7255,8 +7261,15 @@ class ProxyConfig:
if not isinstance(litellm_settings, dict):
return
for key, value in litellm_settings.items():
if key in LITELLM_SETTINGS_SAFE_DB_OVERRIDES:
setattr(litellm, key, value)
if key not in LITELLM_SETTINGS_SAFE_DB_OVERRIDES:
continue
if key == "prometheus_emit_input_sequence_length_label":
if isinstance(value, bool):
setattr(litellm, key, value)
elif isinstance(value, str) and (normalized_value := str_to_bool(value)) is not None:
setattr(litellm, key, normalized_value)
continue
setattr(litellm, key, value)
async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient):
"""
@ -9007,14 +9020,44 @@ class ProxyStartupEvent:
max_budget,
)
@staticmethod
async def resolve_store_model_in_db(prisma_client: PrismaClient | None, configured: bool) -> bool:
if get_secret_bool("STORE_MODEL_IN_DB", configured) or configured:
return True
if prisma_client is None:
return False
try:
db_general_settings: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first(
where={"param_name": "general_settings"}
)
except Exception as e: # noqa: BLE001 # a config-row read failure must not block proxy startup
verbose_proxy_logger.debug("Failed to check DB for store_model_in_db: %s", str(e))
return False
if db_general_settings is None or db_general_settings.param_value is None:
return False
db_value: Final = db_general_settings.param_value.get("store_model_in_db")
if db_value is True or (isinstance(db_value, str) and db_value.lower() == "true"):
verbose_proxy_logger.info("store_model_in_db=True loaded from DB, overriding config/env")
return True
return False
@classmethod
def _initialize_startup_logging(
async def _initialize_startup_logging(
cls,
llm_router: Router | None,
proxy_logging_obj: ProxyLogging,
redis_usage_cache: RedisCache | None,
):
prisma_client: PrismaClient | None = None,
should_load_db_litellm_settings: bool = False,
proxy_config_obj: ProxyConfig | None = None,
) -> None:
"""Initialize logging and alerting on startup"""
if should_load_db_litellm_settings and prisma_client is not None and proxy_config_obj is not None:
try:
await proxy_config_obj.apply_safe_litellm_settings_overrides_from_db(prisma_client=prisma_client)
except Exception as e: # noqa: BLE001 # a config-row read failure must not block proxy startup
verbose_proxy_logger.warning("Could not read litellm_settings from the database: %s", e)
## COST TRACKING ##
cost_tracking()
@ -9545,23 +9588,7 @@ class ProxyStartupEvent:
prisma_client.spend_logs_queue_monitor_task = monitor_task # rebind-ok: the client owns its monitor handle
### ADD NEW MODELS ###
store_model_in_db = get_secret_bool("STORE_MODEL_IN_DB", store_model_in_db) or store_model_in_db
# If store_model_in_db is still False, check DB for override.
# This breaks the chicken-and-egg where DB has store_model_in_db=True
# but YAML config has False.
if store_model_in_db is not True and prisma_client is not None:
try:
_db_gs_record: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first(
where={"param_name": "general_settings"}
)
if _db_gs_record is not None and isinstance(_db_gs_record.param_value, dict):
_db_val: Final = _db_gs_record.param_value.get("store_model_in_db")
if _db_val is True or (isinstance(_db_val, str) and _db_val.lower() == "true"):
store_model_in_db = True
verbose_proxy_logger.info("store_model_in_db=True loaded from DB, overriding config/env")
except Exception as e:
verbose_proxy_logger.debug("Failed to check DB for store_model_in_db: %s", str(e))
store_model_in_db = await cls.resolve_store_model_in_db(prisma_client=prisma_client, configured=store_model_in_db)
config_reload_interval_seconds = proxy_config_reload_interval_seconds
if not isinstance(config_reload_interval_seconds, int) or config_reload_interval_seconds <= 0:
@ -17120,6 +17147,13 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: Final[dict[str, GeneralSettingsUILiteLLMFie
"forgiving it. Applies to key, user, team, team member, org, tag and end-user budgets."
),
},
"prometheus_emit_input_sequence_length_label": { # mutable-ok: nested registry literal; LIT002 exempts only TypedDict-annotated top-level literals
"type": "Boolean",
"description": (
"Break latency and time-to-first-token metrics into input token length buckets. "
"Takes effect on the next proxy restart."
),
},
"max_ui_session_budget": {
"type": "Dollar",
"default": 1.0,

View file

@ -154,6 +154,20 @@ LATENCY_BUCKETS: Final = (
float("inf"),
)
INPUT_SEQUENCE_LENGTH_BUCKETS: Final = (
(1_000, "0-1k"),
(4_000, "1k-4k"),
(16_000, "4k-16k"),
(64_000, "16k-64k"),
(float("inf"), "64k+"),
)
def get_input_sequence_length_bucket(prompt_tokens: int | None) -> str:
token_count: Final = prompt_tokens if isinstance(prompt_tokens, int) and prompt_tokens >= 0 else 0
return next(label for upper, label in INPUT_SEQUENCE_LENGTH_BUCKETS if token_count < upper)
# Batch jobs can run for minutes to hours; buckets span 1 min → 24 h.
BATCH_DURATION_BUCKETS: Final = (
60.0,
@ -205,6 +219,7 @@ class UserAPIKeyLabelNames(Enum):
MCP_TOOL_NAME = "mcp_tool_name"
MCP_SERVER_NAME = "mcp_server_name"
SERVICE_TIER = "service_tier"
INPUT_SEQUENCE_LENGTH = "input_sequence_length"
DEFINED_PROMETHEUS_METRICS = Literal[
@ -857,6 +872,13 @@ class PrometheusMetricLabels:
"litellm_images_generated_metric",
}
)
_input_sequence_length_metrics: ClassVar[frozenset[str]] = frozenset(
{
"litellm_llm_api_latency_metric",
"litellm_llm_api_time_to_first_token_metric",
"litellm_request_total_latency_metric",
}
)
# Managed batch metrics
_batch_user_labels = [
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
@ -962,6 +984,13 @@ class PrometheusMetricLabels:
if label not in default_labels and label not in custom_labels:
custom_labels.append(label)
if (
label_name in PrometheusMetricLabels._input_sequence_length_metrics
and litellm.prometheus_emit_input_sequence_length_label is True
and UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value not in custom_labels
):
custom_labels.append(UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value)
return default_labels + custom_labels
@ -1015,6 +1044,7 @@ class UserAPIKeyLabelValues:
mcp_tool_name: str | None = None
mcp_server_name: str | None = None
service_tier: str | None = None
input_sequence_length: str | None = None
# Added for test compatibility.
def __init__(self, **kwargs: Any) -> None:

View file

@ -0,0 +1,171 @@
import datetime
from collections.abc import Mapping
from typing import Final, cast
import pytest
from prometheus_client import REGISTRY
from prometheus_client.samples import Sample
import litellm
from litellm.integrations.prometheus import PrometheusLogger
from litellm.types.integrations.prometheus import (
PrometheusMetricLabels,
UserAPIKeyLabelNames,
UserAPIKeyLabelValues,
get_input_sequence_length_bucket,
)
from litellm.types.utils import StandardLoggingPayload
LATENCY_METRICS: Final = (
"litellm_llm_api_latency_metric",
"litellm_llm_api_time_to_first_token_metric",
"litellm_request_total_latency_metric",
)
FLAG: Final = "prometheus_emit_input_sequence_length_label"
def _clear_prometheus_registry() -> None:
for collector in list(REGISTRY._collector_to_names): # pyright: ignore[reportPrivateUsage]
REGISTRY.unregister(collector)
@pytest.fixture(autouse=True)
def isolated_registry(monkeypatch: pytest.MonkeyPatch):
_clear_prometheus_registry()
monkeypatch.setattr(litellm, FLAG, False)
yield
_clear_prometheus_registry()
@pytest.mark.parametrize("metric", LATENCY_METRICS)
def test_input_sequence_length_label_is_opt_in(monkeypatch: pytest.MonkeyPatch, metric: str):
assert UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value not in PrometheusMetricLabels.get_labels(metric)
monkeypatch.setattr(litellm, FLAG, True)
assert UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value in PrometheusMetricLabels.get_labels(metric)
def test_input_sequence_length_label_stays_off_non_latency_metrics(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, FLAG, True)
assert UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value not in PrometheusMetricLabels.get_labels(
"litellm_proxy_total_requests_metric"
)
@pytest.mark.parametrize(
"prompt_tokens, expected",
[
(None, "0-1k"),
(0, "0-1k"),
(999, "0-1k"),
(1_000, "1k-4k"),
(3_999, "1k-4k"),
(4_000, "4k-16k"),
(15_999, "4k-16k"),
(16_000, "16k-64k"),
(63_999, "16k-64k"),
(64_000, "64k+"),
(10_000_000, "64k+"),
(-1, "0-1k"),
],
)
def test_input_sequence_length_bucket_boundaries(prompt_tokens: int | None, expected: str):
assert get_input_sequence_length_bucket(prompt_tokens) == expected
def test_user_api_key_label_values_carries_input_sequence_length():
values: Final = UserAPIKeyLabelValues(input_sequence_length="4k-16k")
assert values.input_sequence_length == "4k-16k"
assert values.model_dump()["input_sequence_length"] == "4k-16k"
def _latency_bucket_samples() -> tuple[Sample, ...]:
return tuple(
sample
for metric in REGISTRY.collect()
for sample in metric.samples
if sample.name.endswith("_bucket") and any(name in sample.name for name in LATENCY_METRICS)
)
def _standard_logging_payload(now: datetime.datetime, prompt_tokens: int) -> StandardLoggingPayload:
return cast(
StandardLoggingPayload,
{
"id": "t",
"call_type": "completion",
"response_cost": 0.001,
"status": "success",
"total_tokens": prompt_tokens + 20,
"prompt_tokens": prompt_tokens,
"completion_tokens": 20,
"startTime": now - datetime.timedelta(seconds=3),
"endTime": now,
"completionStartTime": now - datetime.timedelta(seconds=1),
"model": "gpt-4o-mini",
"model_id": "model-123",
"model_group": "gpt-4o-mini",
"api_base": "https://api.openai.com",
"custom_llm_provider": "openai",
"request_tags": [],
"stream": True,
"metadata": {
"user_api_key_hash": "h",
"user_api_key_alias": "a",
"user_api_key_team_id": "t",
"user_api_key_team_alias": "ta",
"user_api_key_user_id": "u",
"user_api_key_user_email": "e@x.com",
"user_api_key_org_id": None,
"user_api_key_org_alias": None,
"requester_metadata": None,
"user_api_key_end_user_id": None,
"usage_object": None,
},
"hidden_params": {"litellm_overhead_time_ms": None, "additional_headers": None},
},
)
def _success_kwargs(now: datetime.datetime, prompt_tokens: int) -> Mapping[str, object]:
return {
"model": "gpt-4o-mini",
"litellm_params": {"metadata": {}},
"standard_logging_object": _standard_logging_payload(now, prompt_tokens),
"stream": True,
"start_time": now - datetime.timedelta(seconds=3),
"api_call_start_time": now - datetime.timedelta(seconds=2),
"completion_start_time": now - datetime.timedelta(seconds=1),
"end_time": now,
}
@pytest.mark.asyncio
@pytest.mark.parametrize("flag_at_request_time", (True, False))
async def test_logger_emits_bucket_from_its_startup_label_set(
monkeypatch: pytest.MonkeyPatch, flag_at_request_time: bool
):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, True)
logger: Final = PrometheusLogger()
monkeypatch.setattr(litellm, FLAG, flag_at_request_time)
await logger.async_log_success_event(dict(_success_kwargs(now, prompt_tokens=4_000)), None, now, now)
samples: Final = _latency_bucket_samples()
assert samples
assert all(sample.labels["input_sequence_length"] == "4k-16k" for sample in samples)
@pytest.mark.asyncio
async def test_logger_built_with_flag_off_emits_no_bucket_label(monkeypatch: pytest.MonkeyPatch):
now: Final = datetime.datetime.now()
logger: Final = PrometheusLogger()
monkeypatch.setattr(litellm, FLAG, True)
await logger.async_log_success_event(dict(_success_kwargs(now, prompt_tokens=4_000)), None, now, now)
samples: Final = _latency_bucket_samples()
assert samples
assert all("input_sequence_length" not in sample.labels for sample in samples)

View file

@ -17,17 +17,15 @@ Pins covered:
from __future__ import annotations
import asyncio
import inspect
import json
import logging
import os
from collections.abc import Awaitable, Callable
from typing import List, Optional, Union
from unittest.mock import AsyncMock, MagicMock, patch
from typing import Final, Optional, Union
from unittest.mock import AsyncMock, MagicMock, call
import pytest
from fastapi import FastAPI
from pydantic import BaseModel
from typing_extensions import TypedDict
@ -749,6 +747,63 @@ async def test_proxy_startup_event_invalid_missing_app_arg_raises():
pass
@pytest.mark.asyncio
@pytest.mark.parametrize(
("configured", "db_value", "expected"),
(
(True, None, True),
(False, True, True),
(False, "true", True),
(False, False, False),
(False, None, False),
),
)
async def test_resolve_store_model_in_db_uses_config_or_db(
monkeypatch: pytest.MonkeyPatch, configured: bool, db_value: object, expected: bool
):
monkeypatch.setattr(ps, "get_secret_bool", lambda name, default: default)
db_record: Final = None if db_value is None else MagicMock(param_value={"store_model_in_db": db_value})
prisma_client: Final = MagicMock()
prisma_client.db.litellm_config.find_first = AsyncMock(return_value=db_record)
result: Final = await ProxyStartupEvent.resolve_store_model_in_db(
prisma_client=prisma_client, configured=configured
)
assert result is expected
assert prisma_client.db.litellm_config.find_first.await_count == (0 if configured else 1)
@pytest.mark.asyncio
async def test_startup_logging_applies_db_settings_before_callback_init(monkeypatch: pytest.MonkeyPatch):
events: Final = MagicMock()
proxy_config: Final = MagicMock()
async def apply_db_settings(prisma_client: object) -> None:
events.db_settings()
proxy_config.apply_safe_litellm_settings_overrides_from_db = apply_db_settings
proxy_logging: Final = MagicMock()
proxy_logging.startup_event.side_effect = lambda **kwargs: events.callback_init()
prisma_client: Final = MagicMock()
prisma_client.db.litellm_config.find_first = AsyncMock(
return_value=MagicMock(param_value={"store_model_in_db": True})
)
monkeypatch.setattr(ps, "cost_tracking", MagicMock())
monkeypatch.setattr(ps, "get_secret_bool", lambda name, default: default)
await ProxyStartupEvent._initialize_startup_logging(
llm_router=None,
proxy_logging_obj=proxy_logging,
redis_usage_cache=None,
prisma_client=prisma_client,
should_load_db_litellm_settings=True,
proxy_config_obj=proxy_config,
)
assert events.method_calls == [call.db_settings(), call.callback_init()]
def test_otel_global_provider_published_after_callback_init():
"""The OTel V2 global-provider publish must run after callback
initialization in ``proxy_startup_event``.

View file

@ -2912,6 +2912,27 @@ async def test_ProxyConfig_add_deployment_applies_db_router_settings(monkeypatch
fake_router.update_settings.assert_called_once_with(routing_strategy="latency-based-routing")
@pytest.mark.asyncio
@pytest.mark.parametrize(
("db_value", "expected"),
((True, True), (False, False), ("true", True), ("false", False), ("invalid", False)),
)
async def test_apply_safe_litellm_settings_overrides_normalizes_input_sequence_length_flag(
monkeypatch: pytest.MonkeyPatch, db_value: object, expected: bool
):
from litellm.proxy import proxy_server
config_record = SimpleNamespace(
param_value={"prometheus_emit_input_sequence_length_label": db_value}
)
monkeypatch.setattr(proxy_server, "get_config_param", AsyncMock(return_value=config_record))
monkeypatch.setattr(litellm, "prometheus_emit_input_sequence_length_label", False)
await ProxyConfig().apply_safe_litellm_settings_overrides_from_db(prisma_client=MagicMock())
assert litellm.prometheus_emit_input_sequence_length_label is expected
# ---------------------------------------------------------------------------
# ProxyConfig._add_general_settings_from_db_config
# ---------------------------------------------------------------------------

View file

@ -10168,6 +10168,7 @@ def test_get_config_list_includes_anthropic_prompt_caching_fields(monkeypatch):
assert fields["enable_anthropic_prompt_caching"]["field_tab"] == "prompt_caching"
assert fields["anthropic_prompt_caching_ttl"]["field_tab"] == "prompt_caching"
assert fields["budget_exceeded_throttle_percentage"]["field_tab"] is None
assert fields["prometheus_emit_input_sequence_length_label"]["field_type"] == "Boolean"
finally:
app.dependency_overrides.clear()