Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,25 @@ async def send_reminder() -> None:

A value outside the 0-900 range (or not an integer) raises `InvalidDelaySecondsError` when the task is kicked.

## FIFO queues

A queue whose name ends in `.fifo` is treated as a FIFO queue automatically, matching SQS's own naming rule (`SQSQueue`'s `is_fifo` field only needs to be set to override that default, and the queue's name must still end in `.fifo` for the broker to accept it as FIFO).

```python
from taskiq_sqs import SQSBroker
from taskiq_sqs.types import SQSQueue

broker = SQSBroker(queues=SQSQueue(name="my-queue.fifo"))

@broker.task(group_id="orders") # defaults to the task name if not set
async def process_order() -> None:
...
```

- `group_id` picks the message's `MessageGroupId` (required by SQS for every FIFO message); it defaults to the task's name.
- `deduplication_id` sets `MessageDeduplicationId`; if not set, the queue must have content-based deduplication enabled, or SQS rejects the message.
- The `delay` label (see [Delayed tasks](#delayed-tasks)) is not supported on FIFO queues — SQS only allows delay to be configured on the queue itself, not per message — and raises `FifoDelayNotSupportedError` if used.

## 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.
Expand Down
43 changes: 41 additions & 2 deletions src/taskiq_sqs/broker.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,14 @@
from taskiq.message import BrokerMessage

from taskiq_sqs import constants
from taskiq_sqs.exceptions import BrokerInitError, InvalidDelaySecondsError, UnknownQueueError
from taskiq_sqs.exceptions import (
BrokerInitError,
FifoDelayNotSupportedError,
InvalidDelaySecondsError,
InvalidMessageDeduplicationIdError,
InvalidMessageGroupIdError,
UnknownQueueError,
)
from taskiq_sqs.types import SQSQueue


Expand Down Expand Up @@ -71,6 +78,12 @@ def _normalize_queues(queues: SQSQueue | Sequence[SQSQueue]) -> list[SQSQueue]:
raise BrokerInitError(
details=f"WaitTimeSeconds for queue '{queue['name']}' can be no greater than 20 or less than 0",
)
ends_with_fifo_suffix = queue["name"].endswith(".fifo")
if "is_fifo" in queue and queue["is_fifo"] != ends_with_fifo_suffix:
raise BrokerInitError(
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",
)
return queue_list

def _resolve_queue(self, queue_name: str | None) -> SQSQueue:
Expand Down Expand Up @@ -129,20 +142,31 @@ async def _get_queue_url(self, queue_name: str) -> str:
async def _build_kick_kwargs(
self,
message: BrokerMessage,
queue: SQSQueue,
queue_url: str,
) -> dict[str, Any]:
"""Build the kwargs for the SQS client kick method.

This function can be extended by the end user to add additional kwargs in the message delivery.
:param message: BrokerMessage object.
:param queue: the queue the message will be sent to.
:param queue_url: URL of the queue the message will be sent to.
"""
kwargs: dict[str, Any] = {
"queue_url": queue_url,
"message_body": message.message.decode("utf-8"),
}
is_fifo = queue.get("is_fifo", queue["name"].endswith(".fifo"))
if constants.SQS_DELAY_SECONDS_LABEL in message.labels:
if is_fifo:
raise FifoDelayNotSupportedError(queue_name=queue["name"])
kwargs["delay_seconds"] = self._validate_delay_seconds(message.labels[constants.SQS_DELAY_SECONDS_LABEL])
if is_fifo:
group_id = message.labels.get(constants.SQS_MESSAGE_GROUP_ID_LABEL, message.task_name)
kwargs["message_group_id"] = self._validate_message_group_id(group_id)
if constants.SQS_MESSAGE_DEDUPLICATION_ID_LABEL in message.labels:
deduplication_id = message.labels[constants.SQS_MESSAGE_DEDUPLICATION_ID_LABEL]
kwargs["message_deduplication_id"] = self._validate_message_deduplication_id(deduplication_id)
return kwargs

@staticmethod
Expand All @@ -153,11 +177,26 @@ def _validate_delay_seconds(delay_seconds: Any) -> int:
raise InvalidDelaySecondsError(delay_seconds=delay_seconds, max_delay_seconds=constants.MAX_DELAY_SECONDS)
return delay_seconds

@staticmethod
def _validate_message_group_id(group_id: Any) -> str:
if not isinstance(group_id, str) or not (1 <= len(group_id) <= constants.MAX_FIFO_ID_LENGTH):
raise InvalidMessageGroupIdError(group_id=group_id, max_length=constants.MAX_FIFO_ID_LENGTH)
return group_id

@staticmethod
def _validate_message_deduplication_id(deduplication_id: Any) -> str:
if not isinstance(deduplication_id, str) or not (1 <= len(deduplication_id) <= constants.MAX_FIFO_ID_LENGTH):
raise InvalidMessageDeduplicationIdError(
deduplication_id=deduplication_id,
max_length=constants.MAX_FIFO_ID_LENGTH,
)
return deduplication_id

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_url)
kwargs = await self._build_kick_kwargs(message, queue, queue_url)
with self._handle_exceptions(queue["name"]):
await self._sqs_client.send_message(**kwargs)

Expand Down
3 changes: 3 additions & 0 deletions src/taskiq_sqs/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,12 @@
MAX_WAIT_TIME_SECONDS: Final[int] = 20
MAX_NUMBER_OF_MESSAGES: Final[int] = 10
MAX_DELAY_SECONDS: Final[int] = 900
MAX_FIFO_ID_LENGTH: Final[int] = 128

SQS_MAX_MESSAGE_SIZE_BYTES: Final[int] = 262_144
DEFAULT_S3_OFFLOAD_THRESHOLD_BYTES: Final[int] = 200_000

SQS_QUEUE_LABEL: Final[str] = "queue_name"
SQS_DELAY_SECONDS_LABEL: Final[str] = "delay"
SQS_MESSAGE_GROUP_ID_LABEL: Final[str] = "group_id"
SQS_MESSAGE_DEDUPLICATION_ID_LABEL: Final[str] = "deduplication_id"
28 changes: 28 additions & 0 deletions src/taskiq_sqs/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,3 +62,31 @@ class InvalidDelaySecondsError(BaseTaskiqSQSError):
__template__ = "delay must be an integer between 0 and {max_delay_seconds}, got {delay_seconds!r}"
delay_seconds: object
max_delay_seconds: int


class FifoDelayNotSupportedError(BaseTaskiqSQSError):
"""Error if a per-message delay is requested on a FIFO queue.

FIFO queues don't support delaying individual messages; delay can only be configured on the queue itself.
"""

__template__ = "Per-message delay is not supported on FIFO queue '{queue_name}'; set it on the queue itself"
queue_name: str


class InvalidMessageGroupIdError(BaseTaskiqSQSError):
"""Error if a message's group_id label is invalid for a FIFO queue."""

__template__ = "group_id must be a non-empty string of at most {max_length} characters, got {group_id!r}"
group_id: object
max_length: int


class InvalidMessageDeduplicationIdError(BaseTaskiqSQSError):
"""Error if a message's deduplication_id label is invalid for a FIFO queue."""

__template__ = (
"deduplication_id must be a non-empty string of at most {max_length} characters, got {deduplication_id!r}"
)
deduplication_id: object
max_length: int
2 changes: 2 additions & 0 deletions src/taskiq_sqs/types/queue.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,10 @@ class SQSQueue(TypedDict):
name: The SQS queue name.
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.
"""

name: str
max_number_of_messages: NotRequired[int]
wait_time_seconds: NotRequired[int]
is_fifo: NotRequired[bool]
5 changes: 3 additions & 2 deletions tests/benchmarks/test_broker.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,9 @@

@pytest.mark.benchmark
async def test_build_kick_kwargs(bench_broker: SQSBroker, broker_message: BrokerMessage) -> None:
queue_url = await bench_broker._get_queue_url(bench_broker._default_queue_name)
await bench_broker._build_kick_kwargs(broker_message, queue_url)
queue = bench_broker._resolve_queue(None)
queue_url = await bench_broker._get_queue_url(queue["name"])
await bench_broker._build_kick_kwargs(broker_message, queue, queue_url)


@pytest.mark.benchmark
Expand Down
33 changes: 31 additions & 2 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,8 +121,12 @@ async def sqs_client(aws_credentials: AWSCredentials) -> AsyncGenerator[capo_sqs
await client.__aexit__(None, None, None)


async def _create_queue(sqs_client: capo_sqs.AsyncSQSClient, name: str) -> str:
response = await sqs_client.create_queue(queue_name=name)
async def _create_queue(
sqs_client: capo_sqs.AsyncSQSClient,
name: str,
attributes: dict[str, str] | None = None,
) -> str:
response = await sqs_client.create_queue(queue_name=name, attributes=attributes)
queue_url = response.get("queue_url")
assert queue_url is not None
return queue_url
Expand All @@ -142,6 +146,17 @@ async def sqs_second_queue(sqs_client: capo_sqs.AsyncSQSClient) -> AsyncGenerato
await sqs_client.delete_queue(queue_url=queue_url)


@pytest.fixture
async def fifo_sqs_queue(sqs_client: capo_sqs.AsyncSQSClient) -> AsyncGenerator[str, Any]:
queue_url = await _create_queue(
sqs_client,
f"{QUEUE_NAME}-fifo-{uuid.uuid4().hex}.fifo",
attributes={"FifoQueue": "true", "ContentBasedDeduplication": "true"},
)
yield queue_url
await sqs_client.delete_queue(queue_url=queue_url)


def _queue_name_from_url(queue_url: str) -> str:
return queue_url.rsplit("/", maxsplit=1)[-1]

Expand All @@ -162,6 +177,20 @@ async def sqs_broker(
await broker.shutdown()


@pytest.fixture
async def fifo_sqs_broker(
aws_credentials: AWSCredentials,
fifo_sqs_queue: str,
) -> AsyncGenerator[SQSBroker, Any]:
broker = SQSBroker(
queues=SQSQueue(name=_queue_name_from_url(fifo_sqs_queue)),
**aws_credentials,
)
await broker.startup()
yield broker
await broker.shutdown()


@pytest.fixture
async def multiqueue_sqs_broker(
aws_credentials: AWSCredentials,
Expand Down
144 changes: 142 additions & 2 deletions tests/test_broker_kick.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,20 @@
from tests.conftest import _queue_name_from_url

from taskiq_sqs import SQSBroker
from taskiq_sqs.constants import SQS_DELAY_SECONDS_LABEL, SQS_QUEUE_LABEL
from taskiq_sqs.exceptions import BrokerInitError, InvalidDelaySecondsError, UnknownQueueError
from taskiq_sqs.constants import (
SQS_DELAY_SECONDS_LABEL,
SQS_MESSAGE_DEDUPLICATION_ID_LABEL,
SQS_MESSAGE_GROUP_ID_LABEL,
SQS_QUEUE_LABEL,
)
from taskiq_sqs.exceptions import (
BrokerInitError,
FifoDelayNotSupportedError,
InvalidDelaySecondsError,
InvalidMessageDeduplicationIdError,
InvalidMessageGroupIdError,
UnknownQueueError,
)


async def test_when_kick_called__than_message_should_be_published_to_queue(
Expand Down Expand Up @@ -99,3 +111,131 @@ async def test_when_kick_called_with_invalid_delay_label__then_should_raise_an_e

with pytest.raises(InvalidDelaySecondsError):
await sqs_broker.kick(broker_message)


async def test_when_kick_called_on_standard_queue__then_no_fifo_attributes_are_sent(
sqs_broker: SQSBroker,
sqs_client: capo_sqs.AsyncSQSClient,
sqs_queue: str,
broker_message: BrokerMessage,
) -> None:
await sqs_broker.kick(broker_message)

response = await sqs_client.receive_message(
queue_url=sqs_queue,
message_system_attribute_names=["MessageGroupId"],
)
messages = response.get("messages", [])
assert len(messages) == 1
assert "MessageGroupId" not in messages[0].get("attributes", {})


async def test_when_kick_called_on_fifo_queue_without_group_id_label__then_task_name_is_used(
fifo_sqs_broker: SQSBroker,
sqs_client: capo_sqs.AsyncSQSClient,
fifo_sqs_queue: str,
broker_message: BrokerMessage,
) -> None:
await fifo_sqs_broker.kick(broker_message)

response = await sqs_client.receive_message(
queue_url=fifo_sqs_queue,
message_system_attribute_names=["MessageGroupId"],
)
messages = response.get("messages", [])
assert len(messages) == 1
assert messages[0].get("attributes", {}).get("MessageGroupId") == broker_message.task_name


async def test_when_kick_called_on_fifo_queue_with_group_id_label__then_it_is_used(
fifo_sqs_broker: SQSBroker,
sqs_client: capo_sqs.AsyncSQSClient,
fifo_sqs_queue: str,
broker_message: BrokerMessage,
) -> None:
broker_message.labels[SQS_MESSAGE_GROUP_ID_LABEL] = "custom-group"

await fifo_sqs_broker.kick(broker_message)

response = await sqs_client.receive_message(
queue_url=fifo_sqs_queue,
message_system_attribute_names=["MessageGroupId"],
)
messages = response.get("messages", [])
assert len(messages) == 1
assert messages[0].get("attributes", {}).get("MessageGroupId") == "custom-group"


async def test_when_kick_called_on_fifo_queue_with_deduplication_id_label__then_it_is_used(
fifo_sqs_broker: SQSBroker,
sqs_client: capo_sqs.AsyncSQSClient,
fifo_sqs_queue: str,
broker_message: BrokerMessage,
) -> None:
broker_message.labels[SQS_MESSAGE_DEDUPLICATION_ID_LABEL] = "custom-dedup-id"

await fifo_sqs_broker.kick(broker_message)

response = await sqs_client.receive_message(
queue_url=fifo_sqs_queue,
message_system_attribute_names=["MessageDeduplicationId"],
)
messages = response.get("messages", [])
assert len(messages) == 1
assert messages[0].get("attributes", {}).get("MessageDeduplicationId") == "custom-dedup-id"


async def test_when_kick_called_with_delay_on_fifo_queue__then_should_raise_an_error(
fifo_sqs_broker: SQSBroker,
broker_message: BrokerMessage,
) -> None:
broker_message.labels[SQS_DELAY_SECONDS_LABEL] = 5

with pytest.raises(FifoDelayNotSupportedError):
await fifo_sqs_broker.kick(broker_message)


@pytest.mark.parametrize("group_id", ["", "x" * 129, 123, None])
async def test_when_kick_called_on_fifo_queue_with_invalid_group_id__then_should_raise_an_error(
fifo_sqs_broker: SQSBroker,
broker_message: BrokerMessage,
group_id: object,
) -> None:
broker_message.labels[SQS_MESSAGE_GROUP_ID_LABEL] = group_id

with pytest.raises(InvalidMessageGroupIdError):
await fifo_sqs_broker.kick(broker_message)


@pytest.mark.parametrize("deduplication_id", ["", "x" * 129, 123, None])
async def test_when_kick_called_on_fifo_queue_with_invalid_deduplication_id__then_should_raise_an_error(
fifo_sqs_broker: SQSBroker,
broker_message: BrokerMessage,
deduplication_id: object,
) -> None:
broker_message.labels[SQS_MESSAGE_DEDUPLICATION_ID_LABEL] = deduplication_id

with pytest.raises(InvalidMessageDeduplicationIdError):
await fifo_sqs_broker.kick(broker_message)


async def test_when_multiple_messages_kicked_to_same_group__then_order_is_preserved(
fifo_sqs_broker: SQSBroker,
sqs_client: capo_sqs.AsyncSQSClient,
fifo_sqs_queue: str,
) -> None:
for i in range(3):
message = BrokerMessage(
task_id=f"task-{i}",
task_name="ordered_task",
message=f"message-{i}".encode(),
labels={
SQS_MESSAGE_GROUP_ID_LABEL: "same-group",
SQS_MESSAGE_DEDUPLICATION_ID_LABEL: f"dedup-{i}",
},
)
await fifo_sqs_broker.kick(message)

response = await sqs_client.receive_message(queue_url=fifo_sqs_queue, max_number_of_messages=3)
bodies = [message.get("body") for message in response.get("messages", [])]
assert bodies == ["message-0", "message-1", "message-2"]
Loading