diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index ca7539131f60e..497af0d9509e2 100644 --- a/airflow-core/src/airflow/dag_processing/processor.py +++ b/airflow-core/src/airflow/dag_processing/processor.py @@ -733,6 +733,7 @@ def wait(self) -> int: raise NotImplementedError(f"Don't call wait on {type(self).__name__} objects") def close(self): + self.cleanup_sockets_after_kill() try: self.logger_filehandle.close() except OSError: diff --git a/airflow-core/tests/unit/dag_processing/test_manager.py b/airflow-core/tests/unit/dag_processing/test_manager.py index 74da895bca30a..139ccf430cb8d 100644 --- a/airflow-core/tests/unit/dag_processing/test_manager.py +++ b/airflow-core/tests/unit/dag_processing/test_manager.py @@ -23,6 +23,7 @@ import os import random import re +import selectors import shutil import signal import textwrap @@ -1187,6 +1188,72 @@ def test_create_process_subprocess_logs_to_stdout( _, kwargs = mock_start.call_args assert kwargs["subprocess_logs_to_stdout"] is expected_subprocess_logs_to_stdout + def test_terminate_orphan_processes_kills_then_closes_processor(self): + manager = DagFileProcessorManager(max_runs=1) + processor, _ = self.mock_processor() + file_info = DagFileInfo( + bundle_name="testing", rel_path=Path("removed.py"), bundle_path=TEST_DAGS_FOLDER + ) + manager._processors = {file_info: processor} + + call_order: list[str] = [] + processor.close = mock.Mock(side_effect=lambda: call_order.append("close")) + + with mock.patch.object( + type(processor), "kill", side_effect=lambda *_args, **_kwargs: call_order.append("kill") + ): + manager.terminate_orphan_processes(present=set()) + + assert call_order == ["kill", "close"] + + def test_terminate_orphan_processes_does_not_dispatch_request_frames_after_kill(self): + manager = DagFileProcessorManager(max_runs=1) + processor, _ = self.mock_processor() + request_sock, request_peer = socketpair() + real_selector = selectors.DefaultSelector() + try: + processor.selector = real_selector + processor._open_sockets[request_sock] = "requests" + + file_info = DagFileInfo( + bundle_name="testing", rel_path=Path("removed.py"), bundle_path=TEST_DAGS_FOLDER + ) + manager._processors = {file_info: processor} + + request_handler = mock.Mock(return_value=False) + + def on_close(sock): + real_selector.unregister(sock) + + real_selector.register(request_sock, selectors.EVENT_READ, (request_handler, on_close)) + + with mock.patch.object(type(processor), "kill"): + manager.terminate_orphan_processes(present=set()) + + request_handler.assert_not_called() + with pytest.raises((KeyError, ValueError)): + real_selector.get_key(request_sock) + finally: + real_selector.close() + request_peer.close() + + def test_kill_timed_out_processors_kills_then_closes_processor(self): + manager = DagFileProcessorManager(max_runs=1, processor_timeout=5) + start_time = time.monotonic() - manager.processor_timeout - 1 + processor, _ = self.mock_processor(start_time=start_time) + file_info = DagFileInfo(bundle_name="testing", rel_path=Path("abc.txt"), bundle_path=TEST_DAGS_FOLDER) + manager._processors = {file_info: processor} + + call_order: list[str] = [] + processor.close = mock.Mock(side_effect=lambda: call_order.append("close")) + + with mock.patch.object( + type(processor), "kill", side_effect=lambda *_args, **_kwargs: call_order.append("kill") + ): + manager._kill_timed_out_processors() + + assert call_order == ["kill", "close"] + def test_kill_timed_out_processors_kill(self): manager = DagFileProcessorManager(max_runs=1, processor_timeout=5) # Set start_time to ensure timeout occurs: start_time = current_time - (timeout + 1) = always (timeout + 1) seconds diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 447f7c7594de6..908f66dcec52c 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -1007,6 +1007,46 @@ def _cleanup_open_sockets(self): self.selector.close() self.stdin.close() + def cleanup_sockets_after_kill(self) -> None: + """Drain log-bearing sockets, then close every remaining socket after a forced kill.""" + for sock, socket_type in list(self._open_sockets.items()): + try: + key = self.selector.get_key(sock) + except KeyError: + key = None + + if key is not None: + socket_handler, on_close = key.data + try: + if socket_type != "requests": + sock.setblocking(False) + while True: + try: + if not socket_handler(sock): + break + except (BlockingIOError, InterruptedError, OSError): + break + + if on_close is not None: + on_close(sock) + else: + with suppress(KeyError): + self.selector.unregister(sock) + self._open_sockets.pop(sock, None) + except Exception: + log.exception( + "Failed to clean up killed subprocess socket", + pid=self.pid, + socket_type=socket_type, + ) + with suppress(KeyError): + self.selector.unregister(sock) + self._open_sockets.pop(sock, None) + with suppress(OSError, ValueError): + sock.close() + + self._open_sockets.clear() + def kill( self, signal_to_send: signal.Signals = signal.SIGINT, @@ -2350,7 +2390,17 @@ def process_log_messages_from_subprocess( if level := NAME_TO_LEVEL.get(event.pop("level")): msg = event.pop("event", None) for target in loggers: - target.log(level, msg, **event) + _log_to_target(target, level, msg, **event) + + +def _log_to_target(target: FilteringBoundLogger, level: int, msg: str | None, **event) -> None: + rendered_msg = msg if msg is not None else "" + try: + target.log(level, rendered_msg, **event) + except ValueError as e: + if "closed file" not in str(e): + raise + log.debug("Dropped log line for closed logger handle", level=level, logger=event.get("logger")) def forward_to_log( @@ -2365,7 +2415,7 @@ def forward_to_log( except UnicodeDecodeError: msg = line.decode("ascii", errors="replace") for log in target_loggers: - log.log(level, msg, logger=logger) + _log_to_target(log, level, msg, logger=logger) def ensure_secrets_backend_loaded() -> list[BaseSecretsBackend]: diff --git a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py index 870b2e6deedca..4be9e3b2e9b38 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py +++ b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py @@ -49,6 +49,7 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from opentelemetry.trace import get_current_span from pytest_unordered import unordered +from structlog.typing import FilteringBoundLogger from task_sdk import FAKE_BUNDLE, make_client from uuid6 import uuid7 @@ -166,6 +167,7 @@ ProcessTracker, _make_process_nondumpable, _remote_logging_conn, + forward_to_log, in_process_api_server, make_buffered_socket_reader, process_log_messages_from_subprocess, @@ -3946,6 +3948,120 @@ def test_process_log_messages_from_subprocess(monkeypatch, caplog): ] +@pytest.mark.parametrize( + "error_message", + ["write to closed file", "I/O operation on closed file"], +) +def test_process_log_messages_closed_logger_is_skipped(error_message): + closed_logger = mock.Mock(spec=FilteringBoundLogger) + closed_logger.log.side_effect = ValueError(error_message) + + good_logger = mock.Mock(spec=FilteringBoundLogger) + + def fake_reconfigure(logger, *args, **kwargs): + return logger + + with ( + mock.patch( + "airflow.sdk.execution_time.supervisor.reconfigure_logger", + side_effect=fake_reconfigure, + ), + mock.patch.object(supervisor.log, "debug") as mock_debug, + ): + gen = process_log_messages_from_subprocess(loggers=(closed_logger, good_logger)) + next(gen) + + gen.send(b'{"level": "info", "event": "hello"}\n') + gen.send(b'{"level": "info", "event": "world"}\n') + + assert good_logger.log.call_count == 2 + assert mock_debug.call_count == 2 + + +def test_forward_to_log_closed_logger_is_skipped(): + closed_logger = mock.Mock(spec=FilteringBoundLogger) + closed_logger.log.side_effect = ValueError("I/O operation on closed file") + good_logger = mock.Mock(spec=FilteringBoundLogger) + + with mock.patch.object(supervisor.log, "debug") as mock_debug: + gen = forward_to_log((closed_logger, good_logger), logger="task.stdout", level=logging.INFO) + next(gen) + gen.send(b"hello\n") + gen.send(b"world\n") + + assert good_logger.log.call_count == 2 + good_logger.log.assert_any_call(logging.INFO, "hello", logger="task.stdout") + good_logger.log.assert_any_call(logging.INFO, "world", logger="task.stdout") + assert mock_debug.call_count == 2 + + +def test_process_log_messages_unexpected_value_error_is_reraised(): + """A ValueError unrelated to a closed file handle must propagate, not be silently swallowed.""" + buggy_logger = mock.Mock(spec=FilteringBoundLogger) + buggy_logger.log.side_effect = ValueError("unexpected formatting bug") + + def fake_reconfigure(log, *args, **kwargs): + return log + + with mock.patch( + "airflow.sdk.execution_time.supervisor.reconfigure_logger", + side_effect=fake_reconfigure, + ): + gen = process_log_messages_from_subprocess(loggers=(buggy_logger,)) + next(gen) + + with pytest.raises(ValueError, match="unexpected formatting bug"): + gen.send(b'{"level": "info", "event": "test"}\n') + + +def test_cleanup_sockets_after_kill_drains_logs_but_not_requests(mocker): + request_read, request_write = socket.socketpair() + stdout_read, stdout_write = socket.socketpair() + log_read, log_write = socket.socketpair() + + subprocess = ActivitySubprocess( + process_log=mocker.MagicMock(), + id=TI_ID, + pid=12345, + stdin=stdout_write, + client=mocker.Mock(), + process=mocker.Mock(), + ) + selector = selectors.DefaultSelector() + subprocess.selector = selector + + request_handler = mock.Mock(return_value=False) + stdout_handler = mock.Mock(return_value=False) + log_handler = mock.Mock(return_value=False) + + def on_close(sock): + selector.unregister(sock) + subprocess._open_sockets.pop(sock, None) + + try: + subprocess._open_sockets[request_read] = "requests" + subprocess._open_sockets[stdout_read] = "stdout" + subprocess._open_sockets[log_read] = "logs" + + selector.register(request_read, selectors.EVENT_READ, (request_handler, on_close)) + selector.register(stdout_read, selectors.EVENT_READ, (stdout_handler, on_close)) + selector.register(log_read, selectors.EVENT_READ, (log_handler, on_close)) + + subprocess.cleanup_sockets_after_kill() + + request_handler.assert_not_called() + stdout_handler.assert_called_once_with(stdout_read) + log_handler.assert_called_once_with(log_read) + assert not subprocess._open_sockets + with pytest.raises((KeyError, ValueError)): + selector.get_key(request_read) + finally: + selector.close() + request_write.close() + stdout_write.close() + log_write.close() + + def test_reinit_supervisor_comms(monkeypatch, client_with_ti_start, caplog): def subprocess_main(): # This is run in the subprocess!