fix: address CodeRabbit review on PR #758

- db: roll back the session on a retryable checkout failure before
  retrying — a failed autobegin can leave it pending-rollback, making the
  next db.connection() raise instead of re-checking-out cleanly. Cheap
  Python-side cleanup when no connection was bound.
- metrics: guard DBPoolCollector.collect() so a pool-read/import hiccup
  can't raise and abort the whole /metrics scrape (Prometheus drops ALL
  metrics if any collector raises) — log and fall back to empty.
This commit is contained in:
Vineeth Voruganti 2026-06-01 00:07:50 -04:00
parent a5c43b8dc1
commit 8118e4a024
3 changed files with 36 additions and 8 deletions

View File

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

View File

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

View File

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