This is an automated email from the ASF dual-hosted git repository. jason810496 pushed a commit to branch jason/lang-sdk-e2e/03e-lang-sdk-dag-bag in repository https://gitbox.apache.org/repos/asf/airflow.git
commit fa2b6523f71a490658e19070850ac2d19627c71d Author: ZHE YOU LIU <[email protected]> AuthorDate: Thu Oct 1 14:52:07 2026 +0000 Parse coordinator-claimed Dag files with their runtime in the Dag processor The Dag processor parses a file a coordinator's Dag importer claims with LangSDKDagFileProcessorProcess, in the new lang_sdk_processor module. Its Python child finds the coordinator, reports the runtime's schema version over fd 0 and execs the runtime, which connects back to two listeners the manager owns and answers the parse request. Each returned Dag gets its unset settings from the Airflow config and must pass validate_serialized_dag. A failed start, a missing result, an invalid frame or message, or an invalid Dag is an import error. A runtime still running 5s after its result is killed, and the result is kept. Callbacks for the file are dropped. The file's Dag source is read with its Dag importer. A Dag bag cannot parse such a file yet: its importer reports an import error, and sync_bag_to_db leaves the file's import errors to the Dag processor. --- airflow-core/src/airflow/dag_processing/dagbag.py | 19 +- .../src/airflow/dag_processing/importer_routing.py | 59 +++ .../airflow/dag_processing/lang_sdk_processor.py | 493 +++++++++++++++++++ airflow-core/src/airflow/dag_processing/manager.py | 29 +- .../tests/unit/dag_processing/fake_lang_sdk.py | 127 +++++ .../tests/unit/dag_processing/test_dagbag.py | 20 + .../unit/dag_processing/test_importer_routing.py | 38 ++ .../unit/dag_processing/test_lang_sdk_processor.py | 533 +++++++++++++++++++++ .../tests/unit/dag_processing/test_manager.py | 137 +++++- generated/known_sdk_imports_in_core.txt | 1 + .../src/airflow/sdk/coordinators/_dag_importer.py | 90 ++++ .../src/airflow/sdk/coordinators/_subprocess.py | 9 + .../task_sdk/coordinators/test_dag_importer.py | 70 +++ .../tests/task_sdk/coordinators/test_subprocess.py | 12 + 14 files changed, 1626 insertions(+), 11 deletions(-) diff --git a/airflow-core/src/airflow/dag_processing/dagbag.py b/airflow-core/src/airflow/dag_processing/dagbag.py index 295eb52536b..4e77a94e65d 100644 --- a/airflow-core/src/airflow/dag_processing/dagbag.py +++ b/airflow-core/src/airflow/dag_processing/dagbag.py @@ -33,6 +33,7 @@ from airflow import settings from airflow._shared.timezones import timezone from airflow.configuration import conf from airflow.dag_processing.bundles.local import LocalDagBundle +from airflow.dag_processing.importer_routing import get_claiming_coordinator from airflow.exceptions import ( AirflowClusterPolicyError, AirflowClusterPolicySkipDag, @@ -585,16 +586,30 @@ def sync_bag_to_db( version_data: dict[str, Any] | None = None, session: Session = NEW_SESSION, ) -> None: - """Save attributes about list of DAG to the DB.""" + """ + Save attributes about list of DAG to the DB. + + Files that a Lang-SDK runtime parses are left out, with their import errors: the Dag processor + stores those. + """ from airflow.dag_processing.collection import update_dag_parsing_results_in_db - import_errors = {(bundle_name, rel_path): error for rel_path, error in dagbag.import_errors.items()} + def is_parsed_by_runtime(rel_path: str) -> bool: + return get_claiming_coordinator(Path(dagbag.bundle_path or "", rel_path), bundle_name) is not None + + import_errors = { + (bundle_name, rel_path): error + for rel_path, error in dagbag.import_errors.items() + if not is_parsed_by_runtime(rel_path) + } # Build the set of all files that were parsed and include files with import errors # in case they are not in parsed_definitions files_parsed = set(import_errors) if dagbag.bundle_path: for rel_path in dagbag.parsed_definitions: + if is_parsed_by_runtime(rel_path): + continue files_parsed.add((bundle_name, rel_path)) # A definition nested in an archive also clears the archive's own discovery errors. if enclosing_file := find_enclosing_file(Path(dagbag.bundle_path, rel_path)): diff --git a/airflow-core/src/airflow/dag_processing/importer_routing.py b/airflow-core/src/airflow/dag_processing/importer_routing.py new file mode 100644 index 00000000000..956723eae26 --- /dev/null +++ b/airflow-core/src/airflow/dag_processing/importer_routing.py @@ -0,0 +1,59 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Route Dag files to the process that parses them, by the bundle's Dag importer registry.""" + +from __future__ import annotations + +import logging +import os +from pathlib import Path +from typing import TYPE_CHECKING + +from airflow.sdk.coordinators._dag_importer import CoordinatorDagImporter # noqa: SDK001 +from airflow.sdk.importers import DagImporterRegistry, get_importer_registry # noqa: SDK001 + +if TYPE_CHECKING: + from airflow.sdk.coordinators._subprocess import SubprocessCoordinator # noqa: SDK001 + +log = logging.getLogger(__name__) + + +def _get_registry(bundle_name: str | None) -> DagImporterRegistry | None: + try: + return get_importer_registry(bundle_name) + except Exception: + log.exception("Cannot build the Dag importer registry for bundle %s", bundle_name) + return None + + +def get_claiming_coordinator( + path: str | os.PathLike[str], bundle_name: str | None +) -> SubprocessCoordinator | None: + """ + Return the coordinator whose runtime parses ``path``, or ``None`` when a Python child parses it. + + A runtime parses the file when its importer is a coordinator's Dag importer. + """ + if (registry := _get_registry(bundle_name)) is None: + return None + try: + importer = registry.get_importer(Path(path)) + except Exception: + log.exception("Cannot load the Dag importer for %s", path) + return None + return importer.coordinator if isinstance(importer, CoordinatorDagImporter) else None diff --git a/airflow-core/src/airflow/dag_processing/lang_sdk_processor.py b/airflow-core/src/airflow/dag_processing/lang_sdk_processor.py new file mode 100644 index 00000000000..877af756ed6 --- /dev/null +++ b/airflow-core/src/airflow/dag_processing/lang_sdk_processor.py @@ -0,0 +1,493 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Parse a native Lang-SDK Dag file with the runtime of the coordinator whose Dag importer claims it.""" + +from __future__ import annotations + +import functools +import os +import selectors +import signal +import time +from pathlib import Path +from socket import socket +from typing import TYPE_CHECKING, Annotated, ClassVar, Literal, cast, get_args + +import attrs +import msgspec +from pydantic import BaseModel, Field, TypeAdapter +from uuid6 import uuid7 + +from airflow.dag_processing.importer_routing import get_claiming_coordinator +from airflow.dag_processing.processor import ( + BaseDagFileProcessorProcess, + DagFileParseRequest, + DagFileParsingResult, + ToManager, +) +from airflow.exceptions import DeserializationError +from airflow.sdk.coordinators._subprocess import _is_connection_from_pid, _start_server +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.comms import CommsDecoder, ErrorResponse, MaskSecret, _RequestFrame +from airflow.sdk.execution_time.supervisor import ( + ResponseSent, + length_prefixed_frame_reader, + make_buffered_socket_reader, + process_log_messages_from_subprocess, + register_request_method, +) +from airflow.sdk.importers import DagSourceCode, FilesystemDagDefinition +from airflow.serialization.serialized_objects import DagSerialization, LazyDeserializedDAG + +if TYPE_CHECKING: + from collections.abc import Generator + + from structlog.typing import FilteringBoundLogger + + from airflow.sdk.api.client import Client + from airflow.sdk.execution_time.supervisor import RequestHandler, RequestResult + from airflow.typing_compat import Self + +# How long a runtime may keep running after its parse result, as Node does while a handle stays open. +_EXIT_GRACE_PERIOD = 5.0 + + +# StartLangSDKRuntime and LangSDKRuntimeSchemaVersion pass only between the manager and its forked +# child before the exec, so they are not part of the supervisor schema the runtimes speak. + + +class StartLangSDKRuntime(BaseModel): + """Ask the parse child to exec the runtime that parses *file*.""" + + file: str + bundle_path: Path + bundle_name: str + dag_file_rel_path: str + comm_address: tuple[str, int] + logs_address: tuple[str, int] + type: Literal["StartLangSDKRuntime"] = "StartLangSDKRuntime" + + +class LangSDKRuntimeSchemaVersion(BaseModel): + """The supervisor schema version of the runtime the parse child is about to exec.""" + + schema_version: str | None + type: Literal["LangSDKRuntimeSchemaVersion"] = "LangSDKRuntimeSchemaVersion" + + +def _start_runtime_entrypoint() -> None: + """Exec the runtime that parses the file named by the start request, or report why it cannot start.""" + os.environ["_AIRFLOW_PROCESS_CONTEXT"] = "client" + # fd 0 becomes the runtime's stdin, so the request channel moves to a close-on-exec copy. + comms = CommsDecoder[StartLangSDKRuntime, LangSDKRuntimeSchemaVersion | DagFileParsingResult]( + socket=socket(fileno=os.dup(0)), + body_decoder=TypeAdapter(StartLangSDKRuntime), + ) + devnull = os.open(os.devnull, os.O_RDONLY) + os.dup2(devnull, 0) + os.close(devnull) + + msg = comms._get_response() + if not isinstance(msg, StartLangSDKRuntime): + raise RuntimeError(f"Required first message to be a StartLangSDKRuntime, it was {msg}") + + def report_schema_version(schema_version: str | None) -> None: + comms.send(LangSDKRuntimeSchemaVersion(schema_version=schema_version)) + + try: + coordinator = get_claiming_coordinator(msg.file, msg.bundle_name) + if coordinator is None: + raise RuntimeError(f"No coordinator's Dag importer claims {msg.file}") + coordinator.parse_dag( + path=Path(msg.file), + bundle_path=msg.bundle_path, + comm_address=msg.comm_address, + logs_address=msg.logs_address, + report_schema_version=report_schema_version, + ) + except Exception as e: + comms.send( + DagFileParsingResult( + fileloc=msg.file, + serialized_dags=[], + import_errors={ + msg.dag_file_rel_path: f"Cannot start the Lang-SDK runtime: {type(e).__name__}: {e}" + }, + ) + ) + + +_Channel = Literal["comm", "logs"] + + [email protected](kw_only=True) +class LangSDKDagFileProcessorProcess(BaseDagFileProcessorProcess): + """ + Parse a native Lang-SDK Dag file with its coordinator's runtime. + + The forked parse child finds the coordinator, reports the runtime's schema version and execs the + runtime. The runtime connects back to two listeners this process owns and answers the + ``DagFileParseRequest`` itself, so the request is sent once it has connected. + """ + + client: Client | None = None # type: ignore[assignment] + """Answers the runtime's requests; without one, as in a Dag bag, those that need it get an error.""" + + decoder = TypeAdapter( + Annotated[LangSDKRuntimeSchemaVersion | get_args(ToManager)[0], Field(discriminator="type")] + ) + + _listeners: dict[_Channel, socket] + _parse_request: DagFileParseRequest + _runtime_schema_version: str | None = attrs.field(default=None, init=False) + _schema_version_reported: bool = attrs.field(default=False, init=False) + _parsing_result_monotonic: float | None = attrs.field(default=None, init=False) + _unverified_connections: list[tuple[socket, _Channel]] = attrs.field(factory=list, init=False) + + @classmethod + def start( # type: ignore[override] + cls, + *, + path: str | os.PathLike[str], + bundle_path: Path, + bundle_name: str, + dag_file_rel_path: str, + **kwargs, + ) -> Self: + listeners: dict[_Channel, socket] = {"comm": _start_server(), "logs": _start_server()} + try: + for listener in listeners.values(): + listener.setblocking(False) + parse_request = DagFileParseRequest( + file=os.fspath(path), bundle_path=bundle_path, bundle_name=bundle_name + ) + proc = super().start( + target=_start_runtime_entrypoint, + use_exec=supervisor._should_use_exec(), + new_process_group=True, + bundle_name=bundle_name, + dag_file_rel_path=dag_file_rel_path, + listeners=listeners, + parse_request=parse_request, + **kwargs, + ) + except BaseException: + for listener in listeners.values(): + listener.close() + raise + for channel, listener in listeners.items(): + proc._open_sockets[listener] = f"{channel}-listener" + proc.selector.register( + listener, + selectors.EVENT_READ, + (functools.partial(proc._accept_connection, channel=channel), proc._on_socket_closed), + ) + proc.send_msg( + StartLangSDKRuntime( + file=parse_request.file, + bundle_path=bundle_path, + bundle_name=bundle_name, + dag_file_rel_path=dag_file_rel_path, + comm_address=listeners["comm"].getsockname()[:2], + logs_address=listeners["logs"].getsockname()[:2], + ), + request_id=0, + ) + return proc + + @classmethod + def run( + cls, + *, + path: str | os.PathLike[str], + bundle_path: Path, + bundle_name: str, + dag_file_rel_path: str, + timeout: float | None, + logger: FilteringBoundLogger, + ) -> DagFileParsingResult: + """ + Parse *path* outside the Dag processor and wait for the result. + + There is no API client, so each request of the runtime that needs one gets an error. + + :raises TimeoutError: if the parse does not finish within *timeout* seconds. The runtime is + killed. + """ + with selectors.DefaultSelector() as selector: + proc = cls.start( + id=uuid7(), + path=path, + bundle_path=bundle_path, + bundle_name=bundle_name, + dag_file_rel_path=dag_file_rel_path, + selector=selector, + logger=logger, + ) + try: + while not proc.is_ready: + wait = 0.1 + if timeout is not None: + if (remaining := proc.start_time + timeout - time.monotonic()) <= 0: + raise TimeoutError( + f"The Lang-SDK runtime did not parse {os.fspath(path)} within {timeout}s" + ) + wait = min(wait, remaining) + proc._service_subprocess(max_wait_time=wait) + except BaseException: + proc._kill_runtime() + raise + finally: + proc.close() + return cast("DagFileParsingResult", proc.parsing_result) + + def _accept_connection(self, listener: socket, *, channel: _Channel) -> bool: + try: + conn, _ = listener.accept() + except (BlockingIOError, InterruptedError): + return True + conn.setblocking(True) + self._unverified_connections.append((conn, channel)) + self._verify_connections() + return True + + def _verify_connections(self) -> None: + """ + Use each accepted connection once it is confirmed to come from the runtime. + + A connection that is not visible yet stays pending and is checked again on the next + ``is_ready`` poll, so the caller's loop never waits here. + """ + pending = [] + for conn, channel in self._unverified_connections: + if channel not in self._listeners: + # The runtime already connected this channel. + conn.close() + continue + try: + owned = _is_connection_from_pid(conn, self.pid) + except OSError: + conn.close() + continue + if not owned: + pending.append((conn, channel)) + continue + self._close_listener(channel) + if channel == "comm": + self._register_comm(conn) + else: + self._register_logs(conn) + self._unverified_connections = pending + + def _close_listener(self, channel: _Channel) -> None: + if (listener := self._listeners.pop(channel, None)) is not None: + self._on_socket_closed(listener) + listener.close() + + def _close_listeners(self) -> None: + """Close the listeners of a runtime that did not connect, and connections never verified.""" + for channel in list(self._listeners): + self._close_listener(channel) + for conn, _ in self._unverified_connections: + conn.close() + self._unverified_connections = [] + + def _register_comm(self, conn: socket) -> None: + self.stdin = conn + self._open_sockets[conn] = "requests" + read_frame, on_close = length_prefixed_frame_reader( + self._handle_valid_requests(), on_close=self._on_socket_closed + ) + + def read_valid_frame(sock: socket) -> bool: + # A frame that does not decode would otherwise escape the Dag processor's selector loop. + try: + return read_frame(sock) + except msgspec.DecodeError as e: + self._fail_on_invalid_message(f"The Lang-SDK runtime sent an invalid frame: {e}") + return False + + self.selector.register(conn, selectors.EVENT_READ, (read_valid_frame, on_close)) + # The parse child reports the version and waits for the reply before it execs the runtime, + # so the version is known here. It is set only now, so the child's messages are not migrated. + self._subprocess_schema_version = self._runtime_schema_version + self.send_msg(self._parse_request, request_id=0) + + def _handle_valid_requests(self) -> Generator[None, _RequestFrame, None]: + """ + Pass each request on to ``handle_requests``, or kill the runtime at one that does not validate. + + ``handle_requests`` would only log such a request, and the runtime would wait for a reply. + """ + requests = self.handle_requests(self.process_log) + next(requests) + while True: + frame = yield + try: + self.decoder.validate_python(self._deserialize_request(frame.body)) + except ValueError as e: + self._fail_on_invalid_message( + f"The Lang-SDK runtime sent a message that does not validate: {e}" + ) + return + requests.send(frame) + + def _fail_on_invalid_message(self, message: str) -> None: + """Kill the runtime; *message* is the import error unless a parse result was already received.""" + if self.parsing_result is None: + self._set_import_error(message) + else: + self.process_log.warning( + "Ignoring an invalid message from the Lang-SDK runtime after its parse result", error=message + ) + self._kill_runtime() + + def _register_logs(self, conn: socket) -> None: + self._open_sockets[conn] = "logs" + self.selector.register( + conn, + selectors.EVENT_READ, + make_buffered_socket_reader( + process_log_messages_from_subprocess(self._get_target_loggers()), + on_close=self._on_socket_closed, + ), + ) + + def _set_import_error(self, message: str) -> None: + self.parsing_result = DagFileParsingResult( + fileloc=self._parse_request.file, + serialized_dags=[], + import_errors={self.dag_file_rel_path: message}, + ) + + def _handle_runtime_schema_version( + self, msg: LangSDKRuntimeSchemaVersion, log: FilteringBoundLogger, req_id: int + ) -> RequestResult | ResponseSent: + if self._schema_version_reported: + self._reject_request(msg, log, req_id) + return ResponseSent.ALREADY_SENT + self._runtime_schema_version = msg.schema_version + self._schema_version_reported = True + return None, {} + + def _handle_parsing_result( + self, msg: DagFileParsingResult, log: FilteringBoundLogger, req_id: int + ) -> RequestResult | ResponseSent: + if self.parsing_result is not None: + log.warning("Ignoring another parse result from the Lang-SDK runtime", fileloc=msg.fileloc) + self.send_msg( + None, + request_id=req_id, + error=ErrorResponse(detail={"message": "A parse result was already received"}), + ) + return ResponseSent.ALREADY_SENT + import_errors = dict(msg.import_errors or {}) + serialized_dags = [] + for dag in msg.serialized_dags: + DagSerialization.fill_config_defaults(dag.data) + try: + DagSerialization.validate_serialized_dag(dag.data) + except DeserializationError as e: + message = f"Cannot load the serialized Dag: {e}" + self.process_log.warning(message) + previous = import_errors.get(self.dag_file_rel_path) + import_errors[self.dag_file_rel_path] = f"{previous}\n{message}" if previous else message + continue + serialized_dags.append(dag) + self.parsing_result = msg.model_copy( + update={ + "serialized_dags": serialized_dags, + "import_errors": import_errors or None, + "dag_source_codes": self._read_dag_source_codes(serialized_dags), + } + ) + self._parsing_result_monotonic = time.monotonic() + return None, {} + + def _read_dag_source_codes(self, serialized_dags: list[LazyDeserializedDAG]) -> dict[str, DagSourceCode]: + """ + Read the file's source with its Dag importer, for the fileloc of each Dag. + + A binary artifact cannot be read as text, so a source that cannot be read is a placeholder. + """ + if not serialized_dags: + return {} + file = self._parse_request.file + try: + coordinator = get_claiming_coordinator(file, self.bundle_name) + if coordinator is None or (importer := coordinator.get_dag_importer()) is None: + raise RuntimeError(f"No coordinator's Dag importer claims {file}") + source = importer.get_source_code(FilesystemDagDefinition(Path(file))) + except Exception as e: + self.process_log.warning("Cannot read the Dag source", fileloc=file, error=str(e)) + source = DagSourceCode(f"Cannot read the source of {self.dag_file_rel_path}: {e}", "text") + return {dag.data["dag"].get("fileloc", file): source for dag in serialized_dags} + + _request_handlers: ClassVar[dict[type[BaseModel], RequestHandler[LangSDKDagFileProcessorProcess]]] = { + **BaseDagFileProcessorProcess._common_request_handlers, + **dict([register_request_method(LangSDKRuntimeSchemaVersion, _handle_runtime_schema_version)]), + } + + def _handle_request(self, msg, log: FilteringBoundLogger, req_id: int) -> None: + if self.client is None and not isinstance( + msg, (DagFileParsingResult, LangSDKRuntimeSchemaVersion, MaskSecret) + ): + self.send_msg( + None, + request_id=req_id, + error=ErrorResponse( + detail={"message": f"{type(msg).__name__} is answered only in the Dag processor"} + ), + ) + return + super()._handle_request(msg, log, req_id) + + @property + def is_ready(self) -> bool: + self._verify_connections() + if ( + self._parsing_result_monotonic is not None + and self._exit_code is None + and time.monotonic() - self._parsing_result_monotonic > _EXIT_GRACE_PERIOD + ): + self.process_log.warning("The Lang-SDK runtime did not exit after its parse result; killing it") + self._kill_runtime() + if self._check_subprocess_exit() is None: + return False + self._close_listeners() + if not super().is_ready: + return False + if self.parsing_result is None: + self._set_import_error( + f"The Lang-SDK runtime exited with code {self._exit_code} without a parse result" + ) + return True + + def _kill_runtime(self) -> None: + """Kill the runtime and wait for it, without servicing its sockets, whose handler may have failed.""" + if self._exit_code is not None: + return + try: + self._signal_subprocess(signal.SIGKILL) + self._exit_code = self._process.wait(timeout=None) + except (self._process.ProcessNotFound, ProcessLookupError): + self._exit_code = -1 + + def close(self) -> None: + # A listener has nothing to drain, and cleanup would call its accept handler forever. + self._close_listeners() + super().close() diff --git a/airflow-core/src/airflow/dag_processing/manager.py b/airflow-core/src/airflow/dag_processing/manager.py index 8c4fd0588eb..7311e510382 100644 --- a/airflow-core/src/airflow/dag_processing/manager.py +++ b/airflow-core/src/airflow/dag_processing/manager.py @@ -55,7 +55,13 @@ from airflow.dag_processing.bundles.base import ( ) from airflow.dag_processing.bundles.manager import DagBundlesManager from airflow.dag_processing.collection import update_dag_parsing_results_in_db -from airflow.dag_processing.processor import DagFileParsingResult, DagFileProcessorProcess +from airflow.dag_processing.importer_routing import get_claiming_coordinator +from airflow.dag_processing.lang_sdk_processor import LangSDKDagFileProcessorProcess +from airflow.dag_processing.processor import ( + BaseDagFileProcessorProcess, + DagFileParsingResult, + DagFileProcessorProcess, +) from airflow.models.asset import remove_references_to_deleted_dags from airflow.models.dag import DagModel from airflow.models.dagbag import DagPriorityParsingRequest @@ -261,7 +267,7 @@ class DagFileProcessorManager(LoggingMixin): _multi_team: bool = attrs.field(factory=lambda: conf.getboolean("core", "multi_team"), init=False) _bundle_name_to_team_name: dict[str, str | None] = attrs.field(factory=dict, init=False) - _processors: dict[DagFileInfo, DagFileProcessorProcess] = attrs.field(factory=dict, init=False) + _processors: dict[DagFileInfo, BaseDagFileProcessorProcess] = attrs.field(factory=dict, init=False) _parsing_start_time: float | None = attrs.field(default=None, init=False) _num_run: int = attrs.field(default=0, init=False) @@ -1260,7 +1266,7 @@ class DagFileProcessorManager(LoggingMixin): def handle_parsing_result( self, file: DagFileInfo, - proc: DagFileProcessorProcess, + proc: BaseDagFileProcessorProcess, *, session: Session = NEW_SESSION, ) -> None: @@ -1444,19 +1450,17 @@ class DagFileProcessorManager(LoggingMixin): client.base_url = "http://in-process.invalid./" return client - def _create_process(self, dag_file: DagFileInfo) -> DagFileProcessorProcess: + def _create_process(self, dag_file: DagFileInfo) -> BaseDagFileProcessorProcess: id = uuid7() callback_to_execute_for_file = self._callback_to_execute.pop(dag_file, []) logger, logger_filehandle = self._get_logger_for_dag_file(dag_file) - - return DagFileProcessorProcess.start( + kwargs: dict[str, Any] = dict( id=id, path=dag_file.absolute_path, bundle_path=cast("Path", dag_file.bundle_path), bundle_name=dag_file.bundle_name, dag_file_rel_path=str(dag_file.rel_path), - callbacks=callback_to_execute_for_file, selector=self.selector, logger=logger, logger_filehandle=logger_filehandle, @@ -1464,6 +1468,17 @@ class DagFileProcessorManager(LoggingMixin): client=self.client, ) + if get_claiming_coordinator(dag_file.absolute_path, dag_file.bundle_name) is not None: + if callback_to_execute_for_file: + self.log.warning( + "Dropping %d callbacks for %s: Lang-SDK runtimes do not run callbacks", + len(callback_to_execute_for_file), + dag_file.rel_path, + ) + return LangSDKDagFileProcessorProcess.start(**kwargs) + + return DagFileProcessorProcess.start(callbacks=callback_to_execute_for_file, **kwargs) + def _start_new_processes(self): """Start more processors if we have enough slots and files to process.""" bundle_to_team = self._get_team_names({file.bundle_name for file in self._file_queue}) diff --git a/airflow-core/tests/unit/dag_processing/fake_lang_sdk.py b/airflow-core/tests/unit/dag_processing/fake_lang_sdk.py new file mode 100644 index 00000000000..53a7cef75f6 --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/fake_lang_sdk.py @@ -0,0 +1,127 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""A coordinator that claims ``.native`` Dag files, and a runtime a test can play in the parse child.""" + +from __future__ import annotations + +import contextlib +import json +import socket +from pathlib import Path +from typing import TYPE_CHECKING, Any +from unittest import mock + +import attrs +from pydantic import TypeAdapter + +from airflow.dag_processing.processor import ( + DagFileParseRequest, + DagFileParsingResult, + ToDagProcessor, + ToManager, +) +from airflow.sdk.coordinators._dag_importer import CoordinatorDagImporter +from airflow.sdk.coordinators._subprocess import SubprocessCoordinator +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.comms import CommsDecoder +from airflow.sdk.importers import DagSourceCode, reset_importer_registry + +from tests_common.test_utils.config import conf_vars + +if TYPE_CHECKING: + from collections.abc import Callable, Iterator, Sequence + + [email protected](kw_only=True) +class FakeCoordinator(SubprocessCoordinator): + """ + Claim ``.native`` files; the file's JSON names the command that parses it. + + ``argv`` is the command, ``schema_version`` its schema version, and ``command_error`` an error to + raise instead. + """ + + def _build_parse_dag_command(self, *, path: Path) -> tuple[list[str], str | None]: + spec = json.loads(path.read_text()) + if error := spec.get("command_error"): + raise FileNotFoundError(error) + return spec.get("argv", ["/bin/false"]), spec.get("schema_version") + + @classmethod + def get_dag_importer_class(cls) -> type[FakeCoordinatorDagImporter]: + return FakeCoordinatorDagImporter + + +class FakeCoordinatorDagImporter(CoordinatorDagImporter): + artifact_suffix = ".native" + supported_extensions = [".native"] + + def get_source_code(self, definition) -> DagSourceCode: + return DagSourceCode(definition.read_text(), "fake") + + [email protected] +def fake_coordinator(**kwargs: Any) -> Iterator[None]: + """ + Configure a ``FakeCoordinator``, with fresh coordinators and registries inside and after the block. + + The parse child is a bare fork even on macOS, so it sees the test's ``parse_dag`` patch and config. + """ + spec = {"fake": {"classpath": f"{__name__}.FakeCoordinator", "kwargs": kwargs}} + reset_importer_registry() + try: + with ( + conf_vars({("sdk", "coordinators"): json.dumps(spec)}), + mock.patch.object(supervisor, "_should_use_exec", return_value=False), + ): + yield + finally: + reset_importer_registry() + + +def write_native_file(path: Path, **spec: Any) -> Path: + path.write_text(json.dumps(spec)) + return path + + +def play_runtime( + reply: Callable[[DagFileParseRequest, CommsDecoder], DagFileParsingResult | None], + *, + schema_version: str | None = None, + log_lines: Sequence[dict[str, Any]] = (), +) -> Callable[..., None]: + """ + Return a ``parse_dag`` that plays the runtime in the parse child instead of exec'ing one. + + It connects back as a runtime does, writes *log_lines* to its logs channel and answers the parse + request with what *reply* returns; ``None`` sends no result. + """ + + def parse_dag(coordinator, *, comm_address, logs_address, report_schema_version, **kwargs) -> None: + report_schema_version(schema_version) + comm = socket.create_connection(comm_address) + logs = socket.create_connection(logs_address) + for line in log_lines: + logs.sendall(json.dumps(line).encode() + b"\n") + comms = CommsDecoder[ToDagProcessor, ToManager](socket=comm, body_decoder=TypeAdapter(ToDagProcessor)) + request = comms._get_response() + assert isinstance(request, DagFileParseRequest) + if (result := reply(request, comms)) is not None: + comms.send(result) + + return parse_dag diff --git a/airflow-core/tests/unit/dag_processing/test_dagbag.py b/airflow-core/tests/unit/dag_processing/test_dagbag.py index 54b06ce75aa..ecfaef5612d 100644 --- a/airflow-core/tests/unit/dag_processing/test_dagbag.py +++ b/airflow-core/tests/unit/dag_processing/test_dagbag.py @@ -46,6 +46,7 @@ from airflow.exceptions import UnknownExecutorException from airflow.executors.executor_loader import ExecutorLoader from airflow.models.dag import DagModel from airflow.models.dagwarning import DagWarning, DagWarningType +from airflow.models.errors import ParseImportError from airflow.models.pool import Pool from airflow.models.serialized_dag import SerializedDagModel from airflow.sdk import DAG, BaseOperator @@ -64,6 +65,7 @@ from tests_common.pytest_plugin import AIRFLOW_ROOT_PATH from tests_common.test_utils import db from tests_common.test_utils.config import conf_vars from unit import cluster_policies +from unit.dag_processing.fake_lang_sdk import fake_coordinator, write_native_file from unit.models import TEST_DAGS_FOLDER pytestmark = pytest.mark.db_test @@ -1530,3 +1532,21 @@ class TestBundlePathSysPath: assert str(tmp_path) not in dag.description assert sys.path == syspath_before + + +def test_sync_bag_to_db_leaves_native_files_to_the_dag_processor(tmp_path, session, testing_dag_bundle): + db.clear_db_import_errors() + write_native_file(tmp_path / "dags.native") + session.add(ParseImportError(bundle_name="testing", filename="dags.native", stacktrace="stored")) + session.commit() + + with fake_coordinator(): + dagbag = DagBag(dag_folder=os.fspath(tmp_path), bundle_path=tmp_path, bundle_name="testing") + sync_bag_to_db(dagbag, "testing", None, session=session) + + assert dagbag.import_errors == { + "dags.native": "A native Lang-SDK Dag is parsed only by the Dag processor" + } + assert {(e.filename, e.stacktrace) for e in session.scalars(select(ParseImportError))} == { + ("dags.native", "stored") + } diff --git a/airflow-core/tests/unit/dag_processing/test_importer_routing.py b/airflow-core/tests/unit/dag_processing/test_importer_routing.py new file mode 100644 index 00000000000..982fdc0cb68 --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_importer_routing.py @@ -0,0 +1,38 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from unittest import mock + +from airflow.dag_processing.importer_routing import get_claiming_coordinator + +from unit.dag_processing.fake_lang_sdk import FakeCoordinator, fake_coordinator + + +def test_get_claiming_coordinator_returns_the_coordinator_of_its_importer(tmp_path): + with fake_coordinator(): + coordinator = get_claiming_coordinator(tmp_path / "dags.native", "testing") + others = [get_claiming_coordinator(tmp_path / name, "testing") for name in ("dags.other", "dag.py")] + + assert isinstance(coordinator, FakeCoordinator) + assert others == [None, None] + + [email protected]("airflow.dag_processing.importer_routing._get_registry", autospec=True, return_value=None) +def test_get_claiming_coordinator_without_a_registry(mock_registry, tmp_path): + assert get_claiming_coordinator(tmp_path / "dags.native", "testing") is None diff --git a/airflow-core/tests/unit/dag_processing/test_lang_sdk_processor.py b/airflow-core/tests/unit/dag_processing/test_lang_sdk_processor.py new file mode 100644 index 00000000000..4a5fd758681 --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_lang_sdk_processor.py @@ -0,0 +1,533 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import contextlib +import os +import selectors +import signal +import socket +import sys +import time +import uuid +from pathlib import Path +from unittest.mock import ANY, MagicMock, patch + +import psutil +import pytest +import structlog + +from airflow.configuration import conf +from airflow.dag_processing.lang_sdk_processor import ( + LangSDKDagFileProcessorProcess, + LangSDKRuntimeSchemaVersion, +) +from airflow.dag_processing.processor import DagFileParseRequest, DagFileParsingResult +from airflow.sdk import DAG, BaseOperator +from airflow.sdk.api.client import Client +from airflow.sdk.api.datamodels._generated import VariableResponse +from airflow.sdk.exceptions import AirflowRuntimeError +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.comms import GetVariable, MaskSecret, _RequestFrame +from airflow.sdk.importers import DagSourceCode +from airflow.serialization.serialized_objects import DagSerialization, LazyDeserializedDAG + +from tests_common.test_utils.config import conf_vars +from unit.dag_processing.fake_lang_sdk import ( + FakeCoordinator, + fake_coordinator, + play_runtime, + write_native_file, +) + +# The oldest supervisor schema version, so the parse request is downgraded. +OLDEST_SCHEMA_VERSION = "2026-06-16" + + +def _serialize_dag(dag_id: str, description: str | None = None) -> LazyDeserializedDAG: + with DAG(dag_id, schedule=None, description=description) as dag: + BaseOperator(task_id="extract") + return LazyDeserializedDAG(data=DagSerialization.to_dict(dag)) + + +def _reply_with(*dags: LazyDeserializedDAG, **result): + def reply(request: DagFileParseRequest, comms) -> DagFileParsingResult: + return DagFileParsingResult(fileloc=request.file, serialized_dags=list(dags), **result) + + return reply + + +def _get_open_fds() -> set[int]: + # Without /proc, as on macOS, this is empty, so the fd leak checks pass trivially. + return {int(fd) for fd in os.listdir("/proc/self/fd")} if os.path.isdir("/proc/self/fd") else set() + + [email protected](autouse=True) +def _coordinator(): + with fake_coordinator(): + yield + + +def _start(tmp_path, selector, *, client: Client | None = None, **spec) -> LangSDKDagFileProcessorProcess: + return LangSDKDagFileProcessorProcess.start( + id=uuid.uuid4(), + path=write_native_file(tmp_path / "dag.native", **spec), + bundle_path=tmp_path, + bundle_name="testing", + dag_file_rel_path="dag.native", + selector=selector, + logger=structlog.get_logger(), + client=client or MagicMock(spec=Client), + ) + + [email protected] +def parse(tmp_path): + """Parse ``dag.native`` as the Dag processor does, and check that nothing is left open.""" + + def _parse(**kwargs) -> LangSDKDagFileProcessorProcess: + fds_before = _get_open_fds() + with selectors.DefaultSelector() as selector: + proc = _start(tmp_path, selector, **kwargs) + deadline = time.monotonic() + 30 + while not proc.is_ready: + assert time.monotonic() < deadline, "the Lang-SDK parse did not finish" + proc._service_subprocess(max_wait_time=0.1) + assert selector.get_map() == {} + proc.close() + assert _get_open_fds() <= fds_before + return proc + + return _parse + + +def _send_an_invalid_frame(request, comms) -> None: + comms.socket.sendall(bytes.fromhex("00000003c1c1c1")) + time.sleep(60) + + +class TestLangSDKDagFileProcessorProcess: + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_parses_the_dags_the_runtime_returns(self, mock_parse_dag, parse, tmp_path, cap_structlog): + mock_parse_dag.side_effect = play_runtime( + _reply_with(_serialize_dag("native_dag")), + schema_version=OLDEST_SCHEMA_VERSION, + log_lines=[{"event": "Parsing the bundle", "level": "info"}], + ) + + proc = parse() + + assert proc.parsing_result.fileloc == os.fspath(tmp_path / "dag.native") + assert proc.parsing_result.import_errors is None + assert [dag.dag_id for dag in proc.parsing_result.serialized_dags] == ["native_dag"] + assert list(proc.parsing_result.dag_source_codes.values()) == [ + DagSourceCode((tmp_path / "dag.native").read_text(), "fake") + ] + assert proc._subprocess_schema_version == OLDEST_SCHEMA_VERSION + assert "Parsing the bundle" in cap_structlog + + @patch("airflow.dag_processing.lang_sdk_processor._is_connection_from_pid", autospec=True) + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_a_connection_is_used_once_it_is_verified(self, mock_parse_dag, mock_owned, tmp_path): + mock_parse_dag.side_effect = play_runtime(_reply_with(_serialize_dag("native_dag"))) + mock_owned.return_value = False + with selectors.DefaultSelector() as selector: + proc = _start(tmp_path, selector) + child_stdin = proc.stdin + deadline = time.monotonic() + 30 + while len(proc._unverified_connections) < 2: + assert proc.stdin is child_stdin, "an unverified connection was used" + assert time.monotonic() < deadline, "the runtime did not connect" + proc._service_subprocess(max_wait_time=0.1) + assert proc.stdin is child_stdin + + mock_owned.return_value = True + while not proc.is_ready: + assert time.monotonic() < deadline, "the Lang-SDK parse did not finish" + proc._service_subprocess(max_wait_time=0.1) + proc.close() + + assert [dag.dag_id for dag in proc.parsing_result.serialized_dags] == ["native_dag"] + + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_requests_are_answered_by_the_client(self, mock_parse_dag, parse): + def reply(request, comms): + variable = comms.send(GetVariable(key="native_var")) + return _reply_with(_serialize_dag("native_dag", description=variable.value))(request, comms) + + mock_parse_dag.side_effect = play_runtime(reply) + client = MagicMock(spec=Client) + client.variables = MagicMock() + client.variables.get.return_value = VariableResponse(key="native_var", value="from-db") + + proc = parse(client=client) + + [dag] = proc.parsing_result.serialized_dags + assert dag.data["dag"]["description"] == "from-db" + + @pytest.mark.parametrize( + ("change", "error"), + [ + pytest.param( + {"max_active_runs": "many"}, + "Dag 'broken_dag' does not match the schema: 'many' is not of type 'number'", + id="schema", + ), + pytest.param( + {"timetable": {"__type": "no.such.Timetable", "__var": {}}}, + "Dag 'broken_dag' cannot be deserialized: TimetableNotRegistered: " + "Timetable class 'no.such.Timetable' is not registered", + id="deserialize", + ), + ], + ) + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_a_dag_that_does_not_validate_is_an_import_error(self, mock_parse_dag, parse, change, error): + broken = _serialize_dag("broken_dag") + broken.data["dag"].update(change) + mock_parse_dag.side_effect = play_runtime(_reply_with(broken, _serialize_dag("good_dag"))) + + proc = parse() + + assert [dag.dag_id for dag in proc.parsing_result.serialized_dags] == ["good_dag"] + [message] = proc.parsing_result.import_errors.values() + assert message.startswith(f"Cannot load the serialized Dag: {error}") + + @conf_vars( + { + ("core", "max_active_tasks_per_dag"): "7", + ("core", "max_active_runs_per_dag"): "3", + ("scheduler", "catchup_by_default"): "True", + } + ) + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_a_dag_setting_left_unset_is_filled_from_the_config(self, mock_parse_dag, parse): + dag = _serialize_dag("native_dag") + del dag.data["dag"]["max_active_tasks"], dag.data["dag"]["catchup"] + dag.data["dag"]["max_active_runs"] = 16 + mock_parse_dag.side_effect = play_runtime(_reply_with(dag)) + + proc = parse() + + [stored] = proc.parsing_result.serialized_dags + assert proc.parsing_result.import_errors is None + assert stored.data["dag"]["max_active_tasks"] == 7 + assert stored.data["dag"]["max_active_runs"] == 16 + assert stored.data["dag"]["catchup"] is True + + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_a_dag_with_a_cycle_is_an_import_error(self, mock_parse_dag, parse): + cyclic = _serialize_dag("cyclic_dag") + [task] = cyclic.data["dag"]["tasks"] + task["__var"]["downstream_task_ids"] = ["extract"] + mock_parse_dag.side_effect = play_runtime(_reply_with(cyclic)) + + proc = parse() + + assert proc.parsing_result.serialized_dags == [] + assert proc.parsing_result.import_errors == { + "dag.native": "Cannot load the serialized Dag: Dag 'cyclic_dag' has a cycle through task 'extract'" + } + + @pytest.mark.parametrize( + ("spec", "reply", "error"), + [ + pytest.param( + {"command_error": "no runtime"}, + None, + "Cannot start the Lang-SDK runtime: FileNotFoundError: no runtime", + id="command-not-resolved", + ), + pytest.param( + {"argv": ["/no/such/runtime"]}, + None, + "Cannot start the Lang-SDK runtime: FileNotFoundError: " + "[Errno 2] No such file or directory: '/no/such/runtime'", + id="exec-failed", + ), + pytest.param( + {"argv": ["/bin/sh", "-c", "exit 3"]}, + None, + "The Lang-SDK runtime exited with code 3 without a parse result", + id="exits-before-connecting", + ), + pytest.param( + {}, + lambda request, comms: None, + "The Lang-SDK runtime exited with code 0 without a parse result", + id="exits-without-a-result", + ), + pytest.param( + {}, + _send_an_invalid_frame, + "The Lang-SDK runtime sent an invalid frame: MessagePack data is malformed: " + "invalid opcode '\\xc1' (byte 0)", + id="invalid-frame", + ), + ], + ) + def test_a_failed_parse_is_an_import_error(self, parse, spec, reply, error): + with ( + patch.object(FakeCoordinator, "parse_dag", autospec=True, side_effect=play_runtime(reply)) + if reply + else contextlib.nullcontext() + ): + proc = parse(**spec) + + assert proc.parsing_result.serialized_dags == [] + assert proc.parsing_result.import_errors == {"dag.native": error} + + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_a_message_that_does_not_validate_is_an_import_error(self, mock_parse_dag, parse): + def reply(request, comms): + body = { + "type": "DagFileParsingResult", + "fileloc": request.file, + "serialized_dags": [{"data": "not a dict"}], + } + comms.socket.sendall(_RequestFrame(id=1, body=body).as_bytes()) + time.sleep(60) + + mock_parse_dag.side_effect = play_runtime(reply) + + proc = parse() + + [message] = proc.parsing_result.import_errors.values() + assert message.startswith("The Lang-SDK runtime sent a message that does not validate: ") + assert "DagFileParsingResult.serialized_dags.0.data\n Input should be a valid dictionary" in message + assert proc._exit_code == -signal.SIGKILL + + @patch.object( + FakeCoordinator, "parse_dag", autospec=True, side_effect=play_runtime(_send_an_invalid_frame) + ) + def test_killing_the_runtime_is_not_reported_as_out_of_memory(self, mock_parse_dag, parse, cap_structlog): + proc = parse() + + assert proc._exit_code == -signal.SIGKILL + assert not any("Likely out of memory" in str(entry.get("event")) for entry in cap_structlog.entries) + + @patch("airflow.dag_processing.lang_sdk_processor._EXIT_GRACE_PERIOD", 0.5) + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_a_runtime_that_runs_on_after_its_result_is_killed(self, mock_parse_dag, parse, cap_structlog): + def reply(request, comms): + comms.send(_reply_with(_serialize_dag("native_dag"))(request, comms)) + time.sleep(60) + + mock_parse_dag.side_effect = play_runtime(reply) + + proc = parse() + + assert [dag.dag_id for dag in proc.parsing_result.serialized_dags] == ["native_dag"] + assert proc._exit_code == -signal.SIGKILL + assert "The Lang-SDK runtime did not exit after its parse result; killing it" in cap_structlog + + @pytest.mark.skipif(not Path("/proc/self/fd").is_dir(), reason="reads /proc") + @pytest.mark.parametrize("use_exec", [False, True], ids=["fork", "spawn"]) + def test_the_runtime_inherits_only_its_standard_streams(self, monkeypatch, tmp_path, use_exec): + if use_exec: + # The spawned interpreter finds the coordinator again from its environment. + monkeypatch.setattr(supervisor, "_should_use_exec", lambda: True) + monkeypatch.setenv("PYTHONPATH", os.pathsep.join(sys.path)) + monkeypatch.setenv("AIRFLOW__SDK__COORDINATORS", conf.get("sdk", "coordinators")) + with selectors.DefaultSelector() as selector: + proc = _start(tmp_path, selector, argv=["/bin/sh", "-c", "exec sleep 30"]) + deadline = time.monotonic() + 30 + while psutil.Process(proc.pid).name() != "sleep": + assert time.monotonic() < deadline, "the runtime did not start" + proc._service_subprocess(max_wait_time=0.1) + fd_dir = Path(f"/proc/{proc.pid}/fd") + fds = {fd.name: os.readlink(fd) for fd in fd_dir.iterdir()} + proc.kill(signal.SIGKILL) + proc.close() + + assert sorted(fds) == ["0", "1", "2"] + assert fds["0"] == "/dev/null" + + +class TestRun: + @staticmethod + def _run(tmp_path, *, timeout: float | None = 30) -> DagFileParsingResult: + return LangSDKDagFileProcessorProcess.run( + path=write_native_file(tmp_path / "dag.native"), + bundle_path=tmp_path, + bundle_name="testing", + dag_file_rel_path="dag.native", + timeout=timeout, + logger=structlog.get_logger(), + ) + + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_requests_get_an_error(self, mock_parse_dag, tmp_path): + def reply(request, comms): + with pytest.raises(AirflowRuntimeError) as ctx: + comms.send(GetVariable(key="native_var")) + description = ctx.value.error.detail["message"] + return _reply_with(_serialize_dag("native_dag", description=description))(request, comms) + + mock_parse_dag.side_effect = play_runtime(reply) + + result = self._run(tmp_path) + + assert result.serialized_dags[0].data["dag"]["description"] == ( + "GetVariable is answered only in the Dag processor" + ) + + @patch("airflow.sdk.execution_time.request_handlers.mask_secret", autospec=True) + @patch.object(FakeCoordinator, "parse_dag", autospec=True) + def test_a_secret_is_masked_without_a_client(self, mock_parse_dag, mock_mask_secret, tmp_path): + def reply(request, comms): + comms.send(MaskSecret(value="native-secret", name="native_conn")) + return _reply_with(_serialize_dag("native_dag"))(request, comms) + + mock_parse_dag.side_effect = play_runtime(reply) + + result = self._run(tmp_path) + + assert [dag.dag_id for dag in result.serialized_dags] == ["native_dag"] + mock_mask_secret.assert_called_once_with("native-secret", "native_conn") + + @pytest.mark.parametrize("connected", [False, True], ids=["before-connecting", "after-connecting"]) + def test_a_parse_past_its_timeout_is_killed(self, tmp_path, connected): + # A runtime that never connects leaves both listeners open when it is killed. + write_native_file(tmp_path / "dag.native", argv=["/bin/sh", "-c", "exec sleep 60"]) + fds_before = _get_open_fds() + + with ( + patch.object( + FakeCoordinator, + "parse_dag", + autospec=True, + side_effect=play_runtime(lambda request, comms: time.sleep(60)), + ) + if connected + else contextlib.nullcontext(), + patch.object( + LangSDKDagFileProcessorProcess, + "close", + autospec=True, + side_effect=LangSDKDagFileProcessorProcess.close, + ) as mock_close, + pytest.raises(TimeoutError, match=r"did not parse .*dag\.native within 1s"), + ): + LangSDKDagFileProcessorProcess.run( + path=tmp_path / "dag.native", + bundle_path=tmp_path, + bundle_name="testing", + dag_file_rel_path="dag.native", + timeout=1, + logger=structlog.get_logger(), + ) + + [proc] = [c.args[0] for c in mock_close.call_args_list] + assert proc._exit_code == -9 + assert not proc._open_sockets + assert _get_open_fds() <= fds_before + + +def _make_process(**kwargs) -> LangSDKDagFileProcessorProcess: + return LangSDKDagFileProcessorProcess( + id=uuid.uuid4(), + pid=1, + stdin=MagicMock(), + process=MagicMock(), + process_log=MagicMock(), + selector=MagicMock(), + bundle_name="testing", + dag_file_rel_path="dag.native", + listeners={}, + parse_request=DagFileParseRequest( + file="/b/dag.native", bundle_path=Path("/b"), bundle_name="testing" + ), + **kwargs, + ) + + +def test_a_dag_source_that_cannot_be_read_is_a_placeholder(): + proc = _make_process() + dag = _serialize_dag("native_dag") + + proc._handle_request(DagFileParsingResult(fileloc="/b/dag.native", serialized_dags=[dag]), MagicMock(), 1) + + source = proc.parsing_result.dag_source_codes[dag.data["dag"]["fileloc"]] + assert source.language == "text" + assert source.source_code.startswith("Cannot read the source of dag.native: [Errno 2] No such file") + + [email protected](LangSDKDagFileProcessorProcess, "send_msg", autospec=True) +def test_the_first_parse_result_wins(mock_send_msg): + proc = _make_process() + first = DagFileParsingResult(fileloc="/b/dag.native", serialized_dags=[_serialize_dag("first")]) + second = DagFileParsingResult(fileloc="/b/dag.native", serialized_dags=[_serialize_dag("second")]) + + proc._handle_request(first, MagicMock(), 1) + proc._handle_request(second, MagicMock(), 2) + + assert [dag.dag_id for dag in proc.parsing_result.serialized_dags] == ["first"] + assert mock_send_msg.call_args.kwargs["error"].detail == { + "message": "A parse result was already received" + } + + [email protected]( + "invalid_frame", + [ + pytest.param(bytes.fromhex("00000003c1c1c1"), id="does-not-decode"), + pytest.param( + _RequestFrame( + id=2, + body={ + "type": "DagFileParsingResult", + "fileloc": "/b/dag.native", + "serialized_dags": [{"data": "not a dict"}], + }, + ).as_bytes(), + id="does-not-validate", + ), + ], +) [email protected](LangSDKDagFileProcessorProcess, "_kill_runtime", autospec=True) [email protected](LangSDKDagFileProcessorProcess, "send_msg", autospec=True) +def test_an_invalid_message_after_the_parse_result_keeps_it(mock_send_msg, mock_kill_runtime, invalid_frame): + proc = _make_process() + result = DagFileParsingResult(fileloc="/b/dag.native", serialized_dags=[_serialize_dag("native_dag")]) + runtime, conn = socket.socketpair() + with runtime, conn: + proc._register_comm(conn) + read_frame, _ = proc.selector.register.call_args.args[2] + runtime.sendall(_RequestFrame(id=1, body=result.model_dump(mode="json")).as_bytes()) + assert read_frame(conn) + runtime.sendall(invalid_frame) + assert not read_frame(conn) + + assert [dag.dag_id for dag in proc.parsing_result.serialized_dags] == ["native_dag"] + assert proc.parsing_result.import_errors is None + proc.process_log.warning.assert_called_with( + "Ignoring an invalid message from the Lang-SDK runtime after its parse result", error=ANY + ) + mock_kill_runtime.assert_called_once_with(proc) + + [email protected](LangSDKDagFileProcessorProcess, "send_msg", autospec=True) +def test_the_schema_version_is_reported_once(mock_send_msg): + proc = _make_process() + + proc._handle_request(LangSDKRuntimeSchemaVersion(schema_version=OLDEST_SCHEMA_VERSION), MagicMock(), 1) + proc._handle_request(LangSDKRuntimeSchemaVersion(schema_version=None), MagicMock(), 2) + + assert proc._runtime_schema_version == OLDEST_SCHEMA_VERSION + assert mock_send_msg.call_args.kwargs["error"].detail["message"] == "Unhandled request" diff --git a/airflow-core/tests/unit/dag_processing/test_manager.py b/airflow-core/tests/unit/dag_processing/test_manager.py index a3624e1d4a1..2ea33b386a0 100644 --- a/airflow-core/tests/unit/dag_processing/test_manager.py +++ b/airflow-core/tests/unit/dag_processing/test_manager.py @@ -51,6 +51,7 @@ from airflow.dag_processing.bundles.base import BaseDagBundle, BundleVersion from airflow.dag_processing.bundles.manager import DagBundlesManager from airflow.dag_processing.collection import update_dag_parsing_results_in_db from airflow.dag_processing.dagbag import DagBag +from airflow.dag_processing.lang_sdk_processor import LangSDKDagFileProcessorProcess from airflow.dag_processing.manager import ( BundleState, DagFileInfo, @@ -71,9 +72,9 @@ from airflow.models.dagcode import DagCode from airflow.models.serialized_dag import SerializedDagModel from airflow.models.team import Team from airflow.providers.standard.operators.empty import EmptyOperator -from airflow.sdk import DAG as SdkDAG +from airflow.sdk import DAG as SdkDAG, BaseOperator from airflow.sdk.importers import DagSourceCode -from airflow.serialization.serialized_objects import LazyDeserializedDAG +from airflow.serialization.serialized_objects import DagSerialization, LazyDeserializedDAG from airflow.utils.net import get_hostname from airflow.utils.session import create_session @@ -90,6 +91,12 @@ from tests_common.test_utils.db import ( clear_db_serialized_dags, clear_db_teams, ) +from unit.dag_processing.fake_lang_sdk import ( + FakeCoordinator, + fake_coordinator, + play_runtime, + write_native_file, +) from unit.models import TEST_DAGS_FOLDER pytestmark = pytest.mark.db_test @@ -1549,6 +1556,132 @@ class TestDagFileProcessorManager: _, kwargs = mock_start.call_args assert kwargs["subprocess_logs_to_stdout"] is expected_subprocess_logs_to_stdout + @mock.patch.object(DagFileProcessorManager, "_get_logger_for_dag_file", autospec=True) + def test_create_process_parses_a_coordinator_file_with_its_runtime(self, mock_get_logger, tmp_path): + mock_get_logger.return_value = (MagicMock(), MagicMock()) + dag_file = DagFileInfo(bundle_name="testing", rel_path=Path("dags.native"), bundle_path=tmp_path) + + with ( + fake_coordinator(), + mock.patch.object(LangSDKDagFileProcessorProcess, "start", autospec=True) as mock_start, + ): + manager = DagFileProcessorManager(max_runs=1) + manager._create_process(dag_file) + + kwargs = mock_start.call_args.kwargs + assert (kwargs["path"], kwargs["dag_file_rel_path"]) == (tmp_path / "dags.native", "dags.native") + assert kwargs["client"] is manager.client + + @pytest.mark.parametrize("rel_path", ["my_dag.py", "dags.fake"]) + @mock.patch.object(DagFileProcessorManager, "_get_logger_for_dag_file", autospec=True) + def test_create_process_keeps_the_python_parse_for_other_files(self, mock_get_logger, rel_path, tmp_path): + mock_get_logger.return_value = (MagicMock(), MagicMock()) + dag_file = DagFileInfo(bundle_name="testing", rel_path=Path(rel_path), bundle_path=tmp_path) + + with ( + fake_coordinator(), + mock.patch.object(DagFileProcessorProcess, "start", autospec=True) as mock_start, + mock.patch.object(LangSDKDagFileProcessorProcess, "start", autospec=True) as mock_lang_sdk_start, + ): + DagFileProcessorManager(max_runs=1)._create_process(dag_file) + + mock_start.assert_called_once() + mock_lang_sdk_start.assert_not_called() + + @mock.patch.object(DagFileProcessorManager, "_get_logger_for_dag_file", autospec=True) + def test_create_process_drops_callbacks_for_a_coordinator_file(self, mock_get_logger, tmp_path, caplog): + mock_get_logger.return_value = (MagicMock(), MagicMock()) + dag_file = DagFileInfo(bundle_name="testing", rel_path=Path("dags.native"), bundle_path=tmp_path) + callback = DagCallbackRequest( + filepath="dags.native", + dag_id="native_dag", + run_id="run", + bundle_name="testing", + bundle_version=None, + is_failure_callback=True, + ) + + with ( + fake_coordinator(), + mock.patch.object(LangSDKDagFileProcessorProcess, "start", autospec=True) as mock_start, + caplog.at_level(logging.WARNING, logger="airflow.dag_processing.manager"), + ): + manager = DagFileProcessorManager(max_runs=1) + manager._callback_to_execute[dag_file] = [callback] + manager._create_process(dag_file) + + assert "callbacks" not in mock_start.call_args.kwargs + assert dag_file not in manager._callback_to_execute + assert [r.getMessage() for r in caplog.records if r.levelno == logging.WARNING] == [ + "Dropping 1 callbacks for dags.native: Lang-SDK runtimes do not run callbacks" + ] + + @mock.patch.object(FakeCoordinator, "parse_dag", autospec=True) + @mock.patch.object( + DagFileProcessorManager, "_find_files_in_bundle", autospec=True, return_value=[Path("good.native")] + ) + def test_coordinator_files_are_persisted( + self, mock_find_files, mock_parse_dag, tmp_path, configure_testing_dag_bundle + ): + def reply(request, comms): + with SdkDAG("native_dag", schedule=None) as dag: + BaseOperator(task_id="extract") + data = DagSerialization.to_dict(dag) + data["dag"].update(fileloc=request.file, relative_fileloc="good.native") + return DagFileParsingResult( + fileloc=request.file, serialized_dags=[LazyDeserializedDAG(data=data)] + ) + + mock_parse_dag.side_effect = play_runtime(reply) + write_native_file(tmp_path / "good.native") + + with fake_coordinator(), configure_testing_dag_bundle(tmp_path): + DagFileProcessorManager(max_runs=1, processor_timeout=60).run() + + with create_session() as session: + serialized_dag = session.scalar( + select(SerializedDagModel).where(SerializedDagModel.dag_id == "native_dag") + ) + dag_code = session.scalar(select(DagCode).where(DagCode.dag_id == "native_dag")) + + assert serialized_dag.data["dag"]["tasks"][0]["__var"]["task_id"] == "extract" + assert dag_code.source_code == (tmp_path / "good.native").read_text() + + @mock.patch.object(FakeCoordinator, "parse_dag", autospec=True) + @mock.patch.object( + DagFileProcessorManager, + "_find_files_in_bundle", + autospec=True, + return_value=[Path("garbage.native"), Path("python_dag.py")], + ) + def test_an_invalid_frame_does_not_stop_other_files_parsing( + self, mock_find_files, mock_parse_dag, tmp_path, configure_testing_dag_bundle + ): + def reply(request, comms): + comms.socket.sendall(bytes.fromhex("00000003c1c1c1")) + time.sleep(60) + + mock_parse_dag.side_effect = play_runtime(reply) + write_native_file(tmp_path / "garbage.native") + (tmp_path / "python_dag.py").write_text( + "from airflow.sdk import DAG\nfrom airflow.sdk.bases.operator import BaseOperator\n\n" + 'with DAG("python_dag", schedule=None):\n BaseOperator(task_id="task")\n' + ) + + with fake_coordinator(), configure_testing_dag_bundle(tmp_path): + manager = DagFileProcessorManager(max_runs=1, processor_timeout=60) + manager.run() + + with create_session() as session: + dag_ids = session.scalars(select(SerializedDagModel.dag_id)).all() + import_errors = session.scalars(select(ParseImportError)).all() + + assert dag_ids == ["python_dag"] + [import_error] = import_errors + assert import_error.filename == "garbage.native" + assert import_error.stacktrace.startswith("The Lang-SDK runtime sent an invalid frame: ") + assert manager.selector.get_map() == {} + def test_terminate_orphan_processes_kills_then_closes_processor(self): manager = DagFileProcessorManager(max_runs=1) processor, _ = self.mock_processor() diff --git a/generated/known_sdk_imports_in_core.txt b/generated/known_sdk_imports_in_core.txt index b906ba25855..1172b79aac8 100644 --- a/generated/known_sdk_imports_in_core.txt +++ b/generated/known_sdk_imports_in_core.txt @@ -4,6 +4,7 @@ airflow-core/src/airflow/cli/commands/task_command.py::7 airflow-core/src/airflow/cli/commands/triggerer_command.py::1 airflow-core/src/airflow/configuration.py::1 airflow-core/src/airflow/dag_processing/dagbag.py::3 +airflow-core/src/airflow/dag_processing/lang_sdk_processor.py::7 airflow-core/src/airflow/dag_processing/manager.py::4 airflow-core/src/airflow/dag_processing/processor.py::14 airflow-core/src/airflow/exceptions.py::1 diff --git a/task-sdk/src/airflow/sdk/coordinators/_dag_importer.py b/task-sdk/src/airflow/sdk/coordinators/_dag_importer.py new file mode 100644 index 00000000000..048c209b5e0 --- /dev/null +++ b/task-sdk/src/airflow/sdk/coordinators/_dag_importer.py @@ -0,0 +1,90 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""The Dag importer a coordinator hands out for its native Dag files.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, ClassVar + +from airflow.sdk.importers.base import ( + AbstractDagImporter, + DagImportError, + DagImportResult, + FilesystemDagDefinition, + find_file_dag_definitions, +) + +if TYPE_CHECKING: + from collections.abc import Iterator + from pathlib import Path + + from airflow.dag_processing.bundles.base import BaseDagBundle # noqa: SDK002 + from airflow.sdk.coordinators._subprocess import SubprocessCoordinator + from airflow.sdk.importers.base import DagDefinition + + +class CoordinatorDagImporter(AbstractDagImporter[FilesystemDagDefinition]): + """ + Claim the native Dag files of a coordinator's artifacts, which the coordinator's runtime parses. + + The Dag processor does not call :meth:`import_definition`: it runs the runtime itself and stores + the Dags the runtime serialized. A Dag bag, such as a CLI command's, reports such a file as an + import error. + + Subclasses set :attr:`artifact_suffix` and :attr:`supported_extensions`, and implement + :meth:`get_source_code`. + """ + + artifact_suffix: ClassVar[str] + """The file name suffix of the artifacts this importer claims, such as ``.min.mjs``.""" + + supported_extensions: list[str] + """ + The extensions a registry routes to this importer, such as ``[".mjs"]`` for ``.min.mjs`` artifacts. + + A registry routes by the last suffix alone. + """ + + def __init__(self, *, coordinator: SubprocessCoordinator) -> None: + self.coordinator = coordinator + + def can_handle(self, definition: DagDefinition | str | Path) -> bool: + return str(definition).endswith(self.artifact_suffix) + + def list_dag_definitions( + self, bundle: BaseDagBundle, *, safe_mode: bool = True + ) -> Iterator[FilesystemDagDefinition]: + for definition in find_file_dag_definitions(bundle.path, self.supported_extensions): + if definition.path.name.endswith(self.artifact_suffix) and self.might_contain_dag( + definition, safe_mode + ): + yield definition + + def import_definition( + self, definition: FilesystemDagDefinition, bundle: BaseDagBundle + ) -> DagImportResult: + """Report that only the Dag processor parses *definition*.""" + return DagImportResult( + definition=definition, + errors=[ + DagImportError( + source_reference=repr(definition), + message="A native Lang-SDK Dag is parsed only by the Dag processor", + ) + ], + ) diff --git a/task-sdk/src/airflow/sdk/coordinators/_subprocess.py b/task-sdk/src/airflow/sdk/coordinators/_subprocess.py index 0a23d4f830b..0c41bf82503 100644 --- a/task-sdk/src/airflow/sdk/coordinators/_subprocess.py +++ b/task-sdk/src/airflow/sdk/coordinators/_subprocess.py @@ -62,6 +62,7 @@ if TYPE_CHECKING: from airflow.dag_processing.bundles.base import BaseDagBundle # noqa: SDK002 from airflow.sdk.api.client import Client from airflow.sdk.api.datamodels._generated import TaskInstance + from airflow.sdk.coordinators._dag_importer import CoordinatorDagImporter Tracked = TypeVar("Tracked", socket.socket, subprocess.Popen) @@ -530,6 +531,14 @@ class SubprocessCoordinator(BaseCoordinator): configured_roots=[str(root) for root in self._configured_roots], ) + @classmethod + def get_dag_importer_class(cls) -> type[CoordinatorDagImporter] | None: + return None + + def get_dag_importer(self) -> CoordinatorDagImporter | None: + importer_cls = self.get_dag_importer_class() + return None if importer_cls is None else importer_cls(coordinator=self) + @classmethod def get_parsed_bundles(cls, kwargs: Mapping[str, Any]) -> frozenset[str] | None: if cls._explicit_root_kwarg is not None and kwargs.get(cls._explicit_root_kwarg): diff --git a/task-sdk/tests/task_sdk/coordinators/test_dag_importer.py b/task-sdk/tests/task_sdk/coordinators/test_dag_importer.py new file mode 100644 index 00000000000..abbada9cba7 --- /dev/null +++ b/task-sdk/tests/task_sdk/coordinators/test_dag_importer.py @@ -0,0 +1,70 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from airflow.sdk.coordinators._dag_importer import CoordinatorDagImporter +from airflow.sdk.coordinators._subprocess import SubprocessCoordinator +from airflow.sdk.importers import DagSourceCode, FilesystemDagDefinition + + +class _BundleImporter(CoordinatorDagImporter): + artifact_suffix = ".min.mjs" + supported_extensions = [".mjs"] + + def get_source_code(self, definition) -> DagSourceCode: + return DagSourceCode("", "typescript") + + def might_contain_dag(self, definition, safe_mode: bool) -> bool: + return not definition.path.name.startswith("skip") + + [email protected] +def importer() -> _BundleImporter: + return _BundleImporter(coordinator=MagicMock(spec=SubprocessCoordinator)) + + [email protected](("path", "expected"), [("dags/main.min.mjs", True), ("dags/main.mjs", False)]) +def test_handles_only_its_artifacts(importer, path, expected): + assert importer.can_handle(path) is expected + + +def test_lists_only_its_artifacts(importer, tmp_path): + for name in ("main.min.mjs", "helper.mjs", "skip.min.mjs"): + (tmp_path / name).write_text("") + + definitions = list(importer.list_dag_definitions(SimpleNamespace(name="testing", path=tmp_path))) + + assert [d.path.name for d in definitions] == ["main.min.mjs"] + + +def test_import_definition_reports_that_only_the_dag_processor_parses_it(importer, tmp_path): + bundle_file = tmp_path / "main.min.mjs" + bundle_file.write_text("") + definition = FilesystemDagDefinition(bundle_file) + + result = importer.import_definition(definition, SimpleNamespace(name="testing", path=tmp_path)) + + assert result.dags == [] + assert [error.message for error in result.errors] == [ + "A native Lang-SDK Dag is parsed only by the Dag processor" + ] diff --git a/task-sdk/tests/task_sdk/coordinators/test_subprocess.py b/task-sdk/tests/task_sdk/coordinators/test_subprocess.py index 6de738fb894..6be4820fc5b 100644 --- a/task-sdk/tests/task_sdk/coordinators/test_subprocess.py +++ b/task-sdk/tests/task_sdk/coordinators/test_subprocess.py @@ -987,6 +987,18 @@ class TestServesBundle: assert _StubSubprocessCoordinator.get_parsed_bundles(kwargs) == expected +class TestGetDagImporter: + def test_returns_none_without_an_importer_class(self): + assert _StubSubprocessCoordinator(command=["x"]).get_dag_importer() is None + + def test_builds_the_importer_class_for_the_coordinator(self): + importer_cls = MagicMock() + with patch.object(_StubSubprocessCoordinator, "get_dag_importer_class", return_value=importer_cls): + coordinator = _StubSubprocessCoordinator(command=["x"]) + assert coordinator.get_dag_importer() is importer_cls.return_value + importer_cls.assert_called_once_with(coordinator=coordinator) + + class TestInitRootSource: """Execute-time root resolution, dispatched on the classified mode."""
