Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 4 additions & 5 deletions db_retry/retriable.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
4 changes: 4 additions & 0 deletions tests/test_retriable.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
25 changes: 25 additions & 0 deletions tests/test_retry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading