From 0bb26c3f1b087ed3dde217c129e28377cc115aa1 Mon Sep 17 00:00:00 2001 From: ohadgur Date: Mon, 9 Mar 2026 17:49:46 +0200 Subject: [PATCH] feat(proxy): add Prisma DB pool and engine health metrics to Prometheus (#22655) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 * 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 * 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 * ci: retrigger CI checks Co-Authored-By: Claude Opus 4.6 * 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 * 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 * fix: include "unknown" in _PG_STATES for stale label clearing Co-Authored-By: Claude Opus 4.6 * 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 * 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 * 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 * fix: use defensive null coalescing for query_raw row values Co-Authored-By: Claude Opus 4.6 * 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 --------- Co-authored-by: Claude Opus 4.6 --- docs/my-website/docs/proxy/prometheus.md | 19 +- litellm/proxy/db/prisma_metrics_collector.py | 180 +++++++++ litellm/proxy/utils.py | 85 +++- litellm/types/integrations/prometheus.py | 11 + .../proxy/db/test_prisma_metrics_collector.py | 369 ++++++++++++++++++ 5 files changed, 642 insertions(+), 22 deletions(-) create mode 100644 litellm/proxy/db/prisma_metrics_collector.py create mode 100644 tests/test_litellm/proxy/db/test_prisma_metrics_collector.py diff --git a/docs/my-website/docs/proxy/prometheus.md b/docs/my-website/docs/proxy/prometheus.md index d8f0d83b59d..dd9e52355be 100644 --- a/docs/my-website/docs/proxy/prometheus.md +++ b/docs/my-website/docs/proxy/prometheus.md @@ -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 diff --git a/litellm/proxy/db/prisma_metrics_collector.py b/litellm/proxy/db/prisma_metrics_collector.py new file mode 100644 index 00000000000..d60887aa6c7 --- /dev/null +++ b/litellm/proxy/db/prisma_metrics_collector.py @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 2f9d27568e3..d44f5a07482 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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, diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 0856d8a6f9b..8bc2171c9f2 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -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) diff --git a/tests/test_litellm/proxy/db/test_prisma_metrics_collector.py b/tests/test_litellm/proxy/db/test_prisma_metrics_collector.py new file mode 100644 index 00000000000..43ef39d4337 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_prisma_metrics_collector.py @@ -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