mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(prometheus): start the series cap over on a one-worker restart and audit it live
A proxy with one worker and an operator-set PROMETHEUS_MULTIPROC_DIR now drops litellm's admission files at boot, so a restart frees every slot there the way it already does with several workers. A cap or TTL that is not a number greater than 0 (a bool, a non-numeric string, an empty value) is ignored with the startup warning instead of breaking the logger The integration cells drive the cap on every endpoint through the OpenAI and Anthropic SDKs and raw httpx, streaming and not, plus gauges, cache hits, failures, both workers of one instance, the TTL on one worker and its warning on two, ignored settings, excluded labels on the fallback counters, a null cap, /config/update, a concurrent burst scraped mid-flight, a provider outage, a killed worker, and restarts with one and two workers
This commit is contained in:
parent
77dd85c663
commit
87b092c83d
9 changed files with 1255 additions and 24 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
362
tests/integration/_support/prometheus_series.py
Normal file
362
tests/integration/_support/prometheus_series.py
Normal file
|
|
@ -0,0 +1,362 @@
|
|||
"""Rig and readers for the Prometheus series-cap cells: a capped proxy, keys that fill the cap, and the
|
||||
scrape, the multiprocess sample files, and the spend log read back per request."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
import uuid
|
||||
from collections import Counter
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import yaml
|
||||
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from integration._support.responses_vendor import ResponsesVendor, same_response
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from prometheus_client.mmap_dict import MmapedDict
|
||||
from prometheus_client.parser import text_string_to_metric_families
|
||||
from pydantic import JsonValue
|
||||
|
||||
REQUESTS: Final = "litellm_requests_metric_total"
|
||||
PROXY_REQUESTS: Final = "litellm_proxy_total_requests_metric_total"
|
||||
PROXY_FAILURES: Final = "litellm_proxy_failed_requests_metric_total"
|
||||
CACHE_HITS: Final = "litellm_cache_hits_metric_total"
|
||||
REMAINING_REQUESTS: Final = "litellm_remaining_api_key_requests_for_model"
|
||||
SUCCESSFUL_FALLBACKS: Final = "litellm_deployment_successful_fallbacks_total"
|
||||
FAILED_FALLBACKS: Final = "litellm_deployment_failed_fallbacks_total"
|
||||
OVERFLOW: Final = "other"
|
||||
USER_AGENT: Final = "litellm-series-cap-audit/1"
|
||||
AGENT_HEADERS: Final[Mapping[str, str]] = MappingProxyType({"User-Agent": USER_AGENT})
|
||||
PROVIDER_OUTAGE: Final = "synthetic provider outage"
|
||||
_SERIES_ATTRIBUTES: Final = frozenset({"le", "pid"})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Call:
|
||||
"""One request the cells can follow end to end: the call id the proxy keeps as the spend log's request id
|
||||
and the marker the scripted provider echoes in its answer."""
|
||||
|
||||
call_id: str
|
||||
marker: str
|
||||
|
||||
@classmethod
|
||||
def new(cls) -> Call:
|
||||
identity: Final = uuid.uuid4()
|
||||
return cls(str(identity), identity.hex)
|
||||
|
||||
@property
|
||||
def text(self) -> str:
|
||||
return f"say marker-{self.marker}"
|
||||
|
||||
@property
|
||||
def answer(self) -> str:
|
||||
return f"answer marker-{self.marker}"
|
||||
|
||||
@property
|
||||
def message(self) -> dict[str, str]:
|
||||
return {"role": "user", "content": self.text}
|
||||
|
||||
@property
|
||||
def headers(self) -> dict[str, str]:
|
||||
return {"x-litellm-call-id": self.call_id}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Key:
|
||||
token: str
|
||||
alias: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Provider:
|
||||
vendor: ResponsesVendor
|
||||
outage: threading.Event
|
||||
failing_models: frozenset[str]
|
||||
|
||||
def respond(self, request: Request) -> Reply:
|
||||
if self.outage.is_set() or self._failing(request):
|
||||
return Reply(status=500, body=json.dumps({"error": {"message": PROVIDER_OUTAGE}}).encode())
|
||||
return self.vendor.respond(request)
|
||||
|
||||
def _failing(self, request: Request) -> bool:
|
||||
if request.method != "POST" or not self.failing_models:
|
||||
return False
|
||||
body: Final = json.loads(request.body)
|
||||
return isinstance(body, dict) and body.get("model") in self.failing_models
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Sample:
|
||||
family: str
|
||||
kind: str
|
||||
name: str
|
||||
labels: Mapping[str, str]
|
||||
value: float
|
||||
|
||||
def identity(self) -> tuple[tuple[str, str], ...]:
|
||||
"""The label set that makes this a series of its own: the histogram bucket and the multiprocess pid are
|
||||
attributes of one series, not separate ones."""
|
||||
return tuple(sorted((name, value) for name, value in self.labels.items() if name not in _SERIES_ATTRIBUTES))
|
||||
|
||||
def is_overflow(self) -> bool:
|
||||
identity: Final = self.identity()
|
||||
return bool(identity) and all(value == OVERFLOW for _, value in identity)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CapRig:
|
||||
proxy: OwnedProxy
|
||||
scenario: Scenario
|
||||
model: str
|
||||
provider: Wire
|
||||
outage: threading.Event
|
||||
warm: tuple[Key, ...]
|
||||
prom_dir: Path
|
||||
|
||||
@property
|
||||
def gateway(self) -> Gateway:
|
||||
return self.proxy.gateway
|
||||
|
||||
@property
|
||||
def base_url(self) -> str:
|
||||
return str(self.gateway.client.base_url).rstrip("/")
|
||||
|
||||
@property
|
||||
def openai_base(self) -> str:
|
||||
return self.base_url + "/v1"
|
||||
|
||||
@property
|
||||
def warm_aliases(self) -> frozenset[str]:
|
||||
return frozenset(key.alias for key in self.warm)
|
||||
|
||||
def key(self, cell: str) -> Key:
|
||||
alias: Final = f"{cell}-{uuid.uuid4().hex[:12]}"
|
||||
return Key(self.scenario.key(key_alias=alias), alias)
|
||||
|
||||
def chat(self, key: Key, call: Call) -> httpx.Response:
|
||||
return chat_once(self.base_url, key, self.model, call)
|
||||
|
||||
|
||||
def chat_once(base_url: str, key: Key, model: str, call: Call) -> httpx.Response:
|
||||
with httpx.Client(base_url=base_url, timeout=60, trust_env=False) as client:
|
||||
return client.post(
|
||||
"/v1/chat/completions",
|
||||
json={"model": model, "messages": [call.message]},
|
||||
headers={**AGENT_HEADERS, **call.headers, "Authorization": f"Bearer {key.token}"},
|
||||
)
|
||||
|
||||
|
||||
def series_cap_config(
|
||||
directory: Path,
|
||||
settings: Mapping[str, JsonValue],
|
||||
*,
|
||||
model_list: Sequence[Mapping[str, JsonValue]] = (),
|
||||
router_settings: Mapping[str, JsonValue] | None = None,
|
||||
) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
assert isinstance(config, dict)
|
||||
config["litellm_settings"] = {**config["litellm_settings"], "callbacks": ["prometheus"], **settings}
|
||||
config["router_settings"] = {**config["router_settings"], "num_retries": 0, **(router_settings or {})}
|
||||
if model_list:
|
||||
config["model_list"] = [dict(entry) for entry in model_list]
|
||||
path: Final = directory / "series-cap.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
@contextmanager
|
||||
def series_cap_rig(
|
||||
directory: Path,
|
||||
settings: Mapping[str, JsonValue],
|
||||
*,
|
||||
workers: int,
|
||||
warm_keys: int,
|
||||
multiproc_dir: Path | None = None,
|
||||
failing_models: frozenset[str] = frozenset(),
|
||||
deployments: Callable[[str], Sequence[Mapping[str, JsonValue]]] | None = None,
|
||||
router_settings: Mapping[str, JsonValue] | None = None,
|
||||
) -> Iterator[CapRig]:
|
||||
outage: Final = threading.Event()
|
||||
double: Final = Provider(ResponsesVendor(), outage, failing_models)
|
||||
prom_dir: Final = multiproc_dir if multiproc_dir is not None else directory / "prom"
|
||||
prom_dir.mkdir(exist_ok=True)
|
||||
shared_samples: Final = multiproc_dir is not None or workers > 1
|
||||
with ExitStack() as stack:
|
||||
gateway: Final = stack.enter_context(gateway_from_environment())
|
||||
provider: Final = stack.enter_context(wire_server(double.respond))
|
||||
config: Final = series_cap_config(
|
||||
directory,
|
||||
settings,
|
||||
model_list=deployments(provider.url) if deployments is not None else (),
|
||||
router_settings=router_settings,
|
||||
)
|
||||
owned: Final = stack.enter_context(
|
||||
owned_proxy_process(
|
||||
gateway,
|
||||
directory,
|
||||
{"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)} if shared_samples else {},
|
||||
config=config,
|
||||
remove_environment=() if shared_samples else ("PROMETHEUS_MULTIPROC_DIR",),
|
||||
workers=workers,
|
||||
)
|
||||
)
|
||||
scenario: Final = stack.enter_context(owned.gateway.scenario())
|
||||
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
|
||||
rig: Final = CapRig(owned, scenario, model, provider, outage, (), prom_dir)
|
||||
warm: Final = tuple(rig.key("warm") for _ in range(warm_keys))
|
||||
for key in warm:
|
||||
_warm_up(rig, key)
|
||||
eventually(
|
||||
lambda: alias_values(scrape(owned.gateway), REQUESTS),
|
||||
lambda seen: all(key.alias in seen for key in warm),
|
||||
seconds=60,
|
||||
)
|
||||
yield CapRig(owned, scenario, model, provider, outage, warm, prom_dir)
|
||||
|
||||
|
||||
def _warm_up(rig: CapRig, key: Key) -> None:
|
||||
response: Final = rig.chat(key, Call.new())
|
||||
assert response.status_code == 200, response.text
|
||||
|
||||
|
||||
def _samples(text: str) -> Iterator[Sample]:
|
||||
for family in text_string_to_metric_families(text):
|
||||
for sample in family.samples:
|
||||
yield Sample(
|
||||
family.name, family.type, sample.name, MappingProxyType(dict(sample.labels)), float(sample.value)
|
||||
)
|
||||
|
||||
|
||||
def scrape(gateway: Gateway) -> tuple[Sample, ...]:
|
||||
response: Final = gateway.client.request(
|
||||
"GET", "/metrics", headers={"Authorization": f"Bearer {gateway.key}"}, follow_redirects=True
|
||||
)
|
||||
assert response.status_code == 200, f"GET /metrics: {response.status_code} {response.text[:300]}"
|
||||
return tuple(_samples(response.text))
|
||||
|
||||
|
||||
def alias_values(samples: Sequence[Sample], name: str) -> frozenset[str]:
|
||||
"""The key aliases holding a series of their own on the metric; the shared overflow series is not one."""
|
||||
return frozenset(
|
||||
sample.labels["api_key_alias"]
|
||||
for sample in samples
|
||||
if sample.name == name and sample.labels.get("api_key_alias") not in (None, OVERFLOW)
|
||||
)
|
||||
|
||||
|
||||
def alias_total(samples: Sequence[Sample], name: str, alias: str) -> float:
|
||||
return sum(
|
||||
sample.value for sample in samples if sample.name == name and sample.labels.get("api_key_alias") == alias
|
||||
)
|
||||
|
||||
|
||||
def overflow_total(samples: Sequence[Sample], name: str) -> float:
|
||||
return sum(sample.value for sample in samples if sample.name == name and sample.is_overflow())
|
||||
|
||||
|
||||
def label_values(samples: Sequence[Sample]) -> frozenset[str]:
|
||||
return frozenset(chain.from_iterable(sample.labels.values() for sample in samples))
|
||||
|
||||
|
||||
def gauge_samples(samples: Sequence[Sample]) -> tuple[Sample, ...]:
|
||||
return tuple(sample for sample in samples if sample.kind == "gauge")
|
||||
|
||||
|
||||
def series_per_family(samples: Sequence[Sample]) -> Mapping[str, int]:
|
||||
"""How many series of their own each metric family holds, the shared `other` series left out."""
|
||||
owned: Final = frozenset(
|
||||
(sample.family, sample.identity()) for sample in samples if sample.identity() and not sample.is_overflow()
|
||||
)
|
||||
return MappingProxyType(dict(Counter(family for family, _ in owned)))
|
||||
|
||||
|
||||
def families_over(samples: Sequence[Sample], cap: int) -> tuple[tuple[str, int], ...]:
|
||||
"""Every metric family holding more series of its own than the cap allows."""
|
||||
return tuple(sorted((family, count) for family, count in series_per_family(samples).items() if count > cap))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpendRow:
|
||||
request_id: str
|
||||
status: str
|
||||
|
||||
|
||||
def spend_rows(alias: str) -> tuple[SpendRow, ...]:
|
||||
"""Every spend log row the key wrote: a success row carries the response id the caller received, a failure
|
||||
row the call id the caller sent."""
|
||||
rows: Final = read_rows(
|
||||
"SELECT request_id, status FROM \"LiteLLM_SpendLogs\" WHERE metadata->>'user_api_key_alias' = %s",
|
||||
(alias,),
|
||||
)
|
||||
return tuple(SpendRow(str(row["request_id"]), str(row["status"])) for row in rows)
|
||||
|
||||
|
||||
def expect_spend_rows(
|
||||
alias: str, response_ids: Sequence[str], call_ids: Sequence[str] = (), earlier: Sequence[SpendRow] = ()
|
||||
) -> None:
|
||||
"""One new row per request on top of the rows the key already had: successes found by the response id the
|
||||
caller got, failures by their call id."""
|
||||
expected: Final = len(earlier) + len(response_ids) + len(call_ids)
|
||||
rows: Final = eventually(lambda: spend_rows(alias), lambda found: len(found) >= expected, seconds=70)
|
||||
fresh: Final = tuple(row for row in rows if row not in earlier)
|
||||
assert len(rows) == expected and len(fresh) == len(response_ids) + len(call_ids), (rows, earlier)
|
||||
for response_id in response_ids:
|
||||
assert any(row.status == "success" and same_response(row.request_id, response_id) for row in fresh), (
|
||||
response_id,
|
||||
fresh,
|
||||
)
|
||||
for call_id in call_ids:
|
||||
assert any(row.status == "failure" and row.request_id == call_id for row in fresh), (call_id, fresh)
|
||||
|
||||
|
||||
def sse_data(text: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
"""The JSON payload of every `data:` frame in a server-sent event stream, the `[DONE]` sentinel left out."""
|
||||
payloads: Final = tuple(
|
||||
line.removeprefix("data:").strip() for line in text.splitlines() if line.startswith("data:")
|
||||
)
|
||||
return tuple(object_value(json.loads(payload)) for payload in payloads if payload and payload != "[DONE]")
|
||||
|
||||
|
||||
def received_markers(provider: Wire) -> tuple[str, ...]:
|
||||
return tuple(chain.from_iterable(_markers_in(request.body) for request in provider.drain()))
|
||||
|
||||
|
||||
def _markers_in(body: bytes) -> tuple[str, ...]:
|
||||
return tuple(part[:32].decode() for part in body.split(b"marker-")[1:])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WorkerSamples:
|
||||
pid: int
|
||||
aliases: frozenset[str]
|
||||
overflow: float
|
||||
|
||||
|
||||
def worker_samples(prom_dir: Path, name: str) -> tuple[WorkerSamples, ...]:
|
||||
return tuple(_worker_samples(path, name) for path in sorted(prom_dir.glob("counter_*.db")))
|
||||
|
||||
|
||||
def _worker_samples(path: Path, name: str) -> WorkerSamples:
|
||||
pid: Final = int(path.stem.rsplit("_", 1)[1])
|
||||
rows: Final = tuple(_counter_rows(path, name))
|
||||
return WorkerSamples(
|
||||
pid,
|
||||
frozenset(labels["api_key_alias"] for labels, _ in rows if labels.get("api_key_alias") not in (None, OVERFLOW)),
|
||||
sum(value for labels, value in rows if all(label == OVERFLOW for label in labels.values())),
|
||||
)
|
||||
|
||||
|
||||
def _counter_rows(path: Path, name: str) -> Iterator[tuple[Mapping[str, str], float]]:
|
||||
for key, value, *_ in MmapedDict.read_all_values_from_file(str(path)):
|
||||
_, sample_name, labels, _ = json.loads(key)
|
||||
if sample_name == name:
|
||||
yield labels, float(value)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
453
tests/integration/observability/test_prometheus_series_cap.py
Normal file
453
tests/integration/observability/test_prometheus_series_cap.py
Normal file
|
|
@ -0,0 +1,453 @@
|
|||
"""Prometheus series cap on the live proxy: label sets past prometheus_metrics_max_series_per_metric share one
|
||||
`other` series on every labeled counter and histogram and stay out of the gauges, idle series expire under
|
||||
prometheus_metrics_ttl_seconds in single-process mode only, and a setting that is not a positive number is
|
||||
ignored with a warning instead of silencing the metrics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Iterator, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from integration._support.client import eventually, object_value, string_value
|
||||
from integration._support.prometheus_series import (
|
||||
AGENT_HEADERS,
|
||||
CACHE_HITS,
|
||||
FAILED_FALLBACKS,
|
||||
OVERFLOW,
|
||||
PROVIDER_OUTAGE,
|
||||
PROXY_FAILURES,
|
||||
PROXY_REQUESTS,
|
||||
REMAINING_REQUESTS,
|
||||
REQUESTS,
|
||||
SUCCESSFUL_FALLBACKS,
|
||||
Call,
|
||||
CapRig,
|
||||
Key,
|
||||
Sample,
|
||||
WorkerSamples,
|
||||
alias_values,
|
||||
chat_once,
|
||||
expect_spend_rows,
|
||||
families_over,
|
||||
gauge_samples,
|
||||
label_values,
|
||||
overflow_total,
|
||||
received_markers,
|
||||
scrape,
|
||||
series_cap_rig,
|
||||
sse_data,
|
||||
worker_samples,
|
||||
)
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark = pytest.mark.timeout(240)
|
||||
|
||||
CAP: Final = 3
|
||||
TTL_SECONDS: Final = 2
|
||||
CLEANUP_SECONDS: Final = 1
|
||||
PRIMARY: Final = "primary"
|
||||
FALLBACK: Final = "fallback"
|
||||
TTL_IGNORED_WARNING: Final = "prometheus_metrics_ttl_seconds is ignored while PROMETHEUS_MULTIPROC_DIR is set"
|
||||
|
||||
|
||||
def _grew(before: Sequence[Sample], after: Sequence[Sample], name: str, by: int) -> bool:
|
||||
return overflow_total(after, name) - overflow_total(before, name) >= by
|
||||
|
||||
|
||||
def _overflowed(rig: CapRig, key: Key, before: Sequence[Sample], requests: int) -> tuple[Sample, ...]:
|
||||
"""The scrape once the key's requests landed on `other` for both request counters, or as soon as the key got
|
||||
a series of its own, so the caller's assertion fails fast on a proxy without the cap."""
|
||||
return eventually(
|
||||
lambda: scrape(rig.gateway),
|
||||
lambda after: (
|
||||
(_grew(before, after, REQUESTS, requests) and _grew(before, after, PROXY_REQUESTS, requests))
|
||||
or key.alias in alias_values(after, REQUESTS)
|
||||
),
|
||||
seconds=60,
|
||||
)
|
||||
|
||||
|
||||
def _expect_other(
|
||||
rig: CapRig, key: Key, calls: Sequence[Call], response_ids: Sequence[str], before: Sequence[Sample]
|
||||
) -> None:
|
||||
samples: Final = _overflowed(rig, key, before, len(calls))
|
||||
assert key.alias not in alias_values(samples, REQUESTS) | alias_values(samples, PROXY_REQUESTS), key.alias
|
||||
assert not families_over(samples, CAP), families_over(samples, CAP)
|
||||
assert overflow_total(samples, REQUESTS) - overflow_total(before, REQUESTS) == len(calls)
|
||||
expect_spend_rows(key.alias, response_ids)
|
||||
markers: Final = received_markers(rig.provider)
|
||||
assert all(call.marker in markers for call in calls), (calls, markers)
|
||||
|
||||
|
||||
def _bearer(key: Key, call: Call) -> dict[str, str]:
|
||||
return {**AGENT_HEADERS, **call.headers, "Authorization": f"Bearer {key.token}"}
|
||||
|
||||
|
||||
class TestCapped:
|
||||
def test_openai_sync_chat_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
|
||||
"""H1: two OpenAI SDK chat completions from a fourth key count on `other` and keep their spend rows."""
|
||||
key: Final = capped.key("h1")
|
||||
calls: Final = (Call.new(), Call.new())
|
||||
before: Final = scrape(capped.gateway)
|
||||
client: Final = openai.OpenAI(
|
||||
base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
|
||||
)
|
||||
completions: Final = tuple(
|
||||
client.chat.completions.create(model=capped.model, messages=[call.message], extra_headers=call.headers)
|
||||
for call in calls
|
||||
)
|
||||
assert tuple(completion.choices[0].message.content for completion in completions) == tuple(
|
||||
call.answer for call in calls
|
||||
)
|
||||
_expect_other(capped, key, calls, tuple(completion.id for completion in completions), before)
|
||||
|
||||
async def test_openai_async_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
|
||||
"""H2: a streamed AsyncOpenAI chat completion from a fourth key counts on `other` once the stream ends."""
|
||||
key: Final = capped.key("h2")
|
||||
call: Final = Call.new()
|
||||
before: Final = scrape(capped.gateway)
|
||||
client: Final = openai.AsyncOpenAI(
|
||||
base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
|
||||
)
|
||||
stream: Final = await client.chat.completions.create(
|
||||
model=capped.model, messages=[call.message], stream=True, extra_headers=call.headers
|
||||
)
|
||||
chunks: Final = tuple([chunk async for chunk in stream])
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == call.answer
|
||||
ids: Final = frozenset(chunk.id for chunk in chunks)
|
||||
assert len(ids) == 1, ids
|
||||
_expect_other(capped, key, (call,), tuple(ids), before)
|
||||
|
||||
def test_anthropic_sync_messages_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
|
||||
"""H3: an Anthropic SDK /v1/messages call from a fourth key counts on `other`."""
|
||||
key: Final = capped.key("h3")
|
||||
call: Final = Call.new()
|
||||
before: Final = scrape(capped.gateway)
|
||||
client: Final = anthropic.Anthropic(
|
||||
base_url=capped.base_url, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
|
||||
)
|
||||
message: Final = client.messages.create(
|
||||
model=capped.model, max_tokens=64, messages=[call.message], extra_headers=call.headers
|
||||
)
|
||||
assert "".join(block.text for block in message.content if block.type == "text") == call.answer
|
||||
_expect_other(capped, key, (call,), (message.id,), before)
|
||||
|
||||
async def test_anthropic_async_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
|
||||
"""H4: a streamed AsyncAnthropic /v1/messages call from a fourth key counts on `other`."""
|
||||
key: Final = capped.key("h4")
|
||||
call: Final = Call.new()
|
||||
before: Final = scrape(capped.gateway)
|
||||
client: Final = anthropic.AsyncAnthropic(
|
||||
base_url=capped.base_url, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
|
||||
)
|
||||
async with client.messages.stream(
|
||||
model=capped.model, max_tokens=64, messages=[call.message], extra_headers=call.headers
|
||||
) as stream:
|
||||
final: Final = await stream.get_final_message()
|
||||
assert "".join(block.text for block in final.content if block.type == "text") == call.answer
|
||||
_expect_other(capped, key, (call,), (final.id,), before)
|
||||
|
||||
def test_openai_sync_responses_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
|
||||
"""H5: an OpenAI SDK /v1/responses call from a fourth key counts on `other`."""
|
||||
key: Final = capped.key("h5")
|
||||
call: Final = Call.new()
|
||||
before: Final = scrape(capped.gateway)
|
||||
client: Final = openai.OpenAI(
|
||||
base_url=capped.openai_base, api_key=key.token, default_headers=dict(AGENT_HEADERS), max_retries=0
|
||||
)
|
||||
response: Final = client.responses.create(model=capped.model, input=call.text, extra_headers=call.headers)
|
||||
assert response.output_text == call.answer
|
||||
_expect_other(capped, key, (call,), (response.id,), before)
|
||||
|
||||
def test_raw_responses_stream_past_the_cap_counts_on_other(self, capped: CapRig) -> None:
|
||||
"""H6: a raw httpx streamed /v1/responses call from a fourth key counts on `other`."""
|
||||
key: Final = capped.key("h6")
|
||||
call: Final = Call.new()
|
||||
before: Final = scrape(capped.gateway)
|
||||
with httpx.Client(base_url=capped.base_url, timeout=60, trust_env=False) as client:
|
||||
response: Final = client.post(
|
||||
"/v1/responses",
|
||||
json={"model": capped.model, "input": call.text, "stream": True},
|
||||
headers=_bearer(key, call),
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
events: Final = sse_data(response.text)
|
||||
deltas: Final = tuple(event for event in events if event.get("type") == "response.output_text.delta")
|
||||
assert "".join(string_value(event["delta"]) for event in deltas) == call.answer
|
||||
completed: Final = tuple(event for event in events if event.get("type") == "response.completed")
|
||||
assert len(completed) == 1, events
|
||||
response_id: Final = string_value(object_value(completed[0]["response"])["id"])
|
||||
_expect_other(capped, key, (call,), (response_id,), before)
|
||||
|
||||
def test_gauges_never_get_an_other_series(self, capped: CapRig) -> None:
|
||||
"""H7: a fourth key's request leaves no gauge sample for it and no gauge sample labeled `other`."""
|
||||
key: Final = capped.key("h7")
|
||||
before: Final = scrape(capped.gateway)
|
||||
assert capped.chat(key, Call.new()).status_code == 200
|
||||
samples: Final = _overflowed(capped, key, before, 1)
|
||||
assert key.alias not in label_values(samples)
|
||||
gauges: Final = gauge_samples(samples)
|
||||
assert not any(OVERFLOW in gauge.labels.values() for gauge in gauges), gauges
|
||||
for alias in capped.warm_aliases:
|
||||
assert any(
|
||||
gauge.name == REMAINING_REQUESTS and gauge.labels.get("api_key_alias") == alias for gauge in gauges
|
||||
), alias
|
||||
|
||||
def test_cache_hits_past_the_cap_count_on_other(self, capped: CapRig) -> None:
|
||||
"""H8: the cache-hit twin: one populating call, hits from the warm keys, then a fourth key's hit on `other`."""
|
||||
shared: Final = Call.new()
|
||||
first, second, third = capped.warm
|
||||
extra: Final = capped.key("h8")
|
||||
capped.provider.drain()
|
||||
before: Final = scrape(capped.gateway)
|
||||
for key in (first, first, second, third, extra):
|
||||
response = capped.chat(key, shared)
|
||||
assert response.status_code == 200 and shared.answer in response.text, response.text
|
||||
samples: Final = eventually(
|
||||
lambda: scrape(capped.gateway),
|
||||
lambda after: _grew(before, after, CACHE_HITS, 1) or extra.alias in alias_values(after, CACHE_HITS),
|
||||
seconds=60,
|
||||
)
|
||||
assert alias_values(samples, CACHE_HITS) == capped.warm_aliases
|
||||
assert overflow_total(samples, CACHE_HITS) - overflow_total(before, CACHE_HITS) == 1
|
||||
assert received_markers(capped.provider).count(shared.marker) == 1
|
||||
|
||||
def test_both_workers_share_the_admitted_series(self, capped: CapRig) -> None:
|
||||
"""H9: fresh connections reach both workers, and each worker's own sample file names only the warm aliases
|
||||
while counting the fourth key on `other`, since the admitted sets live in the shared directory."""
|
||||
extra: Final = capped.key("h9")
|
||||
|
||||
def send_on_a_fresh_connection() -> tuple[WorkerSamples, ...]:
|
||||
assert capped.chat(extra, Call.new()).status_code == 200
|
||||
return worker_samples(capped.prom_dir, REQUESTS)
|
||||
|
||||
workers: Final = eventually(
|
||||
send_on_a_fresh_connection,
|
||||
lambda found: (
|
||||
sum(1 for worker in found if worker.overflow > 0) >= 2
|
||||
or any(extra.alias in worker.aliases for worker in found)
|
||||
),
|
||||
seconds=90,
|
||||
)
|
||||
assert all(extra.alias not in worker.aliases for worker in workers), workers
|
||||
assert sum(1 for worker in workers if worker.overflow > 0) >= 2, workers
|
||||
assert frozenset().union(*(worker.aliases for worker in workers)) == capped.warm_aliases, workers
|
||||
|
||||
def test_failures_past_the_cap_count_on_other(self, capped: CapRig) -> None:
|
||||
"""F1: provider failures fill the failure counter's cap with the warm keys, a fourth key's lands on `other`."""
|
||||
key: Final = capped.key("f1")
|
||||
call: Final = Call.new()
|
||||
before: Final = scrape(capped.gateway)
|
||||
capped.outage.set()
|
||||
try:
|
||||
for warm in capped.warm:
|
||||
assert capped.chat(warm, Call.new()).status_code == 500
|
||||
response: Final = capped.chat(key, call)
|
||||
finally:
|
||||
capped.outage.clear()
|
||||
assert response.status_code == 500 and PROVIDER_OUTAGE in response.text, response.text
|
||||
samples: Final = eventually(
|
||||
lambda: scrape(capped.gateway),
|
||||
lambda after: _grew(before, after, PROXY_FAILURES, 1) or key.alias in alias_values(after, PROXY_FAILURES),
|
||||
seconds=60,
|
||||
)
|
||||
assert key.alias not in label_values(samples)
|
||||
assert overflow_total(samples, PROXY_FAILURES) - overflow_total(before, PROXY_FAILURES) == 1
|
||||
assert alias_values(samples, PROXY_FAILURES) == capped.warm_aliases
|
||||
expect_spend_rows(key.alias, (), (call.call_id,))
|
||||
|
||||
def test_config_update_cannot_lift_a_yaml_cap(self, capped: CapRig) -> None:
|
||||
"""E1: /config/update refuses the YAML-owned cap, so a fourth key still lands on `other`."""
|
||||
response: Final = capped.gateway.client.post(
|
||||
"/config/update",
|
||||
json={"litellm_settings": {"prometheus_metrics_max_series_per_metric": 50}},
|
||||
headers={"Authorization": f"Bearer {capped.gateway.key}"},
|
||||
)
|
||||
assert response.status_code == 400, response.text
|
||||
key: Final = capped.key("e1")
|
||||
before: Final = scrape(capped.gateway)
|
||||
assert capped.chat(key, Call.new()).status_code == 200
|
||||
samples: Final = _overflowed(capped, key, before, 1)
|
||||
assert key.alias not in label_values(samples)
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def ttl(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
|
||||
with series_cap_rig(
|
||||
tmp_path_factory.mktemp("series-ttl"),
|
||||
{
|
||||
"prometheus_metrics_max_series_per_metric": 2,
|
||||
"prometheus_metrics_ttl_seconds": TTL_SECONDS,
|
||||
"prometheus_metrics_cleanup_interval_seconds": CLEANUP_SECONDS,
|
||||
},
|
||||
workers=1,
|
||||
warm_keys=2,
|
||||
) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
class TestTtl:
|
||||
def test_idle_series_expire_and_free_their_slot(self, ttl: CapRig) -> None:
|
||||
"""T1: a third key lands on `other`; once the idle first key expires, a new key gets its own series."""
|
||||
first, second = ttl.warm
|
||||
extra: Final = ttl.key("t1-extra")
|
||||
assert ttl.chat(extra, Call.new()).status_code == 200
|
||||
samples: Final = eventually(
|
||||
lambda: scrape(ttl.gateway),
|
||||
lambda after: overflow_total(after, REQUESTS) >= 1 or extra.alias in alias_values(after, REQUESTS),
|
||||
seconds=60,
|
||||
)
|
||||
assert extra.alias not in label_values(samples)
|
||||
|
||||
def keep_second_busy() -> tuple[Sample, ...]:
|
||||
assert ttl.chat(second, Call.new()).status_code == 200
|
||||
return scrape(ttl.gateway)
|
||||
|
||||
expired: Final = eventually(keep_second_busy, lambda after: first.alias not in label_values(after), seconds=30)
|
||||
assert second.alias in alias_values(expired, REQUESTS)
|
||||
late: Final = ttl.key("t1-late")
|
||||
assert ttl.chat(late, Call.new()).status_code == 200
|
||||
named: Final = eventually(
|
||||
lambda: scrape(ttl.gateway), lambda after: late.alias in alias_values(after, REQUESTS), seconds=30
|
||||
)
|
||||
assert late.alias in alias_values(named, REQUESTS)
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def ttl_multiproc(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
|
||||
with series_cap_rig(
|
||||
tmp_path_factory.mktemp("series-ttl-multiproc"),
|
||||
{
|
||||
"prometheus_metrics_max_series_per_metric": CAP,
|
||||
"prometheus_metrics_ttl_seconds": TTL_SECONDS,
|
||||
"prometheus_metrics_cleanup_interval_seconds": CLEANUP_SECONDS,
|
||||
},
|
||||
workers=2,
|
||||
warm_keys=3,
|
||||
) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
class TestTtlMultiproc:
|
||||
def test_ttl_is_ignored_with_two_workers_while_the_cap_applies(self, ttl_multiproc: CapRig) -> None:
|
||||
"""M1: with two workers an idle key keeps its series past the TTL, the cap still applies, and the log says so."""
|
||||
first, second, _ = ttl_multiproc.warm
|
||||
deadline: Final = time.monotonic() + 2 * TTL_SECONDS
|
||||
while time.monotonic() < deadline:
|
||||
assert ttl_multiproc.chat(second, Call.new()).status_code == 200
|
||||
assert first.alias in alias_values(scrape(ttl_multiproc.gateway), REQUESTS)
|
||||
extra: Final = ttl_multiproc.key("m1")
|
||||
before: Final = scrape(ttl_multiproc.gateway)
|
||||
assert ttl_multiproc.chat(extra, Call.new()).status_code == 200
|
||||
samples: Final = _overflowed(ttl_multiproc, extra, before, 1)
|
||||
assert extra.alias not in label_values(samples)
|
||||
assert TTL_IGNORED_WARNING in ttl_multiproc.proxy.log.read_text()
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def ignored(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
|
||||
with series_cap_rig(
|
||||
tmp_path_factory.mktemp("series-ignored"),
|
||||
{"prometheus_metrics_max_series_per_metric": "five", "prometheus_metrics_ttl_seconds": ""},
|
||||
workers=1,
|
||||
warm_keys=5,
|
||||
) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
class TestIgnored:
|
||||
def test_settings_that_are_not_positive_numbers_are_ignored_with_a_warning(self, ignored: CapRig) -> None:
|
||||
"""I1: a cap of "five" and an empty TTL leave every key its own series and each warning names its setting."""
|
||||
samples: Final = scrape(ignored.gateway)
|
||||
assert alias_values(samples, REQUESTS) >= ignored.warm_aliases
|
||||
assert not any(sample.is_overflow() for sample in samples)
|
||||
log: Final = ignored.proxy.log.read_text()
|
||||
assert (
|
||||
"prometheus_metrics_max_series_per_metric is ignored because it is not a number greater than 0 (got 'five')"
|
||||
in log
|
||||
)
|
||||
assert "prometheus_metrics_ttl_seconds is ignored because it is not a number greater than 0 (got '')" in log
|
||||
|
||||
|
||||
def _fallback_deployments(provider_url: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
{
|
||||
"model_name": name,
|
||||
"litellm_params": {
|
||||
"model": f"openai/gpt-{name}",
|
||||
"api_base": provider_url + "/v1",
|
||||
"api_key": "synthetic-provider-key",
|
||||
},
|
||||
}
|
||||
for name in (PRIMARY, FALLBACK)
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def excluded(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
|
||||
with series_cap_rig(
|
||||
tmp_path_factory.mktemp("series-excluded"),
|
||||
{"prometheus_metrics_max_series_per_metric": CAP, "prometheus_exclude_labels": ["api_key_alias"]},
|
||||
workers=1,
|
||||
warm_keys=0,
|
||||
failing_models=frozenset({f"gpt-{PRIMARY}"}),
|
||||
deployments=_fallback_deployments,
|
||||
router_settings={"fallbacks": [{PRIMARY: [FALLBACK]}]},
|
||||
) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
class TestExcluded:
|
||||
def test_fallback_counters_drop_excluded_labels(self, excluded: CapRig) -> None:
|
||||
"""X1: the successful and failed fallback counters honor prometheus_exclude_labels like every other metric."""
|
||||
key: Final = excluded.key("x1")
|
||||
call: Final = Call.new()
|
||||
response: Final = chat_once(excluded.base_url, key, PRIMARY, call)
|
||||
assert response.status_code == 200 and call.answer in response.text, response.text
|
||||
after_success: Final = eventually(
|
||||
lambda: scrape(excluded.gateway),
|
||||
lambda samples: any(sample.name == SUCCESSFUL_FALLBACKS for sample in samples),
|
||||
seconds=60,
|
||||
)
|
||||
successes: Final = tuple(sample for sample in after_success if sample.name == SUCCESSFUL_FALLBACKS)
|
||||
assert any(sample.labels.get("fallback_model") == FALLBACK for sample in successes), successes
|
||||
assert all("api_key_alias" not in sample.labels for sample in successes), successes
|
||||
excluded.outage.set()
|
||||
try:
|
||||
failed: Final = chat_once(excluded.base_url, key, PRIMARY, Call.new())
|
||||
finally:
|
||||
excluded.outage.clear()
|
||||
assert failed.status_code == 500, failed.text
|
||||
after_failure: Final = eventually(
|
||||
lambda: scrape(excluded.gateway),
|
||||
lambda samples: any(sample.name == FAILED_FALLBACKS for sample in samples),
|
||||
seconds=60,
|
||||
)
|
||||
failures: Final = tuple(sample for sample in after_failure if sample.name == FAILED_FALLBACKS)
|
||||
assert any(sample.labels.get("fallback_model") == FALLBACK for sample in failures), failures
|
||||
assert all("api_key_alias" not in sample.labels for sample in failures), failures
|
||||
|
||||
|
||||
@pytest.fixture(scope="class")
|
||||
def nocap(tmp_path_factory: pytest.TempPathFactory) -> Iterator[CapRig]:
|
||||
with series_cap_rig(
|
||||
tmp_path_factory.mktemp("series-nocap"),
|
||||
{"prometheus_metrics_max_series_per_metric": None},
|
||||
workers=2,
|
||||
warm_keys=4,
|
||||
) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
class TestNoCap:
|
||||
def test_a_null_cap_keeps_every_series(self, nocap: CapRig) -> None:
|
||||
"""N1: an explicit null cap and a missing TTL leave every key its own series and no `other` series."""
|
||||
samples: Final = scrape(nocap.gateway)
|
||||
assert alias_values(samples, REQUESTS) >= nocap.warm_aliases
|
||||
assert not any(sample.is_overflow() for sample in samples)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue