feat(prometheus): cap series per metric for every labeled metric (#44420)

* feat(prometheus): cap series per metric for every labeled metric

Add prometheus_metrics_max_series_per_metric: per metric and per worker process, the first N label
sets keep a series of their own. Counters and histograms record every later label set on one series
whose labels are all "other", so totals stay exact, and gauges skip it. The cap holds with multiple
workers because it never needs to remove a series.

Add prometheus_metrics_ttl_seconds: a series idle for that long is removed and its slot is freed.
The prometheus client cannot remove a series in multi-process mode, so the TTL is ignored there with
a startup warning.

Both settings are off by default. The end_user caps are unchanged.

* fix(prometheus): share the series cap across workers of one proxy instance

Workers writing to one PROMETHEUS_MULTIPROC_DIR now agree on which label
sets get a series through an append-only admissions file per metric, so a
merged scrape stays at the cap plus `other` instead of growing with every
worker and every worker restart. The two fallback counters now pass their
label names as a keyword so the cap and prometheus_exclude_labels apply to
them, admission and child creation happen under one lock, the test fixture
restores the shared registry, and the `other` label value lives in
constants.py.

* test(prometheus): check emitted labels instead of wrapper types, close the admission match

The exclude-labels test now emits through the spend and provider budget
metrics and checks the scrape keeps all their labels. The admission match
arms end in assert_never so the match is exhaustive.

* fix(prometheus): return the exhaustive-match fallback so every admission arm returns

* fix(prometheus): pick the series tracker with isinstance so every path of _admits returns

* fix(prometheus): skip an admissions line a worker could only write part of

* fix(prometheus): frame each admissions record with newlines so a cut-off record cannot swallow the next

A record a worker could only write part of used to merge with the next worker's record, and both were skipped for one request. Each record is now written between two newlines, so the fragment is a line of its own. The clock fixture in the series tests starts from a constant instead of reading the real clock

* fix(prometheus): ignore a non-positive series cap or TTL with a warning instead of failing the logger

A cap or TTL of 0 or less raised at logger init. The proxy logs that as a non-blocking error and keeps serving, so the result was a running proxy with no Prometheus metrics at all. The setting is now ignored with a startup warning naming it, the same rule the end_user cap already follows for a non-positive value

* fix(prometheus): start the series cap over on a one-worker restart and audit it live

A proxy with one worker and an operator-set PROMETHEUS_MULTIPROC_DIR now drops litellm's admission files at boot, so a restart frees every slot there the way it already does with several workers. A cap or TTL that is not a number greater than 0 (a bool, a non-numeric string, an empty value) is ignored with the startup warning instead of breaking the logger

The integration cells drive the cap on every endpoint through the OpenAI and Anthropic SDKs and raw httpx, streaming and not, plus gauges, cache hits, failures, both workers of one instance, the TTL on one worker and its warning on two, ignored settings, excluded labels on the fallback counters, a null cap, /config/update, a concurrent burst scraped mid-flight, a provider outage, a killed worker, and restarts with one and two workers

* fix(prometheus): wipe an operator-set multiprocess directory on a one-worker boot too

* fix(prometheus): leave the multiprocess directory alone on a setup-only run

A run with --skip_server_startup starts no worker, so it no longer creates or
wipes PROMETHEUS_MULTIPROC_DIR. Wiping there deleted the samples of a proxy
already running against the same directory

* fix(prometheus): free the capped series slots when a gateway or backend container restarts

The component image entrypoint starts uvicorn without the proxy CLI and wiped only the .db sample files at container start, so the admitted-series files of the previous container survived an in-place restart. Every label set seen after the restart was then counted on `other` once the previous container had filled the cap

* test(prometheus): cover a setup-only run and a gateway image restart under the cap

Two integration cells from the audit: a `--skip_server_startup` run pointed at a live
two-worker proxy's operator directory leaves its samples alone, and the gateway image
(`docker/component_entrypoint.sh` running `python -m gateway.launch`) restarted on a kept
PROMETHEUS_MULTIPROC_DIR starts the cap over. The burst cells now wait for every counter
they assert on, since the request and failure counters of one call increment at different
points of the logging callback

* test(prometheus): prove the cap reaches the fallback counters in the X1 cell

* fix(prometheus): ignore a cleanup interval that is not a number of at least 0

A string or negative prometheus_metrics_cleanup_interval_seconds reached the
series tracker unvalidated, so the first labeled emit with a TTL on raised
TypeError inside the callback and recorded no series. The interval is now
validated the way the cap and the TTL are: an invalid value is ignored with a
warning and the default 60 seconds applies. The I2 integration cell drives a
string interval through a live proxy and reads the warning from its log

* test(prometheus): give the restart cells the boot budget of their siblings

C4 and C5 boot two proxies each and hit the file's 240 s budget on a loaded
box; C3 and D1 already carry 420 s

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-06 01:28:44 +00:00 • committed by GitHub
parent b58e2d7175
commit 6f5ca84f69
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
18 changed files with 2260 additions and 84 deletions

View file

@ -3,7 +3,7 @@
# stale samples from a previous container incarnation would be summed into the aggregate
if [ -n "$PROMETHEUS_MULTIPROC_DIR" ]; then
mkdir -p "$PROMETHEUS_MULTIPROC_DIR"
rm -f "$PROMETHEUS_MULTIPROC_DIR"/*.db
rm -f "$PROMETHEUS_MULTIPROC_DIR"/*.db "$PROMETHEUS_MULTIPROC_DIR"/litellm_admitted_series_*
fi
case "$USE_DDTRACE" in

View file

@ -501,6 +501,9 @@ prometheus_user_budget_label_include_email_alias: bool = False
prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000
prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0
prometheus_end_user_metrics_cleanup_interval_seconds: Optional[float] = 60.0
prometheus_metrics_max_series_per_metric: Optional[int] = None
prometheus_metrics_ttl_seconds: Optional[float] = None
prometheus_metrics_cleanup_interval_seconds: Optional[float] = 60.0
disable_add_prefix_to_prompt: bool = False # used by anthropic, to disable adding prefix to prompt
disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior.
public_mcp_servers: Optional[List[str]] = None

View file

@ -1549,6 +1549,8 @@ AZURE_STORAGE_DEFAULT_ENDPOINT_SUFFIX: Final = "core.windows.net"
PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES: Final = int(
os.getenv("PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES", 5)
)
PROMETHEUS_OVERFLOW_SERIES_LABEL_VALUE: Final = "other"
PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX: Final = "litellm_admitted_series_"
CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60))
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))

View file

@ -10,16 +10,21 @@ import sys
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import replace
from datetime import datetime, timedelta
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, TypeVar, cast
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, Protocol, TypeAlias, TypeVar, cast
from pydantic import BaseModel
from typing_extensions import ReadOnly, TypedDict
from pydantic import BaseModel, Field, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict, assert_never
import litellm
from litellm._internal_context import with_service_target
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import PROXY_LLM_PROVIDER_FALLBACK, PROXY_REJECTED_BEFORE_ROUTING_KEY
from litellm.constants import (
PROMETHEUS_OVERFLOW_SERIES_LABEL_VALUE,
PROXY_LLM_PROVIDER_FALLBACK,
PROXY_REJECTED_BEFORE_ROUTING_KEY,
)
from litellm.exceptions import (
validate_rate_limit_category,
validate_rate_limit_type,
@ -31,6 +36,10 @@ from litellm.integrations.prometheus_helpers import (
)
from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import (
BoundedPrometheusSeriesTracker,
PrometheusSeriesLimits,
)
from litellm.integrations.prometheus_helpers.shared_prometheus_series_admissions import (
SharedPrometheusSeriesAdmissions,
)
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
@ -164,30 +173,131 @@ def _customer_budget_metrics_enabled() -> bool:
return litellm.enable_end_user_cost_tracking_prometheus_only is True and not litellm.disable_end_user_cost_tracking
class _ExcludedLabelMetric:
"""Proxies a prometheus metric whose declared ``labelnames`` had globally
excluded labels removed, dropping those labels from every ``labels(...)``
call so the emitted arguments always match the metric's real label set."""
class _LabeledMetric:
"""Proxies a labeled prometheus metric. Globally excluded labels are dropped from every ``labels(...)``
call so the emitted arguments match the metric's real label set. With ``limits.max_series`` set, only that
many label sets get a series of their own: a counter or histogram records every later label set on one
series whose labels are all ``other``, so totals stay exact, and a gauge skips it, since one shared gauge
value would mean nothing. In multi-process mode the tracker is the one the workers share, and ``remove`` does
nothing there, since the prometheus client cannot remove a series."""
__slots__ = (
"_excluded_labels",
"_limits",
"_metric",
"_metric_name",
"_original_labelnames",
"_overflow_child",
"_tracker",
)
def __init__(
self,
metric: MetricWrapperBase,
metric_name: str,
original_labelnames: tuple[str, ...],
excluded_labels: frozenset[str],
tracker: BoundedPrometheusSeriesTracker | SharedPrometheusSeriesAdmissions,
limits: PrometheusSeriesLimits,
shares_overflow_series: bool,
) -> None:
kept_label_count: Final = len(tuple(name for name in original_labelnames if name not in excluded_labels))
self._metric = metric
self._metric_name = metric_name
self._original_labelnames = original_labelnames
self._excluded_labels = excluded_labels
def labels(self, *labelvalues: str, **labelkwargs: str) -> MetricWrapperBase:
values: Final = labelvalues or tuple(labelkwargs[name] for name in self._original_labelnames)
kept_values: Final = tuple(
value for name, value in zip(self._original_labelnames, values) if name not in self._excluded_labels
self._tracker = tracker
self._limits = limits
self._overflow_child: Callable[[], MetricWrapperBase | NoOpMetric] = (
partial(metric.labels, *(PROMETHEUS_OVERFLOW_SERIES_LABEL_VALUE,) * kept_label_count)
if shares_overflow_series
else NoOpMetric
)
def labels(self, *labelvalues: object, **labelkwargs: object) -> MetricWrapperBase | NoOpMetric:
values: Final = labelvalues or tuple(labelkwargs[name] for name in self._original_labelnames)
kept_values: Final = self._kept_values(values)
if not kept_values:
return self._metric
if not self._limits.enabled:
return self._metric.labels(*kept_values)
with self._tracker.lock:
if self._admits(kept_values):
return self._metric.labels(*kept_values)
return self._overflow_child()
def remove(self, *labelvalues: object) -> None:
match self._tracker:
case SharedPrometheusSeriesAdmissions():
pass
case BoundedPrometheusSeriesTracker():
kept_values: Final = self._kept_values(labelvalues)
with self._tracker.lock:
self._tracker.forget_series(self._metric_name, kept_values)
self._metric.remove(*kept_values)
case _:
assert_never(self._tracker)
def _admits(self, kept_values: tuple[str, ...]) -> bool:
if isinstance(self._tracker, SharedPrometheusSeriesAdmissions):
return self._limits.max_series is None or self._tracker.admit_series(
metric_name=self._metric_name, label_values=kept_values, max_series=self._limits.max_series
)
return self._tracker.admit_series(
metric=self._metric, metric_name=self._metric_name, label_values=kept_values, limits=self._limits
)
def _kept_values(self, values: tuple[object, ...]) -> tuple[str, ...]:
return tuple(
str(value) for name, value in zip(self._original_labelnames, values) if name not in self._excluded_labels
)
return self._metric.labels(*kept_values) if kept_values else self._metric
_MetricLike: TypeAlias = "NoOpMetric | _ExcludedLabelMetric | MetricWrapperBase"
_MetricLike: TypeAlias = "NoOpMetric | _LabeledMetric | MetricWrapperBase"
_SeriesLimitT: Final = TypeVar("_SeriesLimitT", int, float)
_POSITIVE_SERIES_CAP: Final[TypeAdapter[int]] = TypeAdapter(Annotated[int, Field(gt=0)])
_POSITIVE_SERIES_TTL: Final[TypeAdapter[float]] = TypeAdapter(Annotated[float, Field(gt=0)])
_SERIES_CLEANUP_INTERVAL: Final[TypeAdapter[float]] = TypeAdapter(Annotated[float, Field(ge=0)])
_DEFAULT_SERIES_CLEANUP_INTERVAL_SECONDS: Final = 60.0
def _number_or_none(value: object, limit: TypeAdapter[_SeriesLimitT]) -> _SeriesLimitT | None:
if isinstance(value, bool):
return None
try:
return limit.validate_python(value)
except ValidationError:
return None
def _positive_or_ignored(setting: str, value: object, limit: TypeAdapter[_SeriesLimitT]) -> _SeriesLimitT | None:
if value is None:
return None
validated: Final = _number_or_none(value, limit)
if validated is not None:
return validated
verbose_logger.warning(
"%s is ignored because it is not a number greater than 0 (got %r). Prometheus metrics are emitted without it",
setting,
value,
)
return None
def _cleanup_interval_or_default(value: object) -> float | None:
if value is None:
return None
validated: Final = _number_or_none(value, _SERIES_CLEANUP_INTERVAL)
if validated is not None:
return validated
verbose_logger.warning(
"prometheus_metrics_cleanup_interval_seconds is ignored because it is not a number of at least 0 (got %r). "
"Idle series are checked every %s seconds",
value,
_DEFAULT_SERIES_CLEANUP_INTERVAL_SECONDS,
)
return _DEFAULT_SERIES_CLEANUP_INTERVAL_SECONDS
def _get_budget_metrics_per_request_timeout() -> float:
@ -301,10 +411,17 @@ class PrometheusLogger(CustomLogger):
_custom_buckets: Final = litellm.prometheus_latency_buckets
self.latency_buckets = tuple(_custom_buckets) if _custom_buckets is not None else LATENCY_BUCKETS
self._bounded_prometheus_series_tracker = BoundedPrometheusSeriesTracker()
_multiproc_dir: Final = os.environ.get("PROMETHEUS_MULTIPROC_DIR")
self._series_cap_tracker = (
BoundedPrometheusSeriesTracker()
if _multiproc_dir is None
else SharedPrometheusSeriesAdmissions(directory=_multiproc_dir)
)
self._series_limits = self._configured_series_limits(multiprocess_mode=_multiproc_dir is not None)
# Create metric factory functions
self._counter_factory = self._create_metric_factory(Counter)
self._gauge_factory = self._create_metric_factory(Gauge)
self._gauge_factory = self._create_metric_factory(Gauge, shares_overflow_series=False)
self._histogram_factory = self._create_metric_factory(Histogram)
self.litellm_proxy_failed_requests_metric = self._counter_factory(
@ -694,13 +811,13 @@ class PrometheusLogger(CustomLogger):
self.litellm_deployment_successful_fallbacks = self._counter_factory(
"litellm_deployment_successful_fallbacks",
"LLM Deployment Analytics - Number of successful fallback requests from primary model -> fallback model",
self.get_labels_for_metric("litellm_deployment_successful_fallbacks"),
labelnames=self.get_labels_for_metric("litellm_deployment_successful_fallbacks"),
)
self.litellm_deployment_failed_fallbacks = self._counter_factory(
"litellm_deployment_failed_fallbacks",
"LLM Deployment Analytics - Number of failed fallback requests from primary model -> fallback model",
self.get_labels_for_metric("litellm_deployment_failed_fallbacks"),
labelnames=self.get_labels_for_metric("litellm_deployment_failed_fallbacks"),
)
# Callback Logging Failure Metrics
@ -1182,27 +1299,55 @@ class PrometheusLogger(CustomLogger):
return metric_name in self.enabled_metrics
def _create_metric_factory(self, metric_class):
def _create_metric_factory(self, metric_class, shares_overflow_series: bool = True):
"""Create a factory function that returns either a real metric or a no-op metric"""
def factory(*args, **kwargs):
# Extract metric name from the first argument or 'name' keyword argument
metric_name: Final = args[0] if args else kwargs.get("name", "")
metric_name: Final = str(args[0] if args else kwargs.get("name", ""))
if not self._is_metric_enabled(metric_name):
return NoOpMetric()
original_labelnames: Final = tuple(kwargs.get("labelnames") or ())
if not (frozenset(original_labelnames) & self.exclude_labels):
kept: Final = tuple(name for name in original_labelnames if name not in self.exclude_labels)
if not original_labelnames or (kept == original_labelnames and not self._series_limits.enabled):
return metric_class(*args, **kwargs)
kept: Final = tuple(name for name in original_labelnames if name not in self.exclude_labels)
kept_kwargs: Final = {**kwargs, "labelnames": kept}
real_metric: Final = metric_class(*args, **kept_kwargs)
return _ExcludedLabelMetric(real_metric, original_labelnames, self.exclude_labels)
return _LabeledMetric(
metric=metric_class(*args, **{**kwargs, "labelnames": kept}),
metric_name=metric_name,
original_labelnames=original_labelnames,
excluded_labels=self.exclude_labels,
tracker=self._series_cap_tracker,
limits=self._series_limits,
shares_overflow_series=shares_overflow_series,
)
return factory
@staticmethod
def _configured_series_limits(multiprocess_mode: bool) -> PrometheusSeriesLimits:
limits: Final = PrometheusSeriesLimits(
max_series=_positive_or_ignored(
"prometheus_metrics_max_series_per_metric",
litellm.prometheus_metrics_max_series_per_metric,
_POSITIVE_SERIES_CAP,
),
ttl_seconds=_positive_or_ignored(
"prometheus_metrics_ttl_seconds", litellm.prometheus_metrics_ttl_seconds, _POSITIVE_SERIES_TTL
),
cleanup_interval_seconds=_cleanup_interval_or_default(litellm.prometheus_metrics_cleanup_interval_seconds),
)
if limits.ttl_seconds is None or not multiprocess_mode:
return limits
verbose_logger.warning(
"prometheus_metrics_ttl_seconds is ignored while PROMETHEUS_MULTIPROC_DIR is set: the prometheus "
"client cannot remove a series in multi-process mode. prometheus_metrics_max_series_per_metric "
"still applies"
)
return replace(limits, ttl_seconds=None)
def get_labels_for_metric(self, metric_name: DEFINED_PROMETHEUS_METRICS) -> list[str]:
"""
Get the labels for a metric, filtered if configured.

View file

@ -2,6 +2,7 @@ from __future__ import annotations
import time
from collections import OrderedDict
from dataclasses import dataclass
from threading import RLock
from typing import Final, Protocol
@ -12,6 +13,17 @@ class _RemovableMetric(Protocol):
def remove(self, *labelvalues: object) -> None: ...
@dataclass(frozen=True, slots=True)
class PrometheusSeriesLimits:
max_series: int | None
ttl_seconds: float | None
cleanup_interval_seconds: float | None
@property
def enabled(self) -> bool:
return self.max_series is not None or self.ttl_seconds is not None
class BoundedPrometheusSeriesTracker:
"""
Tracks Prometheus child series and removes stale/excess labelsets.
@ -49,13 +61,7 @@ class BoundedPrometheusSeriesTracker:
now=now,
cleanup_interval_seconds=cleanup_interval_seconds,
):
expired_label_values: Final = [
tracked_label_values
for tracked_label_values, last_seen in series.items()
if now - last_seen > ttl_seconds
]
for tracked_label_values in expired_label_values:
self._remove_metric_series(metric, series, tracked_label_values)
self._remove_expired_series(metric, series, now, ttl_seconds)
# max_series <= 0 is treated as "unlimited" so a misconfigured zero
# value cannot silently drop every emission for this metric.
@ -66,6 +72,34 @@ class BoundedPrometheusSeriesTracker:
break
del series[tracked_label_values]
def admit_series(
self,
metric: _RemovableMetric,
metric_name: str,
label_values: tuple[str | None, ...],
limits: PrometheusSeriesLimits,
) -> bool:
now: Final = time.monotonic()
with self.lock:
series: Final = self._series.setdefault(metric_name, OrderedDict())
if limits.ttl_seconds is not None and self._should_run_ttl_cleanup(
metric_name=metric_name,
now=now,
cleanup_interval_seconds=limits.cleanup_interval_seconds,
):
self._remove_expired_series(metric, series, now, limits.ttl_seconds)
if label_values not in series and limits.max_series is not None and len(series) >= limits.max_series:
return False
series[label_values] = now
series.move_to_end(label_values)
return True
def forget_series(self, metric_name: str, label_values: tuple[str | None, ...]) -> None:
with self.lock:
self._series.get(metric_name, OrderedDict()).pop(label_values, None)
def remove_series(self, metric: _RemovableMetric, label_values: tuple[str | None, ...]) -> bool:
"""Drop one child series, True when it is gone (removed or never existed)."""
return self._remove_metric_child(metric, label_values)
@ -86,6 +120,19 @@ class BoundedPrometheusSeriesTracker:
return True
return False
def _remove_expired_series(
self,
metric: _RemovableMetric,
series: OrderedDict[tuple[str | None, ...], float],
now: float,
ttl_seconds: float,
) -> None:
expired_label_values: Final = [
tracked_label_values for tracked_label_values, last_seen in series.items() if now - last_seen > ttl_seconds
]
for tracked_label_values in expired_label_values:
self._remove_metric_series(metric, series, tracked_label_values)
def _remove_metric_series(
self,
metric: _RemovableMetric,

View file

@ -0,0 +1,92 @@
from __future__ import annotations
import os
from threading import RLock
from typing import Final
from pydantic import TypeAdapter, ValidationError
from litellm.constants import PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX
_LABEL_VALUES: Final = TypeAdapter(tuple[str, ...])
def _parse_admission(line: bytes) -> tuple[str, ...] | None:
try:
return _LABEL_VALUES.validate_json(line)
except ValidationError:
return None
class _MetricAdmissions:
__slots__ = ("_label_sets", "_max_series", "_path", "_read_offset")
def __init__(self, path: str, max_series: int) -> None:
self._path = path
self._max_series = max_series
self._label_sets: set[tuple[str, ...]] = ( # mutable-ok: a frozenset copy per admission is quadratic in the cap
set()
)
self._read_offset = 0
def admit(self, label_values: tuple[str, ...]) -> bool:
if label_values in self._label_sets:
return True
if self._is_full():
return False
self._read_new_admissions()
if label_values not in self._label_sets and not self._is_full():
self._append(label_values)
self._read_new_admissions()
return label_values in self._label_sets
def _is_full(self) -> bool:
return len(self._label_sets) >= self._max_series
def _append(self, label_values: tuple[str, ...]) -> None:
descriptor: Final = os.open(self._path, os.O_WRONLY | os.O_APPEND | os.O_CREAT, 0o600)
try:
os.write(descriptor, b"\n" + _LABEL_VALUES.dump_json(label_values) + b"\n")
finally:
os.close(descriptor)
def _read_new_admissions(self) -> None:
try:
with open(self._path, "rb") as admissions_file:
admissions_file.seek(self._read_offset)
unread: Final = admissions_file.read()
except FileNotFoundError:
return
complete_lines, newline, _ = unread.rpartition(b"\n")
if not newline:
return
self._read_offset += len(complete_lines) + len(newline)
for label_values in map(_parse_admission, complete_lines.split(b"\n")):
if self._is_full():
return
if label_values is not None:
self._label_sets.add(label_values)
class SharedPrometheusSeriesAdmissions:
"""Picks which label sets get a series when several worker processes write to one
``PROMETHEUS_MULTIPROC_DIR``. Each metric has one append-only file there, and its first ``max_series``
distinct lines are the admitted label sets. Every worker reads the same lines in the same order, so all of
them, including a worker that replaces an exited one, admit the same label sets and a scrape that merges
the workers stays at the cap. Each record sits between two newlines, so a record a worker could only write
part of (the directory ran out of space) is a line of its own that admits nothing for every worker, and
it neither hides the records after it nor runs into the next worker's record."""
def __init__(self, directory: str) -> None:
self._directory = directory
self._admissions: dict[str, _MetricAdmissions] = {} # mutable-ok: one entry per metric, added on first use
self.lock = RLock()
def admit_series(self, metric_name: str, label_values: tuple[str, ...], max_series: int) -> bool:
with self.lock:
if metric_name not in self._admissions:
self._admissions[metric_name] = _MetricAdmissions(
path=os.path.join(self._directory, f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}{metric_name}"),
max_series=max_series,
)
return self._admissions[metric_name].admit(label_values)

View file

@ -1,7 +1,7 @@
"""
Prometheus multiprocess directory cleanup utilities.
Wipes all .db files on startup so workers start with a clean slate.
Wipes all .db files and admitted-series files on startup so workers start with a clean slate.
"""
from __future__ import annotations
@ -12,22 +12,34 @@ import re
from typing import Final
from litellm._logging import verbose_proxy_logger
from litellm.constants import PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX
_LIVE_GAUGE_PID: Final = re.compile(r"gauge_live[a-z]*_(\d+)\.db$")
def wipe_directory(directory: str) -> None:
"""Delete all .db files in the directory. Called once before workers fork."""
files: Final = glob.glob(os.path.join(directory, "*.db"))
deleted = 0
for filepath in files:
try:
os.remove(filepath)
deleted += 1
except OSError as e:
verbose_proxy_logger.warning("Failed to delete stale prometheus file %s: %s", filepath, e)
"""Delete all .db files and admitted-series files in the directory. Called once at boot, before any worker
starts, so a restart frees every capped slot and drops the samples of the workers that exited."""
_remove(directory, (*glob.glob(os.path.join(directory, "*.db")), *_admitted_series_files(directory)))
def _admitted_series_files(directory: str) -> tuple[str, ...]:
return tuple(glob.glob(os.path.join(directory, f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}*")))
def _remove(directory: str, files: tuple[str, ...]) -> None:
deleted: Final = sum(_removed(filepath) for filepath in files)
if deleted:
verbose_proxy_logger.info("Prometheus cleanup: wiped %s stale .db files from %s", deleted, directory)
verbose_proxy_logger.info("Prometheus cleanup: wiped %s stale files from %s", deleted, directory)
def _removed(filepath: str) -> int:
try:
os.remove(filepath)
except OSError as e:
verbose_proxy_logger.warning("Failed to delete stale prometheus file %s: %s", filepath, e)
return 0
return 1
def mark_worker_exit(worker_pid: int) -> None:

View file

@ -687,14 +687,16 @@ class ProxyInitializationHelpers:
"""
import tempfile
if prometheus_metrics_port is None and (
num_workers <= 1 or not ProxyInitializationHelpers._prometheus_callback_configured(litellm_settings)
):
return None
from litellm.proxy.prometheus_cleanup import wipe_directory
configured_dir: Final = os.environ.get("PROMETHEUS_MULTIPROC_DIR") or os.environ.get("prometheus_multiproc_dir")
if prometheus_metrics_port is None and (
num_workers <= 1 or not ProxyInitializationHelpers._prometheus_callback_configured(litellm_settings)
):
if configured_dir:
wipe_directory(configured_dir)
return None
multiproc_dir: Final = configured_dir or os.path.join(tempfile.gettempdir(), "litellm_prometheus_multiproc")
os.environ["PROMETHEUS_MULTIPROC_DIR"] = multiproc_dir
@ -1503,6 +1505,11 @@ def run_server(
os.environ["NUM_WORKERS"] = str(num_workers)
# Skip server startup if requested (after all setup is done)
if skip_server_startup:
print("LiteLLM: Setup complete. Skipping server startup as requested.")
return
# Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups
prometheus_multiproc_dir: Final = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir(
num_workers=num_workers,
@ -1510,11 +1517,6 @@ def run_server(
prometheus_metrics_port=prometheus_metrics_port,
)
# Skip server startup if requested (after all setup is done)
if skip_server_startup:
print("LiteLLM: Setup complete. Skipping server startup as requested.")
return
if prometheus_metrics_port is not None and prometheus_multiproc_dir is not None:
from litellm.proxy.prometheus_metrics_server import MetricsServerStartupError, start_metrics_server_process

View file

@ -180,6 +180,54 @@ def _launch_until_bound(
return _launch_until_bound(command, root, environment, output, attempts - 1)
def _proxy_root() -> Path:
return Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
def _proxy_environment(
gateway: Gateway, overrides: Mapping[str, str], remove_environment: tuple[str, ...]
) -> Mapping[str, str]:
return MappingProxyType(
{
**{
name: value
for name, value in {**os.environ, **proxy_database_environment()}.items()
if name not in remove_environment
},
"LITELLM_MASTER_KEY": gateway.key,
"LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"),
"STORE_MODEL_IN_DB": "True",
**overrides,
}
)
def setup_only_proxy_run(
gateway: Gateway, overrides: Mapping[str, str], *, config: Path, workers: int
) -> subprocess.CompletedProcess[str]:
"""The proxy CLI's `--skip_server_startup` pass (the image's setup step), run to completion with the
environment an owned proxy gets."""
return subprocess.run( # test-quality-ok: the checkout at the working directory is the proxy under test
(
sys.executable,
"-m",
"integration._support.proxy",
"--config",
str(config),
"--num_workers",
str(workers),
*DB_PUSH,
"--skip_server_startup",
),
cwd=_proxy_root(),
env=dict(_proxy_environment(gateway, overrides, ())),
capture_output=True,
text=True,
timeout=300,
check=False,
)
@contextmanager
def owned_proxy_process(
gateway: Gateway,
@ -192,18 +240,8 @@ def owned_proxy_process(
database_setup: tuple[str, ...] = DB_PUSH,
extra_arguments: tuple[str, ...] = (),
) -> Iterator[OwnedProxy]:
root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
environment: Final = {
**{
name: value
for name, value in {**os.environ, **proxy_database_environment()}.items()
if name not in remove_environment
},
"LITELLM_MASTER_KEY": gateway.key,
"LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"),
"STORE_MODEL_IN_DB": "True",
**overrides,
}
root: Final = _proxy_root()
environment: Final = _proxy_environment(gateway, overrides, remove_environment)
output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory)))
output.mkdir(parents=True, exist_ok=True)
command: Final = (
@ -230,6 +268,37 @@ def owned_proxy_process(
_stop(process)
@contextmanager
def owned_gateway_image(
gateway: Gateway, directory: Path, overrides: Mapping[str, str], *, config: Path, workers: int
) -> Iterator[OwnedProxy]:
"""The componentized gateway started the way its image starts it: `docker/component_entrypoint.sh` running
`python -m gateway.launch`, with the config handed over as `CONFIG_FILE_PATH`. It serves the data plane only,
so keys come from a proxy that shares its database."""
root: Final = _proxy_root()
environment: Final = _proxy_environment(gateway, {**overrides, "CONFIG_FILE_PATH": str(config)}, ())
output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory)))
output.mkdir(parents=True, exist_ok=True)
command: Final = (
str(root / "docker" / "component_entrypoint.sh"),
sys.executable,
"-m",
"gateway.launch",
"--workers",
str(workers),
"--host",
"127.0.0.1",
)
launch: Final = _launch_until_bound(command, root, environment, output, _PORT_ATTEMPTS)
try:
with httpx.Client(
base_url=f"http://127.0.0.1:{launch.port}", timeout=15, trust_env=False, limits=GATEWAY_LIMITS
) as client:
yield OwnedProxy(Gateway(client, gateway.key, gateway.upstream_url), launch.process, launch.log)
finally:
_stop(launch.process)
def _is_ready(client: httpx.Client) -> bool:
try:
return client.get("/health/readiness", timeout=2).status_code == 200
@ -245,15 +314,8 @@ def refused_boot_log(
config: Path | None = None,
) -> str:
"""Start the proxy and return its log once it exits non-zero instead of becoming ready."""
root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
environment: Final = {
**os.environ,
**proxy_database_environment(),
"LITELLM_MASTER_KEY": gateway.key,
"LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"),
"STORE_MODEL_IN_DB": "True",
**overrides,
}
root: Final = _proxy_root()
environment: Final = _proxy_environment(gateway, overrides, ())
output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory)))
output.mkdir(parents=True, exist_ok=True)
command: Final = (
@ -336,7 +398,7 @@ class UpstreamSlot:
@contextmanager
def owned_upstream(directory: Path) -> Generator[UpstreamSlot]:
root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
root: Final = _proxy_root()
slot: Final = UpstreamSlot(directory, _free_port(), root)
slot.start()
try:

View file

@ -0,0 +1,362 @@
"""Rig and readers for the Prometheus series-cap cells: a capped proxy, keys that fill the cap, and the
scrape, the multiprocess sample files, and the spend log read back per request."""
from __future__ import annotations
import json
import threading
import uuid
from collections import Counter
from collections.abc import Callable, Iterator, Mapping, Sequence
from contextlib import ExitStack, contextmanager
from dataclasses import dataclass
from itertools import chain
from pathlib import Path
from types import MappingProxyType
from typing import Final
import httpx
import yaml
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
from integration._support.database import read_rows
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.responses_vendor import ResponsesVendor, same_response
from integration._support.wire import Reply, Request, Wire, wire_server
from prometheus_client.mmap_dict import MmapedDict
from prometheus_client.parser import text_string_to_metric_families
from pydantic import JsonValue
REQUESTS: Final = "litellm_requests_metric_total"
PROXY_REQUESTS: Final = "litellm_proxy_total_requests_metric_total"
PROXY_FAILURES: Final = "litellm_proxy_failed_requests_metric_total"
CACHE_HITS: Final = "litellm_cache_hits_metric_total"
REMAINING_REQUESTS: Final = "litellm_remaining_api_key_requests_for_model"
SUCCESSFUL_FALLBACKS: Final = "litellm_deployment_successful_fallbacks_total"
FAILED_FALLBACKS: Final = "litellm_deployment_failed_fallbacks_total"
OVERFLOW: Final = "other"
USER_AGENT: Final = "litellm-series-cap-audit/1"
AGENT_HEADERS: Final[Mapping[str, str]] = MappingProxyType({"User-Agent": USER_AGENT})
PROVIDER_OUTAGE: Final = "synthetic provider outage"
_SERIES_ATTRIBUTES: Final = frozenset({"le", "pid"})
@dataclass(frozen=True, slots=True)
class Call:
"""One request the cells can follow end to end: the call id the proxy keeps as the spend log's request id
and the marker the scripted provider echoes in its answer."""
call_id: str
marker: str
@classmethod
def new(cls) -> Call:
identity: Final = uuid.uuid4()
return cls(str(identity), identity.hex)
@property
def text(self) -> str:
return f"say marker-{self.marker}"
@property
def answer(self) -> str:
return f"answer marker-{self.marker}"
@property
def message(self) -> dict[str, str]:
return {"role": "user", "content": self.text}
@property
def headers(self) -> dict[str, str]:
return {"x-litellm-call-id": self.call_id}
@dataclass(frozen=True, slots=True)
class Key:
token: str
alias: str
@dataclass(frozen=True, slots=True)
class Provider:
vendor: ResponsesVendor
outage: threading.Event
failing_models: frozenset[str]
def respond(self, request: Request) -> Reply:
if self.outage.is_set() or self._failing(request):
return Reply(status=500, body=json.dumps({"error": {"message": PROVIDER_OUTAGE}}).encode())
return self.vendor.respond(request)
def _failing(self, request: Request) -> bool:
if request.method != "POST" or not self.failing_models:
return False
body: Final = json.loads(request.body)
return isinstance(body, dict) and body.get("model") in self.failing_models
@dataclass(frozen=True, slots=True)
class Sample:
family: str
kind: str
name: str
labels: Mapping[str, str]
value: float
def identity(self) -> tuple[tuple[str, str], ...]:
"""The label set that makes this a series of its own: the histogram bucket and the multiprocess pid are
attributes of one series, not separate ones."""
return tuple(sorted((name, value) for name, value in self.labels.items() if name not in _SERIES_ATTRIBUTES))
def is_overflow(self) -> bool:
identity: Final = self.identity()
return bool(identity) and all(value == OVERFLOW for _, value in identity)
@dataclass(frozen=True, slots=True)
class CapRig:
proxy: OwnedProxy
scenario: Scenario
model: str
provider: Wire
outage: threading.Event
warm: tuple[Key, ...]
prom_dir: Path
@property
def gateway(self) -> Gateway:
return self.proxy.gateway
@property
def base_url(self) -> str:
return str(self.gateway.client.base_url).rstrip("/")
@property
def openai_base(self) -> str:
return self.base_url + "/v1"
@property
def warm_aliases(self) -> frozenset[str]:
return frozenset(key.alias for key in self.warm)
def key(self, cell: str) -> Key:
alias: Final = f"{cell}-{uuid.uuid4().hex[:12]}"
return Key(self.scenario.key(key_alias=alias), alias)
def chat(self, key: Key, call: Call) -> httpx.Response:
return chat_once(self.base_url, key, self.model, call)
def chat_once(base_url: str, key: Key, model: str, call: Call) -> httpx.Response:
with httpx.Client(base_url=base_url, timeout=60, trust_env=False) as client:
return client.post(
"/v1/chat/completions",
json={"model": model, "messages": [call.message]},
headers={**AGENT_HEADERS, **call.headers, "Authorization": f"Bearer {key.token}"},
)
def series_cap_config(
directory: Path,
settings: Mapping[str, JsonValue],
*,
model_list: Sequence[Mapping[str, JsonValue]] = (),
router_settings: Mapping[str, JsonValue] | None = None,
) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
assert isinstance(config, dict)
config["litellm_settings"] = {**config["litellm_settings"], "callbacks": ["prometheus"], **settings}
config["router_settings"] = {**config["router_settings"], "num_retries": 0, **(router_settings or {})}
if model_list:
config["model_list"] = [dict(entry) for entry in model_list]
path: Final = directory / "series-cap.yaml"
path.write_text(yaml.safe_dump(config))
return path
@contextmanager
def series_cap_rig(
directory: Path,
settings: Mapping[str, JsonValue],
*,
workers: int,
warm_keys: int,
multiproc_dir: Path | None = None,
failing_models: frozenset[str] = frozenset(),
deployments: Callable[[str], Sequence[Mapping[str, JsonValue]]] | None = None,
router_settings: Mapping[str, JsonValue] | None = None,
) -> Iterator[CapRig]:
outage: Final = threading.Event()
double: Final = Provider(ResponsesVendor(), outage, failing_models)
prom_dir: Final = multiproc_dir if multiproc_dir is not None else directory / "prom"
prom_dir.mkdir(exist_ok=True)
shared_samples: Final = multiproc_dir is not None or workers > 1
with ExitStack() as stack:
gateway: Final = stack.enter_context(gateway_from_environment())
provider: Final = stack.enter_context(wire_server(double.respond))
config: Final = series_cap_config(
directory,
settings,
model_list=deployments(provider.url) if deployments is not None else (),
router_settings=router_settings,
)
owned: Final = stack.enter_context(
owned_proxy_process(
gateway,
directory,
{"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)} if shared_samples else {},
config=config,
remove_environment=() if shared_samples else ("PROMETHEUS_MULTIPROC_DIR",),
workers=workers,
)
)
scenario: Final = stack.enter_context(owned.gateway.scenario())
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
rig: Final = CapRig(owned, scenario, model, provider, outage, (), prom_dir)
warm: Final = tuple(rig.key("warm") for _ in range(warm_keys))
for key in warm:
_warm_up(rig, key)
eventually(
lambda: alias_values(scrape(owned.gateway), REQUESTS),
lambda seen: all(key.alias in seen for key in warm),
seconds=60,
)
yield CapRig(owned, scenario, model, provider, outage, warm, prom_dir)
def _warm_up(rig: CapRig, key: Key) -> None:
response: Final = rig.chat(key, Call.new())
assert response.status_code == 200, response.text
def _samples(text: str) -> Iterator[Sample]:
for family in text_string_to_metric_families(text):
for sample in family.samples:
yield Sample(
family.name, family.type, sample.name, MappingProxyType(dict(sample.labels)), float(sample.value)
)
def scrape(gateway: Gateway) -> tuple[Sample, ...]:
response: Final = gateway.client.request(
"GET", "/metrics", headers={"Authorization": f"Bearer {gateway.key}"}, follow_redirects=True
)
assert response.status_code == 200, f"GET /metrics: {response.status_code} {response.text[:300]}"
return tuple(_samples(response.text))
def alias_values(samples: Sequence[Sample], name: str) -> frozenset[str]:
"""The key aliases holding a series of their own on the metric; the shared overflow series is not one."""
return frozenset(
sample.labels["api_key_alias"]
for sample in samples
if sample.name == name and sample.labels.get("api_key_alias") not in (None, OVERFLOW)
)
def alias_total(samples: Sequence[Sample], name: str, alias: str) -> float:
return sum(
sample.value for sample in samples if sample.name == name and sample.labels.get("api_key_alias") == alias
)
def overflow_total(samples: Sequence[Sample], name: str) -> float:
return sum(sample.value for sample in samples if sample.name == name and sample.is_overflow())
def label_values(samples: Sequence[Sample]) -> frozenset[str]:
return frozenset(chain.from_iterable(sample.labels.values() for sample in samples))
def gauge_samples(samples: Sequence[Sample]) -> tuple[Sample, ...]:
return tuple(sample for sample in samples if sample.kind == "gauge")
def series_per_family(samples: Sequence[Sample]) -> Mapping[str, int]:
"""How many series of their own each metric family holds, the shared `other` series left out."""
owned: Final = frozenset(
(sample.family, sample.identity()) for sample in samples if sample.identity() and not sample.is_overflow()
)
return MappingProxyType(dict(Counter(family for family, _ in owned)))
def families_over(samples: Sequence[Sample], cap: int) -> tuple[tuple[str, int], ...]:
"""Every metric family holding more series of its own than the cap allows."""
return tuple(sorted((family, count) for family, count in series_per_family(samples).items() if count > cap))
@dataclass(frozen=True, slots=True)
class SpendRow:
request_id: str
status: str
def spend_rows(alias: str) -> tuple[SpendRow, ...]:
"""Every spend log row the key wrote: a success row carries the response id the caller received, a failure
row the call id the caller sent."""
rows: Final = read_rows(
"SELECT request_id, status FROM \"LiteLLM_SpendLogs\" WHERE metadata->>'user_api_key_alias' = %s",
(alias,),
)
return tuple(SpendRow(str(row["request_id"]), str(row["status"])) for row in rows)
def expect_spend_rows(
alias: str, response_ids: Sequence[str], call_ids: Sequence[str] = (), earlier: Sequence[SpendRow] = ()
) -> None:
"""One new row per request on top of the rows the key already had: successes found by the response id the
caller got, failures by their call id."""
expected: Final = len(earlier) + len(response_ids) + len(call_ids)
rows: Final = eventually(lambda: spend_rows(alias), lambda found: len(found) >= expected, seconds=70)
fresh: Final = tuple(row for row in rows if row not in earlier)
assert len(rows) == expected and len(fresh) == len(response_ids) + len(call_ids), (rows, earlier)
for response_id in response_ids:
assert any(row.status == "success" and same_response(row.request_id, response_id) for row in fresh), (
response_id,
fresh,
)
for call_id in call_ids:
assert any(row.status == "failure" and row.request_id == call_id for row in fresh), (call_id, fresh)
def sse_data(text: str) -> tuple[dict[str, JsonValue], ...]:
"""The JSON payload of every `data:` frame in a server-sent event stream, the `[DONE]` sentinel left out."""
payloads: Final = tuple(
line.removeprefix("data:").strip() for line in text.splitlines() if line.startswith("data:")
)
return tuple(object_value(json.loads(payload)) for payload in payloads if payload and payload != "[DONE]")
def received_markers(provider: Wire) -> tuple[str, ...]:
return tuple(chain.from_iterable(_markers_in(request.body) for request in provider.drain()))
def _markers_in(body: bytes) -> tuple[str, ...]:
return tuple(part[:32].decode() for part in body.split(b"marker-")[1:])
@dataclass(frozen=True, slots=True)
class WorkerSamples:
pid: int
aliases: frozenset[str]
overflow: float
def worker_samples(prom_dir: Path, name: str) -> tuple[WorkerSamples, ...]:
return tuple(_worker_samples(path, name) for path in sorted(prom_dir.glob("counter_*.db")))
def _worker_samples(path: Path, name: str) -> WorkerSamples:
pid: Final = int(path.stem.rsplit("_", 1)[1])
rows: Final = tuple(_counter_rows(path, name))
return WorkerSamples(
pid,
frozenset(labels["api_key_alias"] for labels, _ in rows if labels.get("api_key_alias") not in (None, OVERFLOW)),
sum(value for labels, value in rows if all(label == OVERFLOW for label in labels.values())),
)
def _counter_rows(path: Path, name: str) -> Iterator[tuple[Mapping[str, str], float]]:
for key, value, *_ in MmapedDict.read_all_values_from_file(str(path)):
_, sample_name, labels, _ = json.loads(key)
if sample_name == name:
yield labels, float(value)

View file

@ -9,6 +9,7 @@ from urllib.parse import urlparse
import pytest
import yaml
from integration._support.otlp_sink import SpanSinks, owned_sinks
from integration._support.prometheus_series import CapRig, series_cap_rig
from pydantic import JsonValue
AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path]
@ -51,3 +52,15 @@ def langfuse_vars(audit_sinks: SpanSinks) -> dict[str, JsonValue]:
"langfuse_secret_key": "sk-lf-audit",
"langfuse_host": audit_sinks.tenant,
}
@pytest.fixture(scope="session")
def capped(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
"""A two-worker proxy capped at three series per metric, with three keys already holding a series each."""
with series_cap_rig(
tmp_path_factory.mktemp("series-cap"),
{"prometheus_metrics_max_series_per_metric": 3},
workers=2,
warm_keys=3,
) as rig:
yield rig

View file

@ -0,0 +1,498 @@
"""Prometheus series cap on the live proxy: label sets past prometheus_metrics_max_series_per_metric share one
`other` series on every labeled counter and histogram and stay out of the gauges, idle series expire under
prometheus_metrics_ttl_seconds in single-process mode only, and a setting that is not a positive number is
ignored with a warning instead of silencing the metrics."""
from __future__ import annotations
import time
from collections.abc import Iterator, Sequence
from typing import Final
import anthropic
import httpx
import openai
import pytest
from integration._support.client import eventually, object_value, string_value
from integration._support.prometheus_series import (
AGENT_HEADERS,
CACHE_HITS,
FAILED_FALLBACKS,
OVERFLOW,
PROVIDER_OUTAGE,
PROXY_FAILURES,
PROXY_REQUESTS,
REMAINING_REQUESTS,
REQUESTS,
SUCCESSFUL_FALLBACKS,
Call,
CapRig,
Key,
Sample,
WorkerSamples,
alias_values,
chat_once,
expect_spend_rows,
families_over,
gauge_samples,
label_values,
overflow_total,
received_markers,
scrape,
series_cap_rig,
sse_data,
worker_samples,
)
from pydantic import JsonValue
pytestmark = pytest.mark.timeout(240)
CAP: Final = 3
TTL_SECONDS: Final = 2
CLEANUP_SECONDS: Final = 1
PRIMARY: Final = "primary"
FALLBACK: Final = "fallback"
TTL_IGNORED_WARNING: Final = "prometheus_metrics_ttl_seconds is ignored while PROMETHEUS_MULTIPROC_DIR is set"
def _grew(before: Sequence[Sample], after: Sequence[Sample], name: str, by: int) -> bool:
return overflow_total(after, name) - overflow_total(before, name) >= by
def _overflowed(rig: CapRig, key: Key, before: Sequence[Sample], requests: int) -> tuple[Sample, ...]:
"""The scrape once the key's requests landed on `other` for both request counters, or as soon as the key got
a series of its own, so the caller's assertion fails fast on a proxy without the cap."""
return eventually(
lambda: scrape(rig.gateway),
lambda after: (
(_grew(before, after, REQUESTS, requests) and _grew(before, after, PROXY_REQUESTS, requests))
or key.alias in alias_values(after, REQUESTS)
),
seconds=60,
)
def _expect_other(
rig: CapRig, key: Key, calls: Sequence[Call], response_ids: Sequence[str], before: Sequence[Sample]
) -> None:
samples: Final = _overflowed(rig, key, before, len(calls))
assert key.alias not in alias_values(samples, REQUESTS) | alias_values(samples, PROXY_REQUESTS), key.alias
assert not families_over(samples, CAP), families_over(samples, CAP)
assert overflow_total(samples, REQUESTS) - overflow_total(before, REQUESTS) == len(calls)
expect_spend_rows(key.alias, response_ids)
markers: Final = received_markers(rig.provider)
assert all(call.marker in markers for call in calls), (calls, markers)
def _bearer(key: Key, call: Call) -> dict[str, str]:
return {**AGENT_HEADERS, **call.headers, "Authorization": f"Bearer {key.token}"}
class TestCapped:
def test_openai_sync_chat_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
"""H1: two OpenAI SDK chat completions from a fourth key count on `other` and keep their spend rows."""
key: Final = capped.key("h1")
calls: Final = (Call.new(), Call.new())
before: Final = scrape(capped.gateway)
client: Final = openai.OpenAI(
base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
)
completions: Final = tuple(
client.chat.completions.create(model=capped.model, messages=[call.message], extra_headers=call.headers)
for call in calls
)
assert tuple(completion.choices[0].message.content for completion in completions) == tuple(
call.answer for call in calls
)
_expect_other(capped, key, calls, tuple(completion.id for completion in completions), before)
async def test_openai_async_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
"""H2: a streamed AsyncOpenAI chat completion from a fourth key counts on `other` once the stream ends."""
key: Final = capped.key("h2")
call: Final = Call.new()
before: Final = scrape(capped.gateway)
client: Final = openai.AsyncOpenAI(
base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
)
stream: Final = await client.chat.completions.create(
model=capped.model, messages=[call.message], stream=True, extra_headers=call.headers
)
chunks: Final = tuple([chunk async for chunk in stream])
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == call.answer
ids: Final = frozenset(chunk.id for chunk in chunks)
assert len(ids) == 1, ids
_expect_other(capped, key, (call,), tuple(ids), before)
def test_anthropic_sync_messages_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
"""H3: an Anthropic SDK /v1/messages call from a fourth key counts on `other`."""
key: Final = capped.key("h3")
call: Final = Call.new()
before: Final = scrape(capped.gateway)
client: Final = anthropic.Anthropic(
base_url=capped.base_url, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
)
message: Final = client.messages.create(
model=capped.model, max_tokens=64, messages=[call.message], extra_headers=call.headers
)
assert "".join(block.text for block in message.content if block.type == "text") == call.answer
_expect_other(capped, key, (call,), (message.id,), before)
async def test_anthropic_async_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
"""H4: a streamed AsyncAnthropic /v1/messages call from a fourth key counts on `other`."""
key: Final = capped.key("h4")
call: Final = Call.new()
before: Final = scrape(capped.gateway)
client: Final = anthropic.AsyncAnthropic(
base_url=capped.base_url, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
)
async with client.messages.stream(
model=capped.model, max_tokens=64, messages=[call.message], extra_headers=call.headers
) as stream:
final: Final = await stream.get_final_message()
assert "".join(block.text for block in final.content if block.type == "text") == call.answer
_expect_other(capped, key, (call,), (final.id,), before)
def test_openai_sync_responses_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
"""H5: an OpenAI SDK /v1/responses call from a fourth key counts on `other`."""
key: Final = capped.key("h5")
call: Final = Call.new()
before: Final = scrape(capped.gateway)
client: Final = openai.OpenAI(
base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
)
response: Final = client.responses.create(model=capped.model, input=call.text, extra_headers=call.headers)
assert response.output_text == call.answer
_expect_other(capped, key, (call,), (response.id,), before)
def test_raw_responses_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
"""H6: a raw httpx streamed /v1/responses call from a fourth key counts on `other`."""
key: Final = capped.key("h6")
call: Final = Call.new()
before: Final = scrape(capped.gateway)
with httpx.Client(base_url=capped.base_url, timeout=60, trust_env=False) as client:
response: Final = client.post(
"/v1/responses",
json={"model": capped.model, "input": call.text, "stream": True},
headers=_bearer(key, call),
)
assert response.status_code == 200, response.text
events: Final = sse_data(response.text)
deltas: Final = tuple(event for event in events if event.get("type") == "response.output_text.delta")
assert "".join(string_value(event["delta"]) for event in deltas) == call.answer
completed: Final = tuple(event for event in events if event.get("type") == "response.completed")
assert len(completed) == 1, events
response_id: Final = string_value(object_value(completed[0]["response"])["id"])
_expect_other(capped, key, (call,), (response_id,), before)
def test_gauges_never_get_an_other_series(self, capped: CapRig) -> None:
"""H7: a fourth key's request leaves no gauge sample for it and no gauge sample labeled `other`."""
key: Final = capped.key("h7")
before: Final = scrape(capped.gateway)
assert capped.chat(key, Call.new()).status_code == 200
samples: Final = _overflowed(capped, key, before, 1)
assert key.alias not in label_values(samples)
gauges: Final = gauge_samples(samples)
assert not any(OVERFLOW in gauge.labels.values() for gauge in gauges), gauges
for alias in capped.warm_aliases:
assert any(
gauge.name == REMAINING_REQUESTS and gauge.labels.get("api_key_alias") == alias for gauge in gauges
), alias
def test_cache_hits_past_the_cap_count_on_other(self, capped: CapRig) -> None:
"""H8: the cache-hit twin: one populating call, hits from the warm keys, then a fourth key's hit on `other`."""
shared: Final = Call.new()
first, second, third = capped.warm
extra: Final = capped.key("h8")
capped.provider.drain()
before: Final = scrape(capped.gateway)
for key in (first, first, second, third, extra):
response = capped.chat(key, shared)
assert response.status_code == 200 and shared.answer in response.text, response.text
samples: Final = eventually(
lambda: scrape(capped.gateway),
lambda after: _grew(before, after, CACHE_HITS, 1) or extra.alias in alias_values(after, CACHE_HITS),
seconds=60,
)
assert alias_values(samples, CACHE_HITS) == capped.warm_aliases
assert overflow_total(samples, CACHE_HITS) - overflow_total(before, CACHE_HITS) == 1
assert received_markers(capped.provider).count(shared.marker) == 1
def test_both_workers_share_the_admitted_series(self, capped: CapRig) -> None:
"""H9: fresh connections reach both workers, and each worker's own sample file names only the warm aliases
while counting the fourth key on `other`, since the admitted sets live in the shared directory."""
extra: Final = capped.key("h9")
def send_on_a_fresh_connection() -> tuple[WorkerSamples, ...]:
assert capped.chat(extra, Call.new()).status_code == 200
return worker_samples(capped.prom_dir, REQUESTS)
workers: Final = eventually(
send_on_a_fresh_connection,
lambda found: (
sum(1 for worker in found if worker.overflow > 0) >= 2
or any(extra.alias in worker.aliases for worker in found)
),
seconds=90,
)
assert all(extra.alias not in worker.aliases for worker in workers), workers
assert sum(1 for worker in workers if worker.overflow > 0) >= 2, workers
assert frozenset().union(*(worker.aliases for worker in workers)) == capped.warm_aliases, workers
def test_failures_past_the_cap_count_on_other(self, capped: CapRig) -> None:
"""F1: provider failures fill the failure counter's cap with the warm keys, a fourth key's lands on `other`."""
key: Final = capped.key("f1")
call: Final = Call.new()
before: Final = scrape(capped.gateway)
capped.outage.set()
try:
for warm in capped.warm:
assert capped.chat(warm, Call.new()).status_code == 500
response: Final = capped.chat(key, call)
finally:
capped.outage.clear()
assert response.status_code == 500 and PROVIDER_OUTAGE in response.text, response.text
samples: Final = eventually(
lambda: scrape(capped.gateway),
lambda after: _grew(before, after, PROXY_FAILURES, 1) or key.alias in alias_values(after, PROXY_FAILURES),
seconds=60,
)
assert key.alias not in label_values(samples)
assert overflow_total(samples, PROXY_FAILURES) - overflow_total(before, PROXY_FAILURES) == 1
assert alias_values(samples, PROXY_FAILURES) == capped.warm_aliases
expect_spend_rows(key.alias, (), (call.call_id,))
def test_config_update_cannot_lift_a_yaml_cap(self, capped: CapRig) -> None:
"""E1: /config/update refuses the YAML-owned cap, so a fourth key still lands on `other`."""
response: Final = capped.gateway.client.post(
"/config/update",
json={"litellm_settings": {"prometheus_metrics_max_series_per_metric": 50}},
headers={"Authorization": f"Bearer {capped.gateway.key}"},
)
assert response.status_code == 400, response.text
key: Final = capped.key("e1")
before: Final = scrape(capped.gateway)
assert capped.chat(key, Call.new()).status_code == 200
samples: Final = _overflowed(capped, key, before, 1)
assert key.alias not in label_values(samples)
@pytest.fixture(scope="class")
def ttl(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
with series_cap_rig(
tmp_path_factory.mktemp("series-ttl"),
{
"prometheus_metrics_max_series_per_metric": 2,
"prometheus_metrics_ttl_seconds": TTL_SECONDS,
"prometheus_metrics_cleanup_interval_seconds": CLEANUP_SECONDS,
},
workers=1,
warm_keys=2,
) as rig:
yield rig
class TestTtl:
def test_idle_series_expire_and_free_their_slot(self, ttl: CapRig) -> None:
"""T1: a third key lands on `other`; once the idle first key expires, a new key gets its own series."""
first, second = ttl.warm
extra: Final = ttl.key("t1-extra")
assert ttl.chat(extra, Call.new()).status_code == 200
samples: Final = eventually(
lambda: scrape(ttl.gateway),
lambda after: overflow_total(after, REQUESTS) >= 1 or extra.alias in alias_values(after, REQUESTS),
seconds=60,
)
assert extra.alias not in label_values(samples)
def keep_second_busy() -> tuple[Sample, ...]:
assert ttl.chat(second, Call.new()).status_code == 200
return scrape(ttl.gateway)
expired: Final = eventually(keep_second_busy, lambda after: first.alias not in label_values(after), seconds=30)
assert second.alias in alias_values(expired, REQUESTS)
late: Final = ttl.key("t1-late")
assert ttl.chat(late, Call.new()).status_code == 200
named: Final = eventually(
lambda: scrape(ttl.gateway), lambda after: late.alias in alias_values(after, REQUESTS), seconds=30
)
assert late.alias in alias_values(named, REQUESTS)
@pytest.fixture(scope="class")
def ttl_multiproc(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
with series_cap_rig(
tmp_path_factory.mktemp("series-ttl-multiproc"),
{
"prometheus_metrics_max_series_per_metric": CAP,
"prometheus_metrics_ttl_seconds": TTL_SECONDS,
"prometheus_metrics_cleanup_interval_seconds": CLEANUP_SECONDS,
},
workers=2,
warm_keys=3,
) as rig:
yield rig
class TestTtlMultiproc:
def test_ttl_is_ignored_with_two_workers_while_the_cap_applies(self, ttl_multiproc: CapRig) -> None:
"""M1: with two workers an idle key keeps its series past the TTL, the cap still applies, and the log says so."""
first, second, _ = ttl_multiproc.warm
deadline: Final = time.monotonic() + 2 * TTL_SECONDS
while time.monotonic() < deadline:
assert ttl_multiproc.chat(second, Call.new()).status_code == 200
assert first.alias in alias_values(scrape(ttl_multiproc.gateway), REQUESTS)
extra: Final = ttl_multiproc.key("m1")
before: Final = scrape(ttl_multiproc.gateway)
assert ttl_multiproc.chat(extra, Call.new()).status_code == 200
samples: Final = _overflowed(ttl_multiproc, extra, before, 1)
assert extra.alias not in label_values(samples)
assert TTL_IGNORED_WARNING in ttl_multiproc.proxy.log.read_text()
@pytest.fixture(scope="class")
def ignored(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
with series_cap_rig(
tmp_path_factory.mktemp("series-ignored"),
{"prometheus_metrics_max_series_per_metric": "five", "prometheus_metrics_ttl_seconds": ""},
workers=1,
warm_keys=5,
) as rig:
yield rig
class TestIgnored:
def test_settings_that_are_not_positive_numbers_are_ignored_with_a_warning(self, ignored: CapRig) -> None:
"""I1: a cap of "five" and an empty TTL leave every key its own series and each warning names its setting."""
samples: Final = scrape(ignored.gateway)
assert alias_values(samples, REQUESTS) >= ignored.warm_aliases
assert not any(sample.is_overflow() for sample in samples)
log: Final = ignored.proxy.log.read_text()
assert (
"prometheus_metrics_max_series_per_metric is ignored because it is not a number greater than 0 (got 'five')"
in log
)
assert "prometheus_metrics_ttl_seconds is ignored because it is not a number greater than 0 (got '')" in log
@pytest.fixture(scope="class")
def ignored_interval(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
with series_cap_rig(
tmp_path_factory.mktemp("series-ignored-interval"),
{
"prometheus_metrics_max_series_per_metric": CAP,
"prometheus_metrics_ttl_seconds": TTL_SECONDS,
"prometheus_metrics_cleanup_interval_seconds": "sixty",
},
workers=1,
warm_keys=CAP,
) as rig:
yield rig
class TestIgnoredInterval:
def test_a_cleanup_interval_that_is_not_a_number_is_ignored_with_a_warning_while_the_cap_and_ttl_apply(
self, ignored_interval: CapRig
) -> None:
"""I2: a cleanup interval of "sixty" next to a TTL is ignored for the default, so the first labeled emit
still counts (it raised inside the logging callback before) and a fourth key lands on `other`."""
key: Final = ignored_interval.key("i2")
call: Final = Call.new()
before: Final = scrape(ignored_interval.gateway)
response: Final = ignored_interval.chat(key, call)
assert response.status_code == 200 and call.answer in response.text, response.text
_expect_other(ignored_interval, key, (call,), (string_value(object_value(response.json())["id"]),), before)
assert (
"prometheus_metrics_cleanup_interval_seconds is ignored because it is not a number of at least 0 "
"(got 'sixty'). Idle series are checked every 60.0 seconds"
) in ignored_interval.proxy.log.read_text()
def _fallback_deployments(provider_url: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(
{
"model_name": name,
"litellm_params": {
"model": f"openai/gpt-{name}",
"api_base": provider_url + "/v1",
"api_key": "synthetic-provider-key",
},
}
for name in (PRIMARY, FALLBACK)
)
@pytest.fixture(scope="class")
def excluded(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
with series_cap_rig(
tmp_path_factory.mktemp("series-excluded"),
{"prometheus_metrics_max_series_per_metric": CAP, "prometheus_exclude_labels": ["api_key_alias"]},
workers=1,
warm_keys=0,
failing_models=frozenset({f"gpt-{PRIMARY}"}),
deployments=_fallback_deployments,
router_settings={"fallbacks": [{PRIMARY: [FALLBACK]}]},
) as rig:
yield rig
def _owned_fallback_series(samples: Sequence[Sample], name: str) -> tuple[Sample, ...]:
return tuple(sample for sample in samples if sample.name == name and not sample.is_overflow())
def _fallback_counter_settled(samples: Sequence[Sample], name: str) -> bool:
"""The counter has admitted the cap and sent the next key to `other`, or has handed out more series than the
cap, which is what a proxy without the cap does and what the caller's assertion then reports."""
owned: Final = _owned_fallback_series(samples, name)
return len(owned) > CAP or (len(owned) == CAP and overflow_total(samples, name) >= 1)
def _expect_capped_fallback_counter(rig: CapRig, name: str) -> None:
samples: Final = eventually(
lambda: scrape(rig.gateway), lambda seen: _fallback_counter_settled(seen, name), seconds=60
)
owned: Final = _owned_fallback_series(samples, name)
assert len(owned) == CAP, owned
assert overflow_total(samples, name) == 1, samples
assert all(sample.labels.get("fallback_model") == FALLBACK for sample in owned), owned
assert len({sample.labels["hashed_api_key"] for sample in owned}) == CAP, owned
assert all("api_key_alias" not in sample.labels for sample in samples if sample.name == name), samples
class TestExcluded:
def test_fallback_counters_are_capped_and_drop_excluded_labels(self, excluded: CapRig) -> None:
"""X1: the successful and failed fallback counters are capped like every other metric (their label names
reached the factory positionally before, so the cap never wrapped them) and keep honoring
prometheus_exclude_labels, which get_labels_for_metric already applied to them."""
keys: Final = tuple(excluded.key("x1") for _ in range(CAP + 1))
for key in keys:
call = Call.new()
response = chat_once(excluded.base_url, key, PRIMARY, call)
assert response.status_code == 200 and call.answer in response.text, response.text
_expect_capped_fallback_counter(excluded, SUCCESSFUL_FALLBACKS)
excluded.outage.set()
try:
for key in keys:
failed = chat_once(excluded.base_url, key, PRIMARY, Call.new())
assert failed.status_code == 500, failed.text
finally:
excluded.outage.clear()
_expect_capped_fallback_counter(excluded, FAILED_FALLBACKS)
@pytest.fixture(scope="class")
def nocap(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
with series_cap_rig(
tmp_path_factory.mktemp("series-nocap"),
{"prometheus_metrics_max_series_per_metric": None},
workers=2,
warm_keys=4,
) as rig:
yield rig
class TestNoCap:
def test_a_null_cap_keeps_every_series(self, nocap: CapRig) -> None:
"""N1: an explicit null cap and a missing TTL leave every key its own series and no `other` series."""
samples: Final = scrape(nocap.gateway)
assert alias_values(samples, REQUESTS) >= nocap.warm_aliases
assert not any(sample.is_overflow() for sample in samples)

View file

@ -0,0 +1,460 @@
"""Prometheus series cap under load: a concurrent burst across every endpoint while /metrics is scraped, a
provider outage between bursts, a worker killed mid-burst, and restarts that wipe or keep the multiprocess
directory."""
from __future__ import annotations
import json
import re
import signal
import threading
from collections.abc import Iterator, Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass
from itertools import cycle, product
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
import httpx
import psutil
import pytest
from integration._support.client import eventually, object_value, string_value
from integration._support.process import owned_gateway_image, setup_only_proxy_run
from integration._support.prometheus_series import (
AGENT_HEADERS,
PROXY_FAILURES,
PROXY_REQUESTS,
REQUESTS,
Call,
CapRig,
Key,
Sample,
SpendRow,
alias_total,
alias_values,
expect_spend_rows,
families_over,
label_values,
overflow_total,
scrape,
series_cap_config,
series_cap_rig,
spend_rows,
sse_data,
worker_samples,
)
from pydantic import JsonValue
pytestmark = pytest.mark.timeout(240)
CAP: Final = 3
CHAT: Final = "/v1/chat/completions"
MESSAGES: Final = "/v1/messages"
RESPONSES: Final = "/v1/responses"
ROUTES: Final = (CHAT, MESSAGES, RESPONSES)
EXTRA_KEYS: Final = 7
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
@dataclass(frozen=True, slots=True)
class Served:
key: Key
call: Call
route: str
streamed: bool
status: int
text: str
@property
def response_id(self) -> str:
assert self.status == 200, self.text
if not self.streamed:
return string_value(object_value(json.loads(self.text))["id"])
events: Final = sse_data(self.text)
match self.route:
case "/v1/messages":
starts: Final = tuple(event for event in events if event.get("type") == "message_start")
return string_value(object_value(starts[0]["message"])["id"])
case "/v1/responses":
completed: Final = tuple(event for event in events if event.get("type") == "response.completed")
return string_value(object_value(completed[0]["response"])["id"])
case _:
ids: Final = frozenset(string_value(event["id"]) for event in events)
assert len(ids) == 1, ids
return next(iter(ids))
def _body(route: str, model: str, call: Call, streamed: bool) -> dict[str, JsonValue]:
match route:
case "/v1/messages":
return {"model": model, "max_tokens": 64, "messages": [call.message], "stream": streamed}
case "/v1/responses":
return {"model": model, "input": call.text, "stream": streamed}
case _:
return {"model": model, "messages": [call.message], "stream": streamed}
def _send(rig: CapRig, key: Key, route: str, streamed: bool, tolerate_transport_errors: bool = False) -> Served:
call: Final = Call.new()
try:
with httpx.Client(base_url=rig.base_url, timeout=60, trust_env=False) as client:
response: Final = client.post(
route,
json=_body(route, rig.model, call, streamed),
headers={**AGENT_HEADERS, **call.headers, "Authorization": f"Bearer {key.token}"},
)
except httpx.TransportError as error:
if not tolerate_transport_errors:
raise
return Served(key, call, route, streamed, 0, repr(error))
return Served(key, call, route, streamed, response.status_code, response.text)
@dataclass(frozen=True, slots=True)
class Plan:
key: Key
route: str
streamed: bool
def _plans(keys: Sequence[Key]) -> tuple[Plan, ...]:
streaming: Final = cycle((False, True))
return tuple(Plan(key, route, next(streaming)) for key, route in product(keys, ROUTES))
def _burst(rig: CapRig, plans: Sequence[Plan], tolerate_transport_errors: bool = False) -> tuple[Served, ...]:
with ThreadPoolExecutor(max_workers=len(plans)) as pool:
return tuple(
pool.map(lambda plan: _send(rig, plan.key, plan.route, plan.streamed, tolerate_transport_errors), plans)
)
def _scrape_until(rig: CapRig, stop: threading.Event, sizes: SimpleQueue[int]) -> None:
while not stop.is_set():
try:
sizes.put(len(scrape(rig.gateway)))
except (AssertionError, httpx.HTTPError):
sizes.put(-1)
def _rows_by_alias(keys: Sequence[Key]) -> Mapping[str, tuple[SpendRow, ...]]:
return MappingProxyType({key.alias: spend_rows(key.alias) for key in keys})
def _landed(before: Sequence[Sample], after: Sequence[Sample], growth: Mapping[str, int]) -> bool:
"""Every counter the cell asserts on has counted its calls on `other`: the request and failure counters of
one call increment at different points of the logging callback, so a scrape between them is not the end
state."""
return all(overflow_total(after, name) - overflow_total(before, name) >= by for name, by in growth.items())
class TestBurst:
def test_concurrent_burst_across_every_endpoint_while_scraping(self, capped: CapRig) -> None:
"""C1: 30 concurrent calls from ten keys across chat, messages, and responses, streamed and not, with
/metrics scraped throughout: every call answers, the warm keys keep their series, every other call counts
on `other`, and every call writes one spend row."""
extra: Final = tuple(capped.key(f"c1-{index}") for index in range(EXTRA_KEYS))
keys: Final = (*capped.warm, *extra)
earlier: Final = _rows_by_alias(keys)
before: Final = scrape(capped.gateway)
stop: Final = threading.Event()
sizes: Final[SimpleQueue[int]] = SimpleQueue()
scraper: Final = threading.Thread(target=_scrape_until, args=(capped, stop, sizes))
scraper.start()
try:
served: Final = _burst(capped, _plans(keys))
finally:
stop.set()
scraper.join()
scrapes: Final = tuple(sizes.get_nowait() for _ in range(sizes.qsize()))
assert scrapes and all(count > 0 for count in scrapes), scrapes
assert all(item.status == 200 and item.call.answer in item.text for item in served), [
(item.route, item.status, item.text[:200]) for item in served if item.status != 200
]
extra_requests: Final = EXTRA_KEYS * len(ROUTES)
off_route: Final = len(capped.warm) * (len(ROUTES) - 1)
growth: Final = {REQUESTS: extra_requests, PROXY_REQUESTS: extra_requests + off_route}
samples: Final = eventually(
lambda: scrape(capped.gateway),
lambda after: (
_landed(before, after, growth) or any(key.alias in alias_values(after, REQUESTS) for key in extra)
),
seconds=90,
)
assert alias_values(samples, REQUESTS) == capped.warm_aliases
assert not families_over(samples, CAP), families_over(samples, CAP)
assert overflow_total(samples, REQUESTS) - overflow_total(before, REQUESTS) == extra_requests
assert alias_values(samples, PROXY_REQUESTS) == capped.warm_aliases
assert overflow_total(samples, PROXY_REQUESTS) - overflow_total(before, PROXY_REQUESTS) == (
extra_requests + off_route
)
for key in keys:
expect_spend_rows(
key.alias,
tuple(item.response_id for item in served if item.key == key),
earlier=earlier[key.alias],
)
def test_outage_between_bursts_counts_every_failure_once(self, capped: CapRig) -> None:
"""C2: a burst answers, the provider goes down for the next burst, and comes back for the last: the warm
keys keep their failure series, the fourth key's failures count on `other`, and every call writes one row."""
extra: Final = capped.key("c2")
keys: Final = (*capped.warm, extra)
earlier: Final = _rows_by_alias(keys)
before: Final = scrape(capped.gateway)
plans: Final = tuple(Plan(key, CHAT, streamed) for key, streamed in product(keys, (False, True)))
first: Final = _burst(capped, plans)
capped.outage.set()
try:
prefill: Final = tuple(_send(capped, warm, CHAT, False) for warm in capped.warm)
down: Final = _burst(capped, plans)
finally:
capped.outage.clear()
last: Final = _burst(capped, plans)
failed: Final = (*prefill, *down)
assert all(item.status == 200 for item in (*first, *last)), [item.status for item in (*first, *last)]
assert all(item.status == 500 for item in failed), [item.status for item in failed]
samples: Final = eventually(
lambda: scrape(capped.gateway),
lambda after: (
_landed(before, after, {PROXY_FAILURES: 2, REQUESTS: 4})
or extra.alias in alias_values(after, PROXY_FAILURES)
),
seconds=90,
)
assert alias_values(samples, PROXY_FAILURES) == capped.warm_aliases
assert not families_over(samples, CAP), families_over(samples, CAP)
assert overflow_total(samples, PROXY_FAILURES) - overflow_total(before, PROXY_FAILURES) == 2
assert overflow_total(samples, REQUESTS) - overflow_total(before, REQUESTS) == 4
for key in keys:
expect_spend_rows(
key.alias,
tuple(item.response_id for item in (*first, *last) if item.key == key),
tuple(item.call.call_id for item in failed if item.key == key),
earlier=earlier[key.alias],
)
def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]:
text: Final = log.read_text()
return tuple(int(pid) for pid in _STARTED_WORKER.findall(text)), text.count("Application startup complete.")
@pytest.mark.timeout(420)
def test_killed_worker_is_replaced_by_one_that_reads_the_same_admissions(tmp_path: Path) -> None:
"""C3: SIGKILL one of two workers mid-burst: the sibling keeps answering, and the replacement worker puts a
fourth key on `other` because the admitted series live in the shared directory, not in the dead process."""
with series_cap_rig(tmp_path, {"prometheus_metrics_max_series_per_metric": CAP}, workers=2, warm_keys=3) as rig:
workers, _ = eventually(
lambda: _worker_startups(rig.proxy.log), lambda found: len(found[0]) == 2 and found[1] == 2, seconds=120
)
extra: Final = tuple(rig.key(f"c3-{index}") for index in range(4))
plans: Final = _plans(extra)
with ThreadPoolExecutor(max_workers=1) as pool:
burst: Final = pool.submit(_burst, rig, plans, True)
victim: Final = psutil.Process(workers[0])
victim.suspend()
victim.send_signal(signal.SIGKILL)
served: Final = burst.result()
answered: Final = tuple(item for item in served if item.status == 200)
assert answered, [(item.status, item.text[:200]) for item in served]
assert all(item.call.answer in item.text for item in answered)
replacement: Final = eventually(
lambda: _worker_startups(rig.proxy.log),
lambda found: len(frozenset(found[0]) - frozenset(workers)) == 1,
seconds=120,
)
(new_pid,) = frozenset(replacement[0]) - frozenset(workers)
late: Final = rig.key("c3-late")
def send_until_the_replacement_counts() -> tuple[Sample, ...]:
assert rig.chat(late, Call.new()).status_code == 200
return scrape(rig.gateway)
samples: Final = eventually(
send_until_the_replacement_counts,
lambda after: (
any(
sample.pid == new_pid and (sample.overflow > 0 or late.alias in sample.aliases)
for sample in worker_samples(rig.prom_dir, REQUESTS)
)
or late.alias in alias_values(after, REQUESTS)
),
seconds=90,
)
assert late.alias not in label_values(samples)
by_pid: Final = {sample.pid: sample for sample in worker_samples(rig.prom_dir, REQUESTS)}
assert by_pid[new_pid].overflow > 0 and by_pid[new_pid].aliases <= rig.warm_aliases, by_pid[new_pid]
@pytest.mark.timeout(420)
def test_restart_with_two_workers_starts_the_cap_over(tmp_path: Path) -> None:
"""C4: a second boot on the same multiprocess directory wipes it: the old keys are gone, three new keys get
their series, and a fourth lands on `other`."""
shared_dir: Final = tmp_path / "prom-shared"
settings: Final = {"prometheus_metrics_max_series_per_metric": CAP}
with series_cap_rig(tmp_path, settings, workers=2, warm_keys=3, multiproc_dir=shared_dir) as first_boot:
old_aliases: Final = first_boot.warm_aliases
assert alias_values(scrape(first_boot.gateway), REQUESTS) == old_aliases
with series_cap_rig(tmp_path, settings, workers=2, warm_keys=3, multiproc_dir=shared_dir) as second_boot:
samples: Final = scrape(second_boot.gateway)
assert alias_values(samples, REQUESTS) == second_boot.warm_aliases
assert not old_aliases & label_values(samples)
extra: Final = second_boot.key("c4")
before: Final = scrape(second_boot.gateway)
assert second_boot.chat(extra, Call.new()).status_code == 200
after: Final = eventually(
lambda: scrape(second_boot.gateway),
lambda now: (
overflow_total(now, REQUESTS) - overflow_total(before, REQUESTS) >= 1
or extra.alias in alias_values(now, REQUESTS)
),
seconds=60,
)
assert extra.alias not in label_values(after)
@pytest.mark.timeout(420)
def test_restart_with_one_worker_and_an_operator_directory_starts_the_cap_over(tmp_path: Path) -> None:
"""C5: one worker, no metrics port, PROMETHEUS_MULTIPROC_DIR set by the operator and kept across a restart:
the directory is wiped at boot the way the multi-worker path wipes it, so the merged scrape shows only the
second boot's three keys and a fourth lands on `other`."""
operator_dir: Final = tmp_path / "prom-operator"
settings: Final = {"prometheus_metrics_max_series_per_metric": CAP}
with series_cap_rig(tmp_path, settings, workers=1, warm_keys=3, multiproc_dir=operator_dir) as first_boot:
old_aliases: Final = first_boot.warm_aliases
assert alias_values(scrape(first_boot.gateway), REQUESTS) == old_aliases
with series_cap_rig(tmp_path, settings, workers=1, warm_keys=3, multiproc_dir=operator_dir) as second_boot:
samples: Final = scrape(second_boot.gateway)
assert alias_values(samples, REQUESTS) == second_boot.warm_aliases
assert not old_aliases & label_values(samples)
extra: Final = second_boot.key("c5")
before: Final = scrape(second_boot.gateway)
assert second_boot.chat(extra, Call.new()).status_code == 200
after: Final = eventually(
lambda: scrape(second_boot.gateway),
lambda now: (
overflow_total(now, REQUESTS) - overflow_total(before, REQUESTS) >= 1
or extra.alias in alias_values(now, REQUESTS)
),
seconds=60,
)
assert extra.alias not in label_values(after)
def test_setup_only_run_leaves_a_live_proxy_samples_alone(tmp_path: Path) -> None:
"""P1: a `--skip_server_startup` run of the proxy CLI (the image's setup step) pointed at a live two-worker
proxy's operator-set `PROMETHEUS_MULTIPROC_DIR` leaves the live samples alone: the warm keys keep their series
and their totals, and a fourth key still lands on `other`."""
operator_dir: Final = tmp_path / "prom-operator"
settings: Final = {"prometheus_metrics_max_series_per_metric": CAP}
with series_cap_rig(tmp_path, settings, workers=2, warm_keys=3, multiproc_dir=operator_dir) as rig:
before: Final = scrape(rig.gateway)
assert alias_values(before, REQUESTS) == rig.warm_aliases
completed: Final = setup_only_proxy_run(
rig.gateway,
{"PROMETHEUS_MULTIPROC_DIR": str(operator_dir)},
config=series_cap_config(tmp_path, settings),
workers=2,
)
assert completed.returncode == 0, completed.stdout[-2000:] + completed.stderr[-2000:]
assert "Skipping server startup" in completed.stdout, completed.stdout[-2000:]
after_setup: Final = scrape(rig.gateway)
assert alias_values(after_setup, REQUESTS) == rig.warm_aliases
assert all(
alias_total(after_setup, REQUESTS, alias) == alias_total(before, REQUESTS, alias)
for alias in rig.warm_aliases
)
extra: Final = rig.key("p1")
assert rig.chat(extra, Call.new()).status_code == 200
after: Final = eventually(
lambda: scrape(rig.gateway),
lambda now: (
overflow_total(now, REQUESTS) - overflow_total(after_setup, REQUESTS) >= 1
or extra.alias in alias_values(now, REQUESTS)
),
seconds=60,
)
assert extra.alias not in label_values(after)
assert alias_values(after, REQUESTS) == rig.warm_aliases
IMAGE_MODEL: Final = "series-cap-image"
def _image_deployment(provider_url: str) -> tuple[dict[str, JsonValue], ...]:
return (
{
"model_name": IMAGE_MODEL,
"litellm_params": {
"model": f"openai/gpt-{IMAGE_MODEL}",
"api_base": provider_url + "/v1",
"api_key": "synthetic-provider-key",
},
},
)
@contextmanager
def _gateway_image(control: CapRig, config: Path, prom_dir: Path) -> Iterator[CapRig]:
"""One container life of the gateway image on `prom_dir`: two workers serving the keys the control plane proxy
mints in the database both read."""
with owned_gateway_image(
control.gateway, config.parent, {"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)}, config=config, workers=2
) as image:
yield CapRig(image, control.scenario, IMAGE_MODEL, control.provider, control.outage, (), prom_dir)
def _fill_the_cap(image: CapRig, cell: str) -> frozenset[str]:
"""Three new keys call once each on a boot that has counted nothing yet, and each gets its own series."""
keys: Final = tuple(image.key(cell) for _ in range(CAP))
assert all(image.chat(key, Call.new()).status_code == 200 for key in keys)
aliases: Final = frozenset(key.alias for key in keys)
samples: Final = eventually(
lambda: scrape(image.gateway),
lambda now: sum(sample.value for sample in now if sample.name == REQUESTS) >= CAP,
seconds=60,
)
assert alias_values(samples, REQUESTS) == aliases, (alias_values(samples, REQUESTS), aliases)
return aliases
@pytest.mark.timeout(420)
def test_gateway_image_restart_on_a_kept_directory_starts_the_cap_over(tmp_path: Path) -> None:
"""D1: the gateway image's launcher (`docker/component_entrypoint.sh` running `python -m gateway.launch`, two
workers) restarted on a kept PROMETHEUS_MULTIPROC_DIR: the entrypoint removes the previous container's samples
and admitted series before the workers fork, so the second boot shows only its own three keys and a fourth
lands on `other`."""
control_dir: Final = tmp_path / "control"
image_dir: Final = tmp_path / "image"
control_dir.mkdir()
image_dir.mkdir()
prom_dir: Final = tmp_path / "prom-image"
with series_cap_rig(control_dir, {}, workers=1, warm_keys=0) as control:
config: Final = series_cap_config(
image_dir,
{"prometheus_metrics_max_series_per_metric": CAP},
model_list=_image_deployment(control.provider.url),
)
with _gateway_image(control, config, prom_dir) as first_boot:
old_aliases: Final = _fill_the_cap(first_boot, "d1-old")
with _gateway_image(control, config, prom_dir) as second_boot:
assert not old_aliases & label_values(scrape(second_boot.gateway))
new_aliases: Final = _fill_the_cap(second_boot, "d1-new")
extra: Final = second_boot.key("d1")
before: Final = scrape(second_boot.gateway)
assert second_boot.chat(extra, Call.new()).status_code == 200
after: Final = eventually(
lambda: scrape(second_boot.gateway),
lambda now: (
overflow_total(now, REQUESTS) - overflow_total(before, REQUESTS) >= 1
or extra.alias in alias_values(now, REQUESTS)
),
seconds=60,
)
assert extra.alias not in label_values(after)
assert alias_values(after, REQUESTS) == new_aliases
assert not old_aliases & label_values(after)

View file

@ -689,9 +689,8 @@ def test_exclude_only_hardcoded_label_drops_all_labels(reset_prometheus_exclude_
def test_exclude_labels_does_not_touch_unrelated_metrics(reset_prometheus_exclude_settings):
"""A metric that never declares the excluded label is left as a plain prometheus metric,
not wrapped, so no behavior changes for it."""
from litellm.integrations.prometheus import _ExcludedLabelMetric
"""A metric that never declares the excluded label keeps every one of its own labels in the scrape."""
from prometheus_client import generate_latest
clear_prometheus_registry()
litellm.prometheus_metrics_config = None
@ -699,10 +698,15 @@ def test_exclude_labels_does_not_touch_unrelated_metrics(reset_prometheus_exclud
litellm.prometheus_exclude_labels = ["guardrail_name"]
logger = PrometheusLogger()
spend_labels = {name: f"{name}-value" for name in PrometheusMetricLabels.get_labels("litellm_spend_metric")}
logger.litellm_spend_metric.labels(**spend_labels).inc(1.5)
logger.litellm_provider_remaining_budget_metric.labels("anthropic").set(5.0)
assert not isinstance(logger.litellm_spend_metric, _ExcludedLabelMetric)
assert not isinstance(logger.litellm_provider_remaining_budget_metric, _ExcludedLabelMetric)
assert isinstance(logger.litellm_guardrail_latency_metric, _ExcludedLabelMetric)
scrape = generate_latest(REGISTRY).decode()
spend_line = next(line for line in scrape.splitlines() if line.startswith("litellm_spend_metric_total{"))
assert all(f'{name}="{value}"' in spend_line for name, value in spend_labels.items())
assert spend_line.endswith(" 1.5")
assert 'litellm_provider_remaining_budget_metric{api_provider="anthropic"} 5.0' in scrape
# ==============================================================================

View file

@ -0,0 +1,442 @@
import logging
import re
from pathlib import Path
from threading import Thread
from typing import Final
import pytest
from prometheus_client import REGISTRY, CollectorRegistry, Counter, generate_latest
import litellm
from litellm.constants import PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX
from litellm.integrations.prometheus import PrometheusLogger, _LabeledMetric, prometheus_label_factory
from litellm.integrations.prometheus_helpers import bounded_prometheus_series_tracker
from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import (
BoundedPrometheusSeriesTracker,
PrometheusSeriesLimits,
)
from litellm.integrations.prometheus_helpers.shared_prometheus_series_admissions import (
SharedPrometheusSeriesAdmissions,
)
from litellm.proxy.prometheus_cleanup import wipe_directory
from litellm.types.integrations.prometheus import UserAPIKeyLabelValues
SERIES_SETTINGS: Final = (
"prometheus_metrics_max_series_per_metric",
"prometheus_metrics_ttl_seconds",
"prometheus_metrics_cleanup_interval_seconds",
"prometheus_exclude_labels",
"prometheus_metrics_config",
"enable_end_user_cost_tracking_prometheus_only",
"prometheus_end_user_metrics_max_series_per_metric",
"prometheus_end_user_metrics_ttl_seconds",
)
def _unregister_everything() -> None:
for collector in list(REGISTRY._collector_to_names):
try:
REGISTRY.unregister(collector)
except Exception:
pass
@pytest.fixture(autouse=True)
def isolated_registry_and_settings(monkeypatch):
collectors_before: Final = tuple(REGISTRY._collector_to_names)
_unregister_everything()
monkeypatch.delenv("PROMETHEUS_MULTIPROC_DIR", raising=False)
for setting in SERIES_SETTINGS:
monkeypatch.setattr(litellm, setting, getattr(litellm, setting))
yield
_unregister_everything()
for collector in collectors_before:
REGISTRY.register(collector)
@pytest.fixture
def clock(monkeypatch):
now: Final = [1_000.0]
monkeypatch.setattr(bounded_prometheus_series_tracker.time, "monotonic", lambda: now[0])
return now
def _scraped_series(sample_name: str, registry: CollectorRegistry = REGISTRY) -> frozenset[str]:
exposition: Final = generate_latest(registry).decode()
return frozenset(line for line in exposition.splitlines() if line.startswith(f"{sample_name}{{"))
def _label_values(series: frozenset[str], label: str) -> frozenset[str]:
pattern: Final = re.compile(rf'[{{,]{label}="([^"]*)"')
return frozenset(match.group(1) for match in map(pattern.search, series) if match is not None)
def _sample_value(series: frozenset[str], label: str, value: str) -> float:
(line,) = (line for line in series if f'{label}="{value}"' in line)
return float(line.rsplit(" ", 1)[1])
def _count_request(logger: PrometheusLogger, user_agent: str) -> None:
PrometheusLogger._inc_labeled_counter(
logger,
logger.litellm_proxy_total_requests_metric,
"litellm_proxy_total_requests_metric",
UserAPIKeyLabelValues(user_agent=user_agent),
)
def _observe_latency(logger: PrometheusLogger, user: str) -> None:
labels: Final = prometheus_label_factory(
supported_enum_labels=logger.get_labels_for_metric("litellm_request_total_latency_metric"),
enum_values=UserAPIKeyLabelValues(user=user),
)
logger.litellm_request_total_latency_metric.labels(**labels).observe(0.5)
def test_label_sets_past_the_cap_are_counted_on_one_other_series():
litellm.prometheus_metrics_max_series_per_metric = 2
litellm.prometheus_metrics_ttl_seconds = None
logger: Final = PrometheusLogger()
for identity in ("one", "two", "three", "one", "four"):
_count_request(logger, f"codex/{identity}")
_observe_latency(logger, f"user-{identity}")
counter_series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(counter_series, "user_agent") == {"codex/one", "codex/two", "other"}
assert _sample_value(counter_series, "user_agent", "codex/one") == 2
assert _sample_value(counter_series, "user_agent", "codex/two") == 1
assert _sample_value(counter_series, "user_agent", "other") == 2
histogram_series: Final = _scraped_series("litellm_request_total_latency_metric_count")
assert _label_values(histogram_series, "user") == {"user-one", "user-two", "other"}
assert _sample_value(histogram_series, "user", "other") == 2
def test_gauge_label_sets_past_the_cap_are_not_emitted():
litellm.prometheus_metrics_max_series_per_metric = 2
litellm.prometheus_metrics_ttl_seconds = None
logger: Final = PrometheusLogger()
for provider in ("openai", "anthropic", "bedrock", "openai"):
logger.track_provider_remaining_budget(provider=provider, spend=1.0, budget_limit=10.0)
series: Final = _scraped_series("litellm_provider_remaining_budget_metric")
assert _label_values(series, "api_provider") == {"openai", "anthropic"}
def test_series_idle_past_the_ttl_are_removed_and_free_their_slot(clock):
litellm.prometheus_metrics_max_series_per_metric = 2
litellm.prometheus_metrics_ttl_seconds = 10.0
litellm.prometheus_metrics_cleanup_interval_seconds = 0.0
logger: Final = PrometheusLogger()
_count_request(logger, "idle-agent")
clock[0] += 9.0
_count_request(logger, "still-fresh-agent")
clock[0] += 2.0
_count_request(logger, "new-agent")
series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(series, "user_agent") == {"still-fresh-agent", "new-agent"}
def test_cap_holds_and_ttl_is_ignored_in_multiprocess_mode(monkeypatch, tmp_path, clock):
monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path))
litellm.prometheus_metrics_max_series_per_metric = 2
litellm.prometheus_metrics_ttl_seconds = 10.0
litellm.prometheus_metrics_cleanup_interval_seconds = 0.0
logger: Final = PrometheusLogger()
_count_request(logger, "first-agent")
_count_request(logger, "second-agent")
clock[0] += 11.0
_count_request(logger, "third-agent")
series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(series, "user_agent") == {"first-agent", "second-agent", "other"}
def test_workers_sharing_a_multiprocess_dir_admit_the_same_label_sets(tmp_path: Path):
first_worker: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
second_worker: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
assert first_worker.admit_series("litellm_requests_metric", ("user-a",), max_series=2)
assert second_worker.admit_series("litellm_requests_metric", ("user-b",), max_series=2)
replacement_worker: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
for worker in (first_worker, second_worker, replacement_worker):
assert worker.admit_series("litellm_requests_metric", ("user-a",), max_series=2)
assert worker.admit_series("litellm_requests_metric", ("user-b",), max_series=2)
assert not worker.admit_series("litellm_requests_metric", ("user-c",), max_series=2)
assert first_worker.admit_series("litellm_spend_metric", ("user-c",), max_series=2)
def test_workers_agree_when_racing_appends_overfill_the_admissions_file(tmp_path: Path):
racing_workers: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
for user in ("user-a", "user-b", "user-c"):
assert racing_workers.admit_series("litellm_requests_metric", (user,), max_series=3)
worker: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
assert not worker.admit_series("litellm_requests_metric", ("user-c",), max_series=2)
assert worker.admit_series("litellm_requests_metric", ("user-a",), max_series=2)
assert worker.admit_series("litellm_requests_metric", ("user-b",), max_series=2)
def test_a_line_another_worker_is_still_writing_is_read_once_it_is_complete(tmp_path: Path):
admissions_file: Final = tmp_path / f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}litellm_requests_metric"
assert SharedPrometheusSeriesAdmissions(directory=str(tmp_path)).admit_series(
"litellm_requests_metric", ("user-a",), max_series=2
)
reader: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
with admissions_file.open("ab") as write_in_progress:
write_in_progress.write(b'["user')
write_in_progress.flush()
assert reader.admit_series("litellm_requests_metric", ("user-a",), max_series=2)
write_in_progress.write(b'-b"]\n')
assert not reader.admit_series("litellm_requests_metric", ("user-c",), max_series=2)
assert reader.admit_series("litellm_requests_metric", ("user-b",), max_series=2)
def test_a_record_cut_short_by_a_full_disk_admits_nothing_and_hides_no_other_record(tmp_path: Path):
admissions_file: Final = tmp_path / f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}litellm_requests_metric"
admissions_file.write_bytes(b'\n["user-a')
writer: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
reader: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
assert writer.admit_series("litellm_requests_metric", ("user-b",), max_series=2)
assert reader.admit_series("litellm_requests_metric", ("user-b",), max_series=2)
assert reader.admit_series("litellm_requests_metric", ("user-c",), max_series=2)
assert writer.admit_series("litellm_requests_metric", ("user-c",), max_series=2)
assert not writer.admit_series("litellm_requests_metric", ("user-a",), max_series=2)
assert not reader.admit_series("litellm_requests_metric", ("user-a",), max_series=2)
def test_wiping_the_multiprocess_dir_frees_every_admitted_slot(tmp_path: Path):
before_restart: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
assert before_restart.admit_series("litellm_requests_metric", ("user-a",), max_series=1)
wipe_directory(str(tmp_path))
after_restart: Final = SharedPrometheusSeriesAdmissions(directory=str(tmp_path))
assert after_restart.admit_series("litellm_requests_metric", ("user-b",), max_series=1)
assert not after_restart.admit_series("litellm_requests_metric", ("user-a",), max_series=1)
def test_eviction_racing_a_new_series_cannot_leave_it_untracked():
registry: Final = CollectorRegistry()
counter: Final = Counter("requests", "requests", labelnames=("user",), registry=registry)
class _EvictedWhileBeingCreated:
def labels(self, *labelvalues: str):
if labelvalues == ("evicted-user",):
eviction.start()
eviction.join(timeout=0.05)
return counter.labels(*labelvalues)
def remove(self, *labelvalues: str) -> None:
counter.remove(*labelvalues)
labeled: Final = _LabeledMetric(
metric=_EvictedWhileBeingCreated(),
metric_name="requests",
original_labelnames=("user",),
excluded_labels=frozenset(),
tracker=BoundedPrometheusSeriesTracker(),
limits=PrometheusSeriesLimits(max_series=1, ttl_seconds=None, cleanup_interval_seconds=None),
shares_overflow_series=True,
)
eviction: Final = Thread(target=labeled.remove, args=("evicted-user",))
labeled.labels("evicted-user").inc()
eviction.join()
labeled.labels("next-user").inc()
assert _label_values(_scraped_series("requests_total", registry), "user") == {"next-user"}
@pytest.mark.parametrize(
"metric_name", ["litellm_deployment_successful_fallbacks", "litellm_deployment_failed_fallbacks"]
)
def test_cap_applies_to_the_fallback_counters(metric_name: str):
litellm.prometheus_metrics_max_series_per_metric = 2
litellm.prometheus_metrics_ttl_seconds = None
logger: Final = PrometheusLogger()
for index in range(4):
PrometheusLogger._inc_labeled_counter(
logger,
getattr(logger, metric_name),
metric_name,
UserAPIKeyLabelValues(fallback_model=f"model-{index}"),
)
series: Final = _scraped_series(f"{metric_name}_total")
assert _label_values(series, "fallback_model") == {"model-0", "model-1", "other"}
assert _sample_value(series, "fallback_model", "other") == 2
def test_end_user_eviction_keeps_the_series_and_its_slot_in_multiprocess_mode(monkeypatch, tmp_path):
monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path))
litellm.enable_end_user_cost_tracking_prometheus_only = True
litellm.prometheus_metrics_config = [
{"group": "end-user-spend", "metrics": ["litellm_spend_metric"], "include_labels": ["end_user"]}
]
litellm.prometheus_end_user_metrics_max_series_per_metric = 2
litellm.prometheus_end_user_metrics_ttl_seconds = None
litellm.prometheus_metrics_max_series_per_metric = 3
litellm.prometheus_metrics_ttl_seconds = None
logger: Final = PrometheusLogger()
for index in range(5):
PrometheusLogger._inc_labeled_counter(
logger,
logger.litellm_spend_metric,
"litellm_spend_metric",
UserAPIKeyLabelValues(end_user=f"end-user-{index}"),
amount=0.01,
)
series: Final = _scraped_series("litellm_spend_metric_total")
assert _label_values(series, "end_user") == {"end-user-0", "end-user-1", "end-user-2", "other"}
def test_cap_applies_under_a_globally_excluded_label():
litellm.prometheus_exclude_labels = ["hook_type"]
litellm.prometheus_metrics_max_series_per_metric = 2
litellm.prometheus_metrics_ttl_seconds = None
logger: Final = PrometheusLogger()
for index in range(4):
logger._record_guardrail_metrics(
guardrail_name=f"guardrail-{index}",
latency_seconds=0.1,
status="success",
error_type=None,
hook_type="pre_call",
)
series: Final = _scraped_series("litellm_guardrail_requests_total")
assert _label_values(series, "guardrail_name") == {"guardrail-0", "guardrail-1", "other"}
assert _sample_value(series, "guardrail_name", "other") == 2
assert all("hook_type" not in line for line in series)
def test_end_user_eviction_frees_a_slot_under_the_cap():
litellm.enable_end_user_cost_tracking_prometheus_only = True
litellm.prometheus_metrics_config = [
{"group": "end-user-spend", "metrics": ["litellm_spend_metric"], "include_labels": ["end_user"]}
]
litellm.prometheus_end_user_metrics_max_series_per_metric = 2
litellm.prometheus_end_user_metrics_ttl_seconds = None
litellm.prometheus_metrics_max_series_per_metric = 3
litellm.prometheus_metrics_ttl_seconds = None
logger: Final = PrometheusLogger()
for index in range(5):
PrometheusLogger._inc_labeled_counter(
logger,
logger.litellm_spend_metric,
"litellm_spend_metric",
UserAPIKeyLabelValues(end_user=f"end-user-{index}"),
amount=0.01,
)
series: Final = _scraped_series("litellm_spend_metric_total")
assert _label_values(series, "end_user") == {"end-user-3", "end-user-4"}
def test_series_stay_unbounded_unless_a_limit_is_configured():
litellm.prometheus_metrics_max_series_per_metric = None
litellm.prometheus_metrics_ttl_seconds = None
logger: Final = PrometheusLogger()
for index in range(5):
_count_request(logger, f"agent-{index}")
series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(series, "user_agent") == {f"agent-{index}" for index in range(5)}
@pytest.mark.parametrize(
("setting", "value"),
[
("prometheus_metrics_max_series_per_metric", 0),
("prometheus_metrics_max_series_per_metric", -5),
("prometheus_metrics_ttl_seconds", 0.0),
("prometheus_metrics_ttl_seconds", -1.0),
("prometheus_metrics_max_series_per_metric", "five"),
("prometheus_metrics_max_series_per_metric", True),
("prometheus_metrics_max_series_per_metric", 2.5),
("prometheus_metrics_ttl_seconds", ""),
],
)
def test_a_series_limit_that_is_not_a_positive_number_is_ignored_with_a_warning_and_metrics_keep_flowing(
setting: str, value: object, clock, caplog
):
litellm.prometheus_metrics_cleanup_interval_seconds = 0.0
setattr(litellm, setting, value)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logger: Final = PrometheusLogger()
for index in range(3):
_count_request(logger, f"agent-{index}")
clock[0] += 100.0
series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(series, "user_agent") == {"agent-0", "agent-1", "agent-2"}
assert setting in caplog.text
@pytest.mark.parametrize("value", ["sixty", -1, True, ""])
def test_a_cleanup_interval_that_is_not_a_number_of_at_least_zero_falls_back_to_the_default_with_a_warning(
value: object, monkeypatch, clock, caplog
):
monkeypatch.setattr(litellm, "prometheus_metrics_max_series_per_metric", 3)
monkeypatch.setattr(litellm, "prometheus_metrics_ttl_seconds", 10.0)
monkeypatch.setattr(litellm, "prometheus_metrics_cleanup_interval_seconds", value)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logger: Final = PrometheusLogger()
_count_request(logger, "agent-0")
clock[0] += 30.0
_count_request(logger, "agent-1")
within_the_interval: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
clock[0] += 31.0
_count_request(logger, "agent-2")
assert "prometheus_metrics_cleanup_interval_seconds" in caplog.text
assert _label_values(within_the_interval, "user_agent") == {"agent-0", "agent-1"}
series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(series, "user_agent") == {"agent-2"}
def test_a_cleanup_interval_written_as_a_numeric_string_is_honored(monkeypatch, clock, caplog):
monkeypatch.setattr(litellm, "prometheus_metrics_max_series_per_metric", 3)
monkeypatch.setattr(litellm, "prometheus_metrics_ttl_seconds", 10.0)
monkeypatch.setattr(litellm, "prometheus_metrics_cleanup_interval_seconds", "0")
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logger: Final = PrometheusLogger()
_count_request(logger, "agent-0")
clock[0] += 30.0
_count_request(logger, "agent-1")
series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(series, "user_agent") == {"agent-1"}
assert "prometheus_metrics_cleanup_interval_seconds" not in caplog.text
def test_a_series_cap_written_as_a_numeric_string_is_honored(caplog):
litellm.prometheus_metrics_max_series_per_metric = "2"
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logger: Final = PrometheusLogger()
for index in range(3):
_count_request(logger, f"agent-{index}")
series: Final = _scraped_series("litellm_proxy_total_requests_metric_total")
assert _label_values(series, "user_agent") == {"agent-0", "agent-1", "other"}
assert "prometheus_metrics_max_series_per_metric" not in caplog.text

View file

@ -15,6 +15,7 @@ from unittest.mock import patch
import pytest
from prometheus_client import CollectorRegistry, multiprocess
from litellm.constants import PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX
from litellm.proxy.prometheus_cleanup import mark_dead_workers, mark_worker_exit, wipe_directory
from litellm.proxy.proxy_cli import ProxyInitializationHelpers
@ -235,3 +236,24 @@ class TestMaybeSetupPrometheusMultiprocDir:
assert result_dir == str(tmp_path)
assert os.environ["PROMETHEUS_MULTIPROC_DIR"] == str(tmp_path)
def test_single_worker_restart_with_an_operator_set_dir_wipes_it(self, tmp_path: Path) -> None:
"""One worker and no metrics server still wipe the operator's directory at boot: the docs promise a
restart frees every capped slot, and the exited worker's samples would otherwise keep the merged scrape
past the cap."""
admitted: Final = tmp_path / f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}litellm_requests_metric"
admitted.write_text('\n["user-a"]\n')
samples: Final = tmp_path / "counter_123.db"
samples.write_bytes(b"operator-owned samples")
with patch.dict(os.environ, {"PROMETHEUS_MULTIPROC_DIR": str(tmp_path)}, clear=False):
os.environ.pop("prometheus_multiproc_dir", None)
result_dir = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir(
num_workers=1,
litellm_settings={"callbacks": ["prometheus"]},
)
assert result_dir is None
assert os.environ["PROMETHEUS_MULTIPROC_DIR"] == str(tmp_path)
assert not admitted.exists()
assert not samples.exists()

View file

@ -566,7 +566,7 @@ class TestProxyInitializationHelpers:
"litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False
)
def test_skip_server_startup(
self, mock_should_update, mock_setup_db, mock_atexit_register, mock_uvicorn_run
self, mock_should_update, mock_setup_db, mock_atexit_register, mock_uvicorn_run, tmp_path: Path
):
from click.testing import CliRunner
@ -587,6 +587,9 @@ class TestProxyInitializationHelpers:
for k, v in os.environ.items()
if k not in ("DATABASE_URL", "DIRECT_URL")
}
clean_env["PROMETHEUS_MULTIPROC_DIR"] = str(tmp_path)
live_proxy_samples = tmp_path / "counter_123.db"
live_proxy_samples.write_bytes(b"samples of a proxy that is still running")
with (
patch.dict(
os.environ,
@ -630,6 +633,7 @@ class TestProxyInitializationHelpers:
), f"exit_code={result.exit_code}, output={result.output}"
assert "Skipping server startup" in result.output
assert "telemetry" not in runner.invoke(run_server, ["--help"]).output
assert live_proxy_samples.exists()
# --- normal startup ---
mock_uvicorn_run.reset_mock()
@ -640,6 +644,7 @@ class TestProxyInitializationHelpers:
result.exit_code == 0
), f"exit_code={result.exit_code}, output={result.output}"
mock_uvicorn_run.assert_called_once()
assert not live_proxy_samples.exists()
@patch("uvicorn.run")
@patch("atexit.register")

View file

@ -10,6 +10,8 @@ from pathlib import Path
import pytest
from litellm.constants import PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX
REPO_ROOT = Path(__file__).resolve().parents[2]
COMPONENT_ENTRYPOINT = REPO_ROOT / "docker" / "component_entrypoint.sh"
PROD_ENTRYPOINT = REPO_ROOT / "docker" / "prod_entrypoint.sh"
@ -229,11 +231,14 @@ def test_gating_matches_the_monolithic_entrypoint_and_get_secret_bool(
def test_wipes_the_prometheus_multiproc_dir_before_uvicorn_forks(tmp_path: Path) -> None:
"""A restarted container inherits the emptyDir of its predecessor, whose worker pids it may reuse, so the
stale .db files must be gone before any worker opens the one carrying its own pid."""
stale .db files must be gone before any worker opens the one carrying its own pid. The admitted-series
files go with them, or the restarted workers would keep counting new label sets on `other` for the
label sets the previous container admitted."""
multiproc_dir = tmp_path / "multiproc"
multiproc_dir.mkdir()
(multiproc_dir / "gauge_livesum_7.db").write_bytes(b"stale")
(multiproc_dir / "counter_7.db").write_bytes(b"stale")
(multiproc_dir / f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}litellm_requests_metric").write_bytes(b"stale")
(multiproc_dir / "keep.txt").write_text("not a sample")
bin_dir = tmp_path / "bin"