From ee60f0813b1bec4f93d774d07da513cab6bcd76e Mon Sep 17 00:00:00 2001 From: vvanglro Date: Wed, 29 Jul 2026 16:50:56 +0800 Subject: [PATCH 1/3] feat(stream): per-task reclaim, prefetch backpressure, NOGROUP self-heal RedisStreamBroker reclaim was driven by a fixed idle_timeout and a broker-side Redis lock, which both double-ran long tasks and recovered short crashed tasks slowly. Replace it with per-task-deadline reclaim: - Resolve reclaim deadline from the message's `timeout` label (via formatter.loads) plus reclaim_timeout_grace, falling back to idle_timeout for messages without a timeout label. - Drop the autoclaim Redis lock; rely on XCLAIM min-idle-time for server-side atomic dedup (unacknowledged_lock_timeout is deprecated and ignored). - Protect only messages held by the current listener instance from reclaim (tracked via a local delivered set), so a worker sharing a consumer_name with a dead predecessor still recovers its pending. - Gate reclaim sweeps with reclaim_interval (default 30s) to avoid scanning pending on every listen iteration. - Enforce prefetch backpressure: do not XREADGROUP while xread_count delivered-but-unacked messages are outstanding. - Claim broker-local buffered entries to an internal abandoned consumer on listener close, so the next reclaim sweep can recover not-yet-yielded messages immediately. - Recreate a missing consumer group on NOGROUP during xpending/xreadgroup. Tests cover timeout-label reclaim, idle_timeout reclaim, shared consumer_name reclaim, prefetch backpressure, buffered-message handoff on listener close, and NOGROUP self-heal. README documents the new reclaim/backpressure/close behavior. --- README.md | 17 ++ taskiq_redis/redis_broker.py | 474 ++++++++++++++++++++++++++++++----- tests/test_broker.py | 286 ++++++++++++++++++++- 3 files changed, 707 insertions(+), 70 deletions(-) diff --git a/README.md b/README.md index 00148af..c7614bf 100644 --- a/README.md +++ b/README.md @@ -98,6 +98,23 @@ Stream brokers use redis [stream type](https://redis.io/docs/latest/develop/data > This broker **supports** acknowledgements and therefore is fine to use in cases when data durability is > required. +`RedisStreamBroker` is the recommended stream broker for single-node redis. It uses a consumer group so every message is delivered to exactly one worker, and pending entries are tracked by redis until they are acknowledged. + +Recovery of messages orphaned by a crashed worker is driven by `XCLAIM`: + +* Messages with a `timeout` label are reclaimed once they have been pending for longer than the task's own timeout plus `reclaim_timeout_grace`. A long-running task is never stolen while it is still within its declared timeout. +* Messages without a `timeout` label fall back to `idle_timeout` (10 minutes by default). +* Reclaim sweeps run at most every `reclaim_interval` (30 seconds by default) to avoid hammering redis on every listen loop. Set it to `0` to scan on every iteration. +* A worker that shares its `consumer_name` with a previous, dead worker can still reclaim that worker's pending messages; only messages actually held by the current listener instance are protected from re-delivery. + +To avoid one worker hoarding the backlog, the listener does not fetch new messages while it already has `xread_count` delivered but unacknowledged messages. Set `xread_count=None` to disable this limit. + +When a listener is closed with messages already fetched from redis but not yet yielded to taskiq, those buffered entries are claimed to an internal `abandoned` consumer and stamped as very idle. The next reclaim sweep can recover them immediately instead of waiting for `idle_timeout`. + +If the consumer group is removed out of band, the broker recreates it on the next `XREADGROUP` so an in-use queue heals instead of spinning on errors. + +The legacy `unacknowledged_lock_timeout` parameter is deprecated and ignored; reclaim correctness now relies on the `XCLAIM` min-idle-time check instead of a broker-side redis lock. + ## RedisAsyncResultBackend configuration RedisAsyncResultBackend parameters: diff --git a/taskiq_redis/redis_broker.py b/taskiq_redis/redis_broker.py index b1ad461..68237d9 100644 --- a/taskiq_redis/redis_broker.py +++ b/taskiq_redis/redis_broker.py @@ -1,11 +1,15 @@ +import asyncio +import time import uuid -from collections.abc import AsyncGenerator, Awaitable, Callable +from collections.abc import AsyncGenerator, Awaitable, Callable, Iterable +from contextlib import suppress from logging import getLogger from typing import ( TYPE_CHECKING, Any, TypeAlias, TypeVar, + cast, ) from redis.asyncio import BlockingConnectionPool, Connection, Redis, ResponseError @@ -18,6 +22,9 @@ logger = getLogger("taskiq.redis_broker") +ABANDONED_CONSUMER = "abandoned" +ABANDONED_IDLE_MS = 10**12 + if TYPE_CHECKING: _BlockingConnectionPool: TypeAlias = BlockingConnectionPool[Connection] # type: ignore @@ -164,6 +171,8 @@ def __init__( unacknowledged_batch_size: int = 100, unacknowledged_lock_timeout: float | None = None, xread_count: int | None = 100, + reclaim_interval: int = 30000, + reclaim_timeout_grace: int = 10000, additional_streams: dict[str, str | int] | None = None, **connection_kwargs: Any, ) -> None: @@ -188,12 +197,19 @@ def __init__( :param approximate: decides wether to trim the stream immediately (False) or later on (True) :param xread_count: number of messages to fetch from the stream at once. + The broker won't make another XREADGROUP call while this listener + already has this many delivered but unacknowledged messages. Set to + None to disable this limit. :param additional_streams: additional streams to read from. Each key is a stream name, value is a consumer id. :param unacknowledged_batch_size: number of unacknowledged messages to fetch. - :param unacknowledged_lock_timeout: time in seconds before auto-releasing - the lock. Useful when the worker crashes or gets killed. - If not set, the lock can remain locked indefinitely. + :param unacknowledged_lock_timeout: deprecated and ignored. Redis' XCLAIM + min-idle-time check is used instead of a broker-side lock. + :param reclaim_interval: milliseconds between timed-out message scans. + Set to 0 to scan on every listen iteration. + :param reclaim_timeout_grace: extra time in milliseconds added to + message's timeout label before it can be reclaimed. Messages without + a timeout label still use idle_timeout. """ super().__init__( url, @@ -215,6 +231,40 @@ def __init__( self.unacknowledged_batch_size = unacknowledged_batch_size self.unacknowledged_lock_timeout = unacknowledged_lock_timeout self.count = xread_count + self.reclaim_interval = reclaim_interval + self.reclaim_timeout_grace = reclaim_timeout_grace + + def _get_available_message_count(self, unacked: int) -> int | None: + """Return how many more messages this listener can reserve.""" + if self.count is None: + return None + return max(0, self.count - unacked) + + def _should_reclaim(self, last_reclaim: float) -> bool: + """Return whether the periodic pending-message sweep should run now.""" + if self.reclaim_interval <= 0: + return True + return time.monotonic() - last_reclaim >= self.reclaim_interval / 1000 + + @staticmethod + def _to_str(value: Any) -> str: + if isinstance(value, bytes): + return value.decode() + return str(value) + + def _message_key(self, stream: Any, message_id: Any) -> tuple[str, str]: + """Build a stable identity for one stream entry.""" + return self._to_str(stream), self._to_str(message_id) + + @staticmethod + def _group_delivered_by_stream( + delivered: set[tuple[str, str]], + ) -> dict[str, list[str]]: + """Group listener-held message ids by their Redis stream key.""" + grouped: dict[str, list[str]] = {} + for stream, message_id in delivered: + grouped.setdefault(stream, []).append(message_id) + return grouped async def _declare_consumer_group(self) -> None: """ @@ -257,69 +307,377 @@ async def kick(self, message: BrokerMessage) -> None: approximate=self.approximate, ) - def _ack_generator(self, id: str, queue_name: str) -> Callable[[], Awaitable[None]]: + def _ack_generator( + self, + id: str, + queue_name: str, + on_ack: Callable[[], None] | None = None, + ) -> Callable[[], Awaitable[None]]: + acked = False + async def _ack() -> None: + nonlocal acked async with Redis(connection_pool=self.connection_pool) as redis_conn: - await redis_conn.xack( - queue_name, - self.consumer_group_name, - id, - ) + try: + await redis_conn.xack( + queue_name, + self.consumer_group_name, + id, + ) + finally: + if not acked and on_ack is not None: + acked = True + on_ack() return _ack - async def listen(self) -> AsyncGenerator[AckableMessage, None]: - """Listen to incoming messages.""" - async with Redis(connection_pool=self.connection_pool) as redis_conn: - while True: - logger.debug("Starting fetching new messages") - fetched = await redis_conn.xreadgroup( + def _get_message_reclaim_timeout(self, message: dict[bytes, bytes]) -> int: + """Resolve how long this message may stay pending before reclaim. + + Taskiq stores execution timeout in the serialized task message labels. + If the payload cannot be decoded or has no timeout label, fall back to + the broker-level idle_timeout to preserve legacy behavior. + """ + raw_data = message.get(b"data") + if raw_data is None: + return self.idle_timeout + + with suppress(Exception): + timeout = self.formatter.loads(raw_data).labels.get("timeout") + if timeout is not None: + return int(float(timeout) * 1000) + self.reclaim_timeout_grace + return self.idle_timeout + + async def _get_pending_message( + self, + redis_conn: Redis, + stream: str, + message_id: str, + ) -> dict[bytes, bytes] | None: + """Fetch the stream entry body for one pending message id.""" + results = await redis_conn.xrange( + stream, + min=message_id, + max=message_id, + count=1, + ) + if not results: + return None + return cast("dict[bytes, bytes]", results[0][1]) + + async def _claim_timed_out_messages( + self, + redis_conn: Redis, + stream: str, + count: int, + protected_message_ids: set[tuple[str, str]], + ) -> list[tuple[str, dict[bytes, bytes]]]: + """Claim pending messages that exceeded their reclaim deadline. + + protected_message_ids contains entries delivered by this listener + instance and not acked yet. We skip those to avoid re-delivering our own + in-flight work, while still allowing a restarted worker with the same + Redis consumer name to recover messages left by its predecessor. + """ + try: + pending = await redis_conn.xpending_range( + stream, + self.consumer_group_name, + min="-", + max="+", + count=self.unacknowledged_batch_size, + idle=0, + ) + except ResponseError as exc: + if "NOGROUP" not in str(exc): + raise + logger.info("Consumer group missing for %s, recreating", stream) + await self._declare_consumer_group() + return [] + + claimed: list[tuple[str, dict[bytes, bytes]]] = [] + for pending_message in pending: + if len(claimed) >= count: + break + + message_id = pending_message["message_id"] + if isinstance(message_id, bytes): + message_id = message_id.decode() + if self._message_key(stream, message_id) in protected_message_ids: + continue + + message = await self._get_pending_message( + redis_conn, + stream, + cast(str, message_id), + ) + if message is None: + await redis_conn.xclaim( + stream, self.consumer_group_name, self.consumer_name, - { - self.queue_name: ">", - **self.additional_streams, # type: ignore[dict-item] - }, - block=self.block, - noack=False, - count=self.count, + min_idle_time=0, + message_ids=[message_id], + justid=True, ) - if not fetched: - continue - for stream, msg_list in fetched: # type: ignore[str-unpack] - for msg_id, msg in msg_list: # type: ignore[str-unpack,union-attr] - logger.debug("Received message: %s", msg) - yield AckableMessage( - data=msg[b"data"], # type: ignore[arg-type,index] - ack=self._ack_generator(id=msg_id, queue_name=stream), # type: ignore[arg-type] - ) - logger.debug("Starting fetching unacknowledged messages") - for stream in [self.queue_name, *self.additional_streams.keys()]: - pipe = redis_conn.pipeline() - lock = pipe.lock( - f"autoclaim:{self.consumer_group_name}:{stream}", - timeout=self.unacknowledged_lock_timeout, - ) - await lock.acquire() - await pipe.xautoclaim( - name=stream, - groupname=self.consumer_group_name, - consumername=self.consumer_name, - min_idle_time=self.idle_timeout, - count=self.unacknowledged_batch_size, - ) - await lock.release() - results = await pipe.execute() - pending = results[1] + continue - logger.debug( - "Found %d pending messages in stream %s", - len(pending[1]), + reclaim_timeout = self._get_message_reclaim_timeout(message) + time_since_delivered = int( + cast(Any, pending_message.get("time_since_delivered", 0)), + ) + if time_since_delivered < reclaim_timeout: + continue + + # XCLAIM rechecks min_idle_time inside Redis, so concurrent workers + # racing for the same overdue message cannot both claim it. + result = await redis_conn.xclaim( + stream, + self.consumer_group_name, + self.consumer_name, + min_idle_time=reclaim_timeout, + message_ids=[message_id], + ) + claimed.extend(cast("list[tuple[str, dict[bytes, bytes]]]", result)) + return claimed + + async def _claim_available_timed_out_messages( + self, + redis_conn: Redis, + count: int, + protected_message_ids: set[tuple[str, str]], + ) -> list[tuple[str, dict[bytes, bytes], str]]: + """Claim overdue messages across all configured streams.""" + claimed: list[tuple[str, dict[bytes, bytes], str]] = [] + for stream in [self.queue_name, *self.additional_streams.keys()]: + remaining_count = count - len(claimed) + if remaining_count <= 0: + break + for msg_id, msg in await self._claim_timed_out_messages( + redis_conn, + stream, + remaining_count, + protected_message_ids, + ): + claimed.append((msg_id, msg, stream)) + return claimed + + async def _read_new_messages( + self, + redis_conn: Redis, + count: int | None, + ) -> Any: + """Read newly added stream messages, recreating a missing group once.""" + try: + return await redis_conn.xreadgroup( + self.consumer_group_name, + self.consumer_name, + { + self.queue_name: ">", + **self.additional_streams, # type: ignore[dict-item] + }, + block=self.block, + noack=False, + count=count, + ) + except ResponseError as exc: + if "NOGROUP" not in str(exc): + raise + logger.info("Consumer group missing, recreating") + await self._declare_consumer_group() + return [] + + async def _abandon_buffered_messages( + self, + redis_conn: Redis, + buffered: list[tuple[str, dict[bytes, bytes], str]], + ) -> None: + """Make fetched but not yet yielded messages immediately reclaimable.""" + message_keys = { + self._message_key(stream, msg_id) for msg_id, _, stream in buffered + } + grouped = self._group_delivered_by_stream(message_keys) + for stream, message_ids in grouped.items(): + try: + await redis_conn.xclaim( + stream, + self.consumer_group_name, + ABANDONED_CONSUMER, + min_idle_time=0, + message_ids=cast(Any, message_ids), + idle=ABANDONED_IDLE_MS, + justid=True, + ) + except ResponseError as exc: + if "NOGROUP" not in str(exc): + logger.warning( + "Failed to abandon messages in stream %s", stream, + exc_info=True, + ) + + def _build_ackable_message( + self, + msg_id: str, + msg: dict[bytes, bytes], + queue_name: str, + on_ack: Callable[[], None] | None = None, + ) -> AckableMessage: + return AckableMessage( + data=msg[b"data"], + ack=self._ack_generator(id=msg_id, queue_name=queue_name, on_ack=on_ack), + ) + + def _build_ackable_messages( + self, + messages: Iterable[tuple[str, dict[bytes, bytes], str]], + make_on_ack: Callable[[tuple[str, str]], Callable[[], None]], + ) -> list[AckableMessage]: + """Convert Redis stream entries to AckableMessage objects. + + The caller adds each entry to the listener-local delivered set right + before yielding it, because entries fetched but not yielded yet can be + abandoned immediately on generator close. + """ + ackable_messages = [] + for msg_id, msg, stream in messages: + logger.debug("Received message: %s", msg) + message_key = self._message_key(stream, msg_id) + ackable_messages.append( + self._build_ackable_message( + msg_id=msg_id, + msg=msg, + queue_name=stream, + on_ack=make_on_ack(message_key), + ), + ) + return ackable_messages + + @staticmethod + def _flatten_fetched_messages( + fetched: Any, + ) -> list[tuple[str, dict[bytes, bytes], str]]: + """Normalize XREADGROUP's grouped response to (id, body, stream) tuples.""" + messages = [] + for stream, msg_list in fetched: + for msg_id, msg in msg_list: + messages.append((msg_id, msg, stream)) + return messages + + async def _build_reclaimed_ackable_messages( + self, + redis_conn: Redis, + count: int, + delivered: set[tuple[str, str]], + ) -> list[tuple[str, dict[bytes, bytes], str]]: + """Claim overdue messages for taskiq's receiver.""" + return await self._claim_available_timed_out_messages( + redis_conn, + count or self.unacknowledged_batch_size, + delivered, + ) + + async def _build_due_reclaimed_messages( + self, + redis_conn: Redis, + count: int, + delivered: set[tuple[str, str]], + last_reclaim: float, + ) -> tuple[float, list[tuple[str, dict[bytes, bytes], str]]]: + """Return overdue pending messages only when the reclaim interval elapsed.""" + if not self._should_reclaim(last_reclaim): + return last_reclaim, [] + return time.monotonic(), await self._build_reclaimed_ackable_messages( + redis_conn, + count, + delivered, + ) + + async def _build_new_ackable_messages( + self, + redis_conn: Redis, + count: int | None, + ) -> list[tuple[str, dict[bytes, bytes], str]]: + """Read fresh stream messages for taskiq's receiver.""" + fetched = await self._read_new_messages(redis_conn, count) + if not fetched: + return [] + return self._flatten_fetched_messages(fetched) + + async def _yield_buffered_messages( + self, + buffered: list[tuple[str, dict[bytes, bytes], str]], + delivered: set[tuple[str, str]], + make_on_ack: Callable[[tuple[str, str]], Callable[[], None]], + ) -> AsyncGenerator[AckableMessage, None]: + """Yield fetched messages and keep only not-yielded entries in buffer.""" + messages = self._build_ackable_messages(buffered, make_on_ack) + while buffered and messages: + msg_id, _, stream = buffered.pop(0) + message = messages.pop(0) + delivered.add(self._message_key(stream, msg_id)) + yield message + + async def listen(self) -> AsyncGenerator[AckableMessage, None]: + """Listen to incoming messages with local prefetch/backpressure.""" + unacked = 0 + # Only entries delivered by this listener instance are protected from + # reclaim. Reusing the same Redis consumer_name after a restart is still + # recoverable because the new listener starts with an empty set. + delivered: set[tuple[str, str]] = set() + buffered: list[tuple[str, dict[bytes, bytes], str]] = [] + slot_freed = asyncio.Event() + last_reclaim = 0.0 + + def on_ack(message_key: tuple[str, str]) -> None: + nonlocal unacked + delivered.discard(message_key) + unacked = max(0, unacked - 1) + slot_freed.set() + + def make_on_ack(message_key: tuple[str, str]) -> Callable[[], None]: + def _on_ack() -> None: + on_ack(message_key) + + return _on_ack + + async with Redis(connection_pool=self.connection_pool) as redis_conn: + try: + while True: + count = self._get_available_message_count(unacked) + if count == 0: + # Do not reserve more stream entries while all local + # prefetch slots are occupied. + await slot_freed.wait() + slot_freed.clear() + continue + last_reclaim, buffered = await self._build_due_reclaimed_messages( + redis_conn, + count or self.unacknowledged_batch_size, + delivered, + last_reclaim, ) - for msg_id, msg in pending[1]: - logger.debug("Received message: %s", msg) - yield AckableMessage( - data=msg[b"data"], - ack=self._ack_generator(id=msg_id, queue_name=stream), + unacked += len(buffered) + async for message in self._yield_buffered_messages( + buffered, + delivered, + make_on_ack, + ): + yield message + count = self._get_available_message_count(unacked) + if count != 0: + logger.debug("Starting fetching new messages") + buffered = await self._build_new_ackable_messages( + redis_conn, + count, ) + unacked += len(buffered) + async for message in self._yield_buffered_messages( + buffered, + delivered, + make_on_ack, + ): + yield message + finally: + if buffered: + await self._abandon_buffered_messages(redis_conn, buffered) diff --git a/tests/test_broker.py b/tests/test_broker.py index 3ee4665..0722ef2 100644 --- a/tests/test_broker.py +++ b/tests/test_broker.py @@ -4,6 +4,7 @@ import pytest from redis.asyncio import Redis from taskiq import AckableMessage, AsyncBroker, BrokerMessage +from taskiq.message import TaskiqMessage from taskiq_redis import ( ListQueueBroker, @@ -14,7 +15,7 @@ RedisStreamClusterBroker, RedisStreamSentinelBroker, ) -from taskiq_redis.redis_broker import RedisStreamBroker +from taskiq_redis.redis_broker import ABANDONED_CONSUMER, RedisStreamBroker def test_no_url_should_raise_typeerror() -> None: @@ -435,11 +436,162 @@ async def test_maxlen_in_sentinel_stream_broker( @pytest.mark.anyio -async def test_unacknowledged_lock_timeout_in_stream_broker( +async def test_stream_broker_reclaims_messages_by_task_timeout( + redis_url: str, +) -> None: + queue_name = uuid.uuid4().hex + consumer_group_name = uuid.uuid4().hex + + first_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=uuid.uuid4().hex, + xread_block=50, + ) + second_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=uuid.uuid4().hex, + xread_block=50, + reclaim_interval=0, + reclaim_timeout_grace=0, + ) + + await first_broker.startup() + await second_broker.startup() + + taskiq_message = TaskiqMessage( + task_id=uuid.uuid4().hex, + task_name=uuid.uuid4().hex, + labels={"timeout": 0.1}, + args=[], + kwargs={}, + ) + broker_message = first_broker.formatter.dumps(taskiq_message) + await first_broker.kick(broker_message) + + first_message = await get_message(first_broker) + assert isinstance(first_message, AckableMessage) + assert first_message.data == broker_message.message + + await asyncio.sleep(0.15) + reclaimed_message = await asyncio.wait_for(get_message(second_broker), timeout=2) + + assert isinstance(reclaimed_message, AckableMessage) + assert reclaimed_message.data == broker_message.message + await reclaimed_message.ack() # type: ignore + + await first_broker.shutdown() + await second_broker.shutdown() + + +@pytest.mark.anyio +async def test_stream_broker_unacked_message_is_reclaimed_by_idle_timeout( + redis_url: str, + valid_broker_message: BrokerMessage, +) -> None: + """A message without a timeout label is reclaimed after idle_timeout. + + The first consumer fetches but never acks; the second consumer (with a + small idle_timeout) reclaims it via XCLAIM. + """ + queue_name = uuid.uuid4().hex + consumer_group_name = uuid.uuid4().hex + + first_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=uuid.uuid4().hex, + xread_block=50, + ) + second_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=uuid.uuid4().hex, + xread_block=50, + idle_timeout=100, + reclaim_interval=0, + reclaim_timeout_grace=0, + ) + + await first_broker.startup() + await second_broker.startup() + await first_broker.kick(valid_broker_message) + + first_message = await get_message(first_broker) + assert isinstance(first_message, AckableMessage) + assert first_message.data == valid_broker_message.message + # Do not ack — leave it pending for the first consumer. + + reclaimed_message = await asyncio.wait_for(get_message(second_broker), timeout=3) + assert isinstance(reclaimed_message, AckableMessage) + assert reclaimed_message.data == valid_broker_message.message + await reclaimed_message.ack() # type: ignore + + await first_broker.shutdown() + await second_broker.shutdown() + + +@pytest.mark.anyio +async def test_stream_broker_reclaims_messages_with_shared_consumer_name( + redis_url: str, + valid_broker_message: BrokerMessage, +) -> None: + """A restarted worker can reclaim its previous consumer's pending messages.""" + queue_name = uuid.uuid4().hex + consumer_group_name = uuid.uuid4().hex + consumer_name = uuid.uuid4().hex + + first_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=consumer_name, + xread_block=50, + ) + second_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=consumer_name, + xread_block=50, + idle_timeout=100, + reclaim_interval=0, + ) + + await first_broker.startup() + await second_broker.startup() + await first_broker.kick(valid_broker_message) + + first_message = await get_message(first_broker) + assert isinstance(first_message, AckableMessage) + assert first_message.data == valid_broker_message.message + + reclaimed_message = await asyncio.wait_for(get_message(second_broker), timeout=3) + assert isinstance(reclaimed_message, AckableMessage) + assert reclaimed_message.data == valid_broker_message.message + await reclaimed_message.ack() # type: ignore + + await first_broker.shutdown() + await second_broker.shutdown() + + +@pytest.mark.anyio +async def test_stream_broker_xread_count_limits_unacked_messages( redis_url: str, valid_broker_message: BrokerMessage, ) -> None: - unacknowledged_lock_timeout = 1 + """The listener does not read more messages while at xread_count capacity.""" queue_name = uuid.uuid4().hex consumer_group_name = uuid.uuid4().hex @@ -448,22 +600,132 @@ async def test_unacknowledged_lock_timeout_in_stream_broker( approximate=False, queue_name=queue_name, consumer_group_name=consumer_group_name, - unacknowledged_lock_timeout=unacknowledged_lock_timeout, + xread_block=50, + xread_count=1, + reclaim_interval=0, ) await broker.startup() await broker.kick(valid_broker_message) + await broker.kick(valid_broker_message) - message = await get_message(broker) - assert isinstance(message, AckableMessage) - assert message.data == valid_broker_message.message + iterator = broker.listen() + first_message = await iterator.__anext__() + assert isinstance(first_message, AckableMessage) + + second_task = asyncio.create_task(iterator.__anext__()) + await asyncio.sleep(0.2) + assert not second_task.done() + + async with Redis(connection_pool=broker.connection_pool) as redis: + pending = await redis.xpending_range( + queue_name, + consumer_group_name, + min="-", + max="+", + count=10, + ) + assert len(pending) == 1 + + await first_message.ack() # type: ignore + second_message = await asyncio.wait_for(second_task, timeout=2) + assert isinstance(second_message, AckableMessage) + await second_message.ack() # type: ignore + + await iterator.aclose() + await broker.shutdown() + + +@pytest.mark.anyio +async def test_stream_broker_abandons_buffered_messages_on_close( + redis_url: str, + valid_broker_message: BrokerMessage, +) -> None: + """Messages fetched but not yielded are handed back on generator close.""" + queue_name = uuid.uuid4().hex + consumer_group_name = uuid.uuid4().hex + + first_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=uuid.uuid4().hex, + xread_block=50, + xread_count=2, + reclaim_interval=0, + ) + second_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name=uuid.uuid4().hex, + xread_block=50, + xread_count=1, + reclaim_interval=0, + ) + + await first_broker.startup() + await second_broker.startup() + await first_broker.kick(valid_broker_message) + await first_broker.kick(valid_broker_message) + + iterator = first_broker.listen() + first_message = await iterator.__anext__() + assert isinstance(first_message, AckableMessage) + + await iterator.aclose() + + async with Redis(connection_pool=first_broker.connection_pool) as redis: + pending = await redis.xpending_range( + queue_name, + consumer_group_name, + min="-", + max="+", + count=10, + ) + + pending_by_consumer = {entry["consumer"] for entry in pending} + assert ABANDONED_CONSUMER.encode() in pending_by_consumer + assert first_broker.consumer_name.encode() in pending_by_consumer + + reclaimed_message = await asyncio.wait_for(get_message(second_broker), timeout=2) + assert isinstance(reclaimed_message, AckableMessage) + assert reclaimed_message.data == valid_broker_message.message + + await first_message.ack() # type: ignore + await reclaimed_message.ack() # type: ignore + await first_broker.shutdown() + await second_broker.shutdown() + + +@pytest.mark.anyio +async def test_stream_broker_recreates_missing_consumer_group( + redis_url: str, + valid_broker_message: BrokerMessage, +) -> None: + queue_name = uuid.uuid4().hex + consumer_group_name = uuid.uuid4().hex + + broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_id="0", + xread_block=50, + reclaim_interval=0, + ) + await broker.startup() async with Redis(connection_pool=broker.connection_pool) as redis: - lock_key = f"autoclaim:{consumer_group_name}:{queue_name}" - await redis.exists(lock_key) - await asyncio.sleep(unacknowledged_lock_timeout + 0.5) + await redis.xgroup_destroy(queue_name, consumer_group_name) + await broker.kick(valid_broker_message) - lock_exists_after_timeout = await redis.exists(lock_key) - assert lock_exists_after_timeout == 0, "Lock should be released after timeout" + message = await asyncio.wait_for(get_message(broker), timeout=2) + assert isinstance(message, AckableMessage) + assert message.data == valid_broker_message.message + await message.ack() # type: ignore await broker.shutdown() From d8531fc8c2870775af605ed4460cfa0aec565681 Mon Sep 17 00:00:00 2001 From: vvanglro Date: Thu, 30 Jul 2026 16:15:26 +0800 Subject: [PATCH 2/3] deprecate Redis stream additional streams --- README.md | 4 +++- taskiq_redis/redis_broker.py | 17 +++++++++++++---- tests/test_broker.py | 12 ++++++++++++ 3 files changed, 28 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index c7614bf..f6fd082 100644 --- a/README.md +++ b/README.md @@ -107,7 +107,9 @@ Recovery of messages orphaned by a crashed worker is driven by `XCLAIM`: * Reclaim sweeps run at most every `reclaim_interval` (30 seconds by default) to avoid hammering redis on every listen loop. Set it to `0` to scan on every iteration. * A worker that shares its `consumer_name` with a previous, dead worker can still reclaim that worker's pending messages; only messages actually held by the current listener instance are protected from re-delivery. -To avoid one worker hoarding the backlog, the listener does not fetch new messages while it already has `xread_count` delivered but unacknowledged messages. Set `xread_count=None` to disable this limit. +For a single-stream broker, `xread_count` also prevents one worker from hoarding the backlog: the listener does not fetch new messages while it already has that many delivered but unacknowledged messages. Set `xread_count=None` to disable this limit. + +`additional_streams` is deprecated and will be removed in a future major release. Configure one `RedisStreamBroker` per stream and run a worker process for each broker instead. This gives each stream an independent consumer, reclaim policy, and prefetch limit. It also avoids Redis `XREADGROUP COUNT` applying separately to every stream in a multi-stream read. When a listener is closed with messages already fetched from redis but not yet yielded to taskiq, those buffered entries are claimed to an internal `abandoned` consumer and stamped as very idle. The next reclaim sweep can recover them immediately instead of waiting for `idle_timeout`. diff --git a/taskiq_redis/redis_broker.py b/taskiq_redis/redis_broker.py index 68237d9..9a05645 100644 --- a/taskiq_redis/redis_broker.py +++ b/taskiq_redis/redis_broker.py @@ -1,6 +1,7 @@ import asyncio import time import uuid +import warnings from collections.abc import AsyncGenerator, Awaitable, Callable, Iterable from contextlib import suppress from logging import getLogger @@ -197,11 +198,11 @@ def __init__( :param approximate: decides wether to trim the stream immediately (False) or later on (True) :param xread_count: number of messages to fetch from the stream at once. - The broker won't make another XREADGROUP call while this listener - already has this many delivered but unacknowledged messages. Set to - None to disable this limit. + For a single-stream broker, it also caps this listener's delivered + but unacknowledged messages. Set to None to disable this limit. :param additional_streams: additional streams to read from. - Each key is a stream name, value is a consumer id. + Each key is a stream name, value is a consumer id. Deprecated: + use one broker and worker process per stream instead. :param unacknowledged_batch_size: number of unacknowledged messages to fetch. :param unacknowledged_lock_timeout: deprecated and ignored. Redis' XCLAIM min-idle-time check is used instead of a broker-side lock. @@ -227,6 +228,14 @@ def __init__( self.maxlen = maxlen self.approximate = approximate self.additional_streams = additional_streams or {} + if self.additional_streams: + warnings.warn( + "additional_streams is deprecated and will be removed in a " + "future major release. Use one RedisStreamBroker and worker " + "process per stream instead.", + DeprecationWarning, + stacklevel=2, + ) self.idle_timeout = idle_timeout self.unacknowledged_batch_size = unacknowledged_batch_size self.unacknowledged_lock_timeout = unacknowledged_lock_timeout diff --git a/tests/test_broker.py b/tests/test_broker.py index 0722ef2..52498f6 100644 --- a/tests/test_broker.py +++ b/tests/test_broker.py @@ -636,6 +636,18 @@ async def test_stream_broker_xread_count_limits_unacked_messages( await broker.shutdown() +def test_stream_broker_additional_streams_is_deprecated() -> None: + """Additional streams warn users to migrate to one broker per stream.""" + with pytest.warns( + DeprecationWarning, + match="additional_streams is deprecated", + ): + RedisStreamBroker( + "redis://localhost:7000", + additional_streams={"secondary": ">"}, + ) + + @pytest.mark.anyio async def test_stream_broker_abandons_buffered_messages_on_close( redis_url: str, From 44606601e2f090a0f8595cef3b95350b21be3da2 Mon Sep 17 00:00:00 2001 From: vvanglro Date: Thu, 30 Jul 2026 17:32:45 +0800 Subject: [PATCH 3/3] fix Redis stream reclaim edge cases --- README.md | 2 - taskiq_redis/redis_broker.py | 136 +++++++++++++++-------------- tests/test_broker.py | 165 ++++++++++++++++++++++++++++------- 3 files changed, 206 insertions(+), 97 deletions(-) diff --git a/README.md b/README.md index f6fd082..9cb07a5 100644 --- a/README.md +++ b/README.md @@ -113,8 +113,6 @@ For a single-stream broker, `xread_count` also prevents one worker from hoarding When a listener is closed with messages already fetched from redis but not yet yielded to taskiq, those buffered entries are claimed to an internal `abandoned` consumer and stamped as very idle. The next reclaim sweep can recover them immediately instead of waiting for `idle_timeout`. -If the consumer group is removed out of band, the broker recreates it on the next `XREADGROUP` so an in-use queue heals instead of spinning on errors. - The legacy `unacknowledged_lock_timeout` parameter is deprecated and ignored; reclaim correctness now relies on the `XCLAIM` min-idle-time check instead of a broker-side redis lock. ## RedisAsyncResultBackend configuration diff --git a/taskiq_redis/redis_broker.py b/taskiq_redis/redis_broker.py index 9a05645..b6658f0 100644 --- a/taskiq_redis/redis_broker.py +++ b/taskiq_redis/redis_broker.py @@ -13,7 +13,13 @@ cast, ) -from redis.asyncio import BlockingConnectionPool, Connection, Redis, ResponseError +from redis.asyncio import ( + BlockingConnectionPool, + Connection, + Redis, + RedisError, + ResponseError, +) from taskiq import AckableMessage from taskiq.abc.broker import AsyncBroker from taskiq.abc.result_backend import AsyncResultBackend @@ -327,16 +333,15 @@ def _ack_generator( async def _ack() -> None: nonlocal acked async with Redis(connection_pool=self.connection_pool) as redis_conn: - try: - await redis_conn.xack( - queue_name, - self.consumer_group_name, - id, - ) - finally: - if not acked and on_ack is not None: - acked = True - on_ack() + await redis_conn.xack( + queue_name, + self.consumer_group_name, + id, + ) + if not acked: + acked = True + if on_ack is not None: + on_ack() return _ack @@ -380,7 +385,8 @@ async def _claim_timed_out_messages( stream: str, count: int, protected_message_ids: set[tuple[str, str]], - ) -> list[tuple[str, dict[bytes, bytes]]]: + pending_start: str, + ) -> tuple[list[tuple[str, dict[bytes, bytes]]], str]: """Claim pending messages that exceeded their reclaim deadline. protected_message_ids contains entries delivered by this listener @@ -388,37 +394,30 @@ async def _claim_timed_out_messages( in-flight work, while still allowing a restarted worker with the same Redis consumer name to recover messages left by its predecessor. """ - try: - pending = await redis_conn.xpending_range( - stream, - self.consumer_group_name, - min="-", - max="+", - count=self.unacknowledged_batch_size, - idle=0, - ) - except ResponseError as exc: - if "NOGROUP" not in str(exc): - raise - logger.info("Consumer group missing for %s, recreating", stream) - await self._declare_consumer_group() - return [] + pending = await redis_conn.xpending_range( + stream, + self.consumer_group_name, + min=pending_start, + max="+", + count=self.unacknowledged_batch_size, + idle=0, + ) claimed: list[tuple[str, dict[bytes, bytes]]] = [] + last_checked_id: str | None = None for pending_message in pending: if len(claimed) >= count: break - message_id = pending_message["message_id"] - if isinstance(message_id, bytes): - message_id = message_id.decode() + message_id = self._to_str(pending_message["message_id"]) + last_checked_id = message_id if self._message_key(stream, message_id) in protected_message_ids: continue message = await self._get_pending_message( redis_conn, stream, - cast(str, message_id), + message_id, ) if message is None: await redis_conn.xclaim( @@ -448,13 +447,20 @@ async def _claim_timed_out_messages( message_ids=[message_id], ) claimed.extend(cast("list[tuple[str, dict[bytes, bytes]]]", result)) - return claimed + next_pending_start = "-" + if ( + len(pending) == self.unacknowledged_batch_size + and last_checked_id is not None + ): + next_pending_start = f"({last_checked_id}" + return claimed, next_pending_start async def _claim_available_timed_out_messages( self, redis_conn: Redis, count: int, protected_message_ids: set[tuple[str, str]], + pending_starts: dict[str, str], ) -> list[tuple[str, dict[bytes, bytes], str]]: """Claim overdue messages across all configured streams.""" claimed: list[tuple[str, dict[bytes, bytes], str]] = [] @@ -462,12 +468,16 @@ async def _claim_available_timed_out_messages( remaining_count = count - len(claimed) if remaining_count <= 0: break - for msg_id, msg in await self._claim_timed_out_messages( - redis_conn, - stream, - remaining_count, - protected_message_ids, - ): + claimed_messages, pending_starts[stream] = ( + await self._claim_timed_out_messages( + redis_conn, + stream, + remaining_count, + protected_message_ids, + pending_starts.get(stream, "-"), + ) + ) + for msg_id, msg in claimed_messages: claimed.append((msg_id, msg, stream)) return claimed @@ -476,25 +486,18 @@ async def _read_new_messages( redis_conn: Redis, count: int | None, ) -> Any: - """Read newly added stream messages, recreating a missing group once.""" - try: - return await redis_conn.xreadgroup( - self.consumer_group_name, - self.consumer_name, - { - self.queue_name: ">", - **self.additional_streams, # type: ignore[dict-item] - }, - block=self.block, - noack=False, - count=count, - ) - except ResponseError as exc: - if "NOGROUP" not in str(exc): - raise - logger.info("Consumer group missing, recreating") - await self._declare_consumer_group() - return [] + """Read newly added stream messages.""" + return await redis_conn.xreadgroup( + self.consumer_group_name, + self.consumer_name, + { + self.queue_name: ">", + **self.additional_streams, # type: ignore[dict-item] + }, + block=self.block, + noack=False, + count=count, + ) async def _abandon_buffered_messages( self, @@ -517,13 +520,12 @@ async def _abandon_buffered_messages( idle=ABANDONED_IDLE_MS, justid=True, ) - except ResponseError as exc: - if "NOGROUP" not in str(exc): - logger.warning( - "Failed to abandon messages in stream %s", - stream, - exc_info=True, - ) + except RedisError: + logger.warning( + "Failed to abandon messages in stream %s", + stream, + exc_info=True, + ) def _build_ackable_message( self, @@ -578,12 +580,14 @@ async def _build_reclaimed_ackable_messages( redis_conn: Redis, count: int, delivered: set[tuple[str, str]], + pending_starts: dict[str, str], ) -> list[tuple[str, dict[bytes, bytes], str]]: """Claim overdue messages for taskiq's receiver.""" return await self._claim_available_timed_out_messages( redis_conn, count or self.unacknowledged_batch_size, delivered, + pending_starts, ) async def _build_due_reclaimed_messages( @@ -592,6 +596,7 @@ async def _build_due_reclaimed_messages( count: int, delivered: set[tuple[str, str]], last_reclaim: float, + pending_starts: dict[str, str], ) -> tuple[float, list[tuple[str, dict[bytes, bytes], str]]]: """Return overdue pending messages only when the reclaim interval elapsed.""" if not self._should_reclaim(last_reclaim): @@ -600,6 +605,7 @@ async def _build_due_reclaimed_messages( redis_conn, count, delivered, + pending_starts, ) async def _build_new_ackable_messages( @@ -637,6 +643,7 @@ async def listen(self) -> AsyncGenerator[AckableMessage, None]: buffered: list[tuple[str, dict[bytes, bytes], str]] = [] slot_freed = asyncio.Event() last_reclaim = 0.0 + pending_starts: dict[str, str] = {} def on_ack(message_key: tuple[str, str]) -> None: nonlocal unacked @@ -665,6 +672,7 @@ def _on_ack() -> None: count or self.unacknowledged_batch_size, delivered, last_reclaim, + pending_starts, ) unacked += len(buffered) async for message in self._yield_buffered_messages( diff --git a/tests/test_broker.py b/tests/test_broker.py index 52498f6..2465f23 100644 --- a/tests/test_broker.py +++ b/tests/test_broker.py @@ -1,8 +1,10 @@ import asyncio import uuid +from contextlib import suppress import pytest from redis.asyncio import Redis +from redis.exceptions import ConnectionError as RedisConnectionError from taskiq import AckableMessage, AsyncBroker, BrokerMessage from taskiq.message import TaskiqMessage @@ -636,6 +638,138 @@ async def test_stream_broker_xread_count_limits_unacked_messages( await broker.shutdown() +@pytest.mark.anyio +async def test_stream_broker_ack_failure_keeps_prefetch_slot( + redis_url: str, + valid_broker_message: BrokerMessage, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A failed XACK must not let the listener reserve more PEL entries.""" + queue_name = uuid.uuid4().hex + consumer_group_name = uuid.uuid4().hex + broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + xread_block=50, + xread_count=1, + reclaim_interval=0, + ) + + await broker.startup() + await broker.kick(valid_broker_message) + await broker.kick(valid_broker_message) + + iterator = broker.listen() + first_message = await iterator.__anext__() + assert isinstance(first_message, AckableMessage) + + async def fail_xack(self: Redis, *args: object, **kwargs: object) -> int: + raise RedisConnectionError("simulated Redis disconnect") + + monkeypatch.setattr(Redis, "xack", fail_xack) + with pytest.raises(RedisConnectionError, match="simulated Redis disconnect"): + await first_message.ack() # type: ignore + + second_task = asyncio.create_task(iterator.__anext__()) + await asyncio.sleep(0.2) + assert not second_task.done() + + second_task.cancel() + with suppress(asyncio.CancelledError): + await second_task + await iterator.aclose() + await broker.shutdown() + + +@pytest.mark.anyio +async def test_stream_broker_reclaim_scan_advances_past_protected_messages( + redis_url: str, + valid_broker_message: BrokerMessage, +) -> None: + """Protected entries at the PEL head do not hide a later orphan forever.""" + queue_name = uuid.uuid4().hex + consumer_group_name = uuid.uuid4().hex + active_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name="active", + xread_block=50, + xread_count=3, + unacknowledged_batch_size=2, + idle_timeout=100, + reclaim_interval=0, + reclaim_timeout_grace=0, + ) + orphan_broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=queue_name, + consumer_group_name=consumer_group_name, + consumer_name="orphan", + xread_block=50, + ) + + await active_broker.startup() + await orphan_broker.startup() + for _ in range(3): + await active_broker.kick(valid_broker_message) + + active_iterator = active_broker.listen() + active_messages = [await active_iterator.__anext__() for _ in range(3)] + assert all(isinstance(message, AckableMessage) for message in active_messages) + + await active_broker.kick(valid_broker_message) + orphan_message = await get_message(orphan_broker) + assert isinstance(orphan_message, AckableMessage) + await asyncio.sleep(0.15) + + await active_messages[0].ack() # type: ignore + reclaimed_message = await asyncio.wait_for(active_iterator.__anext__(), timeout=2) + assert isinstance(reclaimed_message, AckableMessage) + + await active_messages[1].ack() # type: ignore + await active_messages[2].ack() # type: ignore + await reclaimed_message.ack() # type: ignore + await active_iterator.aclose() + await active_broker.shutdown() + await orphan_broker.shutdown() + + +@pytest.mark.anyio +async def test_stream_broker_ignores_redis_error_while_abandoning_buffer( + redis_url: str, + valid_broker_message: BrokerMessage, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Closing a listener remains best-effort when Redis is unavailable.""" + broker = RedisStreamBroker( + url=redis_url, + approximate=False, + queue_name=uuid.uuid4().hex, + consumer_group_name=uuid.uuid4().hex, + xread_block=50, + xread_count=2, + ) + await broker.startup() + await broker.kick(valid_broker_message) + await broker.kick(valid_broker_message) + + iterator = broker.listen() + message = await iterator.__anext__() + assert isinstance(message, AckableMessage) + + async def fail_xclaim(self: Redis, *args: object, **kwargs: object) -> list[object]: + raise RedisConnectionError("simulated Redis disconnect") + + monkeypatch.setattr(Redis, "xclaim", fail_xclaim) + await iterator.aclose() + await broker.shutdown() + + def test_stream_broker_additional_streams_is_deprecated() -> None: """Additional streams warn users to migrate to one broker per stream.""" with pytest.warns( @@ -710,34 +844,3 @@ async def test_stream_broker_abandons_buffered_messages_on_close( await reclaimed_message.ack() # type: ignore await first_broker.shutdown() await second_broker.shutdown() - - -@pytest.mark.anyio -async def test_stream_broker_recreates_missing_consumer_group( - redis_url: str, - valid_broker_message: BrokerMessage, -) -> None: - queue_name = uuid.uuid4().hex - consumer_group_name = uuid.uuid4().hex - - broker = RedisStreamBroker( - url=redis_url, - approximate=False, - queue_name=queue_name, - consumer_group_name=consumer_group_name, - consumer_id="0", - xread_block=50, - reclaim_interval=0, - ) - - await broker.startup() - async with Redis(connection_pool=broker.connection_pool) as redis: - await redis.xgroup_destroy(queue_name, consumer_group_name) - await broker.kick(valid_broker_message) - - message = await asyncio.wait_for(get_message(broker), timeout=2) - assert isinstance(message, AckableMessage) - assert message.data == valid_broker_message.message - await message.ack() # type: ignore - - await broker.shutdown()