diff --git a/src/db.py b/src/db.py index 6c7d5866..c543ce26 100644 --- a/src/db.py +++ b/src/db.py @@ -119,8 +119,11 @@ async def acquire_connection_with_retry(db: AsyncSession, context: str) -> None: in a Sentry span so wait time is visible in traces; on budget exhaustion the original error is reraised after capturing live pool stats to Sentry. - Retrying the same session is safe: no connection is bound until checkout - succeeds, so each attempt re-attempts the checkout cleanly. + Each attempt rolls the session back on a retryable failure before retrying: + a failed checkout can leave the autobegun transaction in a pending-rollback + state, which would make the next ``db.connection()`` raise instead of + re-checking-out cleanly. The rollback is pure Python-side state cleanup when + no connection was bound, so it is cheap and safe. """ with sentry_sdk.start_span(op="db.pool.acquire", name=context): if not settings.DB.CONNECTION_RETRY_ENABLED: @@ -139,7 +142,18 @@ async def acquire_connection_with_retry(db: AsyncSession, context: str) -> None: ): with attempt: attempts += 1 - await db.connection() + try: + await db.connection() + except RETRYABLE_DB_CONNECTION_ERRORS: + # Reset session state so the next attempt starts clean. + try: + await db.rollback() + except Exception: + logger.debug( + "rollback after failed checkout failed", + exc_info=True, + ) + raise except RETRYABLE_DB_CONNECTION_ERRORS as e: _record_acquisition_outcome("exhausted") if settings.SENTRY.ENABLED: diff --git a/src/telemetry/prometheus/metrics.py b/src/telemetry/prometheus/metrics.py index d883816b..1a79c634 100644 --- a/src/telemetry/prometheus/metrics.py +++ b/src/telemetry/prometheus/metrics.py @@ -323,17 +323,25 @@ class DBPoolCollector: self.instance_type: str = instance_type def collect(self) -> Iterator[GaugeMetricFamily]: - # Lazy import to avoid an import cycle at module load (db imports config, - # telemetry is imported widely). Reads the async engine.pool directly. - from src.db import get_pool_stats - namespace = settings.METRICS.NAMESPACE or "" gauge = GaugeMetricFamily( "db_pool_connections", "DB connections held by this instance, by pool state", labels=["namespace", "instance_type", "state"], ) - for state, value in get_pool_stats().items(): + # Fail soft: Prometheus aborts the entire scrape (dropping ALL metrics) + # if any collector raises, so never let a pool/import hiccup here sink + # the whole /metrics response. + try: + # Lazy import to avoid an import cycle at module load (db imports + # config, telemetry is imported widely). Reads engine.pool directly. + from src.db import get_pool_stats + + stats = get_pool_stats() + except Exception: + logger.warning("Failed to collect DB pool stats", exc_info=True) + stats = {} + for state, value in stats.items(): gauge.add_metric([namespace, self.instance_type, state], value) yield gauge diff --git a/tests/test_db_resilience.py b/tests/test_db_resilience.py index 0177373e..696a6c53 100644 --- a/tests/test_db_resilience.py +++ b/tests/test_db_resilience.py @@ -31,12 +31,16 @@ class _FlakyConnSession: self.fail_times: int = fail_times self.always_fail: bool = always_fail self.calls: int = 0 + self.rollback_calls: int = 0 async def connection(self) -> None: self.calls += 1 if self.always_fail or self.calls <= self.fail_times: raise _make_operational_error() + async def rollback(self) -> None: + self.rollback_calls += 1 + def _acq_count(outcome: str) -> float: child = db_connection_acquisitions_counter.labels( @@ -73,6 +77,8 @@ async def test_acquire_retries_then_succeeds_records_retried() -> None: session = _FlakyConnSession(fail_times=2) await acquire_connection_with_retry(session, "request:test") # pyright: ignore[reportArgumentType] assert session.calls == 3 # two failures then success + # Session is reset after each failed checkout before the next attempt. + assert session.rollback_calls == 2 assert _acq_count("retried") == before + 1