From 793306e7a87bb2b9d3da2ea5c77738f8f00cd4b2 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Sun, 6 Sep 2026 02:25:43 -0700 Subject: [PATCH] feat(prometheus): bucket latency by input sequence length --- litellm/__init__.py | 1 + litellm/constants.py | 1 + litellm/integrations/prometheus.py | 1 + litellm/proxy/proxy_server.py | 82 ++++++--- litellm/types/integrations/prometheus.py | 30 +++ ..._prometheus_input_sequence_length_label.py | 171 ++++++++++++++++++ .../proxy/proxy_server/test_lifecycle.py | 63 ++++++- .../proxy/proxy_server/test_proxy_config.py | 21 +++ tests/test_litellm/proxy/test_proxy_server.py | 1 + 9 files changed, 343 insertions(+), 28 deletions(-) create mode 100644 tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py diff --git a/litellm/__init__.py b/litellm/__init__.py index 42c0ea881fd..4ae63716994 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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", diff --git a/litellm/constants.py b/litellm/constants.py index ce744e9c58a..f554f7f07eb 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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"] diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 6766d246894..d7ed16bba9a 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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-"): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ba5714fe950..322b96c1809 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 8498b6f6d00..7350a945964 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -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: diff --git a/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py b/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py new file mode 100644 index 00000000000..10a6b3dd821 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_input_sequence_length_label.py @@ -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) diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index a518c9289ae..7e8ab95a419 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -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``. diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 2babfe432f3..56b3f9cebbf 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 9393ec0f8e6..7773c857d6b 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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()