diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 317b01f3861..9f3baa438ee 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -12,9 +12,9 @@ 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 pydantic import BaseModel, Field, TypeAdapter, ValidationError from typing_extensions import ReadOnly, TypedDict, assert_never import litellm @@ -256,13 +256,27 @@ class _LabeledMetric: _MetricLike: TypeAlias = "NoOpMetric | _LabeledMetric | MetricWrapperBase" _SeriesLimitT: Final = TypeVar("_SeriesLimitT", int, float) +_POSITIVE_SERIES_CAP: Final[TypeAdapter[int]] = TypeAdapter(Annotated[int, Field(gt=0)]) +_POSITIVE_SERIES_TTL: Final[TypeAdapter[float]] = TypeAdapter(Annotated[float, Field(gt=0)]) -def _positive_or_ignored(setting: str, value: _SeriesLimitT | None) -> _SeriesLimitT | None: - if value is None or value > 0: - return value +def _positive_number(value: object, limit: TypeAdapter[_SeriesLimitT]) -> _SeriesLimitT | None: + if isinstance(value, bool): + return None + try: + return limit.validate_python(value) + except ValidationError: + return None + + +def _positive_or_ignored(setting: str, value: object, limit: TypeAdapter[_SeriesLimitT]) -> _SeriesLimitT | None: + if value is None: + return None + validated: Final = _positive_number(value, limit) + if validated is not None: + return validated verbose_logger.warning( - "%s is ignored because it is not greater than 0 (got %s). Prometheus metrics are emitted without it", + "%s is ignored because it is not a number greater than 0 (got %r). Prometheus metrics are emitted without it", setting, value, ) @@ -1299,9 +1313,13 @@ class PrometheusLogger(CustomLogger): 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 + "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 ), - ttl_seconds=_positive_or_ignored("prometheus_metrics_ttl_seconds", litellm.prometheus_metrics_ttl_seconds), cleanup_interval_seconds=litellm.prometheus_metrics_cleanup_interval_seconds, ) if limits.ttl_seconds is None or not multiprocess_mode: diff --git a/litellm/proxy/prometheus_cleanup.py b/litellm/proxy/prometheus_cleanup.py index 391fff1c964..b32dd7d8b75 100644 --- a/litellm/proxy/prometheus_cleanup.py +++ b/litellm/proxy/prometheus_cleanup.py @@ -19,21 +19,34 @@ _LIVE_GAUGE_PID: Final = re.compile(r"gauge_live[a-z]*_(\d+)\.db$") def wipe_directory(directory: str) -> None: """Delete all .db files and admitted-series files in the directory. Called once before workers fork.""" - files: Final = ( - *glob.glob(os.path.join(directory, "*.db")), - *glob.glob(os.path.join(directory, f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}*")), - ) - 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) + _remove(directory, (*glob.glob(os.path.join(directory, "*.db")), *_admitted_series_files(directory))) + + +def wipe_admitted_series(directory: str) -> None: + """Drop only litellm's own admitted-series files, so a restart that keeps an operator-managed directory + (one worker, no separate metrics server) still starts the series cap from an empty set.""" + _remove(directory, _admitted_series_files(directory)) + + +def _admitted_series_files(directory: str) -> tuple[str, ...]: + return tuple(glob.glob(os.path.join(directory, f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}*"))) + + +def _remove(directory: str, files: tuple[str, ...]) -> None: + deleted: Final = sum(_removed(filepath) for filepath in files) if deleted: verbose_proxy_logger.info("Prometheus cleanup: wiped %s stale 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: """Remove prometheus .db files for a dead worker. Called by gunicorn child_exit hook.""" if not os.environ.get("PROMETHEUS_MULTIPROC_DIR"): diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index a9745565799..5b63309507e 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -684,14 +684,16 @@ class ProxyInitializationHelpers: """ import tempfile + from litellm.proxy.prometheus_cleanup import wipe_admitted_series, wipe_directory + + configured_dir: Final = os.environ.get("PROMETHEUS_MULTIPROC_DIR") or os.environ.get("prometheus_multiproc_dir") if prometheus_metrics_port is None and ( num_workers <= 1 or not ProxyInitializationHelpers._prometheus_callback_configured(litellm_settings) ): + if configured_dir: + wipe_admitted_series(configured_dir) return None - from litellm.proxy.prometheus_cleanup import wipe_directory - - configured_dir: Final = os.environ.get("PROMETHEUS_MULTIPROC_DIR") or os.environ.get("prometheus_multiproc_dir") multiproc_dir: Final = configured_dir or os.path.join(tempfile.gettempdir(), "litellm_prometheus_multiproc") os.environ["PROMETHEUS_MULTIPROC_DIR"] = multiproc_dir 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..2dd394307df --- /dev/null +++ b/tests/integration/observability/test_prometheus_series_cap.py @@ -0,0 +1,453 @@ +"""Prometheus series cap on the live proxy: label sets past prometheus_metrics_max_series_per_metric share one +`other` series on every labeled counter and histogram and stay out of the gauges, idle series expire under +prometheus_metrics_ttl_seconds in single-process mode only, and a setting that is not a positive number is +ignored with a warning instead of silencing the metrics.""" + +from __future__ import annotations + +import time +from collections.abc import Iterator, Sequence +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import pytest +from integration._support.client import eventually, object_value, string_value +from integration._support.prometheus_series import ( + AGENT_HEADERS, + CACHE_HITS, + FAILED_FALLBACKS, + OVERFLOW, + PROVIDER_OUTAGE, + PROXY_FAILURES, + PROXY_REQUESTS, + REMAINING_REQUESTS, + REQUESTS, + SUCCESSFUL_FALLBACKS, + Call, + CapRig, + Key, + Sample, + WorkerSamples, + alias_values, + chat_once, + expect_spend_rows, + families_over, + gauge_samples, + label_values, + overflow_total, + received_markers, + scrape, + series_cap_rig, + sse_data, + worker_samples, +) +from pydantic import JsonValue + +pytestmark = pytest.mark.timeout(240) + +CAP: Final = 3 +TTL_SECONDS: Final = 2 +CLEANUP_SECONDS: Final = 1 +PRIMARY: Final = "primary" +FALLBACK: Final = "fallback" +TTL_IGNORED_WARNING: Final = "prometheus_metrics_ttl_seconds is ignored while PROMETHEUS_MULTIPROC_DIR is set" + + +def _grew(before: Sequence[Sample], after: Sequence[Sample], name: str, by: int) -> bool: + return overflow_total(after, name) - overflow_total(before, name) >= by + + +def _overflowed(rig: CapRig, key: Key, before: Sequence[Sample], requests: int) -> tuple[Sample, ...]: + """The scrape once the key's requests landed on `other` for both request counters, or as soon as the key got + a series of its own, so the caller's assertion fails fast on a proxy without the cap.""" + return eventually( + lambda: scrape(rig.gateway), + lambda after: ( + (_grew(before, after, REQUESTS, requests) and _grew(before, after, PROXY_REQUESTS, requests)) + or key.alias in alias_values(after, REQUESTS) + ), + seconds=60, + ) + + +def _expect_other( + rig: CapRig, key: Key, calls: Sequence[Call], response_ids: Sequence[str], before: Sequence[Sample] +) -> None: + samples: Final = _overflowed(rig, key, before, len(calls)) + assert key.alias not in alias_values(samples, REQUESTS) | alias_values(samples, PROXY_REQUESTS), key.alias + assert not families_over(samples, CAP), families_over(samples, CAP) + assert overflow_total(samples, REQUESTS) - overflow_total(before, REQUESTS) == len(calls) + expect_spend_rows(key.alias, response_ids) + markers: Final = received_markers(rig.provider) + assert all(call.marker in markers for call in calls), (calls, markers) + + +def _bearer(key: Key, call: Call) -> dict[str, str]: + return {**AGENT_HEADERS, **call.headers, "Authorization": f"Bearer {key.token}"} + + +class TestCapped: + def test_openai_sync_chat_past_the_cap_counts_on_other(self, capped: CapRig) -> None: + """H1: two OpenAI SDK chat completions from a fourth key count on `other` and keep their spend rows.""" + key: Final = capped.key("h1") + calls: Final = (Call.new(), Call.new()) + before: Final = scrape(capped.gateway) + client: Final = openai.OpenAI( + base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0 + ) + completions: Final = tuple( + client.chat.completions.create(model=capped.model, messages=[call.message], extra_headers=call.headers) + for call in calls + ) + assert tuple(completion.choices[0].message.content for completion in completions) == tuple( + call.answer for call in calls + ) + _expect_other(capped, key, calls, tuple(completion.id for completion in completions), before) + + async def test_openai_async_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None: + """H2: a streamed AsyncOpenAI chat completion from a fourth key counts on `other` once the stream ends.""" + key: Final = capped.key("h2") + call: Final = Call.new() + before: Final = scrape(capped.gateway) + client: Final = openai.AsyncOpenAI( + base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0 + ) + stream: Final = await client.chat.completions.create( + model=capped.model, messages=[call.message], stream=True, extra_headers=call.headers + ) + chunks: Final = tuple([chunk async for chunk in stream]) + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == call.answer + ids: Final = frozenset(chunk.id for chunk in chunks) + assert len(ids) == 1, ids + _expect_other(capped, key, (call,), tuple(ids), before) + + def test_anthropic_sync_messages_past_the_cap_counts_on_other(self, capped: CapRig) -> None: + """H3: an Anthropic SDK /v1/messages call from a fourth key counts on `other`.""" + key: Final = capped.key("h3") + call: Final = Call.new() + before: Final = scrape(capped.gateway) + client: Final = anthropic.Anthropic( + base_url=capped.base_url, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0 + ) + message: Final = client.messages.create( + model=capped.model, max_tokens=64, messages=[call.message], extra_headers=call.headers + ) + assert "".join(block.text for block in message.content if block.type == "text") == call.answer + _expect_other(capped, key, (call,), (message.id,), before) + + async def test_anthropic_async_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None: + """H4: a streamed AsyncAnthropic /v1/messages call from a fourth key counts on `other`.""" + key: Final = capped.key("h4") + call: Final = Call.new() + before: Final = scrape(capped.gateway) + client: Final = anthropic.AsyncAnthropic( + base_url=capped.base_url, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0 + ) + async with client.messages.stream( + model=capped.model, max_tokens=64, messages=[call.message], extra_headers=call.headers + ) as stream: + final: Final = await stream.get_final_message() + assert "".join(block.text for block in final.content if block.type == "text") == call.answer + _expect_other(capped, key, (call,), (final.id,), before) + + def test_openai_sync_responses_past_the_cap_counts_on_other(self, capped: CapRig) -> None: + """H5: an OpenAI SDK /v1/responses call from a fourth key counts on `other`.""" + key: Final = capped.key("h5") + call: Final = Call.new() + before: Final = scrape(capped.gateway) + client: Final = openai.OpenAI( + base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0 + ) + response: Final = client.responses.create(model=capped.model, input=call.text, extra_headers=call.headers) + assert response.output_text == call.answer + _expect_other(capped, key, (call,), (response.id,), before) + + def test_raw_responses_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None: + """H6: a raw httpx streamed /v1/responses call from a fourth key counts on `other`.""" + key: Final = capped.key("h6") + call: Final = Call.new() + before: Final = scrape(capped.gateway) + with httpx.Client(base_url=capped.base_url, timeout=60, trust_env=False) as client: + response: Final = client.post( + "/v1/responses", + json={"model": capped.model, "input": call.text, "stream": True}, + headers=_bearer(key, call), + ) + assert response.status_code == 200, response.text + events: Final = sse_data(response.text) + deltas: Final = tuple(event for event in events if event.get("type") == "response.output_text.delta") + assert "".join(string_value(event["delta"]) for event in deltas) == call.answer + completed: Final = tuple(event for event in events if event.get("type") == "response.completed") + assert len(completed) == 1, events + response_id: Final = string_value(object_value(completed[0]["response"])["id"]) + _expect_other(capped, key, (call,), (response_id,), before) + + def test_gauges_never_get_an_other_series(self, capped: CapRig) -> None: + """H7: a fourth key's request leaves no gauge sample for it and no gauge sample labeled `other`.""" + key: Final = capped.key("h7") + before: Final = scrape(capped.gateway) + assert capped.chat(key, Call.new()).status_code == 200 + samples: Final = _overflowed(capped, key, before, 1) + assert key.alias not in label_values(samples) + gauges: Final = gauge_samples(samples) + assert not any(OVERFLOW in gauge.labels.values() for gauge in gauges), gauges + for alias in capped.warm_aliases: + assert any( + gauge.name == REMAINING_REQUESTS and gauge.labels.get("api_key_alias") == alias for gauge in gauges + ), alias + + def test_cache_hits_past_the_cap_count_on_other(self, capped: CapRig) -> None: + """H8: the cache-hit twin: one populating call, hits from the warm keys, then a fourth key's hit on `other`.""" + shared: Final = Call.new() + first, second, third = capped.warm + extra: Final = capped.key("h8") + capped.provider.drain() + before: Final = scrape(capped.gateway) + for key in (first, first, second, third, extra): + response = capped.chat(key, shared) + assert response.status_code == 200 and shared.answer in response.text, response.text + samples: Final = eventually( + lambda: scrape(capped.gateway), + lambda after: _grew(before, after, CACHE_HITS, 1) or extra.alias in alias_values(after, CACHE_HITS), + seconds=60, + ) + assert alias_values(samples, CACHE_HITS) == capped.warm_aliases + assert overflow_total(samples, CACHE_HITS) - overflow_total(before, CACHE_HITS) == 1 + assert received_markers(capped.provider).count(shared.marker) == 1 + + def test_both_workers_share_the_admitted_series(self, capped: CapRig) -> None: + """H9: fresh connections reach both workers, and each worker's own sample file names only the warm aliases + while counting the fourth key on `other`, since the admitted sets live in the shared directory.""" + extra: Final = capped.key("h9") + + def send_on_a_fresh_connection() -> tuple[WorkerSamples, ...]: + assert capped.chat(extra, Call.new()).status_code == 200 + return worker_samples(capped.prom_dir, REQUESTS) + + workers: Final = eventually( + send_on_a_fresh_connection, + lambda found: ( + sum(1 for worker in found if worker.overflow > 0) >= 2 + or any(extra.alias in worker.aliases for worker in found) + ), + seconds=90, + ) + assert all(extra.alias not in worker.aliases for worker in workers), workers + assert sum(1 for worker in workers if worker.overflow > 0) >= 2, workers + assert frozenset().union(*(worker.aliases for worker in workers)) == capped.warm_aliases, workers + + def test_failures_past_the_cap_count_on_other(self, capped: CapRig) -> None: + """F1: provider failures fill the failure counter's cap with the warm keys, a fourth key's lands on `other`.""" + key: Final = capped.key("f1") + call: Final = Call.new() + before: Final = scrape(capped.gateway) + capped.outage.set() + try: + for warm in capped.warm: + assert capped.chat(warm, Call.new()).status_code == 500 + response: Final = capped.chat(key, call) + finally: + capped.outage.clear() + assert response.status_code == 500 and PROVIDER_OUTAGE in response.text, response.text + samples: Final = eventually( + lambda: scrape(capped.gateway), + lambda after: _grew(before, after, PROXY_FAILURES, 1) or key.alias in alias_values(after, PROXY_FAILURES), + seconds=60, + ) + assert key.alias not in label_values(samples) + assert overflow_total(samples, PROXY_FAILURES) - overflow_total(before, PROXY_FAILURES) == 1 + assert alias_values(samples, PROXY_FAILURES) == capped.warm_aliases + expect_spend_rows(key.alias, (), (call.call_id,)) + + def test_config_update_cannot_lift_a_yaml_cap(self, capped: CapRig) -> None: + """E1: /config/update refuses the YAML-owned cap, so a fourth key still lands on `other`.""" + response: Final = capped.gateway.client.post( + "/config/update", + json={"litellm_settings": {"prometheus_metrics_max_series_per_metric": 50}}, + headers={"Authorization": f"Bearer {capped.gateway.key}"}, + ) + assert response.status_code == 400, response.text + key: Final = capped.key("e1") + before: Final = scrape(capped.gateway) + assert capped.chat(key, Call.new()).status_code == 200 + samples: Final = _overflowed(capped, key, before, 1) + assert key.alias not in label_values(samples) + + +@pytest.fixture(scope="class") +def ttl(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]: + with series_cap_rig( + tmp_path_factory.mktemp("series-ttl"), + { + "prometheus_metrics_max_series_per_metric": 2, + "prometheus_metrics_ttl_seconds": TTL_SECONDS, + "prometheus_metrics_cleanup_interval_seconds": CLEANUP_SECONDS, + }, + workers=1, + warm_keys=2, + ) as rig: + yield rig + + +class TestTtl: + def test_idle_series_expire_and_free_their_slot(self, ttl: CapRig) -> None: + """T1: a third key lands on `other`; once the idle first key expires, a new key gets its own series.""" + first, second = ttl.warm + extra: Final = ttl.key("t1-extra") + assert ttl.chat(extra, Call.new()).status_code == 200 + samples: Final = eventually( + lambda: scrape(ttl.gateway), + lambda after: overflow_total(after, REQUESTS) >= 1 or extra.alias in alias_values(after, REQUESTS), + seconds=60, + ) + assert extra.alias not in label_values(samples) + + def keep_second_busy() -> tuple[Sample, ...]: + assert ttl.chat(second, Call.new()).status_code == 200 + return scrape(ttl.gateway) + + expired: Final = eventually(keep_second_busy, lambda after: first.alias not in label_values(after), seconds=30) + assert second.alias in alias_values(expired, REQUESTS) + late: Final = ttl.key("t1-late") + assert ttl.chat(late, Call.new()).status_code == 200 + named: Final = eventually( + lambda: scrape(ttl.gateway), lambda after: late.alias in alias_values(after, REQUESTS), seconds=30 + ) + assert late.alias in alias_values(named, REQUESTS) + + +@pytest.fixture(scope="class") +def ttl_multiproc(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]: + with series_cap_rig( + tmp_path_factory.mktemp("series-ttl-multiproc"), + { + "prometheus_metrics_max_series_per_metric": CAP, + "prometheus_metrics_ttl_seconds": TTL_SECONDS, + "prometheus_metrics_cleanup_interval_seconds": CLEANUP_SECONDS, + }, + workers=2, + warm_keys=3, + ) as rig: + yield rig + + +class TestTtlMultiproc: + def test_ttl_is_ignored_with_two_workers_while_the_cap_applies(self, ttl_multiproc: CapRig) -> None: + """M1: with two workers an idle key keeps its series past the TTL, the cap still applies, and the log says so.""" + first, second, _ = ttl_multiproc.warm + deadline: Final = time.monotonic() + 2 * TTL_SECONDS + while time.monotonic() < deadline: + assert ttl_multiproc.chat(second, Call.new()).status_code == 200 + assert first.alias in alias_values(scrape(ttl_multiproc.gateway), REQUESTS) + extra: Final = ttl_multiproc.key("m1") + before: Final = scrape(ttl_multiproc.gateway) + assert ttl_multiproc.chat(extra, Call.new()).status_code == 200 + samples: Final = _overflowed(ttl_multiproc, extra, before, 1) + assert extra.alias not in label_values(samples) + assert TTL_IGNORED_WARNING in ttl_multiproc.proxy.log.read_text() + + +@pytest.fixture(scope="class") +def ignored(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]: + with series_cap_rig( + tmp_path_factory.mktemp("series-ignored"), + {"prometheus_metrics_max_series_per_metric": "five", "prometheus_metrics_ttl_seconds": ""}, + workers=1, + warm_keys=5, + ) as rig: + yield rig + + +class TestIgnored: + def test_settings_that_are_not_positive_numbers_are_ignored_with_a_warning(self, ignored: CapRig) -> None: + """I1: a cap of "five" and an empty TTL leave every key its own series and each warning names its setting.""" + samples: Final = scrape(ignored.gateway) + assert alias_values(samples, REQUESTS) >= ignored.warm_aliases + assert not any(sample.is_overflow() for sample in samples) + log: Final = ignored.proxy.log.read_text() + assert ( + "prometheus_metrics_max_series_per_metric is ignored because it is not a number greater than 0 (got 'five')" + in log + ) + assert "prometheus_metrics_ttl_seconds is ignored because it is not a number greater than 0 (got '')" in log + + +def _fallback_deployments(provider_url: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + { + "model_name": name, + "litellm_params": { + "model": f"openai/gpt-{name}", + "api_base": provider_url + "/v1", + "api_key": "synthetic-provider-key", + }, + } + for name in (PRIMARY, FALLBACK) + ) + + +@pytest.fixture(scope="class") +def excluded(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]: + with series_cap_rig( + tmp_path_factory.mktemp("series-excluded"), + {"prometheus_metrics_max_series_per_metric": CAP, "prometheus_exclude_labels": ["api_key_alias"]}, + workers=1, + warm_keys=0, + failing_models=frozenset({f"gpt-{PRIMARY}"}), + deployments=_fallback_deployments, + router_settings={"fallbacks": [{PRIMARY: [FALLBACK]}]}, + ) as rig: + yield rig + + +class TestExcluded: + def test_fallback_counters_drop_excluded_labels(self, excluded: CapRig) -> None: + """X1: the successful and failed fallback counters honor prometheus_exclude_labels like every other metric.""" + key: Final = excluded.key("x1") + call: Final = Call.new() + response: Final = chat_once(excluded.base_url, key, PRIMARY, call) + assert response.status_code == 200 and call.answer in response.text, response.text + after_success: Final = eventually( + lambda: scrape(excluded.gateway), + lambda samples: any(sample.name == SUCCESSFUL_FALLBACKS for sample in samples), + seconds=60, + ) + successes: Final = tuple(sample for sample in after_success if sample.name == SUCCESSFUL_FALLBACKS) + assert any(sample.labels.get("fallback_model") == FALLBACK for sample in successes), successes + assert all("api_key_alias" not in sample.labels for sample in successes), successes + excluded.outage.set() + try: + failed: Final = chat_once(excluded.base_url, key, PRIMARY, Call.new()) + finally: + excluded.outage.clear() + assert failed.status_code == 500, failed.text + after_failure: Final = eventually( + lambda: scrape(excluded.gateway), + lambda samples: any(sample.name == FAILED_FALLBACKS for sample in samples), + seconds=60, + ) + failures: Final = tuple(sample for sample in after_failure if sample.name == FAILED_FALLBACKS) + assert any(sample.labels.get("fallback_model") == FALLBACK for sample in failures), failures + assert all("api_key_alias" not in sample.labels for sample in failures), failures + + +@pytest.fixture(scope="class") +def nocap(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]: + with series_cap_rig( + tmp_path_factory.mktemp("series-nocap"), + {"prometheus_metrics_max_series_per_metric": None}, + workers=2, + warm_keys=4, + ) as rig: + yield rig + + +class TestNoCap: + def test_a_null_cap_keeps_every_series(self, nocap: CapRig) -> None: + """N1: an explicit null cap and a missing TTL leave every key its own series and no `other` series.""" + samples: Final = scrape(nocap.gateway) + assert alias_values(samples, REQUESTS) >= nocap.warm_aliases + assert not any(sample.is_overflow() for sample in samples) 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..3301080e3a2 --- /dev/null +++ b/tests/integration/observability/test_prometheus_series_cap_chaos.py @@ -0,0 +1,332 @@ +"""Prometheus series cap under load: a concurrent burst across every endpoint while /metrics is scraped, a +provider outage between bursts, a worker killed mid-burst, and restarts that wipe or keep the multiprocess +directory.""" + +from __future__ import annotations + +import json +import re +import signal +import threading +from collections.abc import Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from itertools import cycle, product +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final + +import httpx +import psutil +import pytest +from integration._support.client import eventually, object_value, string_value +from integration._support.prometheus_series import ( + AGENT_HEADERS, + PROXY_FAILURES, + PROXY_REQUESTS, + REQUESTS, + Call, + CapRig, + Key, + Sample, + SpendRow, + alias_values, + expect_spend_rows, + families_over, + label_values, + overflow_total, + scrape, + series_cap_rig, + spend_rows, + sse_data, + worker_samples, +) +from pydantic import JsonValue + +pytestmark = pytest.mark.timeout(240) + +CAP: Final = 3 +CHAT: Final = "/v1/chat/completions" +MESSAGES: Final = "/v1/messages" +RESPONSES: Final = "/v1/responses" +ROUTES: Final = (CHAT, MESSAGES, RESPONSES) +EXTRA_KEYS: Final = 7 +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") + + +@dataclass(frozen=True, slots=True) +class Served: + key: Key + call: Call + route: str + streamed: bool + status: int + text: str + + @property + def response_id(self) -> str: + assert self.status == 200, self.text + if not self.streamed: + return string_value(object_value(json.loads(self.text))["id"]) + events: Final = sse_data(self.text) + match self.route: + case "/v1/messages": + starts: Final = tuple(event for event in events if event.get("type") == "message_start") + return string_value(object_value(starts[0]["message"])["id"]) + case "/v1/responses": + completed: Final = tuple(event for event in events if event.get("type") == "response.completed") + return string_value(object_value(completed[0]["response"])["id"]) + case _: + ids: Final = frozenset(string_value(event["id"]) for event in events) + assert len(ids) == 1, ids + return next(iter(ids)) + + +def _body(route: str, model: str, call: Call, streamed: bool) -> dict[str, JsonValue]: + match route: + case "/v1/messages": + return {"model": model, "max_tokens": 64, "messages": [call.message], "stream": streamed} + case "/v1/responses": + return {"model": model, "input": call.text, "stream": streamed} + case _: + return {"model": model, "messages": [call.message], "stream": streamed} + + +def _send(rig: CapRig, key: Key, route: str, streamed: bool, tolerate_transport_errors: bool = False) -> Served: + call: Final = Call.new() + try: + with httpx.Client(base_url=rig.base_url, timeout=60, trust_env=False) as client: + response: Final = client.post( + route, + json=_body(route, rig.model, call, streamed), + headers={**AGENT_HEADERS, **call.headers, "Authorization": f"Bearer {key.token}"}, + ) + except httpx.TransportError as error: + if not tolerate_transport_errors: + raise + return Served(key, call, route, streamed, 0, repr(error)) + return Served(key, call, route, streamed, response.status_code, response.text) + + +@dataclass(frozen=True, slots=True) +class Plan: + key: Key + route: str + streamed: bool + + +def _plans(keys: Sequence[Key]) -> tuple[Plan, ...]: + streaming: Final = cycle((False, True)) + return tuple(Plan(key, route, next(streaming)) for key, route in product(keys, ROUTES)) + + +def _burst(rig: CapRig, plans: Sequence[Plan], tolerate_transport_errors: bool = False) -> tuple[Served, ...]: + with ThreadPoolExecutor(max_workers=len(plans)) as pool: + return tuple( + pool.map(lambda plan: _send(rig, plan.key, plan.route, plan.streamed, tolerate_transport_errors), plans) + ) + + +def _scrape_until(rig: CapRig, stop: threading.Event, sizes: SimpleQueue[int]) -> None: + while not stop.is_set(): + try: + sizes.put(len(scrape(rig.gateway))) + except (AssertionError, httpx.HTTPError): + sizes.put(-1) + + +def _rows_by_alias(keys: Sequence[Key]) -> Mapping[str, tuple[SpendRow, ...]]: + return MappingProxyType({key.alias: spend_rows(key.alias) for key in keys}) + + +class TestBurst: + def test_concurrent_burst_across_every_endpoint_while_scraping(self, capped: CapRig) -> None: + """C1: 30 concurrent calls from ten keys across chat, messages, and responses, streamed and not, with + /metrics scraped throughout: every call answers, the warm keys keep their series, every other call counts + on `other`, and every call writes one spend row.""" + extra: Final = tuple(capped.key(f"c1-{index}") for index in range(EXTRA_KEYS)) + keys: Final = (*capped.warm, *extra) + earlier: Final = _rows_by_alias(keys) + before: Final = scrape(capped.gateway) + stop: Final = threading.Event() + sizes: Final[SimpleQueue[int]] = SimpleQueue() + scraper: Final = threading.Thread(target=_scrape_until, args=(capped, stop, sizes)) + scraper.start() + try: + served: Final = _burst(capped, _plans(keys)) + finally: + stop.set() + scraper.join() + scrapes: Final = tuple(sizes.get_nowait() for _ in range(sizes.qsize())) + assert scrapes and all(count > 0 for count in scrapes), scrapes + assert all(item.status == 200 and item.call.answer in item.text for item in served), [ + (item.route, item.status, item.text[:200]) for item in served if item.status != 200 + ] + extra_requests: Final = EXTRA_KEYS * len(ROUTES) + samples: Final = eventually( + lambda: scrape(capped.gateway), + lambda after: ( + overflow_total(after, REQUESTS) - overflow_total(before, REQUESTS) >= extra_requests + or any(key.alias in alias_values(after, REQUESTS) for key in extra) + ), + seconds=90, + ) + assert alias_values(samples, REQUESTS) == capped.warm_aliases + assert not families_over(samples, CAP), families_over(samples, CAP) + assert overflow_total(samples, REQUESTS) - overflow_total(before, REQUESTS) == extra_requests + assert alias_values(samples, PROXY_REQUESTS) == capped.warm_aliases + off_route: Final = len(capped.warm) * (len(ROUTES) - 1) + assert overflow_total(samples, PROXY_REQUESTS) - overflow_total(before, PROXY_REQUESTS) == ( + extra_requests + off_route + ) + for key in keys: + expect_spend_rows( + key.alias, + tuple(item.response_id for item in served if item.key == key), + earlier=earlier[key.alias], + ) + + def test_outage_between_bursts_counts_every_failure_once(self, capped: CapRig) -> None: + """C2: a burst answers, the provider goes down for the next burst, and comes back for the last: the warm + keys keep their failure series, the fourth key's failures count on `other`, and every call writes one row.""" + extra: Final = capped.key("c2") + keys: Final = (*capped.warm, extra) + earlier: Final = _rows_by_alias(keys) + before: Final = scrape(capped.gateway) + plans: Final = tuple(Plan(key, CHAT, streamed) for key, streamed in product(keys, (False, True))) + first: Final = _burst(capped, plans) + capped.outage.set() + try: + prefill: Final = tuple(_send(capped, warm, CHAT, False) for warm in capped.warm) + down: Final = _burst(capped, plans) + finally: + capped.outage.clear() + last: Final = _burst(capped, plans) + failed: Final = (*prefill, *down) + assert all(item.status == 200 for item in (*first, *last)), [item.status for item in (*first, *last)] + assert all(item.status == 500 for item in failed), [item.status for item in failed] + samples: Final = eventually( + lambda: scrape(capped.gateway), + lambda after: ( + overflow_total(after, PROXY_FAILURES) - overflow_total(before, PROXY_FAILURES) >= 2 + or extra.alias in alias_values(after, PROXY_FAILURES) + ), + seconds=90, + ) + assert alias_values(samples, PROXY_FAILURES) == capped.warm_aliases + assert not families_over(samples, CAP), families_over(samples, CAP) + assert overflow_total(samples, PROXY_FAILURES) - overflow_total(before, PROXY_FAILURES) == 2 + assert overflow_total(samples, REQUESTS) - overflow_total(before, REQUESTS) == 4 + for key in keys: + expect_spend_rows( + key.alias, + tuple(item.response_id for item in (*first, *last) if item.key == key), + tuple(item.call.call_id for item in failed if item.key == key), + earlier=earlier[key.alias], + ) + + +def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]: + text: Final = log.read_text() + return tuple(int(pid) for pid in _STARTED_WORKER.findall(text)), text.count("Application startup complete.") + + +@pytest.mark.timeout(420) +def test_killed_worker_is_replaced_by_one_that_reads_the_same_admissions(tmp_path: Path) -> None: + """C3: SIGKILL one of two workers mid-burst: the sibling keeps answering, and the replacement worker puts a + fourth key on `other` because the admitted series live in the shared directory, not in the dead process.""" + with series_cap_rig(tmp_path, {"prometheus_metrics_max_series_per_metric": CAP}, workers=2, warm_keys=3) as rig: + workers, _ = eventually( + lambda: _worker_startups(rig.proxy.log), lambda found: len(found[0]) == 2 and found[1] == 2, seconds=120 + ) + extra: Final = tuple(rig.key(f"c3-{index}") for index in range(4)) + plans: Final = _plans(extra) + with ThreadPoolExecutor(max_workers=1) as pool: + burst: Final = pool.submit(_burst, rig, plans, True) + victim: Final = psutil.Process(workers[0]) + victim.suspend() + victim.send_signal(signal.SIGKILL) + served: Final = burst.result() + answered: Final = tuple(item for item in served if item.status == 200) + assert answered, [(item.status, item.text[:200]) for item in served] + assert all(item.call.answer in item.text for item in answered) + replacement: Final = eventually( + lambda: _worker_startups(rig.proxy.log), + lambda found: len(frozenset(found[0]) - frozenset(workers)) == 1, + seconds=120, + ) + (new_pid,) = frozenset(replacement[0]) - frozenset(workers) + late: Final = rig.key("c3-late") + + def send_until_the_replacement_counts() -> tuple[Sample, ...]: + assert rig.chat(late, Call.new()).status_code == 200 + return scrape(rig.gateway) + + samples: Final = eventually( + send_until_the_replacement_counts, + lambda after: ( + any( + sample.pid == new_pid and (sample.overflow > 0 or late.alias in sample.aliases) + for sample in worker_samples(rig.prom_dir, REQUESTS) + ) + or late.alias in alias_values(after, REQUESTS) + ), + seconds=90, + ) + assert late.alias not in label_values(samples) + by_pid: Final = {sample.pid: sample for sample in worker_samples(rig.prom_dir, REQUESTS)} + assert by_pid[new_pid].overflow > 0 and by_pid[new_pid].aliases <= rig.warm_aliases, by_pid[new_pid] + + +def test_restart_with_two_workers_starts_the_cap_over(tmp_path: Path) -> None: + """C4: a second boot on the same multiprocess directory wipes it: the old keys are gone, three new keys get + their series, and a fourth lands on `other`.""" + shared_dir: Final = tmp_path / "prom-shared" + settings: Final = {"prometheus_metrics_max_series_per_metric": CAP} + with series_cap_rig(tmp_path, settings, workers=2, warm_keys=3, multiproc_dir=shared_dir) as first_boot: + old_aliases: Final = first_boot.warm_aliases + assert alias_values(scrape(first_boot.gateway), REQUESTS) == old_aliases + with series_cap_rig(tmp_path, settings, workers=2, warm_keys=3, multiproc_dir=shared_dir) as second_boot: + samples: Final = scrape(second_boot.gateway) + assert alias_values(samples, REQUESTS) == second_boot.warm_aliases + assert not old_aliases & label_values(samples) + extra: Final = second_boot.key("c4") + before: Final = scrape(second_boot.gateway) + assert second_boot.chat(extra, Call.new()).status_code == 200 + after: Final = eventually( + lambda: scrape(second_boot.gateway), + lambda now: ( + overflow_total(now, REQUESTS) - overflow_total(before, REQUESTS) >= 1 + or extra.alias in alias_values(now, REQUESTS) + ), + seconds=60, + ) + assert extra.alias not in label_values(after) + + +def test_restart_with_one_worker_and_an_operator_directory_starts_the_cap_over(tmp_path: Path) -> None: + """C5: one worker, no metrics port, PROMETHEUS_MULTIPROC_DIR set by the operator and kept across a restart: + the second boot's three keys get their series and a fourth lands on `other`, because the admitted series + files are dropped at boot even though the operator's sample files are left alone.""" + operator_dir: Final = tmp_path / "prom-operator" + settings: Final = {"prometheus_metrics_max_series_per_metric": CAP} + with series_cap_rig(tmp_path, settings, workers=1, warm_keys=3, multiproc_dir=operator_dir) as first_boot: + old_aliases: Final = first_boot.warm_aliases + assert alias_values(scrape(first_boot.gateway), REQUESTS) == old_aliases + old_pids: Final = frozenset(sample.pid for sample in worker_samples(operator_dir, REQUESTS)) + with series_cap_rig(tmp_path, settings, workers=1, warm_keys=3, multiproc_dir=operator_dir) as second_boot: + fresh: Final = tuple(sample for sample in worker_samples(operator_dir, REQUESTS) if sample.pid not in old_pids) + assert len(fresh) == 1 and fresh[0].aliases == second_boot.warm_aliases, fresh + extra: Final = second_boot.key("c5") + before: Final = scrape(second_boot.gateway) + assert second_boot.chat(extra, Call.new()).status_code == 200 + after: Final = eventually( + lambda: scrape(second_boot.gateway), + lambda now: ( + overflow_total(now, REQUESTS) - overflow_total(before, REQUESTS) >= 1 + or extra.alias in alias_values(now, REQUESTS) + ), + seconds=60, + ) + assert extra.alias not in label_values(after) diff --git a/tests/unit/integrations/test_prometheus_series_cardinality.py b/tests/unit/integrations/test_prometheus_series_cardinality.py index 80de28b6723..1230347bb47 100644 --- a/tests/unit/integrations/test_prometheus_series_cardinality.py +++ b/tests/unit/integrations/test_prometheus_series_cardinality.py @@ -367,10 +367,14 @@ def test_series_stay_unbounded_unless_a_limit_is_configured(): ("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_non_positive_series_limit_is_ignored_with_a_warning_and_metrics_keep_flowing( - setting: str, value: float, clock, caplog +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) @@ -384,3 +388,16 @@ def test_a_non_positive_series_limit_is_ignored_with_a_warning_and_metrics_keep_ series: Final = _scraped_series("litellm_proxy_total_requests_metric_total") assert _label_values(series, "user_agent") == {"agent-0", "agent-1", "agent-2"} assert setting in caplog.text + + +def test_a_series_cap_written_as_a_numeric_string_is_honored(caplog): + litellm.prometheus_metrics_max_series_per_metric = "2" + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + logger: Final = PrometheusLogger() + for index in range(3): + _count_request(logger, f"agent-{index}") + + series: Final = _scraped_series("litellm_proxy_total_requests_metric_total") + assert _label_values(series, "user_agent") == {"agent-0", "agent-1", "other"} + assert "prometheus_metrics_max_series_per_metric" not in caplog.text diff --git a/tests/unit/proxy/test_prometheus_cleanup.py b/tests/unit/proxy/test_prometheus_cleanup.py index 6a1b95c51ff..575e271634c 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,23 @@ class TestMaybeSetupPrometheusMultiprocDir: assert result_dir == str(tmp_path) assert os.environ["PROMETHEUS_MULTIPROC_DIR"] == str(tmp_path) + + def test_single_worker_restart_with_an_operator_set_dir_starts_the_series_cap_over(self, tmp_path: Path) -> None: + """One worker and no metrics server leave the operator's directory alone, except for litellm's own + admitted-series files: the docs promise a restart frees every capped slot.""" + admitted: Final = tmp_path / f"{PROMETHEUS_ADMITTED_SERIES_FILE_PREFIX}litellm_requests_metric" + admitted.write_text('\n["user-a"]\n') + samples: Final = tmp_path / "counter_123.db" + samples.write_bytes(b"operator-owned samples") + with patch.dict(os.environ, {"PROMETHEUS_MULTIPROC_DIR": str(tmp_path)}, clear=False): + os.environ.pop("prometheus_multiproc_dir", None) + + result_dir = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( + num_workers=1, + litellm_settings={"callbacks": ["prometheus"]}, + ) + + assert result_dir is None + assert os.environ["PROMETHEUS_MULTIPROC_DIR"] == str(tmp_path) + assert not admitted.exists() + assert samples.read_bytes() == b"operator-owned samples"