mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(proxy): add Prisma DB pool and engine health metrics to Prometheus (#22655)
* feat(proxy): add Prisma DB pool and engine health metrics to Prometheus Add a PrismaMetricsCollector that periodically queries pg_stat_activity and the Prisma engine process to expose connection pool and engine health as Prometheus gauges/counters. Auto-enabled when prometheus_system is in service_callback. New metrics: - litellm_db_pool_active_connections (Gauge) - litellm_db_pool_idle_connections (Gauge) - litellm_db_pool_total_connections (Gauge) - litellm_db_pool_waiting_connections (Gauge) - litellm_db_engine_up (Gauge) - litellm_db_engine_restarts_total (Counter) Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: address Greptile review feedback - Only increment engine_restarts counter on heavy reconnects (engine actually dead), not lightweight network-blip reconnects - Fix potential KeyError in _get_or_create_gauge/counter fallback path when REGISTRY._names_to_collectors is absent - Rename litellm_db_pool_waiting_connections to litellm_db_pool_lock_waiting_connections to clarify it measures lock contention, not pool slot queuing Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: warn when prometheus_system enabled but watchdog disabled Log a warning when users have prometheus_system in service_callback but PRISMA_HEALTH_WATCHDOG_ENABLED=false, since DB pool and engine metrics won't be collected in that configuration. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * ci: retrigger CI checks Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * refactor: use labeled gauge for DB pool connection metrics Replace 3 separate pool gauges (active, idle, total) with a single `litellm_db_pool_connections` gauge using a `state` label. This is more Prometheus-idiomatic and exposes all pg_stat_activity states (active, idle, idle in transaction, etc.) without ambiguity about what "total" includes. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: address Greptile review — stale labels and fallback re-registration - Zero out known pg_stat_activity states that are absent from the current query result, preventing stale gauge values from persisting. - Simplify _get_or_create_gauge/counter by removing the fallback loop that could re-register an already-registered metric (ValueError). - Add test for stale label clearing across collection cycles. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: include "unknown" in _PG_STATES for stale label clearing Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: collect immediately on start and consolidate into single query - Move sleep to end of loop so metrics appear on /metrics immediately after startup instead of after a 30s delay. - Combine pool state and lock waiting queries into a single SQL query using conditional aggregation, halving per-cycle DB overhead. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: prevent tight spin loop on collection error Move asyncio.sleep outside the try/except so it always executes even when _collect_engine_health() or _collect_pool_metrics() raises. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: add multiprocess_mode to _get_or_create_gauge initialization - Include `multiprocess_mode` parameter to properly support multiprocessing in Gauge creation. - Ensure consistent behavior for labeled and unlabeled Gauges. * fix: handle invalid env var and document watchdog prerequisite - Add try/except ValueError for PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS to prevent proxy startup crash on non-numeric values (e.g. "30s") - Document that DB metrics require both prometheus_system callback and PRISMA_HEALTH_WATCHDOG_ENABLED=true Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix: use defensive null coalescing for query_raw row values Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * test: add invalid env var fallback test and fix mock signature - Add test for non-numeric PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS - Add **kwargs to mock _patched_get_or_create_gauge for forward compat Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
df2e1bca46
commit
0bb26c3f1b
5 changed files with 642 additions and 22 deletions
|
|
@ -561,9 +561,26 @@ Use these metrics to monitor the health of the DB Transaction Queue. Eg. Monitor
|
|||
| `litellm_in_memory_spend_update_queue_size` | In-memory aggregate spend values for keys, users, teams, team members, etc.| In-Memory |
|
||||
| `litellm_redis_spend_update_queue_size` | Redis aggregate spend values for keys, users, teams, etc. | Redis |
|
||||
|
||||
#### DB Connection Pool and Engine Health Metrics
|
||||
|
||||
Monitor PostgreSQL connection pool utilization and Prisma query engine health. These metrics are collected every 30 seconds by default.
|
||||
|
||||
## 🔥 LiteLLM Maintained Grafana Dashboards
|
||||
| Metric Name | Type | Labels | Description |
|
||||
|------------------------------------------|---------|---------|-----------------------------------------------------------|
|
||||
| `litellm_db_pool_connections` | Gauge | `state` | Number of DB connections by state (active, idle, etc.) |
|
||||
| `litellm_db_pool_lock_waiting_connections` | Gauge | | Number of connections blocked on row/table locks |
|
||||
| `litellm_db_engine_up` | Gauge | | Whether the Prisma query engine is alive (1=up, 0=down) |
|
||||
| `litellm_db_engine_restarts_total` | Counter | | Total number of Prisma query engine restarts |
|
||||
|
||||
The `state` label values come from PostgreSQL's `pg_stat_activity.state` column: `active`, `idle`, `idle in transaction`, `idle in transaction (aborted)`, `fastpath function call`, `disabled`.
|
||||
|
||||
**Prerequisites:** Metrics collection requires both:
|
||||
- `prometheus_system` in `service_callback` (see [Monitor System Health](#monitor-system-health))
|
||||
- `PRISMA_HEALTH_WATCHDOG_ENABLED` not set to `false` (default: `true`). If disabled, a warning is logged and no DB metrics are collected.
|
||||
|
||||
The collection interval can be configured via the `PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS` environment variable (default: 30, minimum: 5).
|
||||
|
||||
## 🔥 LiteLLM Maintained Grafana Dashboards
|
||||
|
||||
Link to Grafana Dashboards maintained by LiteLLM
|
||||
|
||||
|
|
|
|||
180
litellm/proxy/db/prisma_metrics_collector.py
Normal file
180
litellm/proxy/db/prisma_metrics_collector.py
Normal file
|
|
@ -0,0 +1,180 @@
|
|||
"""
|
||||
Collects Prisma/PostgreSQL connection pool and engine health metrics
|
||||
and exposes them as Prometheus gauges/counters.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from typing import Optional, Set
|
||||
|
||||
from prometheus_client import REGISTRY, Counter, Gauge
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
|
||||
def _get_or_create_gauge(
|
||||
name: str,
|
||||
description: str,
|
||||
labelnames: Optional[list] = None,
|
||||
multiprocess_mode: str = "max",
|
||||
) -> Gauge:
|
||||
names_to_collectors = getattr(REGISTRY, "_names_to_collectors", None)
|
||||
if names_to_collectors is not None and name in names_to_collectors:
|
||||
return names_to_collectors[name]
|
||||
if labelnames:
|
||||
return Gauge(
|
||||
name, description, labelnames=labelnames, multiprocess_mode=multiprocess_mode
|
||||
)
|
||||
return Gauge(name, description, multiprocess_mode=multiprocess_mode)
|
||||
|
||||
|
||||
def _get_or_create_counter(name: str, description: str) -> Counter:
|
||||
names_to_collectors = getattr(REGISTRY, "_names_to_collectors", None)
|
||||
if names_to_collectors is not None and name in names_to_collectors:
|
||||
return names_to_collectors[name]
|
||||
return Counter(name, description)
|
||||
|
||||
|
||||
_POOL_METRICS_SQL = """
|
||||
SELECT state,
|
||||
count(*) as count,
|
||||
count(*) FILTER (WHERE wait_event_type = 'Lock') as lock_waiting
|
||||
FROM pg_stat_activity
|
||||
WHERE pid != pg_backend_pid() AND datname = current_database() AND usename = current_user
|
||||
GROUP BY state
|
||||
"""
|
||||
|
||||
# All possible pg_stat_activity states — used to zero out stale labels
|
||||
_PG_STATES = [
|
||||
"active",
|
||||
"idle",
|
||||
"idle in transaction",
|
||||
"idle in transaction (aborted)",
|
||||
"fastpath function call",
|
||||
"disabled",
|
||||
"unknown",
|
||||
]
|
||||
|
||||
_MIN_COLLECTION_INTERVAL = 5
|
||||
_DEFAULT_COLLECTION_INTERVAL = 30
|
||||
|
||||
|
||||
class PrismaMetricsCollector:
|
||||
"""Periodically collects DB pool and engine health metrics for Prometheus."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
prisma_client: "litellm.proxy.utils.PrismaClient", # type: ignore[name-defined]
|
||||
collection_interval: Optional[float] = None,
|
||||
) -> None:
|
||||
self.prisma_client = prisma_client
|
||||
|
||||
if collection_interval is not None:
|
||||
self._interval = max(collection_interval, _MIN_COLLECTION_INTERVAL)
|
||||
else:
|
||||
raw = os.environ.get(
|
||||
"PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS",
|
||||
str(_DEFAULT_COLLECTION_INTERVAL),
|
||||
)
|
||||
try:
|
||||
self._interval = max(float(raw), _MIN_COLLECTION_INTERVAL)
|
||||
except ValueError:
|
||||
verbose_proxy_logger.warning(
|
||||
"Invalid PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS=%r; using default %ss",
|
||||
raw,
|
||||
_DEFAULT_COLLECTION_INTERVAL,
|
||||
)
|
||||
self._interval = float(_DEFAULT_COLLECTION_INTERVAL)
|
||||
|
||||
self._task: Optional[asyncio.Task] = None
|
||||
|
||||
# Prometheus metrics
|
||||
self._pool_connections = _get_or_create_gauge(
|
||||
"litellm_db_pool_connections",
|
||||
"Number of DB connections by state",
|
||||
labelnames=["state"],
|
||||
)
|
||||
self._pool_waiting = _get_or_create_gauge(
|
||||
"litellm_db_pool_lock_waiting_connections",
|
||||
"Number of connections blocked on row/table locks in the DB pool",
|
||||
)
|
||||
self._engine_up = _get_or_create_gauge(
|
||||
"litellm_db_engine_up",
|
||||
"Whether the Prisma query engine process is alive (1=up, 0=down)",
|
||||
)
|
||||
self._engine_restarts = _get_or_create_counter(
|
||||
"litellm_db_engine_restarts_total",
|
||||
"Total number of Prisma query engine restarts",
|
||||
)
|
||||
|
||||
def start(self) -> None:
|
||||
"""Start the background collection loop. No-op if already running."""
|
||||
if self._task is not None:
|
||||
return
|
||||
self._task = asyncio.create_task(self._collection_loop())
|
||||
verbose_proxy_logger.info(
|
||||
"Started PrismaMetricsCollector (interval=%ss)", self._interval
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the background collection loop."""
|
||||
if self._task is None:
|
||||
return
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._task = None
|
||||
verbose_proxy_logger.info("Stopped PrismaMetricsCollector")
|
||||
|
||||
async def _collection_loop(self) -> None:
|
||||
while True:
|
||||
try:
|
||||
await self._collect_pool_metrics()
|
||||
self._collect_engine_health()
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("PrismaMetricsCollector loop error: %s", e)
|
||||
try:
|
||||
await asyncio.sleep(self._interval)
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
|
||||
async def _collect_pool_metrics(self) -> None:
|
||||
try:
|
||||
rows = await self.prisma_client.db.query_raw(_POOL_METRICS_SQL)
|
||||
|
||||
seen_states: Set[str] = set()
|
||||
total_lock_waiting = 0
|
||||
for row in rows:
|
||||
state = row.get("state") or "unknown"
|
||||
self._pool_connections.labels(state=state).set(row.get("count") or 0)
|
||||
total_lock_waiting += row.get("lock_waiting") or 0
|
||||
seen_states.add(state)
|
||||
|
||||
# Zero out states absent from this cycle to clear stale values
|
||||
for state in _PG_STATES:
|
||||
if state not in seen_states:
|
||||
self._pool_connections.labels(state=state).set(0)
|
||||
|
||||
self._pool_waiting.set(total_lock_waiting)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"PrismaMetricsCollector failed to collect pool metrics: %s", e
|
||||
)
|
||||
|
||||
def _collect_engine_health(self) -> None:
|
||||
alive = self.prisma_client._is_engine_alive()
|
||||
self._engine_up.set(1 if alive else 0)
|
||||
|
||||
def increment_engine_restarts(self) -> None:
|
||||
"""Increment the engine restart counter. Call from attempt_db_reconnect()."""
|
||||
self._engine_restarts.inc()
|
||||
|
||||
@staticmethod
|
||||
def should_enable() -> bool:
|
||||
"""Check if Prometheus system metrics are enabled."""
|
||||
return "prometheus_system" in litellm.service_callback
|
||||
|
|
@ -105,6 +105,7 @@ from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
|||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.db.log_db_metrics import log_db_metrics
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
from litellm.proxy.db.prisma_metrics_collector import PrismaMetricsCollector
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
|
@ -2045,8 +2046,10 @@ class ProxyLogging:
|
|||
|
||||
## CHECK FOR MODEL-LEVEL GUARDRAILS (cached per-request)
|
||||
if not _guardrail_data_computed:
|
||||
_cached_guardrail_data = _check_and_merge_model_level_guardrails(
|
||||
data=data, llm_router=llm_router
|
||||
_cached_guardrail_data = (
|
||||
_check_and_merge_model_level_guardrails(
|
||||
data=data, llm_router=llm_router
|
||||
)
|
||||
)
|
||||
_guardrail_data_computed = True
|
||||
|
||||
|
|
@ -2316,6 +2319,7 @@ class PrismaClient:
|
|||
self._watching_engine: bool = False
|
||||
self._engine_confirmed_dead: bool = False
|
||||
self._engine_wait_thread: Optional[threading.Thread] = None
|
||||
self._metrics_collector: Optional[PrismaMetricsCollector] = None
|
||||
verbose_proxy_logger.debug("Success - Created Prisma Client")
|
||||
|
||||
def get_request_status(
|
||||
|
|
@ -3637,13 +3641,15 @@ class PrismaClient:
|
|||
probe_pid, _ = os.waitpid(pid, os.WNOHANG)
|
||||
except ChildProcessError:
|
||||
verbose_proxy_logger.debug(
|
||||
"PID %s is not a child process; skipping waitpid watch.", pid,
|
||||
"PID %s is not a child process; skipping waitpid watch.",
|
||||
pid,
|
||||
)
|
||||
return False
|
||||
|
||||
if probe_pid == pid:
|
||||
verbose_proxy_logger.warning(
|
||||
"prisma-query-engine PID %s already dead at watch start.", pid,
|
||||
"prisma-query-engine PID %s already dead at watch start.",
|
||||
pid,
|
||||
)
|
||||
self._engine_confirmed_dead = True
|
||||
self._reap_all_zombies()
|
||||
|
|
@ -3820,11 +3826,17 @@ class PrismaClient:
|
|||
waitpid thread nor pidfd are available.
|
||||
|
||||
"""
|
||||
if self._watching_engine or self._engine_pidfd >= 0 or self._engine_wait_thread is not None:
|
||||
if (
|
||||
self._watching_engine
|
||||
or self._engine_pidfd >= 0
|
||||
or self._engine_wait_thread is not None
|
||||
):
|
||||
return
|
||||
pid = self._get_engine_pid()
|
||||
if pid == 0:
|
||||
verbose_proxy_logger.debug("Could not find prisma-query-engine PID; engine death detection unavailable.")
|
||||
verbose_proxy_logger.debug(
|
||||
"Could not find prisma-query-engine PID; engine death detection unavailable."
|
||||
)
|
||||
return
|
||||
self._engine_pid = pid
|
||||
self._engine_confirmed_dead = False
|
||||
|
|
@ -3833,15 +3845,18 @@ class PrismaClient:
|
|||
pidfd_ok = False if waitpid_ok else self._try_pidfd_watch(pid)
|
||||
if waitpid_ok:
|
||||
verbose_proxy_logger.info(
|
||||
"Watching engine PID %s via waitpid thread.", pid,
|
||||
"Watching engine PID %s via waitpid thread.",
|
||||
pid,
|
||||
)
|
||||
elif pidfd_ok:
|
||||
verbose_proxy_logger.info(
|
||||
"Watching engine PID %s via pidfd.", pid,
|
||||
"Watching engine PID %s via pidfd.",
|
||||
pid,
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.info(
|
||||
"Watching engine PID %s via os.kill polling.", pid,
|
||||
"Watching engine PID %s via os.kill polling.",
|
||||
pid,
|
||||
)
|
||||
self._watching_engine = True
|
||||
asyncio.create_task(self._poll_engine_proc())
|
||||
|
|
@ -3864,7 +3879,9 @@ class PrismaClient:
|
|||
blip -- disconnect, connect, SELECT 1).
|
||||
"""
|
||||
effective_timeout = (
|
||||
timeout_seconds if timeout_seconds is not None else self._db_watchdog_reconnect_timeout_seconds
|
||||
timeout_seconds
|
||||
if timeout_seconds is not None
|
||||
else self._db_watchdog_reconnect_timeout_seconds
|
||||
)
|
||||
|
||||
engine_is_dead = self._engine_confirmed_dead or (
|
||||
|
|
@ -3884,14 +3901,18 @@ class PrismaClient:
|
|||
async def _do_heavy_reconnect() -> None:
|
||||
db_url = os.getenv("DATABASE_URL", "")
|
||||
if not db_url:
|
||||
verbose_proxy_logger.error("DATABASE_URL not set; cannot recreate Prisma client.")
|
||||
verbose_proxy_logger.error(
|
||||
"DATABASE_URL not set; cannot recreate Prisma client."
|
||||
)
|
||||
raise RuntimeError("DATABASE_URL not set")
|
||||
await self.db.recreate_prisma_client(db_url)
|
||||
await self._start_engine_watcher()
|
||||
|
||||
await asyncio.wait_for(_do_heavy_reconnect(), timeout=effective_timeout)
|
||||
else:
|
||||
verbose_proxy_logger.debug("Performing Prisma DB reconnect (engine alive or unknown).")
|
||||
verbose_proxy_logger.debug(
|
||||
"Performing Prisma DB reconnect (engine alive or unknown)."
|
||||
)
|
||||
|
||||
async def _do_direct_reconnect() -> None:
|
||||
try:
|
||||
|
|
@ -3942,6 +3963,9 @@ class PrismaClient:
|
|||
"Attempting Prisma DB reconnect. reason=%s", reason
|
||||
)
|
||||
|
||||
engine_was_dead = self._engine_confirmed_dead or (
|
||||
self._engine_pid > 0 and not self._is_engine_alive()
|
||||
)
|
||||
reconnect_succeeded = False
|
||||
try:
|
||||
await self._run_reconnect_cycle(timeout_seconds=timeout_seconds)
|
||||
|
|
@ -3950,6 +3974,8 @@ class PrismaClient:
|
|||
verbose_proxy_logger.info(
|
||||
"Prisma DB reconnect succeeded. reason=%s", reason
|
||||
)
|
||||
if self._metrics_collector is not None and engine_was_dead:
|
||||
self._metrics_collector.increment_engine_restarts()
|
||||
except Exception as reconnect_err:
|
||||
self._consecutive_reconnect_failures += 1
|
||||
verbose_proxy_logger.error(
|
||||
|
|
@ -3990,7 +4016,9 @@ class PrismaClient:
|
|||
|
||||
if lock_timeout_seconds is None:
|
||||
async with self._db_reconnect_lock:
|
||||
return await self._attempt_reconnect_inside_lock(force, reason, timeout_seconds)
|
||||
return await self._attempt_reconnect_inside_lock(
|
||||
force, reason, timeout_seconds
|
||||
)
|
||||
|
||||
lock_acquired_by_timeout_task = False
|
||||
|
||||
|
|
@ -4039,18 +4067,26 @@ class PrismaClient:
|
|||
return False
|
||||
|
||||
try:
|
||||
return await self._attempt_reconnect_inside_lock(force, reason, timeout_seconds)
|
||||
return await self._attempt_reconnect_inside_lock(
|
||||
force, reason, timeout_seconds
|
||||
)
|
||||
finally:
|
||||
self._db_reconnect_lock.release()
|
||||
|
||||
async def start_db_health_watchdog_task(self) -> None:
|
||||
"""Start background tasks that monitor DB health:
|
||||
- A periodic SELECT 1 probe that triggers reconnect on network/connection failure.
|
||||
- A process-level watcher that detects engine death via waitpid thread, pidfd, or os.kill polling."""
|
||||
- A process-level watcher that detects engine death via waitpid thread, pidfd, or os.kill polling.
|
||||
"""
|
||||
if self._db_health_watchdog_enabled is not True:
|
||||
verbose_proxy_logger.debug(
|
||||
"Prisma DB health watchdog disabled via PRISMA_HEALTH_WATCHDOG_ENABLED"
|
||||
)
|
||||
if PrismaMetricsCollector.should_enable():
|
||||
verbose_proxy_logger.warning(
|
||||
"prometheus_system is enabled but PRISMA_HEALTH_WATCHDOG_ENABLED=false — "
|
||||
"DB pool and engine metrics will not be collected"
|
||||
)
|
||||
return
|
||||
if self._db_health_watchdog_task is not None:
|
||||
return
|
||||
|
|
@ -4066,6 +4102,10 @@ class PrismaClient:
|
|||
)
|
||||
await self._start_engine_watcher()
|
||||
|
||||
if PrismaMetricsCollector.should_enable() and self._metrics_collector is None:
|
||||
self._metrics_collector = PrismaMetricsCollector(self)
|
||||
self._metrics_collector.start()
|
||||
|
||||
async def stop_db_health_watchdog_task(self) -> None:
|
||||
"""Stop DB health watchdog task and engine watcher gracefully."""
|
||||
self._stop_engine_watcher()
|
||||
|
|
@ -4079,6 +4119,10 @@ class PrismaClient:
|
|||
self._db_health_watchdog_task = None
|
||||
verbose_proxy_logger.info("Stopped Prisma DB health watchdog")
|
||||
|
||||
if self._metrics_collector is not None:
|
||||
await self._metrics_collector.stop()
|
||||
self._metrics_collector = None
|
||||
|
||||
async def _db_health_watchdog_loop(self) -> None:
|
||||
while True:
|
||||
try:
|
||||
|
|
@ -4506,9 +4550,9 @@ class ProxyUpdateSpend:
|
|||
:MAX_LOGS_PER_INTERVAL
|
||||
]
|
||||
# Remove the logs we're about to process
|
||||
prisma_client.spend_log_transactions = prisma_client.spend_log_transactions[
|
||||
len(logs_to_process) :
|
||||
]
|
||||
prisma_client.spend_log_transactions = (
|
||||
prisma_client.spend_log_transactions[len(logs_to_process) :]
|
||||
)
|
||||
popped_batch = True
|
||||
if len(logs_to_process) > 0:
|
||||
verbose_proxy_logger.info(
|
||||
|
|
@ -4662,9 +4706,7 @@ async def update_spend_logs_job(
|
|||
return
|
||||
|
||||
async with prisma_client._spend_log_transactions_lock:
|
||||
logs_to_process = prisma_client.spend_log_transactions[
|
||||
:MAX_LOGS_PER_INTERVAL
|
||||
]
|
||||
logs_to_process = prisma_client.spend_log_transactions[:MAX_LOGS_PER_INTERVAL]
|
||||
prisma_client.spend_log_transactions = prisma_client.spend_log_transactions[
|
||||
len(logs_to_process) :
|
||||
]
|
||||
|
|
@ -4682,6 +4724,7 @@ async def update_spend_logs_job(
|
|||
from litellm.proxy.guardrails.usage_tracking import (
|
||||
process_spend_logs_guardrail_usage,
|
||||
)
|
||||
|
||||
await process_spend_logs_guardrail_usage(
|
||||
prisma_client=prisma_client,
|
||||
logs_to_process=logs_to_process,
|
||||
|
|
|
|||
|
|
@ -238,6 +238,11 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
"litellm_llm_api_failed_requests_metric",
|
||||
"litellm_callback_logging_failures_metric",
|
||||
"litellm_in_flight_requests",
|
||||
# Database engine / connection pool metrics
|
||||
"litellm_db_pool_connections",
|
||||
"litellm_db_pool_lock_waiting_connections",
|
||||
"litellm_db_engine_up",
|
||||
"litellm_db_engine_restarts_total",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -618,6 +623,12 @@ class PrometheusMetricLabels:
|
|||
litellm_cache_misses_metric = _cache_metric_labels
|
||||
litellm_cached_tokens_metric = _cache_metric_labels
|
||||
|
||||
# Database engine / connection pool metrics
|
||||
litellm_db_pool_connections: List[str] = ["state"]
|
||||
litellm_db_pool_lock_waiting_connections: List[str] = []
|
||||
litellm_db_engine_up: List[str] = []
|
||||
litellm_db_engine_restarts_total: List[str] = []
|
||||
|
||||
@staticmethod
|
||||
def get_labels(label_name: DEFINED_PROMETHEUS_METRICS) -> List[str]:
|
||||
default_labels = getattr(PrometheusMetricLabels, label_name)
|
||||
|
|
|
|||
369
tests/test_litellm/proxy/db/test_prisma_metrics_collector.py
Normal file
369
tests/test_litellm/proxy/db/test_prisma_metrics_collector.py
Normal file
|
|
@ -0,0 +1,369 @@
|
|||
"""
|
||||
Unit tests for PrismaMetricsCollector.
|
||||
|
||||
All Prometheus metrics are isolated per test using a custom CollectorRegistry
|
||||
to avoid cross-test registration conflicts.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from prometheus_client import CollectorRegistry
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.db.prisma_metrics_collector import (
|
||||
PrismaMetricsCollector,
|
||||
_DEFAULT_COLLECTION_INTERVAL,
|
||||
_MIN_COLLECTION_INTERVAL,
|
||||
)
|
||||
|
||||
|
||||
def _make_prisma_client():
|
||||
"""Create a mock PrismaClient with the interface PrismaMetricsCollector uses."""
|
||||
client = MagicMock()
|
||||
client.db = MagicMock()
|
||||
client.db.query_raw = AsyncMock(return_value=[])
|
||||
client._is_engine_alive = MagicMock(return_value=True)
|
||||
return client
|
||||
|
||||
|
||||
def _make_collector(prisma_client=None, collection_interval=None, registry=None):
|
||||
"""Create a PrismaMetricsCollector with an isolated Prometheus registry.
|
||||
|
||||
Patches the module-level helper functions to use the provided registry,
|
||||
so every test gets its own metric instances.
|
||||
"""
|
||||
if prisma_client is None:
|
||||
prisma_client = _make_prisma_client()
|
||||
if registry is None:
|
||||
registry = CollectorRegistry()
|
||||
|
||||
from prometheus_client import Counter, Gauge
|
||||
|
||||
def _patched_get_or_create_gauge(name, description, labelnames=None, **kwargs):
|
||||
if labelnames:
|
||||
return Gauge(name, description, labelnames=labelnames, registry=registry)
|
||||
return Gauge(name, description, registry=registry)
|
||||
|
||||
def _patched_get_or_create_counter(name, description):
|
||||
return Counter(name, description, registry=registry)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.db.prisma_metrics_collector._get_or_create_gauge",
|
||||
side_effect=_patched_get_or_create_gauge,
|
||||
), patch(
|
||||
"litellm.proxy.db.prisma_metrics_collector._get_or_create_counter",
|
||||
side_effect=_patched_get_or_create_counter,
|
||||
):
|
||||
collector = PrismaMetricsCollector(
|
||||
prisma_client=prisma_client,
|
||||
collection_interval=collection_interval,
|
||||
)
|
||||
|
||||
return collector, registry
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Metric creation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_collector_creates_prometheus_metrics():
|
||||
"""Verify all 4 metrics (pool connections gauge, lock waiting gauge, engine_up gauge, restarts counter) are created."""
|
||||
collector, registry = _make_collector()
|
||||
|
||||
assert collector._pool_connections is not None
|
||||
assert collector._pool_waiting is not None
|
||||
assert collector._engine_up is not None
|
||||
assert collector._engine_restarts is not None
|
||||
|
||||
# Verify names via the registry
|
||||
metric_names = {m.name for m in registry.collect()}
|
||||
expected = {
|
||||
"litellm_db_pool_connections",
|
||||
"litellm_db_pool_lock_waiting_connections",
|
||||
"litellm_db_engine_up",
|
||||
"litellm_db_engine_restarts", # counter exposes _total suffix but name is base
|
||||
}
|
||||
assert expected.issubset(
|
||||
metric_names
|
||||
), f"Missing metrics: {expected - metric_names}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pool metrics collection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collect_pool_metrics_sets_gauges():
|
||||
"""Mock query_raw to return pool stats grouped by state and verify labeled gauge is set."""
|
||||
client = _make_prisma_client()
|
||||
|
||||
pool_rows = [
|
||||
{"state": "active", "count": 5, "lock_waiting": 1},
|
||||
{"state": "idle", "count": 10, "lock_waiting": 0},
|
||||
{"state": "idle in transaction", "count": 3, "lock_waiting": 1},
|
||||
]
|
||||
client.db.query_raw = AsyncMock(return_value=pool_rows)
|
||||
collector, registry = _make_collector(prisma_client=client)
|
||||
|
||||
await collector._collect_pool_metrics()
|
||||
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "active"})
|
||||
== 5
|
||||
)
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "idle"})
|
||||
== 10
|
||||
)
|
||||
assert (
|
||||
registry.get_sample_value(
|
||||
"litellm_db_pool_connections", {"state": "idle in transaction"}
|
||||
)
|
||||
== 3
|
||||
)
|
||||
assert registry.get_sample_value("litellm_db_pool_lock_waiting_connections") == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collect_pool_metrics_handles_empty_result():
|
||||
"""When query_raw returns empty list, known states should be zeroed."""
|
||||
client = _make_prisma_client()
|
||||
client.db.query_raw = AsyncMock(return_value=[])
|
||||
collector, registry = _make_collector(prisma_client=client)
|
||||
|
||||
await collector._collect_pool_metrics()
|
||||
|
||||
# Known states should be zeroed out
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "active"})
|
||||
== 0
|
||||
)
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "idle"}) == 0
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collect_pool_metrics_handles_null_state():
|
||||
"""When pg_stat_activity returns a NULL state, it should be mapped to 'unknown'."""
|
||||
client = _make_prisma_client()
|
||||
pool_rows = [{"state": None, "count": 1, "lock_waiting": 0}]
|
||||
client.db.query_raw = AsyncMock(return_value=pool_rows)
|
||||
collector, registry = _make_collector(prisma_client=client)
|
||||
|
||||
await collector._collect_pool_metrics()
|
||||
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "unknown"})
|
||||
== 1
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collect_pool_metrics_clears_stale_states():
|
||||
"""States present in cycle 1 but absent in cycle 2 should be zeroed out."""
|
||||
client = _make_prisma_client()
|
||||
|
||||
# Cycle 1: active=5
|
||||
pool_rows_1 = [{"state": "active", "count": 5, "lock_waiting": 0}]
|
||||
client.db.query_raw = AsyncMock(return_value=pool_rows_1)
|
||||
collector, registry = _make_collector(prisma_client=client)
|
||||
|
||||
await collector._collect_pool_metrics()
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "active"})
|
||||
== 5
|
||||
)
|
||||
|
||||
# Cycle 2: only idle connections, active should be zeroed
|
||||
pool_rows_2 = [{"state": "idle", "count": 3, "lock_waiting": 0}]
|
||||
client.db.query_raw = AsyncMock(return_value=pool_rows_2)
|
||||
|
||||
await collector._collect_pool_metrics()
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "active"})
|
||||
== 0
|
||||
)
|
||||
assert (
|
||||
registry.get_sample_value("litellm_db_pool_connections", {"state": "idle"}) == 3
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_collect_pool_metrics_handles_query_error():
|
||||
"""When query_raw raises an exception, the collector should log a warning and not crash."""
|
||||
client = _make_prisma_client()
|
||||
client.db.query_raw = AsyncMock(side_effect=RuntimeError("connection lost"))
|
||||
collector, _ = _make_collector(prisma_client=client)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.db.prisma_metrics_collector.verbose_proxy_logger"
|
||||
) as mock_logger:
|
||||
await collector._collect_pool_metrics()
|
||||
mock_logger.warning.assert_called_once()
|
||||
assert "connection lost" in str(mock_logger.warning.call_args)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Engine health
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_collect_engine_health_alive():
|
||||
"""When engine is alive, engine_up gauge should be 1."""
|
||||
client = _make_prisma_client()
|
||||
client._is_engine_alive = MagicMock(return_value=True)
|
||||
collector, registry = _make_collector(prisma_client=client)
|
||||
|
||||
collector._collect_engine_health()
|
||||
|
||||
assert registry.get_sample_value("litellm_db_engine_up") == 1
|
||||
|
||||
|
||||
def test_collect_engine_health_dead():
|
||||
"""When engine is dead, engine_up gauge should be 0."""
|
||||
client = _make_prisma_client()
|
||||
client._is_engine_alive = MagicMock(return_value=False)
|
||||
collector, registry = _make_collector(prisma_client=client)
|
||||
|
||||
collector._collect_engine_health()
|
||||
|
||||
assert registry.get_sample_value("litellm_db_engine_up") == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Engine restart counter
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_increment_engine_restarts():
|
||||
"""Calling increment_engine_restarts N times should result in counter value N."""
|
||||
collector, registry = _make_collector()
|
||||
|
||||
for _ in range(7):
|
||||
collector.increment_engine_restarts()
|
||||
|
||||
assert registry.get_sample_value("litellm_db_engine_restarts_total") == 7
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# should_enable
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_should_enable_true():
|
||||
"""should_enable() returns True when prometheus_system is in service_callback."""
|
||||
original = litellm.service_callback
|
||||
try:
|
||||
litellm.service_callback = ["prometheus_system"]
|
||||
assert PrismaMetricsCollector.should_enable() is True
|
||||
finally:
|
||||
litellm.service_callback = original
|
||||
|
||||
|
||||
def test_should_enable_false():
|
||||
"""should_enable() returns False when service_callback is empty."""
|
||||
original = litellm.service_callback
|
||||
try:
|
||||
litellm.service_callback = []
|
||||
assert PrismaMetricsCollector.should_enable() is False
|
||||
finally:
|
||||
litellm.service_callback = original
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Collection interval configuration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_collection_interval_from_env():
|
||||
"""Interval should be read from PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS env var."""
|
||||
with patch.dict(os.environ, {"PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS": "60"}):
|
||||
collector, _ = _make_collector()
|
||||
assert collector._interval == 60
|
||||
|
||||
|
||||
def test_collection_interval_minimum_enforced():
|
||||
"""Interval below the minimum should be clamped to _MIN_COLLECTION_INTERVAL."""
|
||||
with patch.dict(os.environ, {"PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS": "1"}):
|
||||
collector, _ = _make_collector()
|
||||
assert collector._interval == _MIN_COLLECTION_INTERVAL
|
||||
|
||||
|
||||
def test_collection_interval_constructor_override():
|
||||
"""Explicit collection_interval parameter should take precedence over env."""
|
||||
with patch.dict(os.environ, {"PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS": "999"}):
|
||||
collector, _ = _make_collector(collection_interval=45)
|
||||
assert collector._interval == 45
|
||||
|
||||
|
||||
def test_collection_interval_default():
|
||||
"""Without env var or constructor arg, the default interval is used."""
|
||||
with patch.dict(os.environ, {}, clear=False):
|
||||
# Remove the env var if present
|
||||
env_copy = os.environ.copy()
|
||||
env_copy.pop("PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS", None)
|
||||
with patch.dict(os.environ, env_copy, clear=True):
|
||||
collector, _ = _make_collector()
|
||||
assert collector._interval == _DEFAULT_COLLECTION_INTERVAL
|
||||
|
||||
|
||||
def test_collection_interval_invalid_env_falls_back_to_default():
|
||||
"""Non-numeric PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS should fall back to default."""
|
||||
with patch.dict(os.environ, {"PRISMA_METRICS_COLLECTION_INTERVAL_SECONDS": "30s"}):
|
||||
with patch(
|
||||
"litellm.proxy.db.prisma_metrics_collector.verbose_proxy_logger"
|
||||
) as mock_logger:
|
||||
collector, _ = _make_collector()
|
||||
assert collector._interval == _DEFAULT_COLLECTION_INTERVAL
|
||||
mock_logger.warning.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Start / Stop lifecycle
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_creates_task():
|
||||
"""Calling start() should create a background asyncio task."""
|
||||
collector, _ = _make_collector()
|
||||
|
||||
collector.start()
|
||||
assert collector._task is not None
|
||||
# Clean up
|
||||
await collector.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_idempotent():
|
||||
"""Calling start() twice should not create a second task."""
|
||||
collector, _ = _make_collector()
|
||||
|
||||
collector.start()
|
||||
first_task = collector._task
|
||||
collector.start()
|
||||
assert collector._task is first_task
|
||||
# Clean up
|
||||
await collector.stop()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_cancels_task():
|
||||
"""Calling stop() after start() should cancel the task and set it to None."""
|
||||
collector, _ = _make_collector()
|
||||
|
||||
collector.start()
|
||||
assert collector._task is not None
|
||||
|
||||
await collector.stop()
|
||||
assert collector._task is None
|
||||
Loading…
Add table
Reference in a new issue