diff --git a/airflow/providers/apache/kafka/operators/consume.py b/airflow/providers/apache/kafka/operators/consume.py index 02c9db6556df6..6f1c1ee61e50c 100644 --- a/airflow/providers/apache/kafka/operators/consume.py +++ b/airflow/providers/apache/kafka/operators/consume.py @@ -161,7 +161,8 @@ def execute(self, context) -> Any: batch_size = self.max_batch_size msgs = consumer.consume(num_messages=batch_size, timeout=self.poll_timeout) - messages_left -= len(msgs) + if not self.read_to_end: + messages_left -= len(msgs) if not msgs: # No messages + messages_left is being used. self.log.info("Reached end of log. Exiting.") diff --git a/tests/integration/providers/apache/kafka/operators/test_consume.py b/tests/integration/providers/apache/kafka/operators/test_consume.py index 240b02f9be1ad..861bbbcfab35d 100644 --- a/tests/integration/providers/apache/kafka/operators/test_consume.py +++ b/tests/integration/providers/apache/kafka/operators/test_consume.py @@ -65,7 +65,7 @@ def setup_method(self): extra=json.dumps( { "socket.timeout.ms": 10, - "bootstrap.servers": "localhost:9092", + "bootstrap.servers": "broker:29092", "group.id": f"operator.consumer.test.integration.test_{num}", "enable.auto.commit": False, "auto.offset.reset": "beginning", @@ -135,7 +135,7 @@ def test_consumer_operator_test_3(self): operator = ConsumeFromTopicOperator( kafka_config_id=TOPIC, topics=[TOPIC], - apply_function=_batch_tester, + apply_function_batch=_batch_tester, apply_function_kwargs={"test_string": TOPIC}, task_id="test", poll_timeout=0.0001, diff --git a/tests/providers/apache/kafka/operators/test_consume.py b/tests/providers/apache/kafka/operators/test_consume.py index 178419052699c..0ce1ce6e32eb9 100644 --- a/tests/providers/apache/kafka/operators/test_consume.py +++ b/tests/providers/apache/kafka/operators/test_consume.py @@ -19,6 +19,7 @@ import json import logging from typing import Any +from unittest import mock from airflow.models import Connection @@ -79,3 +80,31 @@ def test_operator_callable(self): # execute the operator (this is essentially a no op as the broker isn't setup) operator.execute(context={}) + + @mock.patch("airflow.providers.apache.kafka.hooks.consume.KafkaConsumerHook.get_consumer") + def test_operator_consume_max(self, mock_get_consumer): + mock_consumer = mock.MagicMock() + + mocked_messages = ["test_messages" for i in range(1001)] + + def mock_consume(num_messages=0, timeout=-1): + nonlocal mocked_messages + if num_messages < 0: + raise Exception("Number of messages needs to be positive") + msg_count = min(num_messages, len(mocked_messages)) + returned_messages = mocked_messages[:msg_count] + mocked_messages = mocked_messages[msg_count:] + return returned_messages + + mock_consumer.consume = mock_consume + mock_get_consumer.return_value = mock_consumer + + operator = ConsumeFromTopicOperator( + kafka_config_id="kafka_d", + topics=["test"], + task_id="test", + poll_timeout=0.0001, + ) + + # execute the operator (this is essentially a no op as we're mocking the consumer) + operator.execute(context={})