From 677815f914db7a2df53142d4a9b50dbed5a32684 Mon Sep 17 00:00:00 2001 From: Dima Anfimov Date: Tue, 15 Sep 2026 23:04:37 +0200 Subject: [PATCH 1/2] feat: add batch send --- README.md | 13 +++ src/taskiq_sqs/broker.py | 85 ++++++++++++++++++++ src/taskiq_sqs/constants.py | 4 + src/taskiq_sqs/types/message.py | 7 ++ src/taskiq_sqs/types/queue.py | 14 ++++ tests/test_broker_kick.py | 137 +++++++++++++++++++++++++++++++- 6 files changed, 259 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index ae8feca..2e85d42 100644 --- a/README.md +++ b/README.md @@ -128,6 +128,19 @@ await process_event.kicker().with_labels(expiry=time.time() + 300).kiq() # disc Expiration is checked by the worker on receipt, not by SQS itself — a message can still sit in the queue past its expiry (e.g. while workers are busy or scaled to zero), it just won't run once picked up. `expiry` must be a non-negative number; anything else raises `InvalidExpiryError` when the task is kicked. +## Message batching + +Set `is_batching_enabled` on a queue to buffer kicked messages in memory and flush them together via `SendMessageBatch` (up to `batch_size` messages, or after `batch_timeout` seconds, whichever comes first) instead of sending each one immediately: + +```python +from taskiq_sqs import SQSBroker +from taskiq_sqs.types import SQSQueue + +broker = SQSBroker( + queues=SQSQueue(name="my-queue", is_batching_enabled=True, batch_size=10, batch_timeout=1.0), +) +``` + ## Offloading large messages to S3 SQS messages are limited to 256 KiB. `S3OffloadMiddleware` transparently uploads task payloads that exceed a configurable threshold to S3 before sending them to the queue, and replaces the message with a reference to the uploaded object. The worker downloads the original payload back from S3 before executing the task, and (by default) removes it from S3 afterwards. diff --git a/src/taskiq_sqs/broker.py b/src/taskiq_sqs/broker.py index 466314e..8a39d25 100644 --- a/src/taskiq_sqs/broker.py +++ b/src/taskiq_sqs/broker.py @@ -13,6 +13,7 @@ from taskiq_sqs import constants from taskiq_sqs.exceptions import BrokerInitError, FifoDelayNotSupportedError, UnknownQueueError from taskiq_sqs.types.message import ( + is_label_true, validate_delay_seconds, validate_expiry, validate_message_deduplication_id, @@ -56,6 +57,8 @@ def __init__( self._default_queue_name = self._queues[0]["name"] self._queues_by_name = {queue["name"]: queue for queue in self._queues} self._queue_urls: dict[str, str] = {} + self._batch_queues: dict[str, asyncio.Queue[dict[str, Any]]] = {} + self._batch_worker_tasks: dict[str, asyncio.Task[None]] = {} @staticmethod def _normalize_queues(queues: SQSQueue | Sequence[SQSQueue]) -> list[SQSQueue]: @@ -106,6 +109,9 @@ async def startup(self) -> None: for queue in self._queues: queue_url = await self._get_queue_url(queue["name"]) logger.info("Resolved queue '%s' URL: %s", queue["name"], queue_url) + if queue.get("is_batching_enabled", False): + self._batch_queues[queue["name"]] = asyncio.Queue() + self._batch_worker_tasks[queue["name"]] = asyncio.create_task(self._batch_worker(queue)) except Exception: await self._sqs_client.__aexit__(None, None, None) raise @@ -114,6 +120,13 @@ async def startup(self) -> None: async def shutdown(self) -> None: """Shuts down the SQS broker.""" + for task in self._batch_worker_tasks.values(): + task.cancel() + await asyncio.gather(*self._batch_worker_tasks.values(), return_exceptions=True) + for queue_name, batch_queue in self._batch_queues.items(): + remaining = self._drain_batch_queue(batch_queue) + if remaining: + await self._send_batch(self._queues_by_name[queue_name], remaining) await self._sqs_client.__aexit__(None, None, None) await super().shutdown() @@ -159,14 +172,86 @@ async def _build_kick_kwargs( } return kwargs + def _should_batch(self, queue: SQSQueue, message: BrokerMessage) -> bool: + if not queue.get("is_batching_enabled", False): + return False + if is_label_true(message.labels.get(constants.SQS_SKIP_BATCHING_LABEL, False)): + return False + return constants.SQS_DELAY_SECONDS_LABEL not in message.labels + async def kick(self, message: BrokerMessage) -> None: """Kick tasks out from current program to configured SQS queue.""" queue = self._resolve_queue(message.labels.get(constants.SQS_QUEUE_LABEL)) queue_url = await self._get_queue_url(queue["name"]) kwargs = await self._build_kick_kwargs(message, queue, queue_url) + if self._should_batch(queue, message): + await self._batch_queues[queue["name"]].put(kwargs) + return with self._handle_exceptions(queue["name"]): await self._sqs_client.send_message(**kwargs) + @staticmethod + def _drain_batch_queue(batch_queue: "asyncio.Queue[dict[str, Any]]") -> list[dict[str, Any]]: + drained = [] + while not batch_queue.empty(): + try: + drained.append(batch_queue.get_nowait()) + except asyncio.QueueEmpty: + break + return drained + + async def _send_batch_to_sqs(self, queue: SQSQueue, queue_url: str, batch: list[dict[str, Any]]) -> None: + entries: list[Any] = [ + {"id": str(index), **{key: value for key, value in kwargs.items() if key != "queue_url"}} + for index, kwargs in enumerate(batch) + ] + with self._handle_exceptions(queue["name"]): + response = await self._sqs_client.send_message_batch(queue_url=queue_url, entries=entries) + for failure in response.get("failed", []): + logger.error( + "Failed to send batched message to queue '%s': %s (%s)", + queue["name"], + failure.get("message"), + failure.get("code"), + ) + + async def _send_batch(self, queue: SQSQueue, batch: list[dict[str, Any]]) -> None: + queue_url = await self._get_queue_url(queue["name"]) + if not queue.get("is_fifo", queue["name"].endswith(".fifo")): + await self._send_batch_to_sqs(queue, queue_url, batch) + return + # Keep each FIFO group's messages together in their own batch call, to preserve their relative order. + groups: dict[str, list[dict[str, Any]]] = {} + for kwargs in batch: + groups.setdefault(kwargs.get("message_group_id", ""), []).append(kwargs) + for group in groups.values(): + await self._send_batch_to_sqs(queue, queue_url, group) + + async def _batch_worker(self, queue: SQSQueue) -> None: + """Buffer kicked messages for queue and flush them together via batch send.""" + batch_queue = self._batch_queues[queue["name"]] + batch_size = queue.get("batch_size", constants.DEFAULT_BATCH_SIZE) + batch_timeout = queue.get("batch_timeout", constants.DEFAULT_BATCH_TIMEOUT) + batch: list[dict[str, Any]] = [] + try: + while True: + batch = [await batch_queue.get()] + deadline = time.monotonic() + batch_timeout + while len(batch) < batch_size: + remaining = deadline - time.monotonic() + if remaining <= 0: + break + try: + batch.append(await asyncio.wait_for(batch_queue.get(), timeout=remaining)) + except TimeoutError: + break + await self._send_batch(queue, batch) + batch = [] + except asyncio.CancelledError: + for kwargs in batch: + batch_queue.put_nowait(kwargs) + raise + def _build_ack_function( self, queue_name: str, diff --git a/src/taskiq_sqs/constants.py b/src/taskiq_sqs/constants.py index 9655b27..ab7a233 100644 --- a/src/taskiq_sqs/constants.py +++ b/src/taskiq_sqs/constants.py @@ -7,6 +7,9 @@ MAX_NUMBER_OF_MESSAGES: Final[int] = 10 MAX_DELAY_SECONDS: Final[int] = 900 MAX_FIFO_ID_LENGTH: Final[int] = 128 +MAX_BATCH_SIZE: Final[int] = 10 +DEFAULT_BATCH_SIZE: Final[int] = 10 +DEFAULT_BATCH_TIMEOUT: Final[float] = 1.0 SQS_MAX_MESSAGE_SIZE_BYTES: Final[int] = 262_144 DEFAULT_S3_OFFLOAD_THRESHOLD_BYTES: Final[int] = 200_000 @@ -16,3 +19,4 @@ SQS_MESSAGE_GROUP_ID_LABEL: Final[str] = "group_id" SQS_MESSAGE_DEDUPLICATION_ID_LABEL: Final[str] = "deduplication_id" SQS_EXPIRY_LABEL: Final[str] = "expiry" +SQS_SKIP_BATCHING_LABEL: Final[str] = "skip_batching" diff --git a/src/taskiq_sqs/types/message.py b/src/taskiq_sqs/types/message.py index 609e86f..a18a4e5 100644 --- a/src/taskiq_sqs/types/message.py +++ b/src/taskiq_sqs/types/message.py @@ -53,3 +53,10 @@ def validate_expiry(expiry: Any) -> float: if value < 0: raise InvalidExpiryError(expiry=expiry) return value + + +def is_label_true(value: Any) -> bool: + """Interpret a boolean-ish message label the way taskiq's own label round-trip does.""" + if isinstance(value, str): + return value.strip().lower() == "true" + return bool(value) diff --git a/src/taskiq_sqs/types/queue.py b/src/taskiq_sqs/types/queue.py index 067cdf0..4c0652b 100644 --- a/src/taskiq_sqs/types/queue.py +++ b/src/taskiq_sqs/types/queue.py @@ -13,12 +13,18 @@ class SQSQueue(TypedDict): max_number_of_messages: Maximum messages to retrieve per poll (1-10). Defaults to 1. wait_time_seconds: Long polling wait time in seconds (0-20). Defaults to 0. is_fifo: Whether this is a FIFO queue. + is_batching_enabled: Whether to buffer kicked messages in memory and flush them via batch send. + batch_size: Maximum messages per batch (1-10). Defaults to 10. + batch_timeout: Maximum seconds to wait for a batch to fill up before flushing it anyway. Defaults to 1.0. """ name: str max_number_of_messages: NotRequired[int] wait_time_seconds: NotRequired[int] is_fifo: NotRequired[bool] + is_batching_enabled: NotRequired[bool] + batch_size: NotRequired[int] + batch_timeout: NotRequired[float] def validate_queue(queue: SQSQueue) -> None: @@ -39,3 +45,11 @@ def validate_queue(queue: SQSQueue) -> None: details=f"Queue '{queue['name']}' has is_fifo={queue['is_fifo']}, but SQS requires FIFO queue " "names to end in '.fifo' and standard queue names not to", ) + batch_size = queue.get("batch_size", constants.DEFAULT_BATCH_SIZE) + if batch_size > constants.MAX_BATCH_SIZE or batch_size < 1: + raise BrokerInitError( + details=f"BatchSize for queue '{queue['name']}' can be no greater than 10 or less than 1", + ) + batch_timeout = queue.get("batch_timeout", constants.DEFAULT_BATCH_TIMEOUT) + if batch_timeout <= 0: + raise BrokerInitError(details=f"BatchTimeout for queue '{queue['name']}' must be greater than 0") diff --git a/tests/test_broker_kick.py b/tests/test_broker_kick.py index 38d1471..2315fee 100644 --- a/tests/test_broker_kick.py +++ b/tests/test_broker_kick.py @@ -1,10 +1,11 @@ import asyncio +from typing import Any import capo_sqs import pytest from taskiq import BrokerMessage -from tests.conftest import _queue_name_from_url +from tests.conftest import AWSCredentials, _queue_name_from_url from taskiq_sqs import SQSBroker from taskiq_sqs.constants import ( @@ -13,6 +14,7 @@ SQS_MESSAGE_DEDUPLICATION_ID_LABEL, SQS_MESSAGE_GROUP_ID_LABEL, SQS_QUEUE_LABEL, + SQS_SKIP_BATCHING_LABEL, ) from taskiq_sqs.exceptions import ( BrokerInitError, @@ -23,6 +25,7 @@ InvalidMessageGroupIdError, UnknownQueueError, ) +from taskiq_sqs.types import SQSQueue async def test_when_kick_called__than_message_should_be_published_to_queue( @@ -311,3 +314,135 @@ async def test_when_kick_called_with_stringified_expiry_label__then_it_is_accept assert len(messages) == 1 attribute = messages[0].get("message_attributes", {}).get(SQS_EXPIRY_LABEL, {}) assert attribute.get("string_value") == "1789505020.5" + + +async def test_when_batching_enabled_and_timeout_elapses__then_batch_is_flushed( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, + sqs_queue: str, +) -> None: + broker = SQSBroker( + queues=SQSQueue(name=_queue_name_from_url(sqs_queue), is_batching_enabled=True, batch_timeout=0.2), + **aws_credentials, + ) + await broker.startup() + try: + message = BrokerMessage(task_id="t1", task_name="t1", message=b"one", labels={}) + await broker.kick(message) + + immediate = await sqs_client.receive_message(queue_url=sqs_queue) + assert not immediate.get("messages") + + await asyncio.sleep(0.4) + + flushed = await sqs_client.receive_message(queue_url=sqs_queue) + assert len(flushed.get("messages", [])) == 1 + finally: + await broker.shutdown() + + +async def test_when_batch_size_reached__then_batch_flushes_without_waiting_for_timeout( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, + sqs_queue: str, +) -> None: + broker = SQSBroker( + queues=SQSQueue( + name=_queue_name_from_url(sqs_queue), + is_batching_enabled=True, + batch_size=2, + batch_timeout=30, + ), + **aws_credentials, + ) + await broker.startup() + try: + await broker.kick(BrokerMessage(task_id="t1", task_name="t1", message=b"one", labels={})) + await broker.kick(BrokerMessage(task_id="t2", task_name="t2", message=b"two", labels={})) + + # batch_size reached, so this must flush well before the 30s batch_timeout + flushed: dict[str, Any] = {} + for _ in range(20): + flushed = await sqs_client.receive_message(queue_url=sqs_queue, max_number_of_messages=2) + if len(flushed.get("messages", [])) == 2: + break + await asyncio.sleep(0.1) + assert len(flushed.get("messages", [])) == 2 + finally: + await broker.shutdown() + + +async def test_when_skip_batching_label_set__then_message_is_sent_immediately( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, + sqs_queue: str, +) -> None: + broker = SQSBroker( + queues=SQSQueue(name=_queue_name_from_url(sqs_queue), is_batching_enabled=True, batch_timeout=30), + **aws_credentials, + ) + await broker.startup() + try: + message = BrokerMessage( + task_id="t1", + task_name="t1", + message=b"urgent", + labels={SQS_SKIP_BATCHING_LABEL: True}, + ) + await broker.kick(message) + + # must be visible well before the 30s batch_timeout + immediate = await sqs_client.receive_message(queue_url=sqs_queue) + assert len(immediate.get("messages", [])) == 1 + finally: + await broker.shutdown() + + +async def test_when_delay_label_set_on_batching_queue__then_message_bypasses_batching( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, + sqs_queue: str, +) -> None: + broker = SQSBroker( + queues=SQSQueue(name=_queue_name_from_url(sqs_queue), is_batching_enabled=True, batch_timeout=30), + **aws_credentials, + ) + await broker.startup() + try: + message = BrokerMessage( + task_id="t1", + task_name="t1", + message=b"delayed", + labels={SQS_DELAY_SECONDS_LABEL: 1}, + ) + await broker.kick(message) + + immediate = await sqs_client.receive_message(queue_url=sqs_queue) + assert not immediate.get("messages") + + # governed by the 1s delay, not the 30s batch_timeout + await asyncio.sleep(1.2) + delayed = await sqs_client.receive_message(queue_url=sqs_queue) + assert len(delayed.get("messages", [])) == 1 + finally: + await broker.shutdown() + + +async def test_when_broker_shuts_down_with_pending_batch__then_it_is_flushed( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, + sqs_queue: str, +) -> None: + broker = SQSBroker( + queues=SQSQueue(name=_queue_name_from_url(sqs_queue), is_batching_enabled=True, batch_timeout=30), + **aws_credentials, + ) + await broker.startup() + + message = BrokerMessage(task_id="t1", task_name="t1", message=b"pending", labels={}) + await broker.kick(message) + + await broker.shutdown() # must flush the pending batch, not just drop it + + response = await sqs_client.receive_message(queue_url=sqs_queue) + assert len(response.get("messages", [])) == 1 From a51b4e61c395429d50a93032f74262a309842acc Mon Sep 17 00:00:00 2001 From: Dima Anfimov Date: Tue, 15 Sep 2026 23:18:32 +0200 Subject: [PATCH 2/2] fix: declare -> is_declare --- README.md | 19 ++++++- src/taskiq_sqs/broker.py | 44 ++++++++++++---- src/taskiq_sqs/exceptions.py | 9 +++- src/taskiq_sqs/middleware.py | 4 +- src/taskiq_sqs/result_backend.py | 4 +- src/taskiq_sqs/types/bucket.py | 9 ++-- src/taskiq_sqs/types/queue.py | 7 ++- tests/benchmarks/test_broker.py | 2 +- tests/test_broker_initialization.py | 82 +++++++++++++++++++++++++++-- tests/test_broker_kick.py | 3 +- tests/test_middleware.py | 2 +- tests/test_result_backend.py | 8 +-- 12 files changed, 161 insertions(+), 32 deletions(-) diff --git a/README.md b/README.md index 2e85d42..41b44fd 100644 --- a/README.md +++ b/README.md @@ -22,7 +22,7 @@ from taskiq_sqs import S3ResultBackend, SQSBroker from taskiq_sqs.types import S3Bucket, SQSQueue broker = SQSBroker( - queues=SQSQueue(name="my-queue"), # specify an existing queue + queues=SQSQueue(name="my-queue"), # by default the broker creates the queue for you if it doesn't exist endpoint_url="http://localhost:4566", aws_region_name="us-east-1", ).with_result_backend( @@ -72,6 +72,23 @@ async def urgent_task() -> None: A worker started against this broker consumes from every configured queue at once. Passing a queue name through the `queue_name` label that isn't configured on the broker raises `UnknownQueueError`. +## Declaring queues + +By default the broker creates a queue on startup if it doesn't exist yet, the same way `S3Bucket` does for buckets. Set `is_declare=False` to require the queue to already exist instead (raises `QueueNotFoundError` if it doesn't). `options` are queue attributes (e.g. `VisibilityTimeout`, `MessageRetentionPeriod`) passed to `CreateQueue`, in AWS's own PascalCase naming, when the queue is declared — they have no effect on a queue that already exists: + +```python +from taskiq_sqs import SQSBroker +from taskiq_sqs.types import SQSQueue + +broker = SQSBroker( + queues=SQSQueue(name="my-queue", options={"VisibilityTimeout": "60", "MessageRetentionPeriod": "86400"}), +) +``` + +FIFO queues get their `FifoQueue` attribute set automatically when declared — no need to include it in `options`. + +`S3Bucket` has the same `options` field, for parameters `CreateBucket` accepts beyond `name` (e.g. `acl`), passed through whenever `S3ResultBackend`/`S3OffloadMiddleware` create the bucket. + ## Delayed tasks Set the `delay` label to delay delivery of a task by that many seconds (0-900, SQS's own limit) diff --git a/src/taskiq_sqs/broker.py b/src/taskiq_sqs/broker.py index 8a39d25..897f85f 100644 --- a/src/taskiq_sqs/broker.py +++ b/src/taskiq_sqs/broker.py @@ -11,7 +11,12 @@ from taskiq.message import BrokerMessage from taskiq_sqs import constants -from taskiq_sqs.exceptions import BrokerInitError, FifoDelayNotSupportedError, UnknownQueueError +from taskiq_sqs.exceptions import ( + BrokerInitError, + FifoDelayNotSupportedError, + QueueNotFoundError, + UnknownQueueError, +) from taskiq_sqs.types.message import ( is_label_true, validate_delay_seconds, @@ -107,7 +112,7 @@ async def startup(self) -> None: await self._sqs_client.__aenter__() try: for queue in self._queues: - queue_url = await self._get_queue_url(queue["name"]) + queue_url = await self._get_queue_url(queue) logger.info("Resolved queue '%s' URL: %s", queue["name"], queue_url) if queue.get("is_batching_enabled", False): self._batch_queues[queue["name"]] = asyncio.Queue() @@ -130,12 +135,29 @@ async def shutdown(self) -> None: await self._sqs_client.__aexit__(None, None, None) await super().shutdown() - async def _get_queue_url(self, queue_name: str) -> str: - if queue_name not in self._queue_urls: - with self._handle_exceptions(queue_name): - result = await self._sqs_client.get_queue_url(queue_name=queue_name) - self._queue_urls[queue_name] = result["queue_url"] - return self._queue_urls[queue_name] + async def _get_queue_url(self, queue: SQSQueue) -> str: + name = queue["name"] + if name not in self._queue_urls: + result: Any + try: + result = await self._sqs_client.get_queue_url(queue_name=name) + except capo_sqs.errors.QueueDoesNotExist as exc: + if not queue.get("is_declare", True): + raise QueueNotFoundError(queue_name=name) from exc + result = await self._create_queue(queue) + except capo_sqs.errors.ServiceError as exc: + raise BrokerInitError(details=exc.code or "") from exc + self._queue_urls[name] = result["queue_url"] + return self._queue_urls[name] + + async def _create_queue(self, queue: SQSQueue) -> Any: + attributes: dict[Any, Any] = dict(queue.get("options", {})) + if queue.get("is_fifo", queue["name"].endswith(".fifo")): + attributes.setdefault("FifoQueue", "true") + try: + return await self._sqs_client.create_queue(queue_name=queue["name"], attributes=attributes or None) + except capo_sqs.errors.ServiceError as exc: + raise BrokerInitError(details=exc.code or "") from exc async def _build_kick_kwargs( self, @@ -182,7 +204,7 @@ def _should_batch(self, queue: SQSQueue, message: BrokerMessage) -> bool: async def kick(self, message: BrokerMessage) -> None: """Kick tasks out from current program to configured SQS queue.""" queue = self._resolve_queue(message.labels.get(constants.SQS_QUEUE_LABEL)) - queue_url = await self._get_queue_url(queue["name"]) + queue_url = await self._get_queue_url(queue) kwargs = await self._build_kick_kwargs(message, queue, queue_url) if self._should_batch(queue, message): await self._batch_queues[queue["name"]].put(kwargs) @@ -216,7 +238,7 @@ async def _send_batch_to_sqs(self, queue: SQSQueue, queue_url: str, batch: list[ ) async def _send_batch(self, queue: SQSQueue, batch: list[dict[str, Any]]) -> None: - queue_url = await self._get_queue_url(queue["name"]) + queue_url = await self._get_queue_url(queue) if not queue.get("is_fifo", queue["name"].endswith(".fifo")): await self._send_batch_to_sqs(queue, queue_url, batch) return @@ -292,7 +314,7 @@ def _is_expired(message: Mapping[str, Any]) -> bool: async def _poll_queue(self, queue: SQSQueue, incoming: "asyncio.Queue[_QueueItem]") -> None: """Continuously receive messages from a single queue and forward them to the shared incoming queue.""" try: - queue_url = await self._get_queue_url(queue["name"]) + queue_url = await self._get_queue_url(queue) while True: with self._handle_exceptions(queue["name"]): results = await self._sqs_client.receive_message( diff --git a/src/taskiq_sqs/exceptions.py b/src/taskiq_sqs/exceptions.py index b9af2ae..71ab6d6 100644 --- a/src/taskiq_sqs/exceptions.py +++ b/src/taskiq_sqs/exceptions.py @@ -29,7 +29,7 @@ class ResultBackendError(BaseTaskiqSQSError): class BucketNotFoundError(BaseTaskiqSQSError): """Error if bucket not found.""" - __template__ = "Bucket '{bucket_name}' not found during initialization and declare=False" + __template__ = "Bucket '{bucket_name}' not found during initialization and is_declare=False" bucket_name: str @@ -56,6 +56,13 @@ class UnknownQueueError(BaseTaskiqSQSError): queue_name: str +class QueueNotFoundError(BaseTaskiqSQSError): + """Error if a queue doesn't exist and is_declare=False.""" + + __template__ = "Queue '{queue_name}' not found during initialization and is_declare=False" + queue_name: str + + class InvalidDelaySecondsError(BaseTaskiqSQSError): """Error if a message's delay label is outside SQS's allowed range.""" diff --git a/src/taskiq_sqs/middleware.py b/src/taskiq_sqs/middleware.py index 7aee90a..89f4693 100644 --- a/src/taskiq_sqs/middleware.py +++ b/src/taskiq_sqs/middleware.py @@ -92,12 +92,12 @@ async def _ensure_bucket_exists(self) -> None: try: await self._s3_client.head_bucket(bucket=self._bucket["name"]) except capo_s3.errors.NotFound: - if not self._bucket.get("declare", True): + if not self._bucket.get("is_declare", True): raise exceptions.BucketNotFoundError(bucket_name=self._bucket["name"]) from None await self._create_bucket() async def _create_bucket(self) -> None: - create_kwargs: dict[str, Any] = {} + create_kwargs: dict[str, Any] = dict(self._bucket.get("options", {})) if self._aws_region and self._aws_region != constants.AWS_DEFAULT_REGION: create_kwargs["create_bucket_configuration"] = {"location_constraint": self._aws_region} with contextlib.suppress(capo_s3.errors.BucketAlreadyOwnedByYou): diff --git a/src/taskiq_sqs/result_backend.py b/src/taskiq_sqs/result_backend.py index a689d25..c5be795 100644 --- a/src/taskiq_sqs/result_backend.py +++ b/src/taskiq_sqs/result_backend.py @@ -73,14 +73,14 @@ async def _ensure_bucket_exists(self) -> None: try: await self._s3_client.head_bucket(bucket=self._bucket["name"]) except capo_s3.errors.NotFound: - if not self._bucket.get("declare", True): + if not self._bucket.get("is_declare", True): raise exceptions.BucketNotFoundError(bucket_name=self._bucket["name"]) from None await self._create_bucket() except capo_s3.errors.ServiceError as exc: raise exceptions.ResultBackendError(code=exc.code) from exc async def _create_bucket(self) -> None: - create_kwargs: dict[str, Any] = {} + create_kwargs: dict[str, Any] = dict(self._bucket.get("options", {})) if self._aws_region and self._aws_region != constants.AWS_DEFAULT_REGION: create_kwargs["create_bucket_configuration"] = {"location_constraint": self._aws_region} with contextlib.suppress(capo_s3.errors.BucketAlreadyOwnedByYou): diff --git a/src/taskiq_sqs/types/bucket.py b/src/taskiq_sqs/types/bucket.py index 9771133..303451b 100644 --- a/src/taskiq_sqs/types/bucket.py +++ b/src/taskiq_sqs/types/bucket.py @@ -1,4 +1,5 @@ -from typing import NotRequired, TypedDict +from collections.abc import Mapping +from typing import Any, NotRequired, TypedDict class S3Bucket(TypedDict): @@ -7,8 +8,10 @@ class S3Bucket(TypedDict): Attributes: name: The name of the bucket. - declare: Whether to create the bucket on startup if it not exists yet. Defaults to True. + is_declare: Whether to create the bucket on startup if it not exists yet. Defaults to True. + options: Extra keyword arguments merged into the create bucket call when the bucket is declared. """ name: str - declare: NotRequired[bool] + is_declare: NotRequired[bool] + options: NotRequired[Mapping[str, Any]] diff --git a/src/taskiq_sqs/types/queue.py b/src/taskiq_sqs/types/queue.py index 4c0652b..929e4b0 100644 --- a/src/taskiq_sqs/types/queue.py +++ b/src/taskiq_sqs/types/queue.py @@ -1,4 +1,5 @@ -from typing import NotRequired, TypedDict +from collections.abc import Mapping +from typing import Any, NotRequired, TypedDict from taskiq_sqs import constants from taskiq_sqs.exceptions import BrokerInitError @@ -16,6 +17,8 @@ class SQSQueue(TypedDict): is_batching_enabled: Whether to buffer kicked messages in memory and flush them via batch send. batch_size: Maximum messages per batch (1-10). Defaults to 10. batch_timeout: Maximum seconds to wait for a batch to fill up before flushing it anyway. Defaults to 1.0. + is_declare: Whether to create the queue on startup if it doesn't exist yet. Defaults to True. + options: Queue attributes passed during queue creation when the queue declaration is enabled. """ name: str @@ -25,6 +28,8 @@ class SQSQueue(TypedDict): is_batching_enabled: NotRequired[bool] batch_size: NotRequired[int] batch_timeout: NotRequired[float] + is_declare: NotRequired[bool] + options: NotRequired[Mapping[str, Any]] def validate_queue(queue: SQSQueue) -> None: diff --git a/tests/benchmarks/test_broker.py b/tests/benchmarks/test_broker.py index 5a0aae9..49071b9 100644 --- a/tests/benchmarks/test_broker.py +++ b/tests/benchmarks/test_broker.py @@ -11,7 +11,7 @@ @pytest.mark.benchmark async def test_build_kick_kwargs(bench_broker: SQSBroker, broker_message: BrokerMessage) -> None: queue = bench_broker._resolve_queue(None) - queue_url = await bench_broker._get_queue_url(queue["name"]) + queue_url = await bench_broker._get_queue_url(queue) await bench_broker._build_kick_kwargs(broker_message, queue, queue_url) diff --git a/tests/test_broker_initialization.py b/tests/test_broker_initialization.py index 0ce08da..bd8720a 100644 --- a/tests/test_broker_initialization.py +++ b/tests/test_broker_initialization.py @@ -1,16 +1,90 @@ +import asyncio +import uuid + +import capo_sqs import pytest +from taskiq import BrokerMessage from tests.conftest import AWSCredentials from taskiq_sqs import SQSBroker -from taskiq_sqs.exceptions import BrokerInitError +from taskiq_sqs.exceptions import BrokerInitError, QueueNotFoundError from taskiq_sqs.types import SQSQueue -async def test_get_queue_url_client_error(aws_credentials: AWSCredentials) -> None: - broker = SQSBroker(queues=SQSQueue(name="nonexistent-queue"), **aws_credentials) - with pytest.raises(BrokerInitError): +async def test_when_queue_missing_and_declare_false__then_startup_raises(aws_credentials: AWSCredentials) -> None: + broker = SQSBroker( + queues=SQSQueue(name=f"declare-false-{uuid.uuid4().hex}", is_declare=False), + **aws_credentials, + ) + with pytest.raises(QueueNotFoundError): + await broker.startup() + + +async def test_when_queue_missing_and_declare_true__then_it_is_created( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, +) -> None: + queue_name = f"declare-true-{uuid.uuid4().hex}" + broker = SQSBroker(queues=SQSQueue(name=queue_name), **aws_credentials) # declare defaults to True + try: await broker.startup() + response = await sqs_client.get_queue_url(queue_name=queue_name) + assert response.get("queue_url") == broker._queue_urls[queue_name] + finally: + await broker.shutdown() + await sqs_client.delete_queue(queue_url=broker._queue_urls[queue_name]) + + +async def test_when_queue_declared_with_options__then_they_become_queue_attributes( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, +) -> None: + queue_name = f"declare-options-{uuid.uuid4().hex}" + broker = SQSBroker( + queues=SQSQueue(name=queue_name, options={"VisibilityTimeout": "1"}), + **aws_credentials, + ) + try: + await broker.startup() + queue_url = broker._queue_urls[queue_name] + await sqs_client.send_message(queue_url=queue_url, message_body="hidden") + + generator = broker.listen() + try: + await asyncio.wait_for(generator.__anext__(), timeout=2) # received, deliberately left unacked + finally: + await generator.aclose() + + # still within the 1s VisibilityTimeout the queue was declared with, so this must not see it yet + immediate = await sqs_client.receive_message(queue_url=queue_url) + assert not immediate.get("messages") + finally: + await broker.shutdown() + await sqs_client.delete_queue(queue_url=broker._queue_urls[queue_name]) + + +async def test_when_fifo_queue_declared__then_fifo_attribute_is_set_automatically( + aws_credentials: AWSCredentials, + sqs_client: capo_sqs.AsyncSQSClient, +) -> None: + queue_name = f"declare-fifo-{uuid.uuid4().hex}.fifo" + broker = SQSBroker(queues=SQSQueue(name=queue_name), **aws_credentials) + try: + await broker.startup() + message = BrokerMessage( + task_id="t1", + task_name="t1", + message=b"x", + labels={"group_id": "g1", "deduplication_id": "d1"}, + ) + await broker.kick(message) + + response = await sqs_client.receive_message(queue_url=broker._queue_urls[queue_name]) + assert len(response.get("messages", [])) == 1 + finally: + await broker.shutdown() + await sqs_client.delete_queue(queue_url=broker._queue_urls[queue_name]) async def test_max_number_of_messages_error(aws_credentials: AWSCredentials) -> None: diff --git a/tests/test_broker_kick.py b/tests/test_broker_kick.py index 2315fee..d48fe32 100644 --- a/tests/test_broker_kick.py +++ b/tests/test_broker_kick.py @@ -1,4 +1,5 @@ import asyncio +import uuid from typing import Any import capo_sqs @@ -46,7 +47,7 @@ async def test_when_during_kick_queue_not_found__then_should_raise_an_error( sqs_broker: SQSBroker, broker_message: BrokerMessage, ) -> None: - sqs_broker._queue_urls[sqs_broker._default_queue_name] = "nonexistent-queue" + sqs_broker._queue_urls[sqs_broker._default_queue_name] = f"not-a-real-queue-url-{uuid.uuid4().hex}" with pytest.raises(BrokerInitError): await sqs_broker.kick(broker_message) diff --git a/tests/test_middleware.py b/tests/test_middleware.py index cf64bce..b6fcd9e 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -125,7 +125,7 @@ async def test_when_bucket_missing_and_declare_false__then_startup_raises( aws_credentials: AWSCredentials, ) -> None: middleware = S3OffloadMiddleware( - bucket=S3Bucket(name="nonexistent-offload-bucket", declare=False), + bucket=S3Bucket(name="nonexistent-offload-bucket", is_declare=False), **aws_credentials, ) with pytest.raises(BucketNotFoundError): diff --git a/tests/test_result_backend.py b/tests/test_result_backend.py index f4c30aa..0bc85ca 100644 --- a/tests/test_result_backend.py +++ b/tests/test_result_backend.py @@ -144,7 +144,7 @@ async def test_when_declare_true_and_bucket_missing__then_bucket_is_created_on_s s3_client: capo_s3.AsyncS3Client, ) -> None: self.backend = S3ResultBackend( - bucket=S3Bucket(name=self.tmp_bucket_name, declare=True), + bucket=S3Bucket(name=self.tmp_bucket_name, is_declare=True), **aws_credentials, ) await self.backend.startup() @@ -157,7 +157,7 @@ async def test_when_declare_false_and_bucket_missing__then_startup_raises( s3_client: capo_s3.AsyncS3Client, ) -> None: backend = S3ResultBackend( - bucket=S3Bucket(name=self.tmp_bucket_name, declare=False), + bucket=S3Bucket(name=self.tmp_bucket_name, is_declare=False), **aws_credentials, ) @@ -173,7 +173,7 @@ async def test_when_declare_false_and_bucket_exists__then_startup_succeeds( s3_bucket: str, ) -> None: self.backend = S3ResultBackend( - bucket=S3Bucket(name=s3_bucket, declare=False), + bucket=S3Bucket(name=s3_bucket, is_declare=False), **aws_credentials, ) await self.backend.startup() @@ -187,7 +187,7 @@ async def test_when_declare_true_and_bucket_already_exists__then_startup_is_idem s3_bucket: str, ) -> None: self.backend = S3ResultBackend( - bucket=S3Bucket(name=s3_bucket, declare=True), + bucket=S3Bucket(name=s3_bucket, is_declare=True), **aws_credentials, ) await self.backend.startup()