pierrejeambrun commented on code in PR #74035:
URL: https://github.com/apache/airflow/pull/74035#discussion_r4182889946


##########
airflow-core/src/airflow/dag_processing/lang_sdk_processor.py:
##########
@@ -0,0 +1,522 @@
+# 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 contextlib import suppress
+from pathlib import Path
+from socket import socket
+from typing import TYPE_CHECKING, Annotated, ClassVar, Literal, get_args
+
+import attrs
+import msgspec
+import psutil
+from pydantic import BaseModel, Field, TypeAdapter
+
+from airflow import settings
+from airflow.dag_processing.dagbag import _get_bundle_team_name, 
_validate_executor_fields
+from airflow.dag_processing.importer_routing import get_claiming_importer
+from airflow.dag_processing.processor import (
+    BaseDagFileProcessorProcess,
+    DagFileParseRequest,
+    DagFileParsingResult,
+    ToManager,
+)
+from airflow.exceptions import DeserializationError
+from airflow.models.pool import Pool
+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, 
_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.enums import Encoding
+from airflow.serialization.serialized_objects import DagSerialization, 
LazyDeserializedDAG
+
+if TYPE_CHECKING:
+    from collections.abc import Generator
+
+    from structlog.typing import FilteringBoundLogger
+
+    from airflow.sdk.execution_time.supervisor import RequestHandler, 
RequestResult
+    from airflow.serialization.definitions.dag import SerializedDAG
+    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
+
+_IMPORT_TIMEOUT_SETTING = "[core] dagbag_import_timeout or the 
get_dagbag_import_timeout policy"
+
+
+# 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)
+        if (importer := get_claiming_importer(msg.file, msg.bundle_name)) is 
None:
+            raise RuntimeError(f"No coordinator Dag importer claims 
{msg.file}")
+        importer.get_parsing_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.
+    """
+
+    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)
+    _group_killed: bool = attrs.field(default=False, 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
+
+    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:
+                deserialized = 
DagSerialization.validate_serialized_dag(dag.data)
+                self._apply_team_rules(dag.data, deserialized)
+            except DeserializationError as e:
+                message = f"Cannot load the serialized Dag: {e}"
+            except Exception as e:
+                message = f"{type(e).__name__}: {e}"
+            else:
+                serialized_dags.append(dag)
+                continue
+            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
+        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 _apply_team_rules(self, data: dict, dag: SerializedDAG) -> None:
+        """
+        Check each task's executor and move tasks in the default pool to the 
team's, as the Dag bag does.
+
+        The bundle's team owns the Dag. *data* is a serialized Dag that 
validates, *dag* is that Dag
+        deserialized, and *data*'s pools are changed in place.
+
+        :raises UnknownExecutorException: if a task's executor is not 
available to the team or globally.
+        """
+        _validate_executor_fields(dag, self.bundle_name)
+        if not (team_name := _get_bundle_team_name(self.bundle_name)):
+            return
+        tasks = {task.task_id: task for task in dag.tasks}
+        for encoded in data["dag"]["tasks"]:
+            task_data = encoded[Encoding.VAR]
+            if tasks[task_data["task_id"]].pool != Pool.DEFAULT_POOL_NAME:
+                continue
+            # A mapped task reads its pool from its partial kwargs first.
+            target = task_data.setdefault("partial_kwargs", {}) if 
task_data.get("_is_mapped") else task_data
+            target["pool"] = Pool.get_default_team_pool_name(team_name)
+
+    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:
+            if (importer := get_claiming_importer(file, self.bundle_name)) is 
None:
+                raise RuntimeError(f"No coordinator 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)]),
+    }
+
+    @property
+    def is_ready(self) -> bool:
+        self._verify_connections()

Review Comment:
   `_verify_connections` should probably be protected by a try catch or 
something. (directly inside or at call site)
   
   In `_accept_connection` it's not a problem cause it's only run in 
   ```
     proc.selector.register(
         listener,
         selectors.EVENT_READ,
         (functools.partial(proc._accept_connection, channel=channel), 
proc._on_socket_closed),
     )
   ```
   
   Which is protected 
   ```
     socket_handler, on_close = key.data
     try:
         need_more = socket_handler(key.fileobj)   # == 
proc._accept_connection(listener, channel=...)
     except (BrokenPipeError, ConnectionResetError):
         need_more = False
   ```
   
   
   But that's implicit, makes it prone to errors.



##########
airflow-core/src/airflow/dag_processing/lang_sdk_processor.py:
##########
@@ -0,0 +1,522 @@
+# 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 contextlib import suppress
+from pathlib import Path
+from socket import socket
+from typing import TYPE_CHECKING, Annotated, ClassVar, Literal, get_args
+
+import attrs
+import msgspec
+import psutil
+from pydantic import BaseModel, Field, TypeAdapter
+
+from airflow import settings
+from airflow.dag_processing.dagbag import _get_bundle_team_name, 
_validate_executor_fields
+from airflow.dag_processing.importer_routing import get_claiming_importer
+from airflow.dag_processing.processor import (
+    BaseDagFileProcessorProcess,
+    DagFileParseRequest,
+    DagFileParsingResult,
+    ToManager,
+)
+from airflow.exceptions import DeserializationError
+from airflow.models.pool import Pool
+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, 
_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.enums import Encoding
+from airflow.serialization.serialized_objects import DagSerialization, 
LazyDeserializedDAG
+
+if TYPE_CHECKING:
+    from collections.abc import Generator
+
+    from structlog.typing import FilteringBoundLogger
+
+    from airflow.sdk.execution_time.supervisor import RequestHandler, 
RequestResult
+    from airflow.serialization.definitions.dag import SerializedDAG
+    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
+
+_IMPORT_TIMEOUT_SETTING = "[core] dagbag_import_timeout or the 
get_dagbag_import_timeout policy"
+
+
+# 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)
+        if (importer := get_claiming_importer(msg.file, msg.bundle_name)) is 
None:
+            raise RuntimeError(f"No coordinator Dag importer claims 
{msg.file}")
+        importer.get_parsing_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.
+    """
+
+    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)
+    _group_killed: bool = attrs.field(default=False, 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
+
+    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:
+                deserialized = 
DagSerialization.validate_serialized_dag(dag.data)
+                self._apply_team_rules(dag.data, deserialized)
+            except DeserializationError as e:
+                message = f"Cannot load the serialized Dag: {e}"
+            except Exception as e:
+                message = f"{type(e).__name__}: {e}"
+            else:
+                serialized_dags.append(dag)
+                continue
+            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
+        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 _apply_team_rules(self, data: dict, dag: SerializedDAG) -> None:
+        """
+        Check each task's executor and move tasks in the default pool to the 
team's, as the Dag bag does.
+
+        The bundle's team owns the Dag. *data* is a serialized Dag that 
validates, *dag* is that Dag
+        deserialized, and *data*'s pools are changed in place.
+
+        :raises UnknownExecutorException: if a task's executor is not 
available to the team or globally.
+        """
+        _validate_executor_fields(dag, self.bundle_name)
+        if not (team_name := _get_bundle_team_name(self.bundle_name)):
+            return
+        tasks = {task.task_id: task for task in dag.tasks}
+        for encoded in data["dag"]["tasks"]:
+            task_data = encoded[Encoding.VAR]
+            if tasks[task_data["task_id"]].pool != Pool.DEFAULT_POOL_NAME:
+                continue
+            # A mapped task reads its pool from its partial kwargs first.
+            target = task_data.setdefault("partial_kwargs", {}) if 
task_data.get("_is_mapped") else task_data
+            target["pool"] = Pool.get_default_team_pool_name(team_name)
+
+    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:
+            if (importer := get_claiming_importer(file, self.bundle_name)) is 
None:
+                raise RuntimeError(f"No coordinator 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)]),
+    }
+
+    @property
+    def is_ready(self) -> bool:
+        self._verify_connections()

Review Comment:
   is_ready calls _verify_connections() directly, which can reach 
_register_comm's send_msg(...) — and that's an unguarded sendall. If the 
runtime dies right after connecting but before this gets to run (which the 
docstring says can happen, deferred to "the next is_ready poll"), wouldn't the 
BrokenPipeError/OSError propagate out of is_ready and take down 
_collect_results's whole loop? _service_processor_sockets already treats this 
exact exception pair as "connection gone" for the _accept_connection path — 
could _verify_connections catch the same here?
   



-- 
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]

Reply via email to