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:
ohadgur 2026-03-09 17:49:46 +02:00 • committed by GitHub
parent df2e1bca46
commit 0bb26c3f1b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 642 additions and 22 deletions

View file

@ -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

View 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

View file

@ -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,

View file

@ -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)

View 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