fix(prometheus): handle unknown input lengths

This commit is contained in:
Yucheng He 2026-09-06 08:02:05 -07:00
parent ea3e2f18d6
commit ba9b29280f
4 changed files with 30 additions and 16 deletions

View file

@ -9033,7 +9033,7 @@ class ProxyStartupEvent:
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:
if db_general_settings is None or not isinstance(db_general_settings.param_value, dict):
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"):

View file

@ -154,6 +154,7 @@ LATENCY_BUCKETS: Final = (
float("inf"),
)
UNKNOWN_INPUT_SEQUENCE_LENGTH: Final = "unknown"
INPUT_SEQUENCE_LENGTH_BUCKETS: Final = (
(1_000, "0-1k"),
(4_000, "1k-4k"),
@ -163,9 +164,10 @@ INPUT_SEQUENCE_LENGTH_BUCKETS: Final = (
)
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)
def get_input_sequence_length_bucket(prompt_tokens: object) -> str:
if not isinstance(prompt_tokens, int) or isinstance(prompt_tokens, bool) or prompt_tokens < 0:
return UNKNOWN_INPUT_SEQUENCE_LENGTH
return next(label for upper, label in INPUT_SEQUENCE_LENGTH_BUCKETS if prompt_tokens < upper)
# Batch jobs can run for minutes to hours; buckets span 1 min → 24 h.
@ -977,21 +979,23 @@ class PrometheusMetricLabels:
custom_labels.append(label)
if label_name in PrometheusMetricLabels._org_label_metrics:
for label in [
for label in (
UserAPIKeyLabelNames.ORG_ID.value,
UserAPIKeyLabelNames.ORG_ALIAS.value,
]:
):
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
input_sequence_length_labels: Final = (
(UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value,)
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
)
else ()
)
return [*default_labels, *custom_labels, *input_sequence_length_labels]
_USER_API_KEY_LABEL_VALUE_INIT_ALIASES: Final[Mapping[str, str]] = MappingProxyType(

View file

@ -55,7 +55,7 @@ def test_input_sequence_length_label_stays_off_non_latency_metrics(monkeypatch:
@pytest.mark.parametrize(
"prompt_tokens, expected",
[
(None, "0-1k"),
(None, "unknown"),
(0, "0-1k"),
(999, "0-1k"),
(1_000, "1k-4k"),
@ -66,7 +66,7 @@ def test_input_sequence_length_label_stays_off_non_latency_metrics(monkeypatch:
(63_999, "16k-64k"),
(64_000, "64k+"),
(10_000_000, "64k+"),
(-1, "0-1k"),
(-1, "unknown"),
],
)
def test_input_sequence_length_bucket_boundaries(prompt_tokens: int | None, expected: str):

View file

@ -774,6 +774,16 @@ async def test_resolve_store_model_in_db_uses_config_or_db(
assert prisma_client.db.litellm_config.find_first.await_count == (0 if configured else 1)
@pytest.mark.asyncio
@pytest.mark.parametrize("param_value", ("legacy", ["store_model_in_db"], 7))
async def test_resolve_store_model_in_db_ignores_non_mapping_row(monkeypatch: pytest.MonkeyPatch, param_value: object):
monkeypatch.setattr(ps, "get_secret_bool", lambda name, default: default)
prisma_client: Final = MagicMock()
prisma_client.db.litellm_config.find_first = AsyncMock(return_value=MagicMock(param_value=param_value))
assert await ProxyStartupEvent.resolve_store_model_in_db(prisma_client=prisma_client, configured=False) is False
@pytest.mark.asyncio
async def test_startup_logging_applies_db_settings_before_callback_init(monkeypatch: pytest.MonkeyPatch):
events: Final = MagicMock()