Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,7 @@ def __init__(
model: BatchSqsTypeModel | None = None,
skip_group_on_error: bool = False,
logger: logging.Logger | None = None,
raise_on_entire_batch_failure: bool = True,
):
"""
Initialize the SqsFifoProcessor.
Expand All @@ -84,12 +85,20 @@ def __init__(
Default is False.
logger: logging.Logger | None
Optional Logger instance to output warnings with tracebacks for failed records.
raise_on_entire_batch_failure: bool
Raise an exception when the entire batch has failed processing.
When set to False, partial failures are reported in the response.

"""
self._skip_group_on_error: bool = skip_group_on_error
self._current_group_id = None
self._failed_group_ids: set[str] = set()
super().__init__(EventType.SQS, model, logger=logger)
super().__init__(
EventType.SQS,
model,
raise_on_entire_batch_failure=raise_on_entire_batch_failure,
logger=logger,
)

def _process_record(self, record):
self._current_group_id = record.get("attributes", {}).get("MessageGroupId")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -500,6 +500,24 @@ def lambda_handler(event, context):
assert result["batchItemFailures"][1]["itemIdentifier"] == third_record.message_id


def test_sqs_fifo_batch_processor_not_raise_when_entire_batch_fails(sqs_event_fifo_factory, record_handler):
first_record = SQSRecord(sqs_event_fifo_factory("fail"))
second_record = SQSRecord(sqs_event_fifo_factory("success"))
event = {"Records": [first_record.raw_event, second_record.raw_event]}

processor = SqsFifoPartialProcessor(raise_on_entire_batch_failure=False)

@batch_processor(record_handler=record_handler, processor=processor)
def lambda_handler(event, context):
return processor.response()

response = lambda_handler(event, {})

assert len(response["batchItemFailures"]) == 2
assert response["batchItemFailures"][0]["itemIdentifier"] == first_record.message_id
assert response["batchItemFailures"][1]["itemIdentifier"] == second_record.message_id


def test_sqs_fifo_batch_processor_middleware_with_skip_group_on_error(sqs_event_fifo_factory, record_handler):
# GIVEN a batch of 5 records with 3 different MessageGroupID
first_record = SQSRecord(sqs_event_fifo_factory("success", "1"))
Expand Down Expand Up @@ -909,7 +927,8 @@ def handler(record):
assert len(warning_records) == 0, "Expected no WARNING logs when logger is None"


def test_sqs_fifo_circuit_breaker_does_not_log(sqs_event_fifo_factory, caplog):
@pytest.mark.parametrize("positional_logger", [False, True])
def test_sqs_fifo_circuit_breaker_does_not_log(sqs_event_fifo_factory, caplog, positional_logger):
failing_record = sqs_event_fifo_factory("fail", "group-1")
short_circuited_record = sqs_event_fifo_factory("would-succeed", "group-1")

Expand All @@ -919,8 +938,10 @@ def handler(record):
return record["body"]

test_logger = logging.getLogger("test_logger")
processor = SqsFifoPartialProcessor(logger=test_logger)
processor.raise_on_entire_batch_failure = False
if positional_logger:
processor = SqsFifoPartialProcessor(None, False, test_logger, raise_on_entire_batch_failure=False)
else:
processor = SqsFifoPartialProcessor(logger=test_logger, raise_on_entire_batch_failure=False)

with caplog.at_level(logging.WARNING, logger="test_logger"):
process_partial_response(
Expand Down