This commit is contained in:
devin-ai-integration[bot] 2026-10-04 23:09:08 +08:00 • committed by GitHub
commit bcf48a2204
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1931 additions and 52 deletions

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

@ -1546,6 +1546,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,114 @@ 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)])
def _positive_number(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 = _positive_number(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 _get_budget_metrics_per_request_timeout() -> float:
@ -301,10 +394,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 +794,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 +1282,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=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,39 @@ 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 before workers fork."""
_remove(directory, (*glob.glob(os.path.join(directory, "*.db")), *_admitted_series_files(directory)))
def wipe_admitted_series(directory: str) -> None:
"""Drop only litellm's own admitted-series files, so a restart that keeps an operator-managed directory
(one worker, no separate metrics server) still starts the series cap from an empty set."""
_remove(directory, _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
from litellm.proxy.prometheus_cleanup import wipe_admitted_series, 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_admitted_series(configured_dir)
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")
multiproc_dir: Final = configured_dir or os.path.join(tempfile.gettempdir(), "litellm_prometheus_multiproc")
os.environ["PROMETHEUS_MULTIPROC_DIR"] = multiproc_dir

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,453 @@
"""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 pathlib import Path
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
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
class TestExcluded:
def test_fallback_counters_drop_excluded_labels(self, excluded: CapRig) -> None:
"""X1: the successful and failed fallback counters honor prometheus_exclude_labels like every other metric."""
key: Final = excluded.key("x1")
call: Final = Call.new()
response: Final = chat_once(excluded.base_url, key, PRIMARY, call)
assert response.status_code == 200 and call.answer in response.text, response.text
after_success: Final = eventually(
lambda: scrape(excluded.gateway),
lambda samples: any(sample.name == SUCCESSFUL_FALLBACKS for sample in samples),
seconds=60,
)
successes: Final = tuple(sample for sample in after_success if sample.name == SUCCESSFUL_FALLBACKS)
assert any(sample.labels.get("fallback_model") == FALLBACK for sample in successes), successes
assert all("api_key_alias" not in sample.labels for sample in successes), successes
excluded.outage.set()
try:
failed: Final = chat_once(excluded.base_url, key, PRIMARY, Call.new())
finally:
excluded.outage.clear()
assert failed.status_code == 500, failed.text
after_failure: Final = eventually(
lambda: scrape(excluded.gateway),
lambda samples: any(sample.name == FAILED_FALLBACKS for sample in samples),
seconds=60,
)
failures: Final = tuple(sample for sample in after_failure if sample.name == FAILED_FALLBACKS)
assert any(sample.labels.get("fallback_model") == FALLBACK for sample in failures), failures
assert all("api_key_alias" not in sample.labels for sample in failures), failures
@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,332 @@
"""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 Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
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.prometheus_series import (
AGENT_HEADERS,
PROXY_FAILURES,
PROXY_REQUESTS,
REQUESTS,
Call,
CapRig,
Key,
Sample,
SpendRow,
alias_values,
expect_spend_rows,
families_over,
label_values,
overflow_total,
scrape,
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})
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)
samples: Final = eventually(
lambda: scrape(capped.gateway),
lambda after: (
overflow_total(after, REQUESTS) - overflow_total(before, REQUESTS) >= extra_requests
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
off_route: Final = len(capped.warm) * (len(ROUTES) - 1)
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: (
overflow_total(after, PROXY_FAILURES) - overflow_total(before, PROXY_FAILURES) >= 2
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]
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)
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 second boot's three keys get their series and a fourth lands on `other`, because the admitted series
files are dropped at boot even though the operator's sample files are left alone."""
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
old_pids: Final = frozenset(sample.pid for sample in worker_samples(operator_dir, REQUESTS))
with series_cap_rig(tmp_path, settings, workers=1, warm_keys=3, multiproc_dir=operator_dir) as second_boot:
fresh: Final = tuple(sample for sample in worker_samples(operator_dir, REQUESTS) if sample.pid not in old_pids)
assert len(fresh) == 1 and fresh[0].aliases == second_boot.warm_aliases, fresh
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)

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,403 @@
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
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,23 @@ 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_starts_the_series_cap_over(self, tmp_path: Path) -> None:
"""One worker and no metrics server leave the operator's directory alone, except for litellm's own
admitted-series files: the docs promise a restart frees every capped slot."""
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 samples.read_bytes() == b"operator-owned samples"