taskqueue-toolkit 0.1.0__tar.gz → 0.2.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (28) hide show
  1. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/PKG-INFO +9 -1
  2. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/README.md +8 -0
  3. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/pyproject.toml +1 -1
  4. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/pyproject.toml.orig +1 -1
  5. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/pubsub.py +8 -0
  6. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/rabbitmq.py +38 -2
  7. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/redis_streams.py +22 -0
  8. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/sns.py +16 -3
  9. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/sqs.py +13 -3
  10. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/task_queue.py +20 -0
  11. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/testing/pubsub.py +26 -9
  12. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/testing/rabbitmq.py +34 -10
  13. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/testing/redis_streams.py +28 -0
  14. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/testing/sns.py +32 -11
  15. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/testing/sqs.py +33 -11
  16. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/LICENSE +0 -0
  17. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/__init__.py +0 -0
  18. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/outbox/__init__.py +0 -0
  19. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/outbox/orm.py +0 -0
  20. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/outbox/relay.py +0 -0
  21. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/outbox/repository.py +0 -0
  22. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/py.typed +0 -0
  23. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/__init__.py +0 -0
  24. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/aws_session.py +0 -0
  25. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/dsn.py +0 -0
  26. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/factory.py +0 -0
  27. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/queue/registry.py +0 -0
  28. {taskqueue_toolkit-0.1.0 → taskqueue_toolkit-0.2.0}/src/taskqueue_toolkit/testing/__init__.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: taskqueue-toolkit
3
- Version: 0.1.0
3
+ Version: 0.2.0
4
4
  Summary: Broker-agnostic async task queue (RabbitMQ, SQS, SNS, Redis Streams, Pub/Sub) with a Postgres outbox pattern
5
5
  Author: Walid Boughdiri
6
6
  Author-email: Walid Boughdiri <walid.boughdiri@gmail.com>
@@ -152,6 +152,12 @@ Each adapter constructor takes an injectable client/connection factory for exact
152
152
 
153
153
  Create one broker instance per test — each owns its own isolated in-memory state, so there's nothing to reset between tests and no risk of one test's queue leaking into another's.
154
154
 
155
+ ## Security
156
+
157
+ `Decoder[T]` runs on bytes received from a broker — in most deployments, a source outside this process's control. Using an unsafe deserializer (`pickle.loads`, `yaml.load` without `SafeLoader`, `eval`, ...) as your `decode` callback means whoever can publish to your queue/topic can run arbitrary code in your consumer. Use a safe format (JSON, msgpack, protobuf, ...) unless every publisher is fully trusted.
158
+
159
+ This package's own code is checked on every push/PR and before every release via [bandit](https://bandit.readthedocs.io/) (unsafe code patterns) and [pip-audit](https://github.com/pypa/pip-audit) (known CVEs in resolved dependencies), alongside ruff, mypy, and the test suite. Found a vulnerability? Please open an issue.
160
+
155
161
  ## Development
156
162
 
157
163
  ```bash
@@ -160,4 +166,6 @@ uv run pytest
160
166
  uv run ruff check .
161
167
  uv run ruff format .
162
168
  uv run mypy .
169
+ uvx bandit -r src/
170
+ uvx pip-audit --path .venv
163
171
  ```
@@ -117,6 +117,12 @@ Each adapter constructor takes an injectable client/connection factory for exact
117
117
 
118
118
  Create one broker instance per test — each owns its own isolated in-memory state, so there's nothing to reset between tests and no risk of one test's queue leaking into another's.
119
119
 
120
+ ## Security
121
+
122
+ `Decoder[T]` runs on bytes received from a broker — in most deployments, a source outside this process's control. Using an unsafe deserializer (`pickle.loads`, `yaml.load` without `SafeLoader`, `eval`, ...) as your `decode` callback means whoever can publish to your queue/topic can run arbitrary code in your consumer. Use a safe format (JSON, msgpack, protobuf, ...) unless every publisher is fully trusted.
123
+
124
+ This package's own code is checked on every push/PR and before every release via [bandit](https://bandit.readthedocs.io/) (unsafe code patterns) and [pip-audit](https://github.com/pypa/pip-audit) (known CVEs in resolved dependencies), alongside ruff, mypy, and the test suite. Found a vulnerability? Please open an issue.
125
+
120
126
  ## Development
121
127
 
122
128
  ```bash
@@ -125,4 +131,6 @@ uv run pytest
125
131
  uv run ruff check .
126
132
  uv run ruff format .
127
133
  uv run mypy .
134
+ uvx bandit -r src/
135
+ uvx pip-audit --path .venv
128
136
  ```
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "taskqueue-toolkit"
3
- version = "0.1.0"
3
+ version = "0.2.0"
4
4
  description = "Broker-agnostic async task queue (RabbitMQ, SQS, SNS, Redis Streams, Pub/Sub) with a Postgres outbox pattern"
5
5
  readme = "README.md"
6
6
  license = "Apache-2.0"
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "taskqueue-toolkit"
3
- version = "0.1.0"
3
+ version = "0.2.0"
4
4
  description = "Broker-agnostic async task queue (RabbitMQ, SQS, SNS, Redis Streams, Pub/Sub) with a Postgres outbox pattern"
5
5
  readme = "README.md"
6
6
  license = "Apache-2.0"
@@ -26,6 +26,7 @@ class PubsubQueuedTask[T]:
26
26
  _subscriber: pubsub_v1.SubscriberClient
27
27
  _subscription_path: str
28
28
  _ack_id: str
29
+ delivery_count: int
29
30
  task: T
30
31
 
31
32
  async def ack(self) -> None:
@@ -154,9 +155,16 @@ class PubsubTaskQueue[T]:
154
155
  timeout=_PULL_TIMEOUT_SECONDS,
155
156
  )
156
157
  for received in response.received_messages:
158
+ # delivery_attempt is only populated when the subscription
159
+ # has a dead-letter policy configured — absent (None)
160
+ # otherwise, in which case this reports 1 rather than
161
+ # silently misreporting every message as never-redelivered
162
+ # forever.
163
+ delivery_count = received.delivery_attempt or 1
157
164
  yield PubsubQueuedTask(
158
165
  _subscriber=self._subscriber,
159
166
  _subscription_path=subscription_path,
160
167
  _ack_id=received.ack_id,
168
+ delivery_count=delivery_count,
161
169
  task=self.decode(received.message.data),
162
170
  )
@@ -6,6 +6,7 @@ from contextlib import suppress
6
6
  from dataclasses import dataclass, field
7
7
 
8
8
  import aio_pika
9
+ from pamqp.common import FieldTable
9
10
 
10
11
  from taskqueue_toolkit.queue.dsn import RabbitMqDsn
11
12
  from taskqueue_toolkit.queue.task_queue import Decoder, Encoder
@@ -15,11 +16,33 @@ logger = logging.getLogger(__name__)
15
16
  Connect = Callable[[str], Awaitable["aio_pika.abc.AbstractRobustConnection"]]
16
17
 
17
18
 
19
+ def _delivery_count(message: aio_pika.abc.AbstractIncomingMessage) -> int:
20
+ """1 on first delivery, incremented on every redelivery (a
21
+ nack(requeue=True), a crashed/disconnected consumer, ...). Sourced from
22
+ x-delivery-count, which RabbitMQ maintains natively for quorum queues —
23
+ see the queue declaration in RabbitMqTaskQueue for why the queue type
24
+ matters here (classic queues don't track this at all; only a
25
+ dead-letter hop shows up there, via x-death, which a plain requeue never
26
+ triggers)."""
27
+ headers = message.headers or {}
28
+ count = headers.get("x-delivery-count")
29
+ if count is None:
30
+ return 1
31
+ # RabbitMQ always sends this header as a long-int field value — the
32
+ # broad FieldValue union is what the type carries at rest, not what
33
+ # actually arrives here.
34
+ return 1 + int(count) # type: ignore[arg-type]
35
+
36
+
18
37
  @dataclass(slots=True)
19
38
  class RabbitMqQueuedTask[T]:
20
39
  _message: aio_pika.abc.AbstractIncomingMessage
21
40
  task: T
22
41
 
42
+ @property
43
+ def delivery_count(self) -> int:
44
+ return _delivery_count(self._message)
45
+
23
46
  async def ack(self) -> None:
24
47
  await self._message.ack()
25
48
 
@@ -36,11 +59,22 @@ class RabbitMqTaskQueue[T]:
36
59
  # connection strategy the caller already has; defaults to the real SDK.
37
60
  connect: Connect = field(default=aio_pika.connect_robust)
38
61
 
62
+ def _queue_arguments(self) -> FieldTable:
63
+ # Quorum queues (not RabbitMQ's classic queue type) are what expose
64
+ # x-delivery-count on redelivery — a classic queue tracks no such
65
+ # counter at all, so delivery_count would always read 1. This is
66
+ # also RabbitMQ's own recommended queue type for task-queue-style
67
+ # workloads (durable, replicated, no classic-queue mirroring), so
68
+ # this isn't a special case adopted only for the counter.
69
+ return {"x-queue-type": "quorum"}
70
+
39
71
  async def publish(self, task: T) -> None:
40
72
  connection = await self.connect(self.dsn.url)
41
73
  try:
42
74
  channel = await connection.channel()
43
- queue = await channel.declare_queue(self.dsn.queue_name, durable=True)
75
+ queue = await channel.declare_queue(
76
+ self.dsn.queue_name, durable=True, arguments=self._queue_arguments()
77
+ )
44
78
  await channel.default_exchange.publish(
45
79
  aio_pika.Message(
46
80
  body=self.encode(task),
@@ -57,7 +91,9 @@ class RabbitMqTaskQueue[T]:
57
91
  try:
58
92
  channel = await connection.channel()
59
93
  await channel.set_qos(prefetch_count=1)
60
- queue = await channel.declare_queue(self.dsn.queue_name, durable=True)
94
+ queue = await channel.declare_queue(
95
+ self.dsn.queue_name, durable=True, arguments=self._queue_arguments()
96
+ )
61
97
 
62
98
  async for message in queue.iterator():
63
99
  yield RabbitMqQueuedTask(_message=message, task=self.decode(message.body))
@@ -29,6 +29,7 @@ class RedisStreamsQueuedTask[T]:
29
29
  _stream: str
30
30
  _group: str
31
31
  _message_id: str
32
+ delivery_count: int
32
33
  task: T
33
34
 
34
35
  async def ack(self) -> None:
@@ -61,6 +62,20 @@ class RedisStreamsTaskQueue[T]:
61
62
  def _client(self) -> redis.Redis:
62
63
  return self.client_factory(self.dsn)
63
64
 
65
+ async def _delivery_count(
66
+ self, client: redis.Redis, stream: str, group: str, message_id: bytes
67
+ ) -> int:
68
+ entries = await client.xpending_range(
69
+ name=stream, groupname=group, min=message_id, max=message_id, count=1
70
+ )
71
+ if not entries:
72
+ # Already ack'd/claimed away between xreadgroup and this lookup
73
+ # — vanishingly unlikely for our single-consumer-per-group
74
+ # usage, but "never delivered before" (1) is the only sane
75
+ # fallback rather than raising.
76
+ return 1
77
+ return int(entries[0]["times_delivered"])
78
+
64
79
  async def _ensure_group(self, client: redis.Redis) -> None:
65
80
  try:
66
81
  await client.xgroup_create(self.dsn.stream_name, self.dsn.group, id="0", mkstream=True)
@@ -106,11 +121,18 @@ class RedisStreamsTaskQueue[T]:
106
121
  )
107
122
  for _stream_name, messages in response:
108
123
  for message_id, fields in messages:
124
+ # xreadgroup itself never reports a delivery
125
+ # count — it's only visible via XPENDING, so
126
+ # this is a second round trip per message.
127
+ delivery_count = await self._delivery_count(
128
+ client, stream, group, message_id
129
+ )
109
130
  yield RedisStreamsQueuedTask(
110
131
  _redis=client,
111
132
  _stream=stream,
112
133
  _group=group,
113
134
  _message_id=message_id.decode(),
135
+ delivery_count=delivery_count,
114
136
  task=self.decode(fields[_FIELD.encode()]),
115
137
  )
116
138
  finally:
@@ -6,10 +6,13 @@ import logging
6
6
  from collections.abc import AsyncIterator, Callable
7
7
  from contextlib import AbstractAsyncContextManager, asynccontextmanager
8
8
  from dataclasses import dataclass, field
9
+ from typing import TYPE_CHECKING
9
10
 
10
11
  import aioboto3
11
- from types_aiobotocore_sns.client import SNSClient
12
- from types_aiobotocore_sqs.client import SQSClient
12
+
13
+ if TYPE_CHECKING:
14
+ from types_aiobotocore_sns.client import SNSClient
15
+ from types_aiobotocore_sqs.client import SQSClient
13
16
 
14
17
  from taskqueue_toolkit.queue.aws_session import aws_session_kwargs
15
18
  from taskqueue_toolkit.queue.dsn import SnsDsn
@@ -20,7 +23,7 @@ logger = logging.getLogger(__name__)
20
23
  _LONG_POLL_WAIT_SECONDS = 10
21
24
  _SUBSCRIBER_QUEUE_SUFFIX = "-subscriber"
22
25
 
23
- ClientFactory = Callable[[SnsDsn], AbstractAsyncContextManager[tuple[SNSClient, SQSClient]]]
26
+ ClientFactory = Callable[[SnsDsn], AbstractAsyncContextManager["tuple[SNSClient, SQSClient]"]]
24
27
 
25
28
 
26
29
  @asynccontextmanager
@@ -62,6 +65,7 @@ class SnsQueuedTask[T]:
62
65
  _client: SQSClient
63
66
  _queue_url: str
64
67
  _receipt_handle: str
68
+ delivery_count: int
65
69
  task: T
66
70
 
67
71
  async def ack(self) -> None:
@@ -151,12 +155,21 @@ class SnsTaskQueue[T]:
151
155
  QueueUrl=queue_url,
152
156
  MaxNumberOfMessages=1,
153
157
  WaitTimeSeconds=_LONG_POLL_WAIT_SECONDS,
158
+ # Not returned unless asked for explicitly — this is the
159
+ # source for delivery_count below. SNS fan-out delivers
160
+ # into a plain SQS subscriber queue, so the same
161
+ # attribute SqsTaskQueue reads applies here too.
162
+ AttributeNames=["ApproximateReceiveCount"],
154
163
  )
155
164
  for message in response.get("Messages", []):
165
+ receive_count = int(
166
+ message.get("Attributes", {}).get("ApproximateReceiveCount", "1")
167
+ )
156
168
  envelope = _unwrap_sns_envelope(message["Body"])
157
169
  yield SnsQueuedTask(
158
170
  _client=sqs,
159
171
  _queue_url=queue_url,
160
172
  _receipt_handle=message["ReceiptHandle"],
173
+ delivery_count=receive_count,
161
174
  task=self.decode(base64.b64decode(envelope)),
162
175
  )
@@ -5,10 +5,12 @@ import logging
5
5
  from collections.abc import AsyncIterator, Callable
6
6
  from contextlib import AbstractAsyncContextManager
7
7
  from dataclasses import dataclass, field
8
- from typing import cast
8
+ from typing import TYPE_CHECKING, cast
9
9
 
10
10
  import aioboto3
11
- from types_aiobotocore_sqs.client import SQSClient
11
+
12
+ if TYPE_CHECKING:
13
+ from types_aiobotocore_sqs.client import SQSClient
12
14
 
13
15
  from taskqueue_toolkit.queue.aws_session import aws_session_kwargs
14
16
  from taskqueue_toolkit.queue.dsn import SqsDsn
@@ -18,7 +20,7 @@ logger = logging.getLogger(__name__)
18
20
 
19
21
  _LONG_POLL_WAIT_SECONDS = 10
20
22
 
21
- ClientFactory = Callable[[SqsDsn], AbstractAsyncContextManager[SQSClient]]
23
+ ClientFactory = Callable[[SqsDsn], AbstractAsyncContextManager["SQSClient"]]
22
24
 
23
25
 
24
26
  def _default_client_factory(dsn: SqsDsn) -> AbstractAsyncContextManager[SQSClient]:
@@ -33,6 +35,7 @@ class SqsQueuedTask[T]:
33
35
  _client: SQSClient
34
36
  _queue_url: str
35
37
  _receipt_handle: str
38
+ delivery_count: int
36
39
  task: T
37
40
 
38
41
  async def ack(self) -> None:
@@ -93,11 +96,18 @@ class SqsTaskQueue[T]:
93
96
  QueueUrl=queue_url,
94
97
  MaxNumberOfMessages=1,
95
98
  WaitTimeSeconds=_LONG_POLL_WAIT_SECONDS,
99
+ # Not returned unless asked for explicitly — this is the
100
+ # source for delivery_count below.
101
+ AttributeNames=["ApproximateReceiveCount"],
96
102
  )
97
103
  for message in response.get("Messages", []):
104
+ receive_count = int(
105
+ message.get("Attributes", {}).get("ApproximateReceiveCount", "1")
106
+ )
98
107
  yield SqsQueuedTask(
99
108
  _client=client,
100
109
  _queue_url=queue_url,
101
110
  _receipt_handle=message["ReceiptHandle"],
111
+ delivery_count=receive_count,
102
112
  task=self.decode(base64.b64decode(message["Body"])),
103
113
  )
@@ -7,6 +7,12 @@ T = TypeVar("T")
7
7
  T_co = TypeVar("T_co", covariant=True)
8
8
 
9
9
  Encoder = Callable[[T], bytes]
10
+ # SECURITY: Decoder runs on bytes received from a broker — a source outside
11
+ # this process's control in most deployments. Using an unsafe deserializer
12
+ # here (pickle.loads, yaml.load without SafeLoader, eval, ...) turns
13
+ # whoever can publish to your queue/topic into someone who can run
14
+ # arbitrary code in your consumer. Prefer a safe format (json, msgpack,
15
+ # protobuf) unless every publisher is fully trusted.
10
16
  Decoder = Callable[[bytes], T]
11
17
 
12
18
 
@@ -22,6 +28,20 @@ class QueuedTask(Protocol[T_co]):
22
28
  @property
23
29
  def task(self) -> T_co: ...
24
30
 
31
+ @property
32
+ def delivery_count(self) -> int:
33
+ """How many times this message has been delivered, including the
34
+ current delivery — 1 the first time a consumer sees it, 2 after one
35
+ requeue, and so on. Sourced from each broker's own native counter
36
+ (RabbitMQ's x-death header behind a dead-letter-to-self queue, SQS's
37
+ ApproximateReceiveCount, Pub/Sub's delivery_attempt, Redis Streams'
38
+ pending-entry delivery count) rather than tracked by this package —
39
+ a consumer that wants to give up after N attempts reads this and
40
+ calls nack(requeue=False) itself; this package has no opinion on
41
+ what N should be, that's a decision about the task's own business
42
+ cost of retrying, not a queueing concern."""
43
+ ...
44
+
25
45
  async def ack(self) -> None:
26
46
  """Mark the task as successfully processed."""
27
47
  ...
@@ -40,6 +40,7 @@ class _Message(NamedTuple):
40
40
  class _ReceivedMessage(NamedTuple):
41
41
  ack_id: str
42
42
  message: _Message
43
+ delivery_attempt: int
43
44
 
44
45
 
45
46
  class _PullResponse(NamedTuple):
@@ -54,11 +55,20 @@ class _PublishFuture:
54
55
  return self._message_id
55
56
 
56
57
 
58
+ @dataclass(slots=True)
59
+ class _Envelope:
60
+ """A message plus its own delivery attempt count — mirrors Pub/Sub's
61
+ delivery_attempt (1 on first delivery, incremented on each redelivery)."""
62
+
63
+ data: bytes
64
+ delivery_attempt: int = 0
65
+
66
+
57
67
  @dataclass(slots=True)
58
68
  class _Subscription:
59
69
  topic_path: str
60
- pending: deque[bytes] = field(default_factory=deque)
61
- in_flight: dict[str, bytes] = field(default_factory=dict)
70
+ pending: deque[_Envelope] = field(default_factory=deque)
71
+ in_flight: dict[str, _Envelope] = field(default_factory=dict)
62
72
 
63
73
 
64
74
  @dataclass(slots=True)
@@ -79,7 +89,7 @@ class _FakePublisherClient:
79
89
  def publish(self, topic_path: str, data: bytes) -> _PublishFuture:
80
90
  for sub in self._broker._subscriptions.values():
81
91
  if sub.topic_path == topic_path:
82
- sub.pending.append(data)
92
+ sub.pending.append(_Envelope(data=data))
83
93
  return _PublishFuture(str(uuid.uuid4()))
84
94
 
85
95
 
@@ -108,10 +118,17 @@ class _FakeSubscriberClient:
108
118
  )
109
119
  received = []
110
120
  for _ in range(min(max_messages, len(sub.pending))):
111
- data = sub.pending.popleft()
121
+ envelope = sub.pending.popleft()
122
+ envelope.delivery_attempt += 1
112
123
  ack_id = str(uuid.uuid4())
113
- sub.in_flight[ack_id] = data
114
- received.append(_ReceivedMessage(ack_id=ack_id, message=_Message(data=data)))
124
+ sub.in_flight[ack_id] = envelope
125
+ received.append(
126
+ _ReceivedMessage(
127
+ ack_id=ack_id,
128
+ message=_Message(data=envelope.data),
129
+ delivery_attempt=envelope.delivery_attempt,
130
+ )
131
+ )
115
132
  return _PullResponse(received_messages=received)
116
133
 
117
134
  def acknowledge(self, *, subscription: str, ack_ids: list[str]) -> None:
@@ -128,9 +145,9 @@ class _FakeSubscriberClient:
128
145
  if sub is None:
129
146
  return
130
147
  for ack_id in ack_ids:
131
- data = sub.in_flight.pop(ack_id, None)
132
- if data is not None and ack_deadline_seconds == 0:
133
- sub.pending.appendleft(data)
148
+ envelope = sub.in_flight.pop(ack_id, None)
149
+ if envelope is not None and ack_deadline_seconds == 0:
150
+ sub.pending.appendleft(envelope)
134
151
 
135
152
 
136
153
  @dataclass(slots=True)
@@ -19,22 +19,34 @@ from dataclasses import dataclass, field
19
19
  __all__ = ["FakeRabbitMqBroker"]
20
20
 
21
21
 
22
+ @dataclass(slots=True)
23
+ class _Envelope:
24
+ """A message plus its own redelivery count — tracked per in-flight copy
25
+ of the message, not by content, so two messages with identical bodies
26
+ (or the same message redelivered many times) each carry their own
27
+ count."""
28
+
29
+ body: bytes
30
+ delivery_count: int = 0
31
+
32
+
22
33
  @dataclass(slots=True)
23
34
  class _Queue:
24
35
  """One named queue's backlog plus messages currently delivered but not
25
36
  yet ack()'d/nack()'d — mirrors a real broker's per-queue state."""
26
37
 
27
- pending: deque[bytes] = field(default_factory=deque)
28
- unacked: dict[int, bytes] = field(default_factory=dict)
38
+ pending: deque[_Envelope] = field(default_factory=deque)
39
+ unacked: dict[int, _Envelope] = field(default_factory=dict)
29
40
  _next_id: int = 0
30
41
 
31
42
  def push(self, body: bytes) -> None:
32
- self.pending.append(body)
43
+ self.pending.append(_Envelope(body=body))
33
44
 
34
45
  def requeue(self, delivery_id: int) -> None:
35
- body = self.unacked.pop(delivery_id, None)
36
- if body is not None:
37
- self.pending.appendleft(body)
46
+ envelope = self.unacked.pop(delivery_id, None)
47
+ if envelope is not None:
48
+ envelope.delivery_count += 1
49
+ self.pending.appendleft(envelope)
38
50
 
39
51
  def drop(self, delivery_id: int) -> None:
40
52
  self.unacked.pop(delivery_id, None)
@@ -47,6 +59,7 @@ class _Queue:
47
59
  @dataclass(slots=True)
48
60
  class _FakeIncomingMessage:
49
61
  body: bytes
62
+ headers: dict[str, int]
50
63
  _delivery_id: int
51
64
  _queue: _Queue
52
65
 
@@ -74,10 +87,19 @@ class _FakeAioPikaQueue:
74
87
  async def __anext__(self) -> _FakeIncomingMessage:
75
88
  while not self._queue.pending:
76
89
  await asyncio.sleep(0)
77
- body = self._queue.pending.popleft()
90
+ envelope = self._queue.pending.popleft()
78
91
  delivery_id = self._queue.next_delivery_id()
79
- self._queue.unacked[delivery_id] = body
80
- return _FakeIncomingMessage(body=body, _delivery_id=delivery_id, _queue=self._queue)
92
+ self._queue.unacked[delivery_id] = envelope
93
+ # Mirrors a real quorum queue's x-delivery-count: absent (no header)
94
+ # on first delivery, present from the first redelivery onward — see
95
+ # RabbitMqTaskQueue's _delivery_count(), which reads this as
96
+ # "1 + count" so a never-redelivered message still reports 1.
97
+ headers = (
98
+ {} if envelope.delivery_count == 0 else {"x-delivery-count": envelope.delivery_count}
99
+ )
100
+ return _FakeIncomingMessage(
101
+ body=envelope.body, headers=headers, _delivery_id=delivery_id, _queue=self._queue
102
+ )
81
103
 
82
104
 
83
105
  @dataclass(slots=True)
@@ -100,7 +122,9 @@ class _FakeChannel:
100
122
  async def set_qos(self, *, prefetch_count: int) -> None:
101
123
  return None
102
124
 
103
- async def declare_queue(self, name: str, *, durable: bool = True) -> _FakeAioPikaQueue:
125
+ async def declare_queue(
126
+ self, name: str, *, durable: bool = True, arguments: dict[str, str] | None = None
127
+ ) -> _FakeAioPikaQueue:
104
128
  queue = self._broker._queue_for(name)
105
129
  self.default_exchange = _FakeExchange(_queue=queue)
106
130
  return _FakeAioPikaQueue(name=name, _queue=queue)
@@ -32,6 +32,10 @@ class _Group:
32
32
  # (a real broker's Pending Entries List) — nack(requeue=True) leaves an
33
33
  # entry here; xreadgroup's "0" branch redelivers from here first.
34
34
  pending_entries: list[tuple[bytes, dict[bytes, bytes]]] = field(default_factory=list)
35
+ # Mirrors a real group's per-entry times_delivered (visible via
36
+ # XPENDING) — 0 before the entry's first delivery, matching a real
37
+ # broker where a never-yet-read entry isn't in the PEL at all.
38
+ times_delivered: dict[bytes, int] = field(default_factory=dict)
35
39
 
36
40
 
37
41
  @dataclass(slots=True)
@@ -89,6 +93,8 @@ class _FakeRedisClient:
89
93
  # per-group, matching this package's single-consumer-per-group
90
94
  # usage pattern.
91
95
  delivered = grp.pending_entries[: count or len(grp.pending_entries)]
96
+ for entry_id, _fields in delivered:
97
+ grp.times_delivered[entry_id] = grp.times_delivered.get(entry_id, 0) + 1
92
98
  return [(name.encode(), delivered)]
93
99
 
94
100
  # ">" branch: deliver new entries the group hasn't seen yet.
@@ -97,6 +103,8 @@ class _FakeRedisClient:
97
103
  if new_entries:
98
104
  stream.cursor[group_key] = cursor + len(new_entries)
99
105
  grp.pending_entries.extend(new_entries)
106
+ for entry_id, _fields in new_entries:
107
+ grp.times_delivered[entry_id] = grp.times_delivered.get(entry_id, 0) + 1
100
108
  return [(name.encode(), new_entries)]
101
109
 
102
110
  # Nothing new — a real broker blocks up to `block` ms; the fake
@@ -113,8 +121,28 @@ class _FakeRedisClient:
113
121
  ids = set(entry_ids)
114
122
  before = len(grp.pending_entries)
115
123
  grp.pending_entries = [e for e in grp.pending_entries if e[0] not in ids]
124
+ for entry_id in ids:
125
+ grp.times_delivered.pop(entry_id, None)
116
126
  return before - len(grp.pending_entries)
117
127
 
128
+ async def xpending_range(
129
+ self,
130
+ name: str,
131
+ groupname: str,
132
+ min: bytes,
133
+ max: bytes,
134
+ count: int,
135
+ consumername: str | None = None,
136
+ ) -> list[dict[str, object]]:
137
+ stream = self._broker._stream_for(name.encode())
138
+ grp = stream.groups.get(groupname.encode())
139
+ if grp is None:
140
+ return []
141
+ times_delivered = grp.times_delivered.get(min)
142
+ if times_delivered is None:
143
+ return []
144
+ return [{"message_id": min, "times_delivered": times_delivered}]
145
+
118
146
  async def delete(self, *names: str) -> int:
119
147
  count = 0
120
148
  for name in names:
@@ -37,10 +37,19 @@ class _Exceptions:
37
37
  QueueDoesNotExist: type[Exception] = QueueDoesNotExist
38
38
 
39
39
 
40
+ @dataclass(slots=True)
41
+ class _Envelope:
42
+ """A message plus its own receive count — mirrors SQS's
43
+ ApproximateReceiveCount on the subscriber queue behind the topic."""
44
+
45
+ body: str
46
+ receive_count: int = 0
47
+
48
+
40
49
  @dataclass(slots=True)
41
50
  class _Queue:
42
- pending: deque[dict[str, str]] = field(default_factory=deque)
43
- in_flight: dict[str, dict[str, str]] = field(default_factory=dict)
51
+ pending: deque[_Envelope] = field(default_factory=deque)
52
+ in_flight: dict[str, _Envelope] = field(default_factory=dict)
44
53
 
45
54
 
46
55
  def _topic_arn(name: str) -> str:
@@ -74,7 +83,7 @@ class _FakeSnsClient:
74
83
  )
75
84
  for queue_url in self._broker._topics.get(TopicArn, set()):
76
85
  queue = self._broker._queues.setdefault(queue_url, _Queue())
77
- queue.pending.append({"Body": envelope})
86
+ queue.pending.append(_Envelope(body=envelope))
78
87
  return {"MessageId": str(uuid.uuid4())}
79
88
 
80
89
 
@@ -105,15 +114,27 @@ class _FakeSqsClient:
105
114
  return None
106
115
 
107
116
  async def receive_message(
108
- self, *, QueueUrl: str, MaxNumberOfMessages: int = 1, WaitTimeSeconds: int = 0
109
- ) -> dict[str, list[dict[str, str]]]:
117
+ self,
118
+ *,
119
+ QueueUrl: str,
120
+ MaxNumberOfMessages: int = 1,
121
+ WaitTimeSeconds: int = 0,
122
+ AttributeNames: list[str] | None = None,
123
+ ) -> dict[str, list[dict[str, object]]]:
110
124
  queue = self._broker._queues.setdefault(QueueUrl, _Queue())
111
125
  messages = []
112
126
  for _ in range(min(MaxNumberOfMessages, len(queue.pending))):
113
- body = queue.pending.popleft()
127
+ envelope = queue.pending.popleft()
128
+ envelope.receive_count += 1
114
129
  receipt_handle = str(uuid.uuid4())
115
- queue.in_flight[receipt_handle] = body
116
- messages.append({"Body": body["Body"], "ReceiptHandle": receipt_handle})
130
+ queue.in_flight[receipt_handle] = envelope
131
+ message: dict[str, object] = {
132
+ "Body": envelope.body,
133
+ "ReceiptHandle": receipt_handle,
134
+ }
135
+ if AttributeNames and "ApproximateReceiveCount" in AttributeNames:
136
+ message["Attributes"] = {"ApproximateReceiveCount": str(envelope.receive_count)}
137
+ messages.append(message)
117
138
  return {"Messages": messages}
118
139
 
119
140
  async def delete_message(self, *, QueueUrl: str, ReceiptHandle: str) -> None:
@@ -124,9 +145,9 @@ class _FakeSqsClient:
124
145
  self, *, QueueUrl: str, ReceiptHandle: str, VisibilityTimeout: int
125
146
  ) -> None:
126
147
  queue = self._broker._queues.setdefault(QueueUrl, _Queue())
127
- body = queue.in_flight.pop(ReceiptHandle, None)
128
- if body is not None and VisibilityTimeout == 0:
129
- queue.pending.appendleft(body)
148
+ envelope = queue.in_flight.pop(ReceiptHandle, None)
149
+ if envelope is not None and VisibilityTimeout == 0:
150
+ queue.pending.appendleft(envelope)
130
151
 
131
152
 
132
153
  class _FakeSnsContext(AbstractAsyncContextManager[tuple["_FakeSnsClient", "_FakeSqsClient"]]):
@@ -33,10 +33,20 @@ class _Exceptions:
33
33
  QueueDoesNotExist: type[Exception] = QueueDoesNotExist
34
34
 
35
35
 
36
+ @dataclass(slots=True)
37
+ class _Envelope:
38
+ """A message plus its own receive count — tracked per in-flight copy of
39
+ the message, mirroring SQS's ApproximateReceiveCount (1 on first
40
+ receive, incremented on each redelivery)."""
41
+
42
+ body: str
43
+ receive_count: int = 0
44
+
45
+
36
46
  @dataclass(slots=True)
37
47
  class _Queue:
38
- pending: deque[dict[str, str]] = field(default_factory=deque)
39
- in_flight: dict[str, dict[str, str]] = field(default_factory=dict)
48
+ pending: deque[_Envelope] = field(default_factory=deque)
49
+ in_flight: dict[str, _Envelope] = field(default_factory=dict)
40
50
 
41
51
 
42
52
  @dataclass(slots=True)
@@ -58,19 +68,31 @@ class _FakeSqsClient:
58
68
 
59
69
  async def send_message(self, *, QueueUrl: str, MessageBody: str) -> dict[str, str]:
60
70
  queue = self._broker._queues.setdefault(QueueUrl, _Queue())
61
- queue.pending.append({"Body": MessageBody})
71
+ queue.pending.append(_Envelope(body=MessageBody))
62
72
  return {"MessageId": str(uuid.uuid4())}
63
73
 
64
74
  async def receive_message(
65
- self, *, QueueUrl: str, MaxNumberOfMessages: int = 1, WaitTimeSeconds: int = 0
66
- ) -> dict[str, list[dict[str, str]]]:
75
+ self,
76
+ *,
77
+ QueueUrl: str,
78
+ MaxNumberOfMessages: int = 1,
79
+ WaitTimeSeconds: int = 0,
80
+ AttributeNames: list[str] | None = None,
81
+ ) -> dict[str, list[dict[str, object]]]:
67
82
  queue = self._broker._queues.setdefault(QueueUrl, _Queue())
68
83
  messages = []
69
84
  for _ in range(min(MaxNumberOfMessages, len(queue.pending))):
70
- body = queue.pending.popleft()
85
+ envelope = queue.pending.popleft()
86
+ envelope.receive_count += 1
71
87
  receipt_handle = str(uuid.uuid4())
72
- queue.in_flight[receipt_handle] = body
73
- messages.append({"Body": body["Body"], "ReceiptHandle": receipt_handle})
88
+ queue.in_flight[receipt_handle] = envelope
89
+ message: dict[str, object] = {
90
+ "Body": envelope.body,
91
+ "ReceiptHandle": receipt_handle,
92
+ }
93
+ if AttributeNames and "ApproximateReceiveCount" in AttributeNames:
94
+ message["Attributes"] = {"ApproximateReceiveCount": str(envelope.receive_count)}
95
+ messages.append(message)
74
96
  return {"Messages": messages}
75
97
 
76
98
  async def delete_message(self, *, QueueUrl: str, ReceiptHandle: str) -> None:
@@ -81,9 +103,9 @@ class _FakeSqsClient:
81
103
  self, *, QueueUrl: str, ReceiptHandle: str, VisibilityTimeout: int
82
104
  ) -> None:
83
105
  queue = self._broker._queues.setdefault(QueueUrl, _Queue())
84
- body = queue.in_flight.pop(ReceiptHandle, None)
85
- if body is not None and VisibilityTimeout == 0:
86
- queue.pending.appendleft(body)
106
+ envelope = queue.in_flight.pop(ReceiptHandle, None)
107
+ if envelope is not None and VisibilityTimeout == 0:
108
+ queue.pending.appendleft(envelope)
87
109
 
88
110
 
89
111
  class _FakeClientContext(AbstractAsyncContextManager["_FakeSqsClient"]):