mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
feat(observability): expose Prisma connection pool saturation metrics
Operators could see latency and database CPU symptoms during an incident but could not tell whether the proxy had run out of connections, because nothing exposed the pool. Ten bounded-cardinality metrics now bridge the Prisma query engine's own pool counters into Prometheus, keeping the time a query spent waiting for a slot separate from the time it spent executing. The configured maximum is derived from the engine as busy + idle rather than parsed out of DATABASE_URL, so no credential is read on this path. Sampling rides on database work instead of a scheduled job, and the exhaustion counter is incremented at the error site, so both keep reporting through the window a delayed exporter would erase. Refs LIT-5435
This commit is contained in:
parent
2bd897a49d
commit
9eaee191d1
14 changed files with 1203 additions and 3 deletions
|
|
@ -1507,6 +1507,12 @@ SPEND_LOG_QUEUE_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_QUEUE_MAX_BYT
|
|||
SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0))
|
||||
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
|
||||
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
|
||||
# Floor on how often the Prisma connection-pool counters are read. The read is a
|
||||
# local HTTP call to the query engine, and it rides along with real DB work, so
|
||||
# this bounds the added load rather than setting a publish cadence.
|
||||
DB_POOL_METRICS_MIN_SAMPLE_INTERVAL_SECONDS: Final = 10.0
|
||||
# Ceiling on how long a pool sample may hold up the database call it rides on.
|
||||
DB_POOL_METRICS_SAMPLE_TIMEOUT_SECONDS: Final = 1.0
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
|
||||
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))
|
||||
RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN", "100")))
|
||||
|
|
|
|||
|
|
@ -58,6 +58,8 @@ from litellm.types.utils import (
|
|||
if TYPE_CHECKING:
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from prometheus_client.metrics import MetricWrapperBase
|
||||
|
||||
from litellm.proxy.db.db_pool_metrics import DBPoolMetricsUpdate
|
||||
else:
|
||||
AsyncIOScheduler = Any
|
||||
|
||||
|
|
@ -691,6 +693,83 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=[],
|
||||
)
|
||||
|
||||
########################################
|
||||
# Database connection pool saturation
|
||||
########################################
|
||||
self.litellm_db_pool_connections_max = self._gauge_factory(
|
||||
"litellm_db_pool_connections_max",
|
||||
"Configured maximum size of the Prisma connection pool, summed across live workers",
|
||||
labelnames=(),
|
||||
multiprocess_mode="livesum",
|
||||
)
|
||||
|
||||
self.litellm_db_pool_connections_busy = self._gauge_factory(
|
||||
"litellm_db_pool_connections_busy",
|
||||
"Pool connections currently checked out by a query, summed across live workers",
|
||||
labelnames=(),
|
||||
multiprocess_mode="livesum",
|
||||
)
|
||||
|
||||
self.litellm_db_pool_connections_idle = self._gauge_factory(
|
||||
"litellm_db_pool_connections_idle",
|
||||
"Pool capacity not currently checked out, summed across live workers",
|
||||
labelnames=(),
|
||||
multiprocess_mode="livesum",
|
||||
)
|
||||
|
||||
self.litellm_db_pool_connections_open = self._gauge_factory(
|
||||
"litellm_db_pool_connections_open",
|
||||
"Connections actually opened to the database, summed across live workers",
|
||||
labelnames=(),
|
||||
multiprocess_mode="livesum",
|
||||
)
|
||||
|
||||
self.litellm_db_pool_pending_acquirers = self._gauge_factory(
|
||||
"litellm_db_pool_pending_acquirers",
|
||||
"Queries waiting for a free pool connection, summed across live workers",
|
||||
labelnames=(),
|
||||
multiprocess_mode="livesum",
|
||||
)
|
||||
|
||||
self.litellm_db_pool_acquire_wait_seconds_total = self._counter_factory(
|
||||
name="litellm_db_pool_acquire_wait_seconds_total",
|
||||
documentation=(
|
||||
"Cumulative seconds spent waiting for a pool connection, excluding query execution. "
|
||||
"Divide by litellm_db_pool_acquire_total for mean acquire latency. Counts only waits that ended in an acquisition, so it understates during exhaustion; pair it with litellm_db_pool_timeouts_total"
|
||||
),
|
||||
labelnames=(),
|
||||
)
|
||||
|
||||
self.litellm_db_pool_acquire_total = self._counter_factory(
|
||||
name="litellm_db_pool_acquire_total",
|
||||
documentation="Total pool connection acquisitions on this worker",
|
||||
labelnames=(),
|
||||
)
|
||||
|
||||
self.litellm_db_query_duration_seconds_total = self._counter_factory(
|
||||
name="litellm_db_query_duration_seconds_total",
|
||||
documentation=(
|
||||
"Cumulative seconds spent executing queries against the database, excluding pool wait. "
|
||||
"Divide by litellm_db_query_total for mean query latency"
|
||||
),
|
||||
labelnames=(),
|
||||
)
|
||||
|
||||
self.litellm_db_query_total = self._counter_factory(
|
||||
name="litellm_db_query_total",
|
||||
documentation="Total queries executed against the database by this worker",
|
||||
labelnames=(),
|
||||
)
|
||||
|
||||
self.litellm_db_pool_timeouts_total = self._counter_factory(
|
||||
name="litellm_db_pool_timeouts_total",
|
||||
documentation=(
|
||||
"Total queries that gave up waiting for a pool connection (prisma P2024). "
|
||||
"Any nonzero rate means the pool is exhausted"
|
||||
),
|
||||
labelnames=(),
|
||||
)
|
||||
|
||||
########################################
|
||||
# MCP Tool Call Metrics
|
||||
########################################
|
||||
|
|
@ -3056,6 +3135,25 @@ class PrometheusLogger(CustomLogger):
|
|||
except Exception as e:
|
||||
verbose_logger.warning("Error recording check batch cost metrics: %s", e)
|
||||
|
||||
def record_db_pool_sample(self, update: DBPoolMetricsUpdate) -> None:
|
||||
"""Publish one reading of this worker's Prisma connection pool."""
|
||||
sample: Final = update.sample
|
||||
self.litellm_db_pool_connections_max.set(sample.max_connections)
|
||||
self.litellm_db_pool_connections_busy.set(sample.busy_connections)
|
||||
self.litellm_db_pool_connections_idle.set(sample.idle_connections)
|
||||
self.litellm_db_pool_connections_open.set(sample.open_connections)
|
||||
self.litellm_db_pool_pending_acquirers.set(update.pending_acquirers)
|
||||
# Deltas are non-negative by construction: DBPoolMetricsSampler zeroes
|
||||
# them when the engine's totals move backwards.
|
||||
self.litellm_db_pool_acquire_total.inc(update.acquire_count_delta)
|
||||
self.litellm_db_pool_acquire_wait_seconds_total.inc(update.acquire_wait_seconds_delta)
|
||||
self.litellm_db_query_total.inc(update.query_count_delta)
|
||||
self.litellm_db_query_duration_seconds_total.inc(update.query_duration_seconds_delta)
|
||||
|
||||
def record_db_pool_timeout(self) -> None:
|
||||
"""Count one query that gave up waiting for a free pool connection."""
|
||||
self.litellm_db_pool_timeouts_total.inc()
|
||||
|
||||
def record_check_batch_cost_error(self, error_type: str):
|
||||
try:
|
||||
self.litellm_check_batch_cost_errors_total.labels(
|
||||
|
|
|
|||
231
litellm/proxy/db/db_pool_metrics.py
Normal file
231
litellm/proxy/db/db_pool_metrics.py
Normal file
|
|
@ -0,0 +1,231 @@
|
|||
"""Bridges the Prisma query engine's connection-pool counters into Prometheus.
|
||||
|
||||
The query engine tracks pool occupancy and, separately, how long a query spent
|
||||
waiting for a pool slot versus executing against the database. That split is the
|
||||
only way to tell "the database is slow" apart from "we ran out of connections",
|
||||
so it is sampled here and re-published under ``litellm_db_pool_*`` names.
|
||||
|
||||
Sampling is driven from the DB call path rather than from a scheduled job: an
|
||||
exporter that runs on a timer stops reporting exactly when the event loop is
|
||||
saturated, which is the window an operator most needs. Requests are what stall
|
||||
during pool exhaustion, and they keep arriving, so a throttled sample taken
|
||||
alongside real DB work stays alive through the incident.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
DB_POOL_METRICS_MIN_SAMPLE_INTERVAL_SECONDS,
|
||||
DB_POOL_METRICS_SAMPLE_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import Metrics
|
||||
|
||||
_MILLISECONDS_PER_SECOND: Final = 1000.0
|
||||
|
||||
_POOL_BUSY_KEY: Final = "prisma_pool_connections_busy"
|
||||
_POOL_IDLE_KEY: Final = "prisma_pool_connections_idle"
|
||||
_POOL_OPEN_KEY: Final = "prisma_pool_connections_open"
|
||||
_PENDING_ACQUIRERS_KEY: Final = "prisma_client_queries_wait"
|
||||
_ACQUIRE_WAIT_HISTOGRAM_KEY: Final = "prisma_client_queries_wait_histogram_ms"
|
||||
_QUERY_DURATION_HISTOGRAM_KEY: Final = "prisma_datasource_queries_duration_histogram_ms"
|
||||
|
||||
|
||||
class SupportsPoolSample(Protocol):
|
||||
"""The one method this module needs from a Prisma client wrapper.
|
||||
|
||||
Deliberately not ``runtime_checkable``: ``PrismaWrapper`` resolves most of
|
||||
its surface through an instance-level ``__getattr__``, which an
|
||||
``isinstance`` check against a protocol does not see, so a runtime check
|
||||
here would reject the very object it exists to accept.
|
||||
"""
|
||||
|
||||
async def get_pool_sample(self) -> DBPoolSample: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DBPoolSample:
|
||||
"""One reading of the engine's pool counters.
|
||||
|
||||
``max_connections`` is derived as ``busy + idle`` rather than parsed out of
|
||||
``DATABASE_URL``. The engine reports ``idle`` as remaining capacity, so the
|
||||
sum is the configured ``connection_limit``, and deriving it here keeps the
|
||||
database credentials out of this path entirely.
|
||||
"""
|
||||
|
||||
busy_connections: float
|
||||
idle_connections: float
|
||||
open_connections: float
|
||||
pending_acquirers: float
|
||||
acquire_wait_seconds_total: float
|
||||
acquire_count_total: int
|
||||
query_duration_seconds_total: float
|
||||
query_count_total: int
|
||||
|
||||
@property
|
||||
def max_connections(self) -> float:
|
||||
return self.busy_connections + self.idle_connections
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DBPoolMetricsUpdate:
|
||||
"""A sample plus the cumulative movement since the previous sample.
|
||||
|
||||
The engine reports totals; Prometheus counters take increments. The deltas
|
||||
are the difference between consecutive samples, and are omitted entirely
|
||||
when the engine's totals move backwards, which happens when the query engine
|
||||
restarts and its counters reset.
|
||||
"""
|
||||
|
||||
sample: DBPoolSample
|
||||
pending_acquirers: float
|
||||
acquire_wait_seconds_delta: float
|
||||
acquire_count_delta: int
|
||||
query_duration_seconds_delta: float
|
||||
query_count_delta: int
|
||||
|
||||
|
||||
def _gauge_value(metrics: Metrics, key: str) -> float:
|
||||
for gauge in metrics.gauges:
|
||||
if gauge.key == key:
|
||||
return float(gauge.value)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _histogram_seconds_and_count(metrics: Metrics, key: str) -> tuple[float, int]:
|
||||
for histogram in metrics.histograms:
|
||||
if histogram.key == key:
|
||||
return (histogram.value.sum / _MILLISECONDS_PER_SECOND, histogram.value.count)
|
||||
return (0.0, 0)
|
||||
|
||||
|
||||
def parse_pool_sample(metrics: Metrics) -> DBPoolSample:
|
||||
acquire_wait_seconds, acquire_count = _histogram_seconds_and_count(metrics, _ACQUIRE_WAIT_HISTOGRAM_KEY)
|
||||
query_duration_seconds, query_count = _histogram_seconds_and_count(metrics, _QUERY_DURATION_HISTOGRAM_KEY)
|
||||
return DBPoolSample(
|
||||
busy_connections=_gauge_value(metrics, _POOL_BUSY_KEY),
|
||||
idle_connections=_gauge_value(metrics, _POOL_IDLE_KEY),
|
||||
open_connections=_gauge_value(metrics, _POOL_OPEN_KEY),
|
||||
pending_acquirers=_gauge_value(metrics, _PENDING_ACQUIRERS_KEY),
|
||||
acquire_wait_seconds_total=acquire_wait_seconds,
|
||||
acquire_count_total=acquire_count,
|
||||
query_duration_seconds_total=query_duration_seconds,
|
||||
query_count_total=query_count,
|
||||
)
|
||||
|
||||
|
||||
class DBPoolMetricsSampler:
|
||||
"""Throttled reader of the engine's pool counters.
|
||||
|
||||
One instance per process. ``maybe_sample`` is safe to call on every DB
|
||||
operation: it returns ``None`` without touching the engine until
|
||||
``min_interval_seconds`` has elapsed since the last reading.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
min_interval_seconds: float = DB_POOL_METRICS_MIN_SAMPLE_INTERVAL_SECONDS,
|
||||
monotonic: Callable[[], float] = time.monotonic,
|
||||
) -> None:
|
||||
self._min_interval_seconds: Final = min_interval_seconds
|
||||
self._monotonic: Final = monotonic
|
||||
self._last_sampled_at: float | None = None
|
||||
self._previous: DBPoolSample | None = None
|
||||
self._pending_baseline: float = 0.0
|
||||
|
||||
def is_due(self) -> bool:
|
||||
"""Whether enough time has passed to justify reading the engine again."""
|
||||
if self._last_sampled_at is None:
|
||||
return True
|
||||
return (self._monotonic() - self._last_sampled_at) >= self._min_interval_seconds
|
||||
|
||||
async def maybe_sample(self, resolve_client: Callable[[], SupportsPoolSample | None]) -> DBPoolMetricsUpdate | None:
|
||||
"""Read the engine's counters, or return ``None`` if not yet due.
|
||||
|
||||
The interval is consumed as soon as an attempt starts, before the client
|
||||
is resolved, so a deployment that has no client to sample throttles the
|
||||
same as one that does rather than retrying on every database call.
|
||||
|
||||
Never raises, and never blocks for long: a wedged query engine costs a
|
||||
bounded wait. A pool sample is diagnostic, and an engine that cannot
|
||||
answer one is already the subject of a louder alarm.
|
||||
|
||||
The client is resolved through a callable rather than passed in so that
|
||||
resolution, which reads proxy module state, happens under the same
|
||||
throttle and the same exception guard as the read itself.
|
||||
"""
|
||||
if not self.is_due():
|
||||
return None
|
||||
self._last_sampled_at = self._monotonic()
|
||||
|
||||
try:
|
||||
client: Final = resolve_client()
|
||||
if client is None:
|
||||
return None
|
||||
sample: Final = await asyncio.wait_for(
|
||||
client.get_pool_sample(), timeout=DB_POOL_METRICS_SAMPLE_TIMEOUT_SECONDS
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a diagnostic read must not fail the database call it rides on
|
||||
verbose_proxy_logger.debug("db pool metrics sample failed: %s", e)
|
||||
return None
|
||||
|
||||
update: Final = self._to_update(sample)
|
||||
self._previous = sample
|
||||
return update
|
||||
|
||||
def _corrected_pending_acquirers(self, sample: DBPoolSample) -> float:
|
||||
"""The engine's waiter gauge, corrected for waiters that timed out.
|
||||
|
||||
``prisma_client_queries_wait`` is decremented when a waiter acquires a
|
||||
connection but not when it gives up, so every pool timeout latches the
|
||||
gauge one higher for the life of the process. Observed directly against
|
||||
a live engine at ``connection_limit=2``: after two bursts that produced
|
||||
seven P2024s, a fully drained pool still reported seven waiters.
|
||||
|
||||
Free capacity is proof that nobody is queued, so any reading taken while
|
||||
a connection is idle is exactly the accumulated latch, and subtracting it
|
||||
recovers the real depth. The baseline re-arms on every idle sample, so
|
||||
repeated exhaustion stays accurate rather than drifting further each time.
|
||||
"""
|
||||
if sample.idle_connections > 0:
|
||||
self._pending_baseline = sample.pending_acquirers
|
||||
return 0.0
|
||||
return max(0.0, sample.pending_acquirers - self._pending_baseline)
|
||||
|
||||
def _to_update(self, sample: DBPoolSample) -> DBPoolMetricsUpdate:
|
||||
pending: Final = self._corrected_pending_acquirers(sample)
|
||||
previous: Final = self._previous
|
||||
engine_restarted: Final = previous is not None and (
|
||||
sample.acquire_count_total < previous.acquire_count_total
|
||||
or sample.query_count_total < previous.query_count_total
|
||||
)
|
||||
if previous is None or engine_restarted:
|
||||
return DBPoolMetricsUpdate(
|
||||
sample=sample,
|
||||
pending_acquirers=pending,
|
||||
acquire_wait_seconds_delta=0.0,
|
||||
acquire_count_delta=0,
|
||||
query_duration_seconds_delta=0.0,
|
||||
query_count_delta=0,
|
||||
)
|
||||
return DBPoolMetricsUpdate(
|
||||
sample=sample,
|
||||
pending_acquirers=pending,
|
||||
acquire_wait_seconds_delta=max(
|
||||
0.0, sample.acquire_wait_seconds_total - previous.acquire_wait_seconds_total
|
||||
),
|
||||
acquire_count_delta=sample.acquire_count_total - previous.acquire_count_total,
|
||||
query_duration_seconds_delta=max(
|
||||
0.0, sample.query_duration_seconds_total - previous.query_duration_seconds_total
|
||||
),
|
||||
query_count_delta=sample.query_count_total - previous.query_count_total,
|
||||
)
|
||||
|
|
@ -123,6 +123,23 @@ class PrismaDBExceptionHandler:
|
|||
|
||||
return type(e) is prisma.errors.DataError
|
||||
|
||||
@staticmethod
|
||||
def is_connection_pool_timeout_error(e: Exception) -> bool:
|
||||
"""True iff the query engine gave up waiting for a free pool slot (P2024).
|
||||
|
||||
Matched on the prisma error code rather than the message, because the
|
||||
message is also keyword-matched as a generic transport timeout elsewhere
|
||||
in this class. The distinction matters to an operator: the database
|
||||
answered fine, the proxy simply had no connection left to ask with, so
|
||||
the fix is pool sizing or shedding load rather than anything database
|
||||
side.
|
||||
|
||||
The ``P####`` codes are prisma's own namespace and only prisma populates
|
||||
``code`` with one, so no type check is needed. Skipping it also keeps
|
||||
this callable on the DB failure path when prisma itself is stubbed.
|
||||
"""
|
||||
return getattr(e, "code", None) == "P2024"
|
||||
|
||||
@staticmethod
|
||||
def is_database_transport_error(e: Exception) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -8,10 +8,78 @@ import asyncio
|
|||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from functools import wraps
|
||||
from typing import Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceTypes
|
||||
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
|
||||
from litellm.proxy.db.db_pool_metrics import DBPoolMetricsSampler
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.proxy.db.db_pool_metrics import SupportsPoolSample
|
||||
|
||||
# One sampler per process. The pool it reads is per-process too, so there is
|
||||
# nothing to key this by.
|
||||
_pool_metrics_sampler: Final = DBPoolMetricsSampler()
|
||||
|
||||
|
||||
def _prometheus_logger() -> "PrometheusLogger | None":
|
||||
"""The active PrometheusLogger, or None when prometheus is not configured.
|
||||
|
||||
Imports lazily. Bare ``import litellm`` does not load the prometheus
|
||||
integration, so hoisting this would make every proxy pay for a module that
|
||||
only prometheus deployments use.
|
||||
"""
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
return PrometheusLogger.get_instance()
|
||||
|
||||
|
||||
def _resolve_pool_client() -> "SupportsPoolSample | None":
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
return None if prisma_client is None else prisma_client.db
|
||||
|
||||
|
||||
async def _sample_db_pool_metrics() -> None:
|
||||
"""Publish a throttled reading of the connection pool, if one is due.
|
||||
|
||||
Runs alongside real database work rather than on a timer: a scheduled
|
||||
exporter goes quiet exactly when the event loop is saturated, which is the
|
||||
window this metric exists to cover.
|
||||
|
||||
Every step is inside the guard, including resolving the client and locating
|
||||
the logger. This runs off a database call that has already succeeded, so
|
||||
nothing here may turn a working query into a failed one.
|
||||
"""
|
||||
try:
|
||||
update: Final = await _pool_metrics_sampler.maybe_sample(_resolve_pool_client)
|
||||
if update is None:
|
||||
return
|
||||
logger: Final = _prometheus_logger()
|
||||
if logger is not None:
|
||||
logger.record_db_pool_sample(update)
|
||||
except Exception as e: # noqa: BLE001 # a metrics failure must never fail the database call it rides on
|
||||
verbose_proxy_logger.debug("db pool metrics publish failed: %s", e)
|
||||
|
||||
|
||||
def _record_db_pool_timeout_if_exhausted(e: Exception) -> None:
|
||||
"""Count a pool exhaustion, without ever displacing the error that caused it.
|
||||
|
||||
The caller re-raises the original exception after this returns. Anything
|
||||
that escaped here would replace a P2024 with a metrics error, during the
|
||||
incident this counter exists to record.
|
||||
"""
|
||||
try:
|
||||
if not PrismaDBExceptionHandler.is_connection_pool_timeout_error(e):
|
||||
return
|
||||
logger: Final = _prometheus_logger()
|
||||
if logger is not None:
|
||||
logger.record_db_pool_timeout()
|
||||
except Exception as metrics_error: # noqa: BLE001 # counting an exhaustion must not mask the exhaustion itself
|
||||
verbose_proxy_logger.debug("db pool timeout metric failed: %s", metrics_error)
|
||||
|
||||
|
||||
def _safe_db_event_metadata(kwargs: dict) -> dict[str, str] | None:
|
||||
|
|
@ -52,6 +120,13 @@ def log_db_metrics(func):
|
|||
end_time: datetime = datetime.now()
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
# Dispatched, never awaited. Awaiting here would add a suspension
|
||||
# point after the query already succeeded, so a client disconnect in
|
||||
# that window would discard a completed result and skip the success
|
||||
# hook below. `is_due` keeps this to one task per sample interval.
|
||||
if _pool_metrics_sampler.is_due():
|
||||
asyncio.create_task(_sample_db_pool_metrics())
|
||||
|
||||
if "PROXY" not in func.__name__:
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.service_logging_obj.async_service_success_hook(
|
||||
|
|
@ -124,6 +199,11 @@ async def _handle_logging_db_exception(
|
|||
) -> None:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
# Counted before the DB-relatedness gate below: a pool timeout is the proxy
|
||||
# running out of connections, so it must be recorded whether or not the
|
||||
# failure is classified as a DB service failure.
|
||||
_record_db_pool_timeout_if_exhausted(e)
|
||||
|
||||
# don't log this as a DB Service Failure, if the DB did not raise an exception
|
||||
if _is_exception_related_to_db(e) is not True:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from datetime import datetime, timedelta
|
|||
from typing import Any, Final, Protocol
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.db.db_pool_metrics import DBPoolSample, parse_pool_sample
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
||||
|
||||
|
|
@ -798,6 +799,17 @@ class PrismaWrapper:
|
|||
seconds_left: Final = (expiration_time - datetime.utcnow()).total_seconds()
|
||||
return seconds_left > self.TOKEN_REFRESH_BUFFER_SECONDS
|
||||
|
||||
async def get_pool_sample(self) -> DBPoolSample:
|
||||
"""This connection pool's current occupancy, as the query engine sees it.
|
||||
|
||||
Declared rather than left to ``__getattr__``, which resolves through an
|
||||
untyped delegation and so cannot satisfy a typed caller, and which also
|
||||
makes the underlying client's identity easy to lose track of. Keeping
|
||||
the engine's raw counter names behind this boundary means the metrics
|
||||
layer never has to know them.
|
||||
"""
|
||||
return parse_pool_sample(await self._original_prisma.get_metrics())
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
"""
|
||||
Proxy attribute access to the underlying Prisma client.
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from collections.abc import Callable
|
|||
from typing import Any, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.db.db_pool_metrics import DBPoolSample
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
|
||||
# Per-model action methods that read from the database. These are routed to
|
||||
|
|
@ -131,6 +132,15 @@ class RoutingPrismaWrapper:
|
|||
reconnect attempts against an already-healthy writer."""
|
||||
self._writer_unavailable = False
|
||||
|
||||
async def get_pool_sample(self) -> DBPoolSample:
|
||||
"""The writer's pool occupancy.
|
||||
|
||||
The writer is the pool that auth, spend writes and every background job
|
||||
contend for, so it is the one that saturates. The reader has its own
|
||||
separate pool that this does not cover; see the metric docs.
|
||||
"""
|
||||
return await self._writer.get_pool_sample()
|
||||
|
||||
def _should_use_reader(self) -> bool:
|
||||
return not self._reader_unavailable
|
||||
|
||||
|
|
|
|||
|
|
@ -119,7 +119,7 @@ EXCEPTION_CLASS: Final = "exception_class"
|
|||
RATE_LIMIT_CATEGORY: Final = "rate_limit_category"
|
||||
RATE_LIMIT_TYPE: Final = "rate_limit_type"
|
||||
STATUS_CODE: Final = "status_code"
|
||||
EXCEPTION_LABELS: Final = [EXCEPTION_STATUS, EXCEPTION_CLASS]
|
||||
EXCEPTION_LABELS: Final = (EXCEPTION_STATUS, EXCEPTION_CLASS)
|
||||
LATENCY_BUCKETS: Final = (
|
||||
0.005,
|
||||
0.01,
|
||||
|
|
@ -273,6 +273,17 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
# MCP tool call metrics
|
||||
"litellm_mcp_tool_calls_total",
|
||||
"litellm_mcp_tool_call_spend_metric",
|
||||
# Database connection pool saturation
|
||||
"litellm_db_pool_connections_max",
|
||||
"litellm_db_pool_connections_busy",
|
||||
"litellm_db_pool_connections_idle",
|
||||
"litellm_db_pool_connections_open",
|
||||
"litellm_db_pool_pending_acquirers",
|
||||
"litellm_db_pool_acquire_wait_seconds_total",
|
||||
"litellm_db_pool_acquire_total",
|
||||
"litellm_db_query_duration_seconds_total",
|
||||
"litellm_db_query_total",
|
||||
"litellm_db_pool_timeouts_total",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -765,6 +776,20 @@ class PrometheusMetricLabels:
|
|||
|
||||
litellm_check_batch_cost_last_run_timestamp: list[str] = []
|
||||
|
||||
# Database connection pool saturation. Deliberately unlabelled: the pool is a
|
||||
# per-worker resource, and any key/team/user label would both be meaningless
|
||||
# here and blow up cardinality on a metric scraped from every pod.
|
||||
litellm_db_pool_connections_max: tuple[str, ...] = ()
|
||||
litellm_db_pool_connections_busy: tuple[str, ...] = ()
|
||||
litellm_db_pool_connections_idle: tuple[str, ...] = ()
|
||||
litellm_db_pool_connections_open: tuple[str, ...] = ()
|
||||
litellm_db_pool_pending_acquirers: tuple[str, ...] = ()
|
||||
litellm_db_pool_acquire_wait_seconds_total: tuple[str, ...] = ()
|
||||
litellm_db_pool_acquire_total: tuple[str, ...] = ()
|
||||
litellm_db_query_duration_seconds_total: tuple[str, ...] = ()
|
||||
litellm_db_query_total: tuple[str, ...] = ()
|
||||
litellm_db_pool_timeouts_total: tuple[str, ...] = ()
|
||||
|
||||
# MCP tool call metrics
|
||||
litellm_mcp_tool_calls_total: list[str] = [
|
||||
UserAPIKeyLabelNames.MCP_TOOL_NAME.value,
|
||||
|
|
@ -834,7 +859,7 @@ class PrometheusMetricLabels:
|
|||
if label not in default_labels and label not in custom_labels:
|
||||
custom_labels.append(label)
|
||||
|
||||
return default_labels + custom_labels
|
||||
return [*default_labels, *custom_labels]
|
||||
|
||||
|
||||
_USER_API_KEY_LABEL_VALUE_INIT_ALIASES: Final[Mapping[str, str]] = MappingProxyType(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,227 @@
|
|||
"""Prometheus surface for the Prisma connection-pool saturation metrics."""
|
||||
|
||||
from typing import get_args
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from prometheus_client import REGISTRY
|
||||
from prisma.errors import DataError
|
||||
|
||||
from litellm.proxy.db.db_pool_metrics import DBPoolMetricsUpdate, DBPoolSample
|
||||
from litellm.types.integrations.prometheus import (
|
||||
DEFINED_PROMETHEUS_METRICS,
|
||||
PrometheusMetricLabels,
|
||||
)
|
||||
|
||||
DB_POOL_METRICS = (
|
||||
"litellm_db_pool_connections_max",
|
||||
"litellm_db_pool_connections_busy",
|
||||
"litellm_db_pool_connections_idle",
|
||||
"litellm_db_pool_connections_open",
|
||||
"litellm_db_pool_pending_acquirers",
|
||||
"litellm_db_pool_acquire_wait_seconds_total",
|
||||
"litellm_db_pool_acquire_total",
|
||||
"litellm_db_query_duration_seconds_total",
|
||||
"litellm_db_query_total",
|
||||
"litellm_db_pool_timeouts_total",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metric", DB_POOL_METRICS)
|
||||
def test_metric_is_registered_so_the_config_and_exclude_lists_accept_it(metric):
|
||||
assert metric in get_args(DEFINED_PROMETHEUS_METRICS)
|
||||
# Asserted on the class attribute rather than get_labels(), which also folds
|
||||
# in whatever custom metadata labels and tags happen to be configured
|
||||
# process-wide and so depends on global state another test may have set.
|
||||
assert getattr(PrometheusMetricLabels, metric) == (), (
|
||||
f"{metric} must stay unlabelled: the pool is a per-worker resource and the ticket "
|
||||
"forbids api-key/team labels on saturation metrics"
|
||||
)
|
||||
|
||||
|
||||
def _update(**kwargs) -> DBPoolMetricsUpdate:
|
||||
sample = DBPoolSample(
|
||||
busy_connections=kwargs.pop("busy", 0.0),
|
||||
idle_connections=kwargs.pop("idle", 0.0),
|
||||
open_connections=kwargs.pop("open_", 0.0),
|
||||
pending_acquirers=kwargs.pop("pending", 0.0),
|
||||
acquire_wait_seconds_total=0.0,
|
||||
acquire_count_total=0,
|
||||
query_duration_seconds_total=0.0,
|
||||
query_count_total=0,
|
||||
)
|
||||
return DBPoolMetricsUpdate(
|
||||
sample=sample,
|
||||
pending_acquirers=sample.pending_acquirers,
|
||||
acquire_wait_seconds_delta=kwargs.pop("wait_delta", 0.0),
|
||||
acquire_count_delta=kwargs.pop("acquire_delta", 0),
|
||||
query_duration_seconds_delta=kwargs.pop("query_seconds_delta", 0.0),
|
||||
query_count_delta=kwargs.pop("query_delta", 0),
|
||||
)
|
||||
|
||||
|
||||
def _clear_registry() -> None:
|
||||
for collector in list(REGISTRY._collector_to_names.keys()):
|
||||
try:
|
||||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def logger():
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
_clear_registry()
|
||||
yield PrometheusLogger()
|
||||
_clear_registry()
|
||||
|
||||
|
||||
def _value(name: str) -> float:
|
||||
return REGISTRY.get_sample_value(name) or 0.0
|
||||
|
||||
|
||||
def test_a_saturated_pool_is_visible_in_the_gauges(logger):
|
||||
logger.record_db_pool_sample(_update(busy=5.0, idle=0.0, open_=5.0, pending=3.0))
|
||||
|
||||
assert _value("litellm_db_pool_connections_busy") == 5.0
|
||||
assert _value("litellm_db_pool_connections_idle") == 0.0
|
||||
assert _value("litellm_db_pool_connections_open") == 5.0
|
||||
assert _value("litellm_db_pool_pending_acquirers") == 3.0
|
||||
assert _value("litellm_db_pool_connections_max") == 5.0, "busy + idle is the configured limit"
|
||||
|
||||
|
||||
def test_pool_wait_and_query_execution_are_counted_separately(logger):
|
||||
logger.record_db_pool_sample(
|
||||
_update(wait_delta=1.5, acquire_delta=6, query_seconds_delta=0.4, query_delta=4)
|
||||
)
|
||||
|
||||
assert _value("litellm_db_pool_acquire_wait_seconds_total") == pytest.approx(1.5)
|
||||
assert _value("litellm_db_pool_acquire_total") == 6.0
|
||||
assert _value("litellm_db_query_duration_seconds_total") == pytest.approx(0.4)
|
||||
assert _value("litellm_db_query_total") == 4.0
|
||||
|
||||
|
||||
def test_counters_accumulate_across_samples(logger):
|
||||
logger.record_db_pool_sample(_update(acquire_delta=6, wait_delta=1.5))
|
||||
logger.record_db_pool_sample(_update(acquire_delta=0, wait_delta=0.0))
|
||||
logger.record_db_pool_sample(_update(acquire_delta=4, wait_delta=0.5))
|
||||
|
||||
assert _value("litellm_db_pool_acquire_total") == 10.0
|
||||
assert _value("litellm_db_pool_acquire_wait_seconds_total") == pytest.approx(2.0)
|
||||
|
||||
|
||||
def test_pool_timeouts_are_counted(logger):
|
||||
logger.record_db_pool_timeout()
|
||||
logger.record_db_pool_timeout()
|
||||
|
||||
assert _value("litellm_db_pool_timeouts_total") == 2.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_broken_metric_never_fails_the_database_call_it_rides_on(logger):
|
||||
"""The publish path sits inline on a DB call, so the safety boundary lives
|
||||
in log_db_metrics rather than in each recorder."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy.db.log_db_metrics import (
|
||||
_record_db_pool_timeout_if_exhausted,
|
||||
_sample_db_pool_metrics,
|
||||
)
|
||||
|
||||
logger.record_db_pool_sample = MagicMock(side_effect=RuntimeError("registry gone"))
|
||||
logger.record_db_pool_timeout = MagicMock(side_effect=RuntimeError("registry gone"))
|
||||
|
||||
prisma_client = MagicMock()
|
||||
|
||||
async def _sample():
|
||||
raise RuntimeError("engine gone")
|
||||
|
||||
prisma_client.db.get_pool_sample = _sample
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.db.log_db_metrics._prometheus_logger", return_value=logger),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
):
|
||||
await _sample_db_pool_metrics()
|
||||
_record_db_pool_timeout_if_exhausted(
|
||||
DataError({"user_facing_error": {"error_code": "P2024", "message": "pool timeout"}})
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_unimportable_prometheus_module_does_not_break_the_db_call():
|
||||
"""_prometheus_logger imports lazily. If that import fails, a database call
|
||||
that already succeeded must still succeed, and a pool timeout must still
|
||||
surface as the original P2024 rather than as an ImportError."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy.db.log_db_metrics import (
|
||||
_record_db_pool_timeout_if_exhausted,
|
||||
_sample_db_pool_metrics,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.db.log_db_metrics._prometheus_logger",
|
||||
side_effect=ImportError("prometheus integration unavailable"),
|
||||
):
|
||||
await _sample_db_pool_metrics()
|
||||
_record_db_pool_timeout_if_exhausted(
|
||||
DataError({"user_facing_error": {"error_code": "P2024", "message": "pool timeout"}})
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metric", DB_POOL_METRICS)
|
||||
def test_the_emitted_series_carries_no_labels(logger, metric):
|
||||
"""The ticket forbids unbounded labels on saturation metrics. Asserted on
|
||||
the constructed metric rather than on PrometheusMetricLabels, which these
|
||||
metrics never consult, so this fails if someone adds a labelname."""
|
||||
assert getattr(logger, metric)._labelnames == ()
|
||||
|
||||
|
||||
POOL_GAUGES = (
|
||||
"litellm_db_pool_connections_max",
|
||||
"litellm_db_pool_connections_busy",
|
||||
"litellm_db_pool_connections_idle",
|
||||
"litellm_db_pool_connections_open",
|
||||
"litellm_db_pool_pending_acquirers",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metric", POOL_GAUGES)
|
||||
def test_pool_gauges_aggregate_across_workers_instead_of_fanning_out_per_pid(logger, metric):
|
||||
"""Under PROMETHEUS_MULTIPROC_DIR a Gauge with no multiprocess_mode defaults
|
||||
to 'all', which stamps an unbounded pid label on every series and never
|
||||
aggregates. Worse, mark_process_dead only reaps gauge_live* files, so a
|
||||
recycled worker leaves its last reading in every later scrape. livesum sums
|
||||
over live workers and drops dead ones, which is what a pod-level pool number
|
||||
means."""
|
||||
assert getattr(logger, metric)._multiprocess_mode == "livesum"
|
||||
|
||||
|
||||
def test_the_gauge_publishes_the_corrected_pending_not_the_raw_engine_reading(logger):
|
||||
"""The engine's waiter gauge latches on pool timeouts, so the sampler
|
||||
corrects it. Publishing sample.pending_acquirers instead would put the
|
||||
latched value on the dashboard."""
|
||||
latched = DBPoolSample(
|
||||
busy_connections=0.0,
|
||||
idle_connections=2.0,
|
||||
open_connections=2.0,
|
||||
pending_acquirers=7.0,
|
||||
acquire_wait_seconds_total=0.0,
|
||||
acquire_count_total=0,
|
||||
query_duration_seconds_total=0.0,
|
||||
query_count_total=0,
|
||||
)
|
||||
logger.record_db_pool_sample(
|
||||
DBPoolMetricsUpdate(
|
||||
sample=latched,
|
||||
pending_acquirers=0.0,
|
||||
acquire_wait_seconds_delta=0.0,
|
||||
acquire_count_delta=0,
|
||||
query_duration_seconds_delta=0.0,
|
||||
query_count_delta=0,
|
||||
)
|
||||
)
|
||||
|
||||
assert _value("litellm_db_pool_pending_acquirers") == 0.0
|
||||
265
tests/test_litellm/proxy/db/test_db_pool_metrics.py
Normal file
265
tests/test_litellm/proxy/db/test_db_pool_metrics.py
Normal file
|
|
@ -0,0 +1,265 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../.."))
|
||||
|
||||
from litellm.proxy.db.db_pool_metrics import DBPoolMetricsSampler, parse_pool_sample
|
||||
|
||||
|
||||
class _Value:
|
||||
def __init__(self, key, value):
|
||||
self.key = key
|
||||
self.value = value
|
||||
self.labels = {}
|
||||
self.description = ""
|
||||
|
||||
|
||||
class _Hist:
|
||||
def __init__(self, total_sum, count):
|
||||
self.sum = total_sum
|
||||
self.count = count
|
||||
self.buckets = []
|
||||
|
||||
|
||||
class _Metrics:
|
||||
"""Stands in for prisma's Metrics model, matching the shape the engine returns."""
|
||||
|
||||
def __init__(self, *, busy=0.0, idle=0.0, open_=0.0, wait=0.0, wait_ms=0.0, waits=0, query_ms=0.0, queries=0):
|
||||
self.counters = []
|
||||
self.gauges = [
|
||||
_Value("prisma_pool_connections_busy", busy),
|
||||
_Value("prisma_pool_connections_idle", idle),
|
||||
_Value("prisma_pool_connections_open", open_),
|
||||
_Value("prisma_client_queries_wait", wait),
|
||||
]
|
||||
self.histograms = [
|
||||
_Value("prisma_client_queries_wait_histogram_ms", _Hist(wait_ms, waits)),
|
||||
_Value("prisma_datasource_queries_duration_histogram_ms", _Hist(query_ms, queries)),
|
||||
]
|
||||
|
||||
|
||||
class _Client:
|
||||
def __init__(self, *metrics, fail=False):
|
||||
self._metrics = list(metrics)
|
||||
self._fail = fail
|
||||
self.calls = 0
|
||||
|
||||
async def get_pool_sample(self):
|
||||
self.calls += 1
|
||||
if self._fail:
|
||||
raise RuntimeError("engine unreachable")
|
||||
return parse_pool_sample(self._metrics[min(self.calls - 1, len(self._metrics) - 1)])
|
||||
|
||||
|
||||
class _Clock:
|
||||
def __init__(self):
|
||||
self.now = 0.0
|
||||
|
||||
def __call__(self):
|
||||
return self.now
|
||||
|
||||
|
||||
def test_max_connections_is_busy_plus_idle_not_parsed_from_the_url():
|
||||
"""The engine reports idle as remaining capacity, so busy + idle is the
|
||||
configured connection_limit. Verified against a live engine at
|
||||
connection_limit=5 under saturation: busy=5, idle=0."""
|
||||
sample = parse_pool_sample(_Metrics(busy=5.0, idle=0.0, open_=5.0))
|
||||
assert sample.max_connections == 5.0
|
||||
|
||||
sample = parse_pool_sample(_Metrics(busy=2.0, idle=3.0, open_=5.0))
|
||||
assert sample.max_connections == 5.0
|
||||
|
||||
|
||||
def test_parse_separates_pool_wait_from_query_execution():
|
||||
sample = parse_pool_sample(_Metrics(wait=3.0, wait_ms=1500.0, waits=6, query_ms=800.0, queries=4))
|
||||
assert sample.pending_acquirers == 3.0
|
||||
assert sample.acquire_wait_seconds_total == 1.5
|
||||
assert sample.acquire_count_total == 6
|
||||
assert sample.query_duration_seconds_total == 0.8
|
||||
assert sample.query_count_total == 4
|
||||
|
||||
|
||||
def test_parse_tolerates_a_metric_the_engine_did_not_report():
|
||||
empty = _Metrics()
|
||||
empty.gauges = []
|
||||
empty.histograms = []
|
||||
sample = parse_pool_sample(empty)
|
||||
assert sample.max_connections == 0.0
|
||||
assert sample.acquire_count_total == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sampler_throttles_until_the_interval_elapses():
|
||||
clock = _Clock()
|
||||
client = _Client(_Metrics(busy=1.0, idle=9.0))
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=10.0, monotonic=clock)
|
||||
|
||||
assert await sampler.maybe_sample(lambda: client) is not None
|
||||
assert client.calls == 1
|
||||
|
||||
clock.now = 9.9
|
||||
assert sampler.is_due() is False
|
||||
assert await sampler.maybe_sample(lambda: client) is None
|
||||
assert client.calls == 1, "engine must not be re-read before the interval elapses"
|
||||
|
||||
clock.now = 10.0
|
||||
assert sampler.is_due() is True
|
||||
assert await sampler.maybe_sample(lambda: client) is not None
|
||||
assert client.calls == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cumulative_totals_are_published_as_deltas():
|
||||
clock = _Clock()
|
||||
client = _Client(
|
||||
_Metrics(wait_ms=1000.0, waits=10, query_ms=500.0, queries=5),
|
||||
_Metrics(wait_ms=2500.0, waits=25, query_ms=900.0, queries=9),
|
||||
)
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=1.0, monotonic=clock)
|
||||
|
||||
first = await sampler.maybe_sample(lambda: client)
|
||||
assert first is not None
|
||||
assert first.acquire_count_delta == 0, "first sample has no predecessor to diff against"
|
||||
assert first.acquire_wait_seconds_delta == 0.0
|
||||
|
||||
clock.now = 5.0
|
||||
second = await sampler.maybe_sample(lambda: client)
|
||||
assert second is not None
|
||||
assert second.acquire_count_delta == 15
|
||||
assert second.acquire_wait_seconds_delta == pytest.approx(1.5)
|
||||
assert second.query_count_delta == 4
|
||||
assert second.query_duration_seconds_delta == pytest.approx(0.4)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_restart_does_not_publish_a_negative_delta():
|
||||
clock = _Clock()
|
||||
client = _Client(
|
||||
_Metrics(wait_ms=9000.0, waits=90, query_ms=4000.0, queries=40),
|
||||
_Metrics(wait_ms=10.0, waits=1, query_ms=5.0, queries=1),
|
||||
)
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=1.0, monotonic=clock)
|
||||
|
||||
await sampler.maybe_sample(lambda: client)
|
||||
clock.now = 5.0
|
||||
after_restart = await sampler.maybe_sample(lambda: client)
|
||||
|
||||
assert after_restart is not None
|
||||
assert after_restart.acquire_count_delta == 0
|
||||
assert after_restart.acquire_wait_seconds_delta == 0.0
|
||||
assert after_restart.query_count_delta == 0
|
||||
assert after_restart.sample.acquire_count_total == 1, "the fresh absolute reading is still reported"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_unreachable_engine_never_raises_into_the_db_call():
|
||||
clock = _Clock()
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=1.0, monotonic=clock)
|
||||
assert await sampler.maybe_sample(lambda: _Client(_Metrics(), fail=True)) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_failed_sample_still_consumes_the_interval():
|
||||
"""Otherwise a wedged engine is re-probed on every single DB call."""
|
||||
clock = _Clock()
|
||||
client = _Client(_Metrics(), fail=True)
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=10.0, monotonic=clock)
|
||||
|
||||
await sampler.maybe_sample(lambda: client)
|
||||
clock.now = 1.0
|
||||
await sampler.maybe_sample(lambda: client)
|
||||
assert client.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_client_that_delegates_via_getattr_is_still_sampled():
|
||||
"""PrismaWrapper resolves its surface through an instance-level __getattr__,
|
||||
so a class-level structural check rejects it while the call itself works.
|
||||
An earlier revision guarded with isinstance() against a runtime_checkable
|
||||
Protocol and silently published zeroes on a live proxy; this pins the shape
|
||||
that regression had."""
|
||||
|
||||
class _Delegating:
|
||||
def __init__(self, inner):
|
||||
self._inner = inner
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._inner, name)
|
||||
|
||||
inner = _Client(_Metrics(busy=4.0, idle=1.0))
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=1.0, monotonic=_Clock())
|
||||
|
||||
update = await sampler.maybe_sample(lambda: _Delegating(inner))
|
||||
|
||||
assert update is not None
|
||||
assert update.sample.max_connections == 5.0
|
||||
assert inner.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_interval_is_consumed_even_when_there_is_no_client_to_sample():
|
||||
"""Otherwise a proxy with prometheus disabled, or one sampled before the DB
|
||||
client exists, never records an attempt and so re-enters this path on every
|
||||
single database call instead of once per interval."""
|
||||
clock = _Clock()
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=10.0, monotonic=clock)
|
||||
|
||||
assert await sampler.maybe_sample(lambda: None) is None
|
||||
assert sampler.is_due() is False, "a no-client attempt must still consume the interval"
|
||||
|
||||
clock.now = 10.0
|
||||
assert sampler.is_due() is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_resolver_that_raises_is_contained():
|
||||
def _boom():
|
||||
raise RuntimeError("proxy_server not importable")
|
||||
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=1.0, monotonic=_Clock())
|
||||
assert await sampler.maybe_sample(_boom) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pool_timeouts_do_not_latch_the_pending_gauge():
|
||||
"""prisma decrements its waiter gauge on acquisition but not on timeout, so
|
||||
every P2024 latches it one higher for the life of the process. Sequence
|
||||
reproduced against a live engine at connection_limit=2: after two bursts
|
||||
producing seven timeouts, a fully drained pool still reported pending=7.
|
||||
Free capacity proves nobody is queued, so an idle reading is the latch."""
|
||||
clock = _Clock()
|
||||
client = _Client(
|
||||
_Metrics(busy=0.0, idle=2.0, wait=0.0), # at rest
|
||||
_Metrics(busy=2.0, idle=0.0, wait=4.0), # saturated, 4 real waiters
|
||||
_Metrics(busy=0.0, idle=2.0, wait=4.0), # drained, 4 latched timeouts
|
||||
_Metrics(busy=2.0, idle=0.0, wait=7.0), # saturated again, 3 real waiters
|
||||
_Metrics(busy=0.0, idle=2.0, wait=7.0), # drained again
|
||||
)
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=1.0, monotonic=clock)
|
||||
|
||||
observed = []
|
||||
for tick in range(5):
|
||||
clock.now = tick * 2.0
|
||||
update = await sampler.maybe_sample(lambda: client)
|
||||
assert update is not None
|
||||
observed.append(update.pending_acquirers)
|
||||
|
||||
assert observed == [0.0, 4.0, 0.0, 3.0, 0.0], (
|
||||
f"expected the latch to be subtracted, got {observed}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_raw_engine_reading_is_still_carried_on_the_sample():
|
||||
"""The correction applies to what is published, not to what was observed."""
|
||||
clock = _Clock()
|
||||
client = _Client(_Metrics(busy=0.0, idle=2.0, wait=4.0))
|
||||
sampler = DBPoolMetricsSampler(min_interval_seconds=1.0, monotonic=clock)
|
||||
|
||||
update = await sampler.maybe_sample(lambda: client)
|
||||
|
||||
assert update is not None
|
||||
assert update.pending_acquirers == 0.0
|
||||
assert update.sample.pending_acquirers == 4.0
|
||||
|
|
@ -582,3 +582,50 @@ def test_is_deadlock_error_matches_postgres_deadlock(error):
|
|||
def test_is_deadlock_error_excludes_non_deadlocks(error):
|
||||
"""Non-deadlock prisma errors, connectivity failures, and non-prisma exceptions are not treated as deadlocks."""
|
||||
assert PrismaDBExceptionHandler.is_deadlock_error(error) is False
|
||||
|
||||
|
||||
def _pool_timeout_error(connection_limit: int = 10, timeout: int = 60) -> DataError:
|
||||
"""The exact payload prisma raises for P2024, captured from a live engine
|
||||
driven to exhaustion with connection_limit=1 and pool_timeout=1."""
|
||||
return DataError(
|
||||
{
|
||||
"user_facing_error": {
|
||||
"error_code": "P2024",
|
||||
"meta": {"connection_limit": connection_limit, "timeout": timeout},
|
||||
"message": (
|
||||
"Timed out fetching a new connection from the connection pool. "
|
||||
f"(Current connection pool timeout: {timeout}, connection limit: {connection_limit})"
|
||||
),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_pool_exhaustion_is_identified_by_its_prisma_code():
|
||||
assert PrismaDBExceptionHandler.is_connection_pool_timeout_error(_pool_timeout_error()) is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"other_error",
|
||||
[
|
||||
DataError({"user_facing_error": {"error_code": "P2002", "message": "unique constraint"}}),
|
||||
PrismaError("can't reach database server"),
|
||||
UniqueViolationError({"user_facing_error": {"error_code": "P2002"}}),
|
||||
httpx.ConnectError("connection refused"),
|
||||
ValueError("unrelated"),
|
||||
],
|
||||
)
|
||||
def test_only_p2024_counts_as_pool_exhaustion(other_error):
|
||||
"""A database that answered and refused the data, and a database that could
|
||||
not be reached at all, are both different problems from having no free
|
||||
connection to ask with. Counting either as exhaustion would send an operator
|
||||
to resize a pool that is fine."""
|
||||
assert PrismaDBExceptionHandler.is_connection_pool_timeout_error(other_error) is False
|
||||
|
||||
|
||||
def test_pool_exhaustion_is_not_mistaken_for_a_reachability_failure():
|
||||
"""P2024's message contains "Timed out", which the transport classifier
|
||||
keyword-matches. The pool predicate must not inherit that ambiguity."""
|
||||
error = _pool_timeout_error()
|
||||
assert PrismaDBExceptionHandler.is_connection_pool_timeout_error(error) is True
|
||||
assert PrismaDBExceptionHandler.is_database_infrastructure_error(error) is False
|
||||
|
|
|
|||
105
tests/test_litellm/proxy/db/test_log_db_metrics.py
Normal file
105
tests/test_litellm/proxy/db/test_log_db_metrics.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from prisma.errors import DataError, UniqueViolationError
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../.."))
|
||||
|
||||
from litellm.proxy.db.log_db_metrics import (
|
||||
_handle_logging_db_exception,
|
||||
_record_db_pool_timeout_if_exhausted,
|
||||
)
|
||||
|
||||
|
||||
def _pool_timeout_error() -> DataError:
|
||||
return DataError(
|
||||
{
|
||||
"user_facing_error": {
|
||||
"error_code": "P2024",
|
||||
"meta": {"connection_limit": 10, "timeout": 60},
|
||||
"message": "Timed out fetching a new connection from the connection pool.",
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_pool_exhaustion_increments_the_timeout_counter():
|
||||
logger = MagicMock()
|
||||
with patch("litellm.proxy.db.log_db_metrics._prometheus_logger", return_value=logger):
|
||||
_record_db_pool_timeout_if_exhausted(_pool_timeout_error())
|
||||
logger.record_db_pool_timeout.assert_called_once()
|
||||
|
||||
|
||||
def test_an_ordinary_db_error_does_not_increment_the_timeout_counter():
|
||||
logger = MagicMock()
|
||||
with patch("litellm.proxy.db.log_db_metrics._prometheus_logger", return_value=logger):
|
||||
_record_db_pool_timeout_if_exhausted(UniqueViolationError({"user_facing_error": {"error_code": "P2002"}}))
|
||||
logger.record_db_pool_timeout.assert_not_called()
|
||||
|
||||
|
||||
def test_no_prometheus_logger_configured_is_not_an_error():
|
||||
with patch("litellm.proxy.db.log_db_metrics._prometheus_logger", return_value=None):
|
||||
_record_db_pool_timeout_if_exhausted(_pool_timeout_error())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pool_exhaustion_is_counted_even_though_it_is_not_a_db_service_failure():
|
||||
"""P2024 is a prisma ``DataError``, which ``_is_exception_related_to_db``
|
||||
classifies as DB-related, but the counter must not depend on that gate:
|
||||
exhaustion is the proxy running out of connections, and it has to be
|
||||
recorded on whichever side of the classification it lands."""
|
||||
logger = MagicMock()
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj.async_service_failure_hook = MagicMock(
|
||||
side_effect=AssertionError("should not be reached when the error is not DB related")
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.db.log_db_metrics._prometheus_logger", return_value=logger),
|
||||
patch("litellm.proxy.db.log_db_metrics._is_exception_related_to_db", return_value=False),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj),
|
||||
):
|
||||
await _handle_logging_db_exception(
|
||||
e=_pool_timeout_error(),
|
||||
func=lambda: None,
|
||||
kwargs={},
|
||||
args=(),
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
logger.record_db_pool_timeout.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_stalled_pool_sample_does_not_delay_the_database_call():
|
||||
"""The sample is dispatched, not awaited. Awaiting it would both add its
|
||||
latency to every sampled query and open a cancellation window after the
|
||||
query had already succeeded, in which a client disconnect would discard a
|
||||
completed result."""
|
||||
import asyncio
|
||||
|
||||
from litellm.proxy.db.log_db_metrics import _pool_metrics_sampler, log_db_metrics
|
||||
|
||||
started = asyncio.Event()
|
||||
|
||||
async def _hang():
|
||||
started.set()
|
||||
await asyncio.sleep(30)
|
||||
|
||||
@log_db_metrics
|
||||
async def fake_db_call(**kwargs):
|
||||
return "rows"
|
||||
|
||||
_pool_metrics_sampler._last_sampled_at = None
|
||||
with patch("litellm.proxy.db.log_db_metrics._sample_db_pool_metrics", _hang):
|
||||
result = await asyncio.wait_for(fake_db_call(table_name="x"), timeout=1.0)
|
||||
|
||||
assert result == "rows"
|
||||
await asyncio.wait_for(started.wait(), timeout=1.0), "the sample must still be dispatched"
|
||||
for task in asyncio.all_tasks():
|
||||
if task is not asyncio.current_task() and task.get_coro().__name__ == "_hang":
|
||||
task.cancel()
|
||||
|
|
@ -215,3 +215,50 @@ def test_db_push_applies_replica_identity_full_when_requested(monkeypatch):
|
|||
|
||||
assert mock_run.call_args[0][0][:3] == ["prisma", "db", "push"]
|
||||
assert applied == [True]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_pool_sample_parses_the_engines_counters():
|
||||
"""The wrapper owns the engine, so it owns the engine's counter names. This
|
||||
pins the delegation and keeps those six key strings out of the metrics
|
||||
layer."""
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
|
||||
class _Value:
|
||||
def __init__(self, key, value):
|
||||
self.key = key
|
||||
self.value = value
|
||||
|
||||
class _Hist:
|
||||
def __init__(self, total, count):
|
||||
self.sum = total
|
||||
self.count = count
|
||||
|
||||
class _Metrics:
|
||||
gauges = [
|
||||
_Value("prisma_pool_connections_busy", 3.0),
|
||||
_Value("prisma_pool_connections_idle", 0.0),
|
||||
_Value("prisma_pool_connections_open", 3.0),
|
||||
_Value("prisma_client_queries_wait", 5.0),
|
||||
]
|
||||
histograms = [
|
||||
_Value("prisma_client_queries_wait_histogram_ms", _Hist(2000.0, 8)),
|
||||
_Value("prisma_datasource_queries_duration_histogram_ms", _Hist(500.0, 20)),
|
||||
]
|
||||
counters = []
|
||||
|
||||
original = MagicMock()
|
||||
original.get_metrics = AsyncMock(return_value=_Metrics())
|
||||
wrapper = PrismaWrapper(original_prisma=original, iam_token_db_auth=False)
|
||||
|
||||
sample = await wrapper.get_pool_sample()
|
||||
|
||||
original.get_metrics.assert_awaited_once()
|
||||
assert sample.busy_connections == 3.0
|
||||
assert sample.max_connections == 3.0
|
||||
assert sample.pending_acquirers == 5.0
|
||||
assert sample.acquire_wait_seconds_total == 2.0
|
||||
assert sample.query_duration_seconds_total == 0.5
|
||||
assert sample.query_count_total == 20
|
||||
|
|
|
|||
|
|
@ -991,3 +991,33 @@ async def test_recreate_keeps_writer_unavailable_when_writer_recreate_fails():
|
|||
await routing.recreate_prisma_client("writer-url")
|
||||
|
||||
assert routing.writer_unavailable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_pool_sample_reads_the_writer_pool():
|
||||
"""Reads route to the replica, but the writer is the pool this PR samples.
|
||||
Pinned so a later change to sample the reader is a deliberate edit rather
|
||||
than a silent one, and so deleting the delegation fails."""
|
||||
from litellm.proxy.db.db_pool_metrics import DBPoolSample
|
||||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||||
|
||||
sample = DBPoolSample(
|
||||
busy_connections=1.0,
|
||||
idle_connections=2.0,
|
||||
open_connections=3.0,
|
||||
pending_acquirers=0.0,
|
||||
acquire_wait_seconds_total=0.0,
|
||||
acquire_count_total=0,
|
||||
query_duration_seconds_total=0.0,
|
||||
query_count_total=0,
|
||||
)
|
||||
writer = MagicMock()
|
||||
writer.get_pool_sample = AsyncMock(return_value=sample)
|
||||
reader = MagicMock()
|
||||
reader.get_pool_sample = AsyncMock()
|
||||
|
||||
wrapper = RoutingPrismaWrapper(writer=writer, reader=reader)
|
||||
|
||||
assert await wrapper.get_pool_sample() is sample
|
||||
writer.get_pool_sample.assert_awaited_once()
|
||||
reader.get_pool_sample.assert_not_awaited()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue