Skip to content
Closed
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 @@ -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
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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]}"
Expand Down Expand Up @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from __future__ import annotations

import json
import threading
from collections.abc import Callable
from unittest.mock import MagicMock, patch

Expand Down Expand Up @@ -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."""
Expand Down