From 34076c312bbf70590c2d97e0f62f7aee0c3a7607 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Fri, 25 Sep 2026 16:12:36 +0300 Subject: [PATCH] fix: retry a raw retriable asyncpg error, so connect-time failures retry on every sqlalchemy --- db_retry/retriable.py | 9 ++++----- tests/test_retriable.py | 4 ++++ tests/test_retry.py | 25 +++++++++++++++++++++++++ 3 files changed, 33 insertions(+), 5 deletions(-) diff --git a/db_retry/retriable.py b/db_retry/retriable.py index c7e154f..176c6bf 100644 --- a/db_retry/retriable.py +++ b/db_retry/retriable.py @@ -6,15 +6,14 @@ def _is_retriable_link(exception: BaseException) -> bool: - return ( - isinstance(exception, DBAPIError) - and exception.orig is not None - and isinstance(exception.orig.__cause__, RETRIABLE_ASYNCPG_ERRORS) + candidate = ( + exception.orig.__cause__ if isinstance(exception, DBAPIError) and exception.orig is not None else exception ) + return isinstance(candidate, RETRIABLE_ASYNCPG_ERRORS) def is_retriable(exception: BaseException) -> bool: - """Walk __cause__/__context__; True if any link is a retriable DBAPIError.""" + """Walk __cause__/__context__; True if any link is a retriable asyncpg error, raw or wrapped in a DBAPIError.""" current: BaseException | None = exception seen: set[int] = set() while current is not None and id(current) not in seen: diff --git a/tests/test_retriable.py b/tests/test_retriable.py index 895a1e3..8dd6a1a 100644 --- a/tests/test_retriable.py +++ b/tests/test_retriable.py @@ -42,6 +42,10 @@ def test_a_statement_of_unknown_outcome_is_never_retriable() -> None: ), pytest.param(_make_dbapi_error(asyncpg.PostgresError()), False, id="non_retriable_postgres_error"), pytest.param(ValueError("not a db error"), False, id="bare_non_dbapi_exception"), + pytest.param(asyncpg.SerializationError(), True, id="raw_serialization_error_40001"), + pytest.param(asyncpg.PostgresConnectionError(), True, id="raw_postgres_connection_error_08000"), + pytest.param(asyncpg.StatementCompletionUnknownError(), False, id="raw_statement_completion_unknown_40003"), + pytest.param(asyncpg.PostgresError(), False, id="raw_non_retriable_postgres_error"), ], ) def test_is_retriable(exception: BaseException, expected: bool) -> None: diff --git a/tests/test_retry.py b/tests/test_retry.py index 2738067..e95d6f0 100644 --- a/tests/test_retry.py +++ b/tests/test_retry.py @@ -144,3 +144,28 @@ async def _always_fails() -> None: assert len(waits) == 3 # noqa: PLR2004 # four attempts, one wait between each pair assert all(wait > 0 for wait in waits) + + +async def test_a_connect_time_connection_error_is_retried(monkeypatch: pytest.MonkeyPatch) -> None: + _record_backoff(monkeypatch) + attempts = 0 + + async def _refusing_creator() -> asyncpg.Connection: + nonlocal attempts + attempts += 1 + msg = "connection refused" + raise asyncpg.PostgresConnectionError(msg) + + engine: typing.Final = sa_async.create_async_engine("postgresql+asyncpg://", async_creator=_refusing_creator) + expected_attempts: typing.Final = 2 + + @postgres_retry(retries=expected_attempts) + async def connect() -> None: + await engine.connect().__aenter__() + + try: + with pytest.raises((asyncpg.PostgresConnectionError, DBAPIError)): + await connect() + finally: + await engine.dispose() + assert attempts == expected_attempts