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
10 changes: 5 additions & 5 deletions CONTEXT.md
Original file line number Diff line number Diff line change
Expand Up @@ -28,8 +28,8 @@ _Avoid_: retry, as a count. The knob is spelled `retries` for callers and cannot
breaking them, but every number in this package counts attempts.

**Primary host**:
The host a `ConnectionPlan`'s first connect attempt is aimed at — which, for a multi-host DSN, is
*every* host in shuffled order, handed to asyncpg to walk itself. "Primary" is a position in the
two-stage connect (one bulk attempt, then host-by-host through `failover`), not a PostgreSQL
replication role; replication role is `target_session_attrs`, where `read-write` selects a writable
node and `prefer-standby` a replica.
The host a connection's first connect attempt is aimed at — which, for a multi-host DSN, is
*every* host, in an order shuffled afresh for each connection, handed to asyncpg to walk itself.
"Primary" is a position in the two-stage connect (one bulk attempt, then host-by-host through
`failover`), not a PostgreSQL replication role; replication role is `target_session_attrs`, where
`read-write` selects a writable node and `prefer-standby` a replica.
46 changes: 17 additions & 29 deletions db_retry/connections.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,8 +22,6 @@
class ConnectionPlan:
connect_args: typing.Mapping[str, typing.Any]
target_session_attrs: SessionAttribute | None
primary_host: str | list[str]
primary_port: int | list[int] | None
failover: tuple[tuple[str, int], ...]


Expand All @@ -42,45 +40,29 @@ def build_connection_plan(url: sqlalchemy.URL) -> ConnectionPlan:
target_session_attrs: SessionAttribute | None = (
SessionAttribute(raw_target_session_attrs) if raw_target_session_attrs else None
)
raw_hosts: str | list[str] = connect_args.pop("host")
raw_ports: int | list[int] | None = connect_args.pop("port", None)
primary_host: str | list[str]
primary_port: int | list[int] | None
failover: tuple[tuple[str, int], ...]
if hosts_and_ports:
random.shuffle(hosts_and_ports)
primary_host = list(map(itemgetter(0), hosts_and_ports))
primary_port = list(map(itemgetter(1), hosts_and_ports))
failover = tuple(hosts_and_ports)
else:
primary_host = raw_hosts
primary_port = raw_ports
failover = ()
del connect_args["host"], connect_args["port"]
return ConnectionPlan(
connect_args=types.MappingProxyType(connect_args),
target_session_attrs=target_session_attrs,
primary_host=primary_host,
primary_port=primary_port,
failover=failover,
failover=tuple(hosts_and_ports),
)


async def _connect(
plan: ConnectionPlan,
host: str | list[str],
port: int | list[int] | None,
timeout: float, # noqa: ASYNC109
**address: str | int | list[str] | list[int],
) -> "ConnectionType":
return await asyncpg.connect(
**plan.connect_args,
host=host,
port=port,
**address,
timeout=timeout,
target_session_attrs=plan.target_session_attrs,
)


def _reshuffled(failover: tuple[tuple[str, int], ...]) -> list[tuple[str, int]]:
def _shuffled(failover: tuple[tuple[str, int], ...]) -> list[tuple[str, int]]:
return random.sample(failover, len(failover))


Expand All @@ -91,17 +73,23 @@ def build_connection_factory(
plan: typing.Final = build_connection_plan(url)

async def _connection_factory() -> "ConnectionType":
if not plan.failover:
return await _connect(plan, timeout)

bulk_order: typing.Final = _shuffled(plan.failover)
try:
return await _connect(plan, plan.primary_host, plan.primary_port, timeout)
return await _connect(
plan,
timeout,
host=list(map(itemgetter(0), bulk_order)),
port=list(map(itemgetter(1), bulk_order)),
)
except TimeoutError:
if not plan.failover:
raise

logger.warning("Failed to fetch asyncpg connection. Trying host by host.")

for host, port in _reshuffled(plan.failover):
for host, port in _shuffled(plan.failover):
try:
return await _connect(plan, host, port, timeout)
return await _connect(plan, timeout, host=host, port=port)
except (TimeoutError, OSError, asyncpg.TargetServerAttributeNotMatched) as exc:
logger.warning("Failed to fetch asyncpg connection from %s, %s", host, exc)
msg: typing.Final = f"None of the hosts match the target attribute requirement {plan.target_session_attrs}"
Expand Down
79 changes: 49 additions & 30 deletions tests/test_connection_factory.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import os
import random
import typing
from unittest import mock

Expand Down Expand Up @@ -70,44 +71,64 @@ async def test_connection_factory_failure_and_success(monkeypatch: pytest.Monkey
assert result is mock_connection


def _six_hosts() -> tuple[list[tuple[str, int]], sqlalchemy.URL]:
hosts_and_ports: typing.Final = [(f"host{n}", 5431 + n) for n in range(1, 7)]
query: typing.Final = "&".join(f"host={host}:{port}" for host, port in hosts_and_ports)
return hosts_and_ports, sqlalchemy.make_url(f"postgresql+asyncpg://user:password@/database?{query}")


def test_build_connection_plan_multihost() -> None:
url: typing.Final = sqlalchemy.make_url(
"postgresql+asyncpg://user:password@/database?host=host1:5432&host=host2:5432&target_session_attrs=read-write"
"postgresql+asyncpg://user:password@/database?host=host1:5432&host=host2:5433&target_session_attrs=read-write"
)
plan: typing.Final[ConnectionPlan] = build_connection_plan(url)
assert set(plan.failover) == {("host1", 5432), ("host2", 5432)}
assert isinstance(plan.primary_host, list)
assert isinstance(plan.primary_port, list)
assert plan.failover == (("host1", 5432), ("host2", 5433))
assert plan.target_session_attrs == SessionAttribute("read-write")
assert "host" not in plan.connect_args
assert "port" not in plan.connect_args
assert "target_session_attrs" not in plan.connect_args


def test_host_and_port_stay_paired_through_the_shuffle() -> None:
"""INVARIANT: no plan ever pairs one host's name with another host's port.
def test_build_connection_plan_is_deterministic() -> None:
hosts_and_ports, url = _six_hosts()
assert build_connection_plan(url) == build_connection_plan(url)
assert build_connection_plan(url).failover == tuple(hosts_and_ports)

Broken by shuffling ``primary_host`` and ``primary_port`` as two independent lists, or by
deriving ``failover`` from a second shuffle of its own rather than from the one that produced
the primary order. Both read as tidier code and both silently mis-pair. Nothing downstream can
catch it: ``_connect`` hands whatever pair it is given straight to asyncpg, so a swap surfaces
as a refused connection that is indistinguishable from a host being down -- on the failover
path, which by definition only runs when hosts are already failing. This is also why the DSN
here gives each host a distinct port; with matching ports a swap is unobservable. Six hosts
rather than two for the same reason: the assertions are order-independent, so a mis-pairing
build can still be let through by a shuffle that happens to come out in step, and six hosts put
that at one run in 720 instead of one in two.

async def test_each_connection_gets_a_fresh_bulk_order(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr("db_retry.connections.random", random.Random(0)) # noqa: S311
connect: typing.Final = mock.AsyncMock(return_value=mock.AsyncMock(spec=asyncpg.Connection))
monkeypatch.setattr("asyncpg.connect", connect)
_, url = _six_hosts()
factory: typing.Final = build_connection_factory(url=url, timeout=1.0)
await factory()
await factory()
first, second = (call.kwargs["host"] for call in connect.await_args_list)
assert sorted(first) == sorted(second)
assert first != second


async def test_host_and_port_stay_paired_through_the_shuffle(monkeypatch: pytest.MonkeyPatch) -> None:
"""INVARIANT: no connect attempt ever pairs one host's name with another host's port.

Broken by shuffling the bulk attempt's hosts and ports as two independent lists, or by drawing
them from different shuffles. Both read as tidier code and both silently mis-pair. Nothing
downstream can catch it: ``_connect`` hands whatever it is given straight to asyncpg, so a swap
surfaces as a refused connection that is indistinguishable from a host being down -- on the
failover path, which by definition only runs when hosts are already failing. This is also why
each host here has a distinct port; with matching ports a swap is unobservable. Six hosts
rather than two for the same reason: a mis-pairing build can still be let through by shuffles
that happen to come out in step, and six hosts put that at one run in 720 instead of one in two.
"""
expected: typing.Final = {(f"host{n}", 5431 + n) for n in range(1, 7)}
hosts: typing.Final = "&".join(f"host={host}:{port}" for host, port in sorted(expected))
plan: typing.Final[ConnectionPlan] = build_connection_plan(
sqlalchemy.make_url(f"postgresql+asyncpg://user:password@/database?{hosts}")
)
assert isinstance(plan.primary_host, list)
assert isinstance(plan.primary_port, list)
assert set(plan.failover) == expected
assert set(zip(plan.primary_host, plan.primary_port, strict=True)) == expected
assert list(zip(plan.primary_host, plan.primary_port, strict=True)) == list(plan.failover)
connect: typing.Final = mock.AsyncMock(side_effect=TimeoutError)
monkeypatch.setattr("asyncpg.connect", connect)
hosts_and_ports, url = _six_hosts()
factory: typing.Final = build_connection_factory(url=url, timeout=1.0)
with pytest.raises(asyncpg.TargetServerAttributeNotMatched):
await factory()
bulk, *failover = connect.await_args_list
assert sorted(zip(bulk.kwargs["host"], bulk.kwargs["port"], strict=True)) == hosts_and_ports
assert sorted((call.kwargs["host"], call.kwargs["port"]) for call in failover) == hosts_and_ports


def test_build_connection_plan_connect_args_is_read_only() -> None:
Expand All @@ -122,9 +143,7 @@ def test_build_connection_plan_single_host() -> None:
url: typing.Final = sqlalchemy.make_url(f"postgresql+asyncpg://user:password@host1:{port}/database")
plan: typing.Final[ConnectionPlan] = build_connection_plan(url)
assert plan.failover == ()
assert plan.primary_host == "host1"
assert plan.primary_port == port
assert plan.connect_args["host"] == "host1"
assert plan.connect_args["port"] == port
assert plan.target_session_attrs is None
assert "host" not in plan.connect_args
assert "port" not in plan.connect_args
assert "target_session_attrs" not in plan.connect_args
Loading