diff --git a/providers/apache/kafka/src/airflow/providers/apache/kafka/plugins/event_producer.py b/providers/apache/kafka/src/airflow/providers/apache/kafka/plugins/event_producer.py index 20f6f1f6e449e..19bfba9cdc1fd 100644 --- a/providers/apache/kafka/src/airflow/providers/apache/kafka/plugins/event_producer.py +++ b/providers/apache/kafka/src/airflow/providers/apache/kafka/plugins/event_producer.py @@ -21,6 +21,7 @@ import logging import os import time +from concurrent.futures import ThreadPoolExecutor from datetime import datetime, timezone from fnmatch import fnmatch from functools import lru_cache @@ -178,21 +179,34 @@ def _reset_state_after_fork() -> None: os.register_at_fork(after_in_child=_reset_state_after_fork) +def _build_producer() -> Producer: + """Resolve the Kafka connection and build a producer from it.""" + # Deferred import: KafkaProducerHook pulls in confluent_kafka.Producer, which is heavy + # and shouldn't be loaded by every Airflow process at plugin import time. + from airflow.providers.apache.kafka.hooks.produce import KafkaProducerHook + + hook_kwargs: dict[str, Any] = {} + if _get_kafka_config_id(): + hook_kwargs["kafka_config_id"] = _get_kafka_config_id() + return KafkaProducerHook(**hook_kwargs).get_producer() + + def _get_producer() -> Producer | None: """Build (once) and return the plugin's Kafka producer, or ``None`` on init failure.""" global _producer if _producer is not None: return _producer - # Deferred import: KafkaProducerHook pulls in confluent_kafka.Producer, which is heavy - # and shouldn't be loaded by every Airflow process at plugin import time. - from airflow.providers.apache.kafka.hooks.produce import KafkaProducerHook - try: - hook_kwargs: dict[str, Any] = {} - if _get_kafka_config_id(): - hook_kwargs["kafka_config_id"] = _get_kafka_config_id() - producer = KafkaProducerHook(**hook_kwargs).get_producer() + # The connection lookup opens a DB session, so it runs on a thread of its own. + # + # settings.Session gives out one session per thread. MetastoreBackend.get_connection + # carries @provide_session, so on the caller's thread it reuses the session that + # thread already has, and closes it on the way out. The scheduler calls the listeners + # in this module while it is moving queued DagRuns to RUNNING. Closing its session + # detaches those DagRuns, and the scheduler crashes the next time it reads one. + with ThreadPoolExecutor(max_workers=1, thread_name_prefix="kafka-event-producer") as pool: + producer = pool.submit(_build_producer).result() except Exception as exc: log.warning("Kafka event producer: failed to initialize producer (%s).", exc) return None diff --git a/providers/apache/kafka/tests/integration/apache/kafka/plugins/test_event_producer.py b/providers/apache/kafka/tests/integration/apache/kafka/plugins/test_event_producer.py index 81cbc4ca931e3..8071d0a0c1e8d 100644 --- a/providers/apache/kafka/tests/integration/apache/kafka/plugins/test_event_producer.py +++ b/providers/apache/kafka/tests/integration/apache/kafka/plugins/test_event_producer.py @@ -64,6 +64,8 @@ class TestEventProducer: dag_folder = os.path.join(test_dir, "dags") KAFKA_CONFIG_ID = "kafka_default" + # This config id lives in the metadata DB only. + DB_KAFKA_CONFIG_ID = "kafka_event_producer_db" # Use a unique topic per run to avoid errors on a re-run, in case # the previous teardown hasn't finished with the topic deletion. TOPIC = f"airflow.events.itest.{uuid.uuid4().hex[:8]}" @@ -128,6 +130,75 @@ def start_components(self): terminate_process(scheduler_process) terminate_process(apiserver_process) + @pytest.fixture + def kafka_connection_in_metadata_db(self): + """Give the plugin a Kafka connection that only the metadata DB has. + + setup_class exports AIRFLOW_CONN_KAFKA_DEFAULT, and the environment secrets backend + answers from that variable without opening a DB session. This test needs the lookup + to reach the metastore backend instead, and that only happens for a connection the + environment does not have. + + The topic is separate too. The other test counts the messages on the class topic, + and the runs here would add to that count. + """ + topic = f"airflow.events.itest.{uuid.uuid4().hex[:8]}" + self._admin.create_topic([(topic, 1, 1)]) + + connection = Connection( + conn_id=self.DB_KAFKA_CONFIG_ID, + conn_type="kafka", + extra=json.dumps(client_config), + ) + subprocess.run( + ["airflow", "connections", "add", self.DB_KAFKA_CONFIG_ID, "--conn-json", connection.as_json()], + check=True, + env=os.environ.copy(), + ) + os.environ["AIRFLOW__KAFKA_EVENT_PRODUCER__KAFKA_CONFIG_ID"] = self.DB_KAFKA_CONFIG_ID + os.environ["AIRFLOW__KAFKA_EVENT_PRODUCER__TOPIC"] = topic + try: + yield + finally: + os.environ["AIRFLOW__KAFKA_EVENT_PRODUCER__TOPIC"] = self.TOPIC + os.environ.pop("AIRFLOW__KAFKA_EVENT_PRODUCER__KAFKA_CONFIG_ID", None) + subprocess.run( + ["airflow", "connections", "delete", self.DB_KAFKA_CONFIG_ID], + check=False, + env=os.environ.copy(), + ) + try: + self._admin.delete_topic([topic]) + except Exception as exc: + log.warning("teardown: failed to delete topic %r: %s", topic, exc) + + @pytest.mark.execution_timeout(300) + def test_scheduler_starts_every_queued_dag_run(self, kafka_connection_in_metadata_db): + """The scheduler starts every queued Dag run, not only the first one. + + Both runs are triggered before the scheduler starts, so a single pass of + _start_queued_dagruns has more than one row to get through. That pass is holding a + DB session. If building the Kafka producer closes it, the rows the pass has not + reached yet are detached, and the scheduler exits. + """ + dag_id = "demo_dag" + run_ids = [unpause_trigger_dag_and_get_run_id(dag_id=dag_id) for _ in range(2)] + + scheduler_process, apiserver_process = start_scheduler() + try: + states = [wait_for_dag_run(dag_id=dag_id, run_id=run_id, max_wait_time=60) for run_id in run_ids] + scheduler_exit_code = scheduler_process.poll() + finally: + terminate_process(scheduler_process) + terminate_process(apiserver_process) + + assert scheduler_exit_code is None, ( + f"the scheduler exited while starting queued Dag runs, exit code {scheduler_exit_code}" + ) + assert states == [State.SUCCESS, State.SUCCESS], ( + f"the scheduler did not finish both queued Dag runs, final states: {states}" + ) + @pytest.mark.execution_timeout(90) def test_dag_run_produces_event_messages(self, start_components): consumer = KafkaConsumerHook(topics=[self.TOPIC], kafka_config_id=self.KAFKA_CONFIG_ID).get_consumer() diff --git a/providers/apache/kafka/tests/unit/apache/kafka/plugins/test_event_producer.py b/providers/apache/kafka/tests/unit/apache/kafka/plugins/test_event_producer.py index 8e4ed0fcf4ee3..cd094009f4627 100644 --- a/providers/apache/kafka/tests/unit/apache/kafka/plugins/test_event_producer.py +++ b/providers/apache/kafka/tests/unit/apache/kafka/plugins/test_event_producer.py @@ -17,6 +17,7 @@ from __future__ import annotations import json +import threading from collections.abc import Callable from unittest.mock import MagicMock, patch @@ -322,6 +323,32 @@ def test_hook_defaults_when_options_unset(self): # hook's own defaults (the "kafka_default" connection) apply. hook_cls_mock.assert_called_once_with() + def test_producer_is_not_built_on_the_calling_thread(self): + """Building the producer must not run on the thread that called the listener. + + Building it resolves the Kafka connection, and settings.Session gives out one + session per thread. The scheduler is what calls these listeners. On its thread + the lookup would reuse the scheduler's own session and close it, detaching the + rows the scheduler is working on. + """ + calling_thread = threading.get_ident() + lookup_thread = {} + + def record_thread(): + lookup_thread["ident"] = threading.get_ident() + return MagicMock() + + # get_producer is where the real hook resolves the connection. + hook_mock = MagicMock() + hook_mock.get_producer.side_effect = record_thread + + with patch(_PRODUCER_HOOK_CLS, return_value=hook_mock): + assert event_producer._get_producer() is not None + + assert lookup_thread["ident"] != calling_thread, ( + "the Kafka connection was resolved on the caller's thread" + ) + class TestCheckTopicExists: """Tests for the topic check and its retry-after-failure cooldown."""