From 995070742c677927863c3041dcb95d5142e0cb4c Mon Sep 17 00:00:00 2001 From: Ash Berlin-Taylor Date: Thu, 13 Mar 2025 14:51:01 +0000 Subject: [PATCH] Show task logs in KubeExecutor on stdout This is needed so that when the KubeExecutor is asked to get logs for a running pod they show up in the pod. This changes the `airflow.sdk.execution_time.execute_workload` entrypoint to a. produce all logs as JSON, and b. ask the SDK to send all output log messages to the top level logger to, so they appear on stdout. In order to make the output a bit nicer this also tidies up/removes some of the logging from dispose_orm so that it doesn't "pollute" the logs (this was required as due to the current hack we have to upload remote logs, we ended up alling dispose_orm at the end and that _wasn't_ JSON formatted). Closes #46894 --- airflow/settings.py | 16 ++++--- .../sdk/execution_time/execute_workload.py | 12 +++-- .../airflow/sdk/execution_time/supervisor.py | 45 +++++++++++++------ .../execution_time/test_supervisor.py | 8 ++-- tests/dag_processing/test_manager.py | 2 +- tests/jobs/test_triggerer_job.py | 2 +- 6 files changed, 58 insertions(+), 27 deletions(-) diff --git a/airflow/settings.py b/airflow/settings.py index d73888d8ac648..40824c463deb2 100644 --- a/airflow/settings.py +++ b/airflow/settings.py @@ -389,8 +389,13 @@ def _session_maker(_engine): Session = scoped_session(NonScopedSession) # https://docs.sqlalchemy.org/en/20/core/pooling.html#using-connection-pools-with-multiprocessing-or-os-fork - os.register_at_fork(after_in_child=lambda: engine.dispose(close=False)) - os.register_at_fork(after_in_child=lambda: async_engine.sync_engine.dispose(close=False)) + def clean_in_fork(): + if engine: + engine.dispose(close=False) + if async_engine: + async_engine.sync_engine.dispose(close=False) + + os.register_at_fork(after_in_child=clean_in_fork) DEFAULT_ENGINE_ARGS = { @@ -480,15 +485,16 @@ def prepare_engine_args(disable_connection_pool=False, pool_class=None): return engine_args -def dispose_orm(): +def dispose_orm(do_log: bool = True): """Properly close pooled database connections.""" global Session, engine, NonScopedSession _globals = globals() - if "engine" not in _globals and "Session" not in _globals: + if _globals.get("engine") is None and _globals.get("Session") is None: return - log.debug("Disposing DB connection pool (PID %s)", os.getpid()) + if do_log: + log.debug("Disposing DB connection pool (PID %s)", os.getpid()) if "Session" in _globals and Session is not None: from sqlalchemy.orm.session import close_all_sessions diff --git a/task-sdk/src/airflow/sdk/execution_time/execute_workload.py b/task-sdk/src/airflow/sdk/execution_time/execute_workload.py index cc7021794b0e9..5fd9d6669b763 100644 --- a/task-sdk/src/airflow/sdk/execution_time/execute_workload.py +++ b/task-sdk/src/airflow/sdk/execution_time/execute_workload.py @@ -42,16 +42,19 @@ def execute_workload(input: str) -> None: from airflow.executors import workloads from airflow.sdk.execution_time.supervisor import supervise from airflow.sdk.log import configure_logging + from airflow.settings import dispose_orm - configure_logging(output=sys.stdout.buffer) + dispose_orm(do_log=False) + + configure_logging(output=sys.stdout.buffer, enable_pretty_log=False) decoder = TypeAdapter[workloads.All](workloads.All) workload = decoder.validate_json(input) if not isinstance(workload, workloads.ExecuteTask): - raise ValueError(f"KubernetesExecutor does not know how to handle {type(workload)}") + raise ValueError(f"We do not know how to handle {type(workload)}") - log.info("Executing workload in Kubernetes", workload=workload) + log.info("Executing workload", workload=workload) supervise( # This is the "wrong" ti type, but it duck types the same. TODO: Create a protocol for this. @@ -61,6 +64,9 @@ def execute_workload(input: str) -> None: token=workload.token, server=conf.get("core", "execution_api_server_url"), log_path=workload.log_path, + # Include the output of the task to stdout too, so that in process logs can be read from via the + # kubeapi as pod logs. + subprocess_logs_to_stdout=True, ) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 36986a3a5f270..e7af31635eef6 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -369,7 +369,10 @@ class WatchedSubprocess: selector: selectors.BaseSelector = attrs.field(factory=selectors.DefaultSelector) - log: FilteringBoundLogger + process_log: FilteringBoundLogger + + subprocess_logs_to_stdout: bool = False + """Duplicate log messages to stdout, or only send them to ``self.process_log``.""" @classmethod def start( @@ -425,7 +428,7 @@ def start( stdin=feed_stdin, process=psutil.Process(pid), requests_fd=requests_fd, - log=logger, + process_log=logger, **constructor_kwargs, ) @@ -445,17 +448,22 @@ def _register_pipe_readers(self, stdout: socket, stderr: socket, requests: socke # alternatives are used automatically) -- this is a way of having "event-based" code, but without # needing full async, to read and process output from each socket as it is received. - self.selector.register(stdout, selectors.EVENT_READ, self._create_socket_handler(self.log, "stdout")) + target_loggers: tuple[FilteringBoundLogger, ...] = (self.process_log,) + if self.subprocess_logs_to_stdout: + target_loggers += (log,) + self.selector.register( + stdout, selectors.EVENT_READ, self._create_socket_handler(target_loggers, channel="stdout") + ) self.selector.register( stderr, selectors.EVENT_READ, - self._create_socket_handler(self.log, "stderr", log_level=logging.ERROR), + self._create_socket_handler(target_loggers, channel="stderr", log_level=logging.ERROR), ) self.selector.register( logs, selectors.EVENT_READ, make_buffered_socket_reader( - process_log_messages_from_subprocess(self.log), on_close=self._on_socket_closed + process_log_messages_from_subprocess(target_loggers), on_close=self._on_socket_closed ), ) self.selector.register( @@ -464,10 +472,10 @@ def _register_pipe_readers(self, stdout: socket, stderr: socket, requests: socke make_buffered_socket_reader(self.handle_requests(log), on_close=self._on_socket_closed), ) - def _create_socket_handler(self, logger, channel, log_level=logging.INFO) -> Callable[[socket], bool]: + def _create_socket_handler(self, loggers, channel, log_level=logging.INFO) -> Callable[[socket], bool]: """Create a socket handler that forwards logs to a logger.""" return make_buffered_socket_reader( - forward_to_log(logger.bind(chan=channel), level=log_level), on_close=self._on_socket_closed + forward_to_log(loggers, chan=channel, level=log_level), on_close=self._on_socket_closed ) def _on_socket_closed(self): @@ -746,7 +754,7 @@ def _upload_logs(self): if self._what else {} ) - upload_to_remote(self.log, log_meta_dict) + upload_to_remote(self.process_log, log_meta_dict) def _monitor_subprocess(self): """ @@ -976,7 +984,9 @@ def cb(sock: socket): return cb -def process_log_messages_from_subprocess(log: FilteringBoundLogger) -> Generator[None, bytes, None]: +def process_log_messages_from_subprocess( + loggers: tuple[FilteringBoundLogger, ...], +) -> Generator[None, bytes, None]: from structlog.stdlib import NAME_TO_LEVEL while True: @@ -1003,10 +1013,16 @@ def process_log_messages_from_subprocess(log: FilteringBoundLogger) -> Generator if exc := event.pop("exception", None): # TODO: convert the dict back to a pretty stack trace event["error_detail"] = exc - log.log(NAME_TO_LEVEL[event.pop("level")], event.pop("event", None), **event) + + level = NAME_TO_LEVEL[event.pop("level")] + msg = event.pop("event", None) + for target in loggers: + target.log(level, msg, **event) -def forward_to_log(target_log: FilteringBoundLogger, level: int) -> Generator[None, bytes, None]: +def forward_to_log( + target_loggers: tuple[FilteringBoundLogger, ...], chan: str, level: int +) -> Generator[None, bytes, None]: while True: buf = yield line = bytes(buf) @@ -1014,10 +1030,10 @@ def forward_to_log(target_log: FilteringBoundLogger, level: int) -> Generator[No line = line.rstrip() try: msg = line.decode("utf-8", errors="replace") - target_log.log(level, msg) except UnicodeDecodeError: msg = line.decode("ascii", errors="replace") - target_log.log(level, msg) + for log in target_loggers: + log.log(level, msg, chan=chan) def supervise( @@ -1029,6 +1045,7 @@ def supervise( server: str | None = None, dry_run: bool = False, log_path: str | None = None, + subprocess_logs_to_stdout: bool = False, client: Client | None = None, ) -> int: """ @@ -1041,6 +1058,7 @@ def supervise( :param server: Base URL of the API server. :param dry_run: If True, execute without actual task execution (simulate run). :param log_path: Path to write logs, if required. + :param subprocess_logs_to_stdout: Should task logs also be sent to stdout via the main logger. :param client: Optional preconfigured client for communication with the server (Mostly for tests). :return: Exit code of the process. """ @@ -1081,6 +1099,7 @@ def supervise( client=client, logger=logger, bundle_info=bundle_info, + subprocess_logs_to_stdout=subprocess_logs_to_stdout, ) exit_code = process.wait() 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 97e49ffc63947..965164d642ab3 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py +++ b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py @@ -516,7 +516,7 @@ def test_heartbeat_failures_handling(self, monkeypatch, mocker, captured_logs, t mock_kill = mocker.patch("airflow.sdk.execution_time.supervisor.WatchedSubprocess.kill") proc = ActivitySubprocess( - log=mocker.MagicMock(), + process_log=mocker.MagicMock(), id=TI_ID, pid=mock_process.pid, stdin=mocker.MagicMock(), @@ -606,7 +606,7 @@ def test_overtime_handling( monkeypatch.setattr(ActivitySubprocess, "TASK_OVERTIME_THRESHOLD", overtime_threshold) mock_watched_subprocess = ActivitySubprocess( - log=mocker.MagicMock(), + process_log=mocker.MagicMock(), id=TI_ID, pid=12345, stdin=mocker.Mock(), @@ -751,7 +751,7 @@ def mock_process(self, mocker): @pytest.fixture def watched_subprocess(self, mocker, mock_process): proc = ActivitySubprocess( - log=mocker.MagicMock(), + process_log=mocker.MagicMock(), id=TI_ID, pid=12345, stdin=mocker.Mock(), @@ -937,7 +937,7 @@ class TestHandleRequest: def watched_subprocess(self, mocker): """Fixture to provide a WatchedSubprocess instance.""" return ActivitySubprocess( - log=mocker.MagicMock(), + process_log=mocker.MagicMock(), id=TI_ID, pid=12345, stdin=BytesIO(), diff --git a/tests/dag_processing/test_manager.py b/tests/dag_processing/test_manager.py index fb3f03bd9b8ef..c65380882c4af 100644 --- a/tests/dag_processing/test_manager.py +++ b/tests/dag_processing/test_manager.py @@ -138,7 +138,7 @@ def mock_processor(self) -> DagFileProcessorProcess: proc.create_time.return_value = time.time() proc.wait.return_value = 0 ret = DagFileProcessorProcess( - log=MagicMock(), + process_log=MagicMock(), id=uuid7(), pid=1234, process=proc, diff --git a/tests/jobs/test_triggerer_job.py b/tests/jobs/test_triggerer_job.py index 9471789b2a8c2..0804068c783f7 100644 --- a/tests/jobs/test_triggerer_job.py +++ b/tests/jobs/test_triggerer_job.py @@ -149,7 +149,7 @@ def builder(job=None): process = mocker.Mock(spec=psutil.Process, pid=10 * job.id + 1) proc = TriggerRunnerSupervisor( - log=mocker.Mock(), + process_log=mocker.Mock(), id=job.id, job=job, pid=process.pid,