From 3d0df4ca15099502bc52bf13ad6bac4cafcea143 Mon Sep 17 00:00:00 2001 From: Dima Anfimov Date: Wed, 16 Sep 2026 16:34:54 +0200 Subject: [PATCH] fix: recover gracefully from connection loss in ack and kick --- README.md | 24 +++++ taskiq_aio_pika/broker.py | 98 ++++++++++++++----- taskiq_aio_pika/retries.py | 17 ++++ tests/test_retries.py | 194 +++++++++++++++++++++++++++++++++++++ 4 files changed, 309 insertions(+), 24 deletions(-) create mode 100644 taskiq_aio_pika/retries.py create mode 100644 tests/test_retries.py diff --git a/README.md b/README.md index c360554..8cb056f 100644 --- a/README.md +++ b/README.md @@ -144,6 +144,30 @@ broker = AioPikaBroker( Once a message has been redelivered more than `delivery_limit` times, RabbitMQ dead-letters it to the broker's dead-letter queue instead of redelivering it again — no application code involved. `delivery_limit` is only supported by quorum queues. See the [RabbitMQ docs](https://www.rabbitmq.com/docs/quorum-queues#poison-message-handling) for details. +## Connection loss and message redelivery + +RabbitMQ deliveries are acknowledged on the specific channel they were delivered on. If the underlying connection drops (a network blip, a broker restart, etc.), any message that was already delivered but not yet acked cannot be acked anymore, even after the connection reconnects — the delivery tag isn't valid on a new channel. `AioPikaBroker` handles this by logging a warning and letting the message go instead of crashing the worker; RabbitMQ automatically requeues the message once the old channel closes. + +The practical consequence is that a task can run **more than once** whenever a connection is lost while the task is in flight — regardless of `ack_time`. This is a property of AMQP itself, not something a broker can paper over, so: + +- Write tasks to be idempotent whenever you can (safe to execute twice with the same effect). +- If you can't, prefer `ack_time="when_received"` (the taskiq default) to shrink the window between delivery and ack, at the cost of losing the message outright if the worker crashes mid-task instead of duplicating it. `ack_time="when_executed"`/`"when_saved"` hold the message unacked for longer (through task execution / result saving), which widens the duplicate-delivery window but guarantees the message isn't lost if the worker itself crashes. + +Publishing (`kick`) is affected too, but differently: right after a connection recovers, there's a short window (order of a second) where the write channel can still be mid-recovery. `AioPikaBroker` retries `kick` automatically in that case — configurable via the `retries` constructor argument: + +```python +from taskiq_aio_pika import AioPikaBroker + +broker = AioPikaBroker( + retries={ + "kick": { + "max_attempts": 4, # set to 0 to disable retrying + "backoff": 0.2, # doubled after each retry + }, + }, +) +``` + ## Custom Queue and Exchange arguments You can pass custom arguments to the underlying RabbitMQ queues and exchange declaration by using the `Queue`/`Exchange` classes from `taskiq_aio_pika`. If you used `faststream` before you are probably familiar with this concept. diff --git a/taskiq_aio_pika/broker.py b/taskiq_aio_pika/broker.py index dc8bb87..fd98992 100644 --- a/taskiq_aio_pika/broker.py +++ b/taskiq_aio_pika/broker.py @@ -1,12 +1,18 @@ import asyncio from collections.abc import AsyncGenerator, Callable from datetime import timedelta +from functools import partial from logging import getLogger from typing import Any, TypeVar import aiormq from aio_pika import DeliveryMode, ExchangeType, Message, connect_robust -from aio_pika.abc import AbstractChannel, AbstractQueue, AbstractRobustConnection +from aio_pika.abc import ( + AbstractChannel, + AbstractIncomingMessage, + AbstractQueue, + AbstractRobustConnection, +) from pamqp.common import FieldTable from taskiq import AckableMessage, AsyncBroker, AsyncResultBackend, BrokerMessage from typing_extensions import Self @@ -19,6 +25,7 @@ ) from taskiq_aio_pika.exchange import Exchange from taskiq_aio_pika.queue import Queue, QueueType +from taskiq_aio_pika.retries import DEFAULT_KICK_RETRIES, RetriesConfig from taskiq_aio_pika.utils import merge_async_iterables _T = TypeVar("_T") @@ -64,6 +71,7 @@ def __init__( delayed_message_exchange: Exchange | None = None, label_for_routing: str = "queue_name", label_for_priority: str = "priority", + retries: RetriesConfig | None = None, **connection_kwargs: Any, ) -> None: """ @@ -82,6 +90,7 @@ def __init__( :param delayed_message_exchange: parameters of exchange that used to send messages with delay. :param label_for_routing: label name to use for routing key selection. :param label_for_priority: label name to use for message priority. + :param retries: per-operation retry policies. :param connection_kwargs: additional keyword arguments, for connect_robust method of aio-pika. """ super().__init__(result_backend, task_id_generator) @@ -96,6 +105,9 @@ def __init__( self._label_for_routing = label_for_routing self._label_for_priority = label_for_priority + self._retries: RetriesConfig = { + "kick": {**DEFAULT_KICK_RETRIES, **(retries or {}).get("kick", {})}, + } self._delay_queue = delay_queue @@ -362,6 +374,7 @@ async def kick(self, message: BrokerMessage) -> None: """ if self.write_channel is None: raise NoStartupError("Please run startup before kicking.") + write_channel = self.write_channel priority = parse_val(int, message.labels.get(self._label_for_priority)) rmq_message = Message( body=message.message, @@ -395,28 +408,65 @@ async def kick(self, message: BrokerMessage) -> None: f"Check routing keys and queue names in broker queues.", ) - if x_delay is None: - exchange = await self.write_channel.get_exchange( - self._exchange.name, - ensure=False, - ) - await exchange.publish(rmq_message, routing_key=routing_key_name) - elif self._delayed_message_exchange_plugin: - rmq_message.headers["x-delay"] = int(x_delay * 1000) - exchange = await self.write_channel.get_exchange( - self._delayed_message_exchange.name, - ) - await exchange.publish(rmq_message, routing_key=routing_key_name) - elif self._delay_queue: - rmq_message.expiration = timedelta(seconds=x_delay) - await self.write_channel.default_exchange.publish( - rmq_message, - routing_key=self._delay_queue.routing_key or self._delay_queue.name, - ) - else: - raise IncorrectRoutingKeyError( - "Delay requested but no delay queue or delayed-message-exchange " - "is configured in the broker.", + async def _publish() -> None: + if x_delay is None: + exchange = await write_channel.get_exchange( + self._exchange.name, + ensure=False, + ) + await exchange.publish(rmq_message, routing_key=routing_key_name) + elif self._delayed_message_exchange_plugin: + rmq_message.headers["x-delay"] = int(x_delay * 1000) + exchange = await write_channel.get_exchange( + self._delayed_message_exchange.name, + ) + await exchange.publish(rmq_message, routing_key=routing_key_name) + elif self._delay_queue: + rmq_message.expiration = timedelta(seconds=x_delay) + await write_channel.default_exchange.publish( + rmq_message, + routing_key=self._delay_queue.routing_key or self._delay_queue.name, + ) + else: + raise IncorrectRoutingKeyError( + "Delay requested but no delay queue or delayed-message-exchange is configured in the broker.", + ) + + await self._publish_with_retry(_publish) + + async def _publish_with_retry(self, _publish: Callable[[], Any]) -> None: + """Run a publish callback, retrying on a transient channel-recovery race.""" + max_attempts = self._retries["kick"]["max_attempts"] + delay = self._retries["kick"]["backoff"] + for attempt in range(max_attempts + 1): + try: + await _publish() + return + except aiormq.exceptions.ChannelInvalidStateError: + if attempt == max_attempts: + raise + logger.warning( + "Publish failed because the write channel was invalidated by a connection recovery race; " + "retrying in %.2fs (attempt %d/%d).", + delay, + attempt + 1, + max_attempts, + ) + await asyncio.sleep(delay) + delay *= 2 + + @staticmethod + async def _safe_ack(message: AbstractIncomingMessage, queue_name: str) -> None: + """Ack a message, tolerating a channel invalidated by a connection loss.""" + try: + await message.ack() + except aiormq.exceptions.ChannelInvalidStateError: + logger.warning( + "Could not ack message (delivery_tag=%s, redelivered=%s) on queue '%s' - the channel was invalidated by" + " a connection loss.", + message.delivery_tag, + message.redelivered, + queue_name, ) async def listen(self) -> AsyncGenerator[AckableMessage, None]: @@ -442,7 +492,7 @@ async def body( async for message in iterator: yield AckableMessage( data=message.body, - ack=message.ack, + ack=partial(self._safe_ack, message, queue.name), ) except (RuntimeError, asyncio.CancelledError): # Suppress errors during iterator cleanup if channel is being closed diff --git a/taskiq_aio_pika/retries.py b/taskiq_aio_pika/retries.py new file mode 100644 index 0000000..eca20cb --- /dev/null +++ b/taskiq_aio_pika/retries.py @@ -0,0 +1,17 @@ +from typing_extensions import TypedDict + + +class KickRetries(TypedDict, total=False): + """Retry policy for `AioPikaBroker.kick`.""" + + max_attempts: int + backoff: float + + +class RetriesConfig(TypedDict, total=False): + """Per-operation retry policies for `AioPikaBroker`.""" + + kick: KickRetries + + +DEFAULT_KICK_RETRIES: KickRetries = {"max_attempts": 4, "backoff": 0.2} diff --git a/tests/test_retries.py b/tests/test_retries.py new file mode 100644 index 0000000..72fc4b4 --- /dev/null +++ b/tests/test_retries.py @@ -0,0 +1,194 @@ +from types import SimpleNamespace +from typing import cast +from unittest.mock import AsyncMock, patch + +import aiormq +import pytest +from aio_pika.abc import AbstractIncomingMessage +from taskiq import BrokerMessage + +from taskiq_aio_pika import AioPikaBroker + + +class TestPublishWithRetry: + async def test_when_publish_succeeds_immediately__it_is_called_only_once( + self, + ) -> None: + broker = AioPikaBroker(retries={"kick": {"max_attempts": 4, "backoff": 0.01}}) + publish = AsyncMock(return_value=None) + + await broker._publish_with_retry(publish) + + assert publish.await_count == 1 + + async def test_when_publish_fails_then_recovers__it_retries_until_success( + self, + ) -> None: + broker = AioPikaBroker(retries={"kick": {"max_attempts": 4, "backoff": 0.01}}) + publish = AsyncMock( + side_effect=[ + aiormq.exceptions.ChannelInvalidStateError("simulated"), + aiormq.exceptions.ChannelInvalidStateError("simulated"), + None, + ], + ) + + await broker._publish_with_retry(publish) + + assert publish.await_count == 3 + + async def test_when_publish_always_fails__it_gives_up_after_max_attempts( + self, + ) -> None: + broker = AioPikaBroker(retries={"kick": {"max_attempts": 2, "backoff": 0.01}}) + publish = AsyncMock( + side_effect=aiormq.exceptions.ChannelInvalidStateError("always failing"), + ) + + with pytest.raises(aiormq.exceptions.ChannelInvalidStateError): + await broker._publish_with_retry(publish) + + # the initial attempt plus `max_attempts` retries + assert publish.await_count == 3 + + async def test_when_max_attempts_is_zero__it_fails_on_first_error_without_retrying( + self, + ) -> None: + broker = AioPikaBroker(retries={"kick": {"max_attempts": 0, "backoff": 0.01}}) + publish = AsyncMock( + side_effect=aiormq.exceptions.ChannelInvalidStateError("simulated"), + ) + + with pytest.raises(aiormq.exceptions.ChannelInvalidStateError): + await broker._publish_with_retry(publish) + + assert publish.await_count == 1 + + async def test_when_publish_raises_other_error__it_is_not_retried(self) -> None: + broker = AioPikaBroker(retries={"kick": {"max_attempts": 4, "backoff": 0.01}}) + publish = AsyncMock(side_effect=ValueError("not a channel-recovery error")) + + with pytest.raises(ValueError, match="not a channel-recovery error"): + await broker._publish_with_retry(publish) + + assert publish.await_count == 1 + + async def test_when_retrying__backoff_doubles_between_attempts( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + broker = AioPikaBroker(retries={"kick": {"max_attempts": 3, "backoff": 0.1}}) + publish = AsyncMock( + side_effect=[ + aiormq.exceptions.ChannelInvalidStateError("simulated"), + aiormq.exceptions.ChannelInvalidStateError("simulated"), + aiormq.exceptions.ChannelInvalidStateError("simulated"), + None, + ], + ) + sleep_delays: list[float] = [] + + async def fake_sleep(delay: float) -> None: + sleep_delays.append(delay) + + monkeypatch.setattr("taskiq_aio_pika.broker.asyncio.sleep", fake_sleep) + + await broker._publish_with_retry(publish) + + assert sleep_delays == [0.1, 0.2, 0.4] + + +class TestSafeAck: + async def test_when_ack_succeeds__nothing_special_happens(self) -> None: + message = SimpleNamespace(ack=AsyncMock(return_value=None)) + await AioPikaBroker._safe_ack( + cast(AbstractIncomingMessage, message), + "some_queue", + ) + message.ack.assert_awaited_once() + + async def test_when_channel_invalid_state_error_raised__it_is_swallowed_and_logged( + self, + caplog: pytest.LogCaptureFixture, + ) -> None: + message = SimpleNamespace( + ack=AsyncMock( + side_effect=aiormq.exceptions.ChannelInvalidStateError("simulated"), + ), + delivery_tag=42, + redelivered=False, + ) + with caplog.at_level("WARNING", logger="taskiq.aio_pika_broker"): + await AioPikaBroker._safe_ack( + cast(AbstractIncomingMessage, message), + "some_queue", + ) + assert any( + "42" in record.message and "some_queue" in record.message + for record in caplog.records + ) + + async def test_when_other_error_raised__it_propagates(self) -> None: + message = SimpleNamespace(ack=AsyncMock(side_effect=ValueError("boom"))) + with pytest.raises(ValueError, match="boom"): + await AioPikaBroker._safe_ack( + cast(AbstractIncomingMessage, message), + "some_queue", + ) + + +class TestKickRetriesIntegration: + async def test_when_write_channel_is_transiently_invalid__kick_retries_and_succeeds( + self, + broker: AioPikaBroker, + ) -> None: + publish = AsyncMock( + side_effect=[ + aiormq.exceptions.ChannelInvalidStateError("simulated recovery race"), + aiormq.exceptions.ChannelInvalidStateError("simulated recovery race"), + None, + ], + ) + broker._retries["kick"]["backoff"] = 0.01 + stub_exchange = SimpleNamespace(publish=publish) + with patch.object( + broker.write_channel, + "get_exchange", + AsyncMock(return_value=stub_exchange), + ): + await broker.kick( + BrokerMessage( + task_id="1", + task_name="t1", + message=b"payload", + labels={}, + ), + ) + assert publish.await_count == 3 + + async def test_when_write_channel_never_recovers__kick_gives_up_and_raises( + self, + broker: AioPikaBroker, + ) -> None: + publish = AsyncMock( + side_effect=aiormq.exceptions.ChannelInvalidStateError("always failing"), + ) + broker._retries["kick"] = {"max_attempts": 2, "backoff": 0.01} + stub_exchange = SimpleNamespace(publish=publish) + with ( + patch.object( + broker.write_channel, + "get_exchange", + AsyncMock(return_value=stub_exchange), + ), + pytest.raises(aiormq.exceptions.ChannelInvalidStateError), + ): + await broker.kick( + BrokerMessage( + task_id="1", + task_name="t1", + message=b"payload", + labels={}, + ), + ) + assert publish.await_count == 3