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
24 changes: 24 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
98 changes: 74 additions & 24 deletions taskiq_aio_pika/broker.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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")
Expand Down Expand Up @@ -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:
"""
Expand All @@ -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)
Expand All @@ -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

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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]:
Expand All @@ -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
Expand Down
17 changes: 17 additions & 0 deletions taskiq_aio_pika/retries.py
Original file line number Diff line number Diff line change
@@ -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}
Loading