diff --git a/docker/component_entrypoint.sh b/docker/component_entrypoint.sh index 173afafe1ad..3cd530e3009 100755 --- a/docker/component_entrypoint.sh +++ b/docker/component_entrypoint.sh @@ -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 diff --git a/litellm/__init__.py b/litellm/__init__.py index fea7a27a5fb..03d8fccb27e 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 diff --git a/litellm/constants.py b/litellm/constants.py index eb95cc12ac7..33c93194ff2 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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)) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 6468dc41ea5..4366f785803 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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. diff --git a/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py index ba7d54fafea..6538b457d44 100644 --- a/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py +++ b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py @@ -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, diff --git a/litellm/integrations/prometheus_helpers/shared_prometheus_series_admissions.py b/litellm/integrations/prometheus_helpers/shared_prometheus_series_admissions.py new file mode 100644 index 00000000000..1df27175445 --- /dev/null +++ b/litellm/integrations/prometheus_helpers/shared_prometheus_series_admissions.py @@ -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) diff --git a/litellm/proxy/prometheus_cleanup.py b/litellm/proxy/prometheus_cleanup.py index c65aaeedfaa..52aa20cea5a 100644 --- a/litellm/proxy/prometheus_cleanup.py +++ b/litellm/proxy/prometheus_cleanup.py @@ -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: diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 9e736757fe2..08dfc79def7 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -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 diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 3fa7f4b0333..e8a0284382c 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -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: diff --git a/tests/integration/_support/prometheus_series.py b/tests/integration/_support/prometheus_series.py new file mode 100644 index 00000000000..3a406a97d11 --- /dev/null +++ b/tests/integration/_support/prometheus_series.py @@ -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) diff --git a/tests/integration/observability/conftest.py b/tests/integration/observability/conftest.py index f09a703f047..c5150877857 100644 --- a/tests/integration/observability/conftest.py +++ b/tests/integration/observability/conftest.py @@ -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 diff --git a/tests/integration/observability/test_prometheus_series_cap.py b/tests/integration/observability/test_prometheus_series_cap.py new file mode 100644 index 00000000000..ecdb6dabfc5 --- /dev/null +++ b/tests/integration/observability/test_prometheus_series_cap.py @@ -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) diff --git a/tests/integration/observability/test_prometheus_series_cap_chaos.py b/tests/integration/observability/test_prometheus_series_cap_chaos.py new file mode 100644 index 00000000000..c2f165ca37e --- /dev/null +++ b/tests/integration/observability/test_prometheus_series_cap_chaos.py @@ -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) diff --git a/tests/unit/enterprise/integrations/test_prometheus.py b/tests/unit/enterprise/integrations/test_prometheus.py index 16d6ff9d9a0..987cfa03a5e 100644 --- a/tests/unit/enterprise/integrations/test_prometheus.py +++ b/tests/unit/enterprise/integrations/test_prometheus.py @@ -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 # ============================================================================== diff --git a/tests/unit/integrations/test_prometheus_series_cardinality.py b/tests/unit/integrations/test_prometheus_series_cardinality.py new file mode 100644 index 00000000000..4b87f0ad995 --- /dev/null +++ b/tests/unit/integrations/test_prometheus_series_cardinality.py @@ -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 diff --git a/tests/unit/proxy/test_prometheus_cleanup.py b/tests/unit/proxy/test_prometheus_cleanup.py index 6a1b95c51ff..c774d376d33 100644 --- a/tests/unit/proxy/test_prometheus_cleanup.py +++ b/tests/unit/proxy/test_prometheus_cleanup.py @@ -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() diff --git a/tests/unit/proxy/test_proxy_cli.py b/tests/unit/proxy/test_proxy_cli.py index 51fbaf2ee9f..8fc807cd8a0 100644 --- a/tests/unit/proxy/test_proxy_cli.py +++ b/tests/unit/proxy/test_proxy_cli.py @@ -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") diff --git a/tests/unit/test_component_entrypoint.py b/tests/unit/test_component_entrypoint.py index 0c2a533b8bc..335c7e54416 100644 --- a/tests/unit/test_component_entrypoint.py +++ b/tests/unit/test_component_entrypoint.py @@ -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"