jason810496 commented on code in PR #74035:
URL: https://github.com/apache/airflow/pull/74035#discussion_r4163180558
##########
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:
Done in 76b203ed12. `_add_callback_to_queue` drops a callback for a native
file before the file is queued, with one warning naming the request type and
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:
Done in 2a253c147f. `before_run()` builds each bundle's registry before
`gc.freeze()`, so a broken `[sdk] coordinators` config shows up at startup.
##########
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:
It's a bridge until #73457 lets importers own their parse process. Without
the fd-0 relay, the runtime has to connect to the manager-side process
directly. e3bd6f0b0b adds this to ADR-0010. #74043 keeps `DagImportResult.dags`
as `list[DAG]`, so that part of the ADR still holds.
##########
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:
Done in 77bd5b757d and f13c43fa99. There is now one Task SDK helper that
returns the claiming coordinator's key, and core wraps it. The message says to
give the Dag's tasks their own queue and map it to that coordinator. The task
is marked failed without retries.
##########
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:
Done in 1697091e40. Once the runtime has exited, the grace path,
`_kill_runtime` and `close()` kill what it left in its process group.
`_kill_runtime` now waits with a timeout. 917db98d4a makes the test wait for
the kill to land.
##########
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:
Done in 4c5f7e2572. The error names `[core] dagbag_import_timeout` (or the
`get_dagbag_import_timeout` policy). In `run()`, before the runtime reports its
version, it names `[dag_processor] dag_file_processor_timeout`.
##########
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:
Done in 778193ac00: `run()`, the client override and the gate move to
#74043. There, ff48b1f19a makes the client optional on the base class, so there
is no `type: ignore`.
##########
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:
I'd keep them private for now. The coordinator plumbing is still an internal
contract between core and the SDK, and core pins the SDK minor version.
77bd5b757d drops the runtime `noqa: SDK001` in importer_routing.py and
regenerates `known_sdk_imports_in_core.txt`.
##########
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:
Done in 77bd5b757d. A case with two coordinators claiming one extension now
asserts the helper returns `None`, without patching `_get_registry`.
##########
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:
Done in fe4b272b1a.
##########
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:
Done in 85c8608063. A fixture now clears import errors before and after the
test.
##########
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:
Done in bfcc16ccfc.
##########
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:
Done in 77bd5b757d. The test only uses `conf_vars` now. Since #74042 makes
`get_dag_importer` the only hook, `NativeCoordinator` overrides 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]