kaxil commented on code in PR #74035:
URL: https://github.com/apache/airflow/pull/74035#discussion_r4159729029
##########
airflow-core/src/airflow/dag_processing/manager.py:
##########
@@ -1444,26 +1450,35 @@ def client(self) -> Client:
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,
subprocess_logs_to_stdout=conf.get("logging",
"dag_processor_log_target") == "stdout",
client=self.client,
)
+ if get_claiming_coordinator(dag_file.absolute_path,
dag_file.bundle_name) is not None:
Review Comment:
A callback for a native file still forces a full parse here.
`_add_callback_to_queue` queues the file with the callback's pinned
`bundle_version` and checkout path, the callbacks are dropped, and
`LangSDKDagFileProcessorProcess` parses that old checkout.
`handle_parsing_result` then persists the result under
`self._bundle_versions[bundle_name]`, the current version. On a Git bundle, a
heartbeat timeout of a native task from an older run (the scheduler sends that
`TaskCallbackRequest` without checking for callbacks) rewrites the current
serialized Dag with the old code until the next regular parse. The Python path
never persists from a callback run, since `_parse_file` returns `None`. Could
native callbacks be dropped in `_add_callback_to_queue`, before the file is
queued, with one log line per request naming its type and the
dag_id/run_id/task_id?
##########
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)
Review Comment:
The manager first builds a bundle's registry here, from `_create_process`,
which runs after `gc.freeze()` in `before_run()`. So the registry, every
coordinator `for_bundle` builds and their modules stay out of the frozen set
that each forked parse child shares. Warming
`get_importer_registry(bundle.name)` per bundle before the freeze would avoid
that, and would also surface a broken `[sdk] coordinators` config once at
startup instead of per file.
##########
task-sdk/src/airflow/sdk/execution_time/task_runner.py:
##########
@@ -1030,6 +1042,17 @@ def parse(what: StartupDetails, log: Logger) ->
RuntimeTaskInstance:
bundle_prepare_ms = int((time.monotonic() - bundle_prepare_start) * 1000)
dag_absolute_path = os.fspath(Path(bundle_instance.path,
what.dag_rel_path))
+ if _is_lang_sdk_dag_file(dag_absolute_path, bundle_info.name):
+ log.error(
+ "A task of a native Lang-SDK Dag cannot run in Python. Route its
queue to the coordinator "
Review Comment:
A few things about this exit. "Route its queue" is risky advice when the
queue is `default`, which is what a task gets when the author forgets to set
one, since mapping `default` sends every Python task on it to the runtime;
"give the Dag's tasks their own queue and route that" is safer.
`_is_lang_sdk_dag_file` throws the coordinator away, so the message can't say
which one to route to, and it duplicates `get_claiming_coordinator`; one shared
helper that returns the coordinator would fix both. The bare `sys.exit(1)` also
lets the task retry `retries` times on a misconfiguration that can't fix itself.
##########
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(
Review Comment:
ADR-0010 makes `import_definition` the integration point and says the
coordinator is reached through the importer, never through the manager's
file-to-process routing. This does the reverse: the manager checks the
importer's type and starts `LangSDKDagFileProcessorProcess` itself. The
no-fd-0-bridge deviation explains part of that, but the routing move is bigger,
and #73457 is heading the other way (importers owning their parse process). Is
routing in the manager the intended end state, or a bridge until #73457 lands?
Either way, could ADR-0010 be amended with this layer? It also still says
`DagImportResult.dags` is `list[DAG]`, which #74043 changes.
##########
airflow-core/src/airflow/dag_processing/lang_sdk_processor.py:
##########
@@ -0,0 +1,521 @@
+# 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 import settings
+from airflow.configuration import conf
+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 schema version and the import timeout of the runtime the parse
child is about to exec."""
+
+ schema_version: str | None
+ import_timeout: float | None = None
+ """Seconds from the start of the parse; ``None`` means no timeout."""
+ type: Literal["LangSDKRuntimeSchemaVersion"] =
"LangSDKRuntimeSchemaVersion"
+
+
+def _get_import_timeout(path: str) -> float | None:
+ """Return the ``get_dagbag_import_timeout`` policy's timeout for *path*;
``None`` means none."""
+ timeout = settings.get_dagbag_import_timeout(path)
+ if not isinstance(timeout, (int, float)):
+ raise TypeError(f"Value ({timeout}) from get_dagbag_import_timeout
must be int or float")
+ return timeout if timeout > 0 else None
+
+
+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,
import_timeout=import_timeout))
+
+ try:
+ # The policy is user code: it runs in this child, where a failure is
only this file's import error.
+ import_timeout = _get_import_timeout(msg.file)
+ 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)
+ _import_timeout: float | 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,
+ 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. The file's import
+ timeout bounds the parse, and ``[dag_processor]
dag_file_processor_timeout`` until the parse child
+ reports it.
+ """
+ processor_timeout = conf.getfloat("dag_processor",
"dag_file_processor_timeout")
+ 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:
+ timeout = proc._import_timeout if
proc._schema_version_reported else processor_timeout
+ if timeout is not None and time.monotonic() -
proc.start_time > timeout:
+ # Unlike is_ready, this does not wait for an exited
runtime's leftover processes,
+ # which can hold its sockets open. close() closes them.
+ proc._time_out(timeout)
+ break
+ proc._service_subprocess(max_wait_time=0.1)
+ 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._import_timeout = msg.import_timeout
+ 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
Review Comment:
If the runtime sends its result and exits while a process it started still
holds stdout or stderr, `is_ready` stays False on `_open_sockets`, and this
grace kill, `_kill_runtime` and the manager's `kill()` at
`dag_file_processor_timeout` all return early because `_exit_code` is set. The
manager then drops the processor without `handle_parsing_result`, so the valid
result is lost and the leftover keeps running, one more per reparse. `run()`
has the same gap, which is why
`test_the_import_timeout_holds_after_the_runtime_exits` calls `os.killpg`
itself. Since the child has its own process group, could the grace path and
`close()` kill the group once the leader has exited? Related: `_kill_runtime`'s
`wait(timeout=None)` runs on the manager loop, so a runtime stuck in
uninterruptible IO stalls the whole processor.
##########
airflow-core/src/airflow/dag_processing/lang_sdk_processor.py:
##########
@@ -0,0 +1,521 @@
+# 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 import settings
+from airflow.configuration import conf
+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
Review Comment:
Core now imports `_is_connection_from_pid` and `_start_server`, both
private, from the Task SDK. Since airflow-core depends on them across the
distribution boundary, should they become public names? Related:
importer_routing.py imports from `airflow.sdk` at runtime with `# noqa:
SDK001`, where the other `SDK001` uses in core are all under `TYPE_CHECKING`.
Regenerating `known_sdk_imports_in_core.txt` for it, as this file got, keeps
that inventory accurate.
##########
airflow-core/src/airflow/dag_processing/manager.py:
##########
@@ -1444,26 +1450,35 @@ def client(self) -> Client:
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(
Review Comment:
Building the kwargs as `dict[str, Any]` turns off mypy's argument checking
for both `start()` calls. Two explicit calls, or a small typed helper for the
shared values, would keep it.
##########
airflow-core/src/airflow/dag_processing/lang_sdk_processor.py:
##########
@@ -0,0 +1,521 @@
+# 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 import settings
+from airflow.configuration import conf
+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 schema version and the import timeout of the runtime the parse
child is about to exec."""
+
+ schema_version: str | None
+ import_timeout: float | None = None
+ """Seconds from the start of the parse; ``None`` means no timeout."""
+ type: Literal["LangSDKRuntimeSchemaVersion"] =
"LangSDKRuntimeSchemaVersion"
+
+
+def _get_import_timeout(path: str) -> float | None:
+ """Return the ``get_dagbag_import_timeout`` policy's timeout for *path*;
``None`` means none."""
+ timeout = settings.get_dagbag_import_timeout(path)
+ if not isinstance(timeout, (int, float)):
+ raise TypeError(f"Value ({timeout}) from get_dagbag_import_timeout
must be int or float")
+ return timeout if timeout > 0 else None
+
+
+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,
import_timeout=import_timeout))
+
+ try:
+ # The policy is user code: it runs in this child, where a failure is
only this file's import error.
+ import_timeout = _get_import_timeout(msg.file)
+ 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]
Review Comment:
The `type: ignore[assignment]` is here because the base from #74040 keeps
`client: Client` required while making `logger_filehandle` optional for the
same client-less path. Nothing in this PR starts the process without a client
(the manager always passes one, and the first caller is the Dag bag in #74043),
so could `run()`, this override and the `None`-client gate in `_handle_request`
move up to #74043? There the base could own the client-less case, rather than a
subclass narrowing what `_common_request_handlers` is typed against.
##########
task-sdk/tests/task_sdk/execution_time/test_task_runner.py:
##########
@@ -902,6 +907,93 @@ def test_parse_module_in_bundle_root(tmp_path: Path,
make_ti_context):
assert ti.task.dag.dag_id == "dag_name"
+class NativeDagImporter(CoordinatorDagImporter):
+ artifact_suffix = ".native"
+ supported_extensions = [".native"]
+
+ def get_source_code(self, definition):
+ return DagSourceCode(source_code=definition.read_text(),
language="native")
+
+
[email protected](kw_only=True)
+class NativeCoordinator(SubprocessCoordinator):
+ """A coordinator whose Dag importer claims ``.native`` files in every
bundle."""
+
+ def get_dag_importer(self):
+ return NativeDagImporter(coordinator=self)
+
+
+@patch("airflow.dag_processing.dagbag.BundleDagBag", autospec=True)
+def test_parse_rejects_a_task_of_a_native_dag(mock_bag, tmp_path: Path,
make_ti_context):
+ tmp_path.joinpath("dag.native").write_text("{}")
+ what = StartupDetails(
+ ti=TaskInstance(
+ id=uuid7(),
+ task_id="a",
+ dag_id="native_dag",
+ run_id="c",
+ try_number=1,
+ dag_version_id=uuid7(),
+ queue="default",
+ ),
+ dag_rel_path="dag.native",
+ bundle_info=BundleInfo(name="my-bundle", version=None),
+ ti_context=make_ti_context(),
+ start_date=timezone.utcnow(),
+ sentry_integration="",
+ )
+ bundle_config = [
+ {
+ "name": "my-bundle",
+ "classpath": "airflow.dag_processing.bundles.local.LocalDagBundle",
+ "kwargs": {"path": str(tmp_path), "refresh_interval": 1},
+ }
+ ]
+ coordinators = {"native": {"classpath": f"{__name__}.NativeCoordinator",
"kwargs": {}}}
+ log = mock.Mock()
+
+ reset_importer_registry()
+ try:
+ with (
+ patch.dict(
+ os.environ,
+ {
+ "AIRFLOW__DAG_PROCESSOR__DAG_BUNDLE_CONFIG_LIST":
json.dumps(bundle_config),
+ "AIRFLOW__SDK__COORDINATORS": json.dumps(coordinators),
+ },
+ ),
+ pytest.raises(SystemExit, match="1"),
+ ):
+ parse(what, log)
+ finally:
+ reset_importer_registry()
+
+ mock_bag.assert_not_called()
+ log.error.assert_called_once_with(
+ "A task of a native Lang-SDK Dag cannot run in Python. Route its queue
to the coordinator "
+ "that parses the Dag, with [sdk] queue_to_coordinator",
+ dag_id="native_dag",
+ task_id="a",
+ queue="default",
+ path="dag.native",
+ )
+
+
[email protected](
+ ("file_name", "expected"),
+ [("dag.native", True), ("dag.py", False), ("dag.pyc", False), ("dags.zip",
False)],
+)
+def test_is_lang_sdk_dag_file(file_name, expected):
Review Comment:
Nothing reaches the `except Exception` in `_is_lang_sdk_dag_file`: flipping
it to `return True` keeps this file green, and that flip would fail every
Python task on a worker whose `[sdk] coordinators` raises (two coordinators
claiming one extension, say). A case with a clashing config asserting `False`
would pin it. `test_get_claiming_coordinator_without_a_registry` has the same
gap, since it patches `_get_registry`, which is where the branch lives.
##########
task-sdk/tests/task_sdk/execution_time/test_task_runner.py:
##########
@@ -902,6 +907,93 @@ def test_parse_module_in_bundle_root(tmp_path: Path,
make_ti_context):
assert ti.task.dag.dag_id == "dag_name"
+class NativeDagImporter(CoordinatorDagImporter):
+ artifact_suffix = ".native"
+ supported_extensions = [".native"]
+
+ def get_source_code(self, definition):
+ return DagSourceCode(source_code=definition.read_text(),
language="native")
+
+
[email protected](kw_only=True)
+class NativeCoordinator(SubprocessCoordinator):
+ """A coordinator whose Dag importer claims ``.native`` files in every
bundle."""
+
+ def get_dag_importer(self):
Review Comment:
Could `NativeCoordinator` override the `get_dag_importer_class` classmethod
rather than `get_dag_importer`? That's the hook real coordinators use, and with
only the instance method the class reports no importer, so
`_check_dag_file_claims` skips it. Also, `conf_vars({("sdk", "coordinators"):
...})` already resets the registry and coordinator manager on enter and exit,
so the `reset_importer_registry()` try/finally and the `patch.dict` of the env
var can go.
##########
airflow-core/src/airflow/dag_processing/lang_sdk_processor.py:
##########
@@ -0,0 +1,521 @@
+# 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 import settings
+from airflow.configuration import conf
+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 schema version and the import timeout of the runtime the parse
child is about to exec."""
+
+ schema_version: str | None
+ import_timeout: float | None = None
+ """Seconds from the start of the parse; ``None`` means no timeout."""
+ type: Literal["LangSDKRuntimeSchemaVersion"] =
"LangSDKRuntimeSchemaVersion"
+
+
+def _get_import_timeout(path: str) -> float | None:
+ """Return the ``get_dagbag_import_timeout`` policy's timeout for *path*;
``None`` means none."""
+ timeout = settings.get_dagbag_import_timeout(path)
+ if not isinstance(timeout, (int, float)):
+ raise TypeError(f"Value ({timeout}) from get_dagbag_import_timeout
must be int or float")
+ return timeout if timeout > 0 else None
+
+
+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,
import_timeout=import_timeout))
+
+ try:
+ # The policy is user code: it runs in this child, where a failure is
only this file's import error.
+ import_timeout = _get_import_timeout(msg.file)
+ 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)
+ _import_timeout: float | 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,
+ 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. The file's import
+ timeout bounds the parse, and ``[dag_processor]
dag_file_processor_timeout`` until the parse child
+ reports it.
+ """
+ processor_timeout = conf.getfloat("dag_processor",
"dag_file_processor_timeout")
+ 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:
+ timeout = proc._import_timeout if
proc._schema_version_reported else processor_timeout
+ if timeout is not None and time.monotonic() -
proc.start_time > timeout:
+ # Unlike is_ready, this does not wait for an exited
runtime's leftover processes,
+ # which can hold its sockets open. close() closes them.
+ proc._time_out(timeout)
+ break
+ proc._service_subprocess(max_wait_time=0.1)
+ 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._import_timeout = msg.import_timeout
+ 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._import_timeout is not None
+ and self.parsing_result is None
+ and self._exit_code is None
+ and time.monotonic() - self.start_time > self._import_timeout
+ ):
+ self._time_out(self._import_timeout)
+ 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 _time_out(self, timeout: float) -> None:
+ if self.parsing_result is None:
+ self._set_import_error(
+ f"The Lang-SDK runtime did not parse
{self._parse_request.file} within {timeout}s"
Review Comment:
This is the whole import error a user sees, and it doesn't name the setting.
Saying it's `[core] dagbag_import_timeout` (or the `get_dagbag_import_timeout`
policy), and in `run()` before the version arrives that it's `[dag_processor]
dag_file_processor_timeout`, would tell them which knob to turn. Once #74043
lands it's also the first native-parse error a CLI user hits.
##########
airflow-core/tests/unit/dag_processing/test_dagbag.py:
##########
@@ -1530,3 +1532,21 @@ def
test_dagbag_no_bundle_path_no_syspath_modification(self, tmp_path):
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"))
Review Comment:
This commits a `ParseImportError` row, and `sync_bag_to_db` commits too, but
the `session` fixture only rolls back, so the row outlives the test. A
`finally` or fixture calling `clear_db_import_errors()` would keep later tests
clean.
##########
airflow-core/tests/unit/dag_processing/test_lang_sdk_processor.py:
##########
@@ -0,0 +1,593 @@
+#
+# 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,
+ _get_import_timeout,
+)
+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.parametrize(
+ ("policy", "error"),
+ [
+ pytest.param(
+ {"side_effect": RuntimeError("policy bug")}, "RuntimeError:
policy bug", id="raises"
+ ),
+ pytest.param(
+ {"return_value": "30"},
+ "TypeError: Value (30) from get_dagbag_import_timeout must be
int or float",
+ id="not-a-number",
+ ),
+ ],
+ )
+ def test_a_failing_import_timeout_policy_is_an_import_error(self, parse,
policy, error):
+ with patch("airflow.settings.get_dagbag_import_timeout",
autospec=True, **policy):
+ proc = parse()
+
+ assert proc.parsing_result.import_errors == {
+ "dag.native": f"Cannot start the Lang-SDK runtime: {error}"
+ }
+
+ @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, **spec) -> DagFileParsingResult:
+ return LangSDKDagFileProcessorProcess.run(
+ path=write_native_file(tmp_path / "dag.native", **spec),
+ bundle_path=tmp_path,
+ bundle_name="testing",
+ dag_file_rel_path="dag.native",
+ 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"])
+ @patch("airflow.settings.get_dagbag_import_timeout", autospec=True,
return_value=1)
+ def test_a_parse_past_the_import_timeout_is_killed(self, mock_timeout,
tmp_path, connected):
+ # A runtime that never connects leaves both listeners open when it is
killed.
+ 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,
+ ):
+ result = self._run(tmp_path, argv=["/bin/sh", "-c", "exec sleep
60"])
+
+ assert result.import_errors == {
+ "dag.native": f"The Lang-SDK runtime did not parse {tmp_path /
'dag.native'} within 1.0s"
+ }
+ [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
+
+ @patch("airflow.settings.get_dagbag_import_timeout", autospec=True,
return_value=1)
+ def test_the_import_timeout_holds_after_the_runtime_exits(self,
mock_timeout, tmp_path):
+ with patch.object(
+ LangSDKDagFileProcessorProcess,
+ "close",
+ autospec=True,
+ side_effect=LangSDKDagFileProcessorProcess.close,
+ ) as mock_close:
+ # The runtime exits, and the process it leaves behind keeps its
output open.
+ result = self._run(tmp_path, argv=["/bin/sh", "-c", "sleep 30 &
exit 0"])
+ [proc] = [c.args[0] for c in mock_close.call_args_list]
+ os.killpg(proc.pid, signal.SIGKILL)
+
+ assert result.import_errors == {
+ "dag.native": f"The Lang-SDK runtime did not parse {tmp_path /
'dag.native'} within 1.0s"
+ }
+ assert proc._exit_code == 0
+ assert not proc._open_sockets
+
+ @conf_vars({("dag_processor", "dag_file_processor_timeout"): "1"})
+ @patch.object(
+ FakeCoordinator,
+ "_build_parse_dag_command",
+ autospec=True,
+ side_effect=lambda self, *, path: time.sleep(60),
+ )
+ def
test_the_dag_file_processor_timeout_applies_until_the_import_timeout_is_reported(
+ self, mock_build_parse_dag_command, tmp_path
+ ):
+ result = self._run(tmp_path)
+
+ assert result.import_errors == {
+ "dag.native": f"The Lang-SDK runtime did not parse {tmp_path /
'dag.native'} within 1.0s"
+ }
+
+
[email protected](("configured", "expected"), [(30, 30), (0.5, 0.5),
(0, None), (-1, None)])
+@patch("airflow.settings.get_dagbag_import_timeout", autospec=True)
+def test_only_a_positive_import_timeout_applies(mock_timeout, configured,
expected):
+ mock_timeout.return_value = configured
+
+ assert _get_import_timeout("/b/dag.native") == expected
+ mock_timeout.assert_called_once_with("/b/dag.native")
+
+
+def _make_process(**kwargs) -> LangSDKDagFileProcessorProcess:
+ return LangSDKDagFileProcessorProcess(
+ id=uuid.uuid4(),
+ pid=1,
+ stdin=MagicMock(),
Review Comment:
`process`, `process_log` and `selector` are bare `MagicMock()`s that the
tests assert through (`process_log.warning.assert_called_with`,
`selector.register.call_args`), so a misspelt method would still pass. A
`spec=` on each would catch that.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]