This is an automated email from the ASF dual-hosted git repository.

ephraimbuddy pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new 9fd0d72bedb Share subprocess request dispatch without changing message 
support (#73290)
9fd0d72bedb is described below

commit 9fd0d72bedb4bcbc496c6485b6b512f21c25e011
Author: Ephraim Anierobi <[email protected]>
AuthorDate: Wed Sep 23 15:00:06 2026 +0100

    Share subprocess request dispatch without changing message support (#73290)
    
    Task and Dag subprocesses duplicate dispatch plumbing, making their 
supported message boundaries and shared behavior harder to keep aligned. Keep 
those boundaries explicit and protect the existing wire behavior as handlers 
evolve.
---
 .../src/airflow/dag_processing/processor.py        | 121 +---
 .../tests/unit/dag_processing/test_processor.py    |  65 +-
 generated/known_sdk_imports_in_core.txt            |   2 +-
 .../airflow/sdk/execution_time/request_handlers.py |   6 +-
 .../src/airflow/sdk/execution_time/supervisor.py   | 717 ++++++++++++++-------
 .../execution_time/_request_registration_types.py  |  41 ++
 .../task_sdk/execution_time/test_supervisor.py     | 168 ++++-
 7 files changed, 768 insertions(+), 352 deletions(-)

diff --git a/airflow-core/src/airflow/dag_processing/processor.py 
b/airflow-core/src/airflow/dag_processing/processor.py
index a3f12670384..fe29a314100 100644
--- a/airflow-core/src/airflow/dag_processing/processor.py
+++ b/airflow-core/src/airflow/dag_processing/processor.py
@@ -71,24 +71,8 @@ from airflow.sdk.execution_time.comms import (
     XComSequenceIndexResult,
     XComSequenceSliceResult,
 )
-from airflow.sdk.execution_time.request_handlers import (
-    handle_delete_variable,
-    handle_get_prev_successful_dag_run,
-    handle_get_previous_dag_run,
-    handle_get_previous_ti,
-    handle_get_task_states,
-    handle_get_ti_count,
-    handle_get_variable_keys,
-    handle_get_xcom,
-    handle_get_xcom_count,
-    handle_get_xcom_sequence_item,
-    handle_get_xcom_sequence_slice,
-    handle_mask_secret,
-    handle_put_variable,
-)
-from airflow.sdk.execution_time.supervisor import WatchedSubprocess
+from airflow.sdk.execution_time.supervisor import WatchedSubprocess, 
register_request_method
 from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance, 
_send_error_email_notification
-from airflow.sdk.log import mask_secret
 from airflow.serialization.serialized_objects import DagSerialization, 
LazyDeserializedDAG
 from airflow.utils.dag_version_inflation_checker import 
check_dag_file_stability
 from airflow.utils.file import iter_airflow_imports
@@ -107,6 +91,7 @@ if TYPE_CHECKING:
     from airflow.sdk.definitions.context import Context
     from airflow.sdk.definitions.dag import DAG
     from airflow.sdk.definitions.mappedoperator import MappedOperator
+    from airflow.sdk.execution_time.supervisor import RequestHandler, 
RequestResult
     from airflow.typing_compat import Self
 
 
@@ -674,76 +659,40 @@ class DagFileProcessorProcess(WatchedSubprocess, 
LoggingMixin):
             log_level=log_level,
         )
 
-    def _handle_request(self, msg: ToManager, log: FilteringBoundLogger, 
req_id: int) -> None:
-        from airflow.sdk.api.datamodels._generated import (
-            ConnectionResponse,
-            VariableResponse,
-        )
-
-        resp: BaseModel | None = None
-        dump_opts: dict[str, bool] = {}
-        if isinstance(msg, DagFileParsingResult):
-            self.parsing_result = msg
-        elif isinstance(msg, GetConnection):
-            conn = self.client.connections.get(msg.conn_id)
-            if isinstance(conn, ConnectionResponse):
-                if conn.password:
-                    mask_secret(conn.password)
-                if conn.extra:
-                    mask_secret(conn.extra)
-                conn_result = ConnectionResult.from_conn_response(conn)
-                resp = conn_result
-                dump_opts = {"exclude_unset": True, "by_alias": True}
-            else:
-                resp = conn
-        elif isinstance(msg, GetVariable):
-            var = self.client.variables.get(msg.key)
-            if isinstance(var, VariableResponse):
-                if var.value:
-                    mask_secret(var.value, var.key)
-                var_result = VariableResult.from_variable_response(var)
-                resp = var_result
-                dump_opts = {"exclude_unset": True}
-            else:
-                resp = var
-        elif isinstance(msg, GetVariableKeys):
-            resp, dump_opts = handle_get_variable_keys(self.client, msg)
-        elif isinstance(msg, PutVariable):
-            resp, dump_opts = handle_put_variable(self.client, msg)
-        elif isinstance(msg, DeleteVariable):
-            resp, dump_opts = handle_delete_variable(self.client, msg)
-        elif isinstance(msg, GetPreviousDagRun):
-            resp, dump_opts = handle_get_previous_dag_run(self.client, msg)
-        elif isinstance(msg, GetPrevSuccessfulDagRun):
-            resp, dump_opts = handle_get_prev_successful_dag_run(self.client, 
self.id)
-        elif isinstance(msg, GetXCom):
-            resp, dump_opts = handle_get_xcom(self.client, msg)
-        elif isinstance(msg, GetXComCount):
-            resp, dump_opts = handle_get_xcom_count(self.client, msg)
-        elif isinstance(msg, GetXComSequenceItem):
-            resp, dump_opts = handle_get_xcom_sequence_item(self.client, msg)
-        elif isinstance(msg, GetXComSequenceSlice):
-            resp, dump_opts = handle_get_xcom_sequence_slice(self.client, msg)
-        elif isinstance(msg, MaskSecret):
-            handle_mask_secret(msg)
-        elif isinstance(msg, GetTICount):
-            resp, dump_opts = handle_get_ti_count(self.client, msg)
-        elif isinstance(msg, GetTaskStates):
-            resp, dump_opts = handle_get_task_states(self.client, msg)
-        elif isinstance(msg, GetPreviousTI):
-            resp, dump_opts = handle_get_previous_ti(self.client, msg)
-        else:
-            log.error("Unhandled request", msg=msg)
-            self.send_msg(
-                None,
-                request_id=req_id,
-                error=ErrorResponse(
-                    detail={"status_code": 400, "message": "Unhandled 
request"},
-                ),
-            )
-            return
+    def _handle_parsing_result(
+        self, msg: DagFileParsingResult, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        self.parsing_result = msg
+        return None, {}
+
+    _request_handlers: ClassVar[dict[type[BaseModel], 
RequestHandler[DagFileProcessorProcess]]] = {
+        **WatchedSubprocess._get_shared_request_handlers(
+            DeleteVariable,
+            GetConnection,
+            GetPrevSuccessfulDagRun,
+            GetPreviousDagRun,
+            GetPreviousTI,
+            GetTICount,
+            GetTaskStates,
+            GetVariable,
+            GetVariableKeys,
+            GetXCom,
+            GetXComCount,
+            GetXComSequenceItem,
+            GetXComSequenceSlice,
+            MaskSecret,
+            PutVariable,
+        ),
+        **dict([register_request_method(DagFileParsingResult, 
_handle_parsing_result)]),
+    }
 
-        self.send_msg(resp, request_id=req_id, error=None, **dump_opts)
+    def _reject_request(self, msg, log: FilteringBoundLogger, req_id: int) -> 
None:
+        log.error("Unhandled request", msg=msg)
+        self.send_msg(
+            None,
+            request_id=req_id,
+            error=ErrorResponse(detail={"status_code": 400, "message": 
"Unhandled request"}),
+        )
 
     @property
     def is_ready(self) -> bool:
diff --git a/airflow-core/tests/unit/dag_processing/test_processor.py 
b/airflow-core/tests/unit/dag_processing/test_processor.py
index d51118f2b61..f3966f3916a 100644
--- a/airflow-core/tests/unit/dag_processing/test_processor.py
+++ b/airflow-core/tests/unit/dag_processing/test_processor.py
@@ -2326,6 +2326,63 @@ class TestDagProcessingMessageTypes:
 
 
 class TestDagFileProcessorProcess:
+    def test_registered_message_types(self):
+        expected = set(typing.get_args(typing.get_args(ToManager)[0]))
+        assert set(DagFileProcessorProcess._request_handlers) == expected
+
+    @pytest.mark.parametrize(
+        "message_type",
+        sorted(set(typing.get_args(typing.get_args(ToManager)[0])) - 
{DagFileParsingResult}, key=str),
+        ids=lambda message_type: message_type.__name__,
+    )
+    def test_reuses_shared_request_handlers(self, message_type):
+        handler = DagFileProcessorProcess._request_handlers[message_type]
+        assert handler is 
supervisor.ActivitySubprocess._request_handlers[message_type]
+        assert handler is 
supervisor.WatchedSubprocess._shared_request_handlers[message_type]
+
+    @patch.object(DagFileProcessorProcess, "send_msg", autospec=True)
+    @pytest.mark.parametrize(
+        "message_type",
+        sorted(
+            set(typing.get_args(typing.get_args(ToSupervisor)[0]))
+            - set(typing.get_args(typing.get_args(ToManager)[0])),
+            key=str,
+        ),
+        ids=lambda message_type: message_type.__name__,
+    )
+    def test_rejects_task_only_messages(self, send_msg, proc, message_type):
+        proc._handle_request(message_type.model_construct(), 
structlog.get_logger(), req_id=42)
+
+        send_msg.assert_called_once_with(
+            proc,
+            None,
+            request_id=42,
+            error=comms.ErrorResponse(detail={"status_code": 400, "message": 
"Unhandled request"}),
+        )
+        assert not proc.client.mock_calls
+
+    @patch.object(DagFileProcessorProcess, "send_msg", autospec=True)
+    def test_dispatch_parsing_result(self, send_msg, proc):
+        result = DagFileParsingResult(fileloc="test_dag.py", 
serialized_dags=[])
+        proc._handle_request(result, structlog.get_logger(), req_id=42)
+
+        assert proc.parsing_result is result
+        send_msg.assert_called_once_with(proc, None, request_id=42, error=None)
+
+    @patch.object(DagFileProcessorProcess, "send_msg", autospec=True)
+    def test_previous_successful_run_uses_process_id(self, send_msg, proc):
+        proc.client.task_instances.get_previous_successful_dagrun.return_value 
= (
+            comms.PrevSuccessfulDagRunResult()
+        )
+        proc._handle_request(
+            comms.GetPrevSuccessfulDagRun(ti_id=uuid.uuid4()), 
structlog.get_logger(), req_id=42
+        )
+
+        
proc.client.task_instances.get_previous_successful_dagrun.assert_called_once_with(proc.id)
+        send_msg.assert_called_once_with(
+            proc, comms.PrevSuccessfulDagRunResult(), request_id=42, 
error=None, exclude_unset=True
+        )
+
     @pytest.fixture
     def proc(self):
         from socket import socketpair
@@ -2388,7 +2445,9 @@ class TestDagFileProcessorProcess:
         )
 
         with (
-            patch("airflow.dag_processing.processor.mask_secret") as 
mock_mask_secret,
+            patch(
+                "airflow.sdk.execution_time.request_handlers.mask_secret", 
autospec=True
+            ) as mock_mask_secret,
             patch.object(DagFileProcessorProcess, "send_msg", autospec=True) 
as mock_send_msg,
         ):
             proc._handle_request(
@@ -2429,7 +2488,9 @@ class TestDagFileProcessorProcess:
         )
 
         with (
-            patch("airflow.dag_processing.processor.mask_secret") as 
mock_mask_secret,
+            patch(
+                "airflow.sdk.execution_time.request_handlers.mask_secret", 
autospec=True
+            ) as mock_mask_secret,
             patch.object(DagFileProcessorProcess, "send_msg", autospec=True) 
as mock_send_msg,
         ):
             proc._handle_request(
diff --git a/generated/known_sdk_imports_in_core.txt 
b/generated/known_sdk_imports_in_core.txt
index 6c212ad6f61..f05d5091514 100644
--- a/generated/known_sdk_imports_in_core.txt
+++ b/generated/known_sdk_imports_in_core.txt
@@ -7,7 +7,7 @@ airflow-core/src/airflow/dag_processing/dagbag.py::1
 airflow-core/src/airflow/dag_processing/importers/base.py::1
 airflow-core/src/airflow/dag_processing/importers/python_importer.py::7
 airflow-core/src/airflow/dag_processing/manager.py::4
-airflow-core/src/airflow/dag_processing/processor.py::15
+airflow-core/src/airflow/dag_processing/processor.py::13
 airflow-core/src/airflow/exceptions.py::1
 airflow-core/src/airflow/executors/base_executor.py::3
 airflow-core/src/airflow/jobs/triggerer_job_runner.py::18
diff --git a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py 
b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py
index a31596b3d2b..902e19447da 100644
--- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py
+++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py
@@ -21,8 +21,10 @@ Shared request handlers for supervised subprocess comms 
channels.
 These functions implement the supervisor-side logic for message types that are
 used by more than one subprocess type (tasks, callbacks, triggerer).  Each
 handler accepts a ``Client`` and a request message and returns
-``(response_model | None, dump_opts)`` so the caller can forward the result
-via ``send_msg``.
+``(response_model | None, dump_opts)``. The caller owns the wire 
acknowledgment:
+asset-state-store mutation handlers return ``(None, {})``, but callers that
+acknowledge those operations with ``OKResponse(ok=True)`` must construct that
+response instead of forwarding the empty result via ``send_msg``.
 """
 
 from __future__ import annotations
diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py 
b/task-sdk/src/airflow/sdk/execution_time/supervisor.py
index 4bdf71d1119..7c28f0aeaca 100644
--- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py
+++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py
@@ -36,9 +36,21 @@ from collections import deque
 from collections.abc import Callable, Generator
 from contextlib import contextmanager, suppress
 from datetime import datetime, timezone
+from enum import Enum, auto
 from http import HTTPStatus
 from socket import socket, socketpair
-from typing import TYPE_CHECKING, Any, BinaryIO, ClassVar, NoReturn, TextIO, 
cast
+from typing import (
+    TYPE_CHECKING,
+    Any,
+    BinaryIO,
+    ClassVar,
+    NoReturn,
+    Protocol,
+    TextIO,
+    TypeAlias,
+    TypeVar,
+    cast,
+)
 from urllib.parse import urlparse
 from uuid import UUID
 
@@ -175,7 +187,95 @@ if TYPE_CHECKING:
     from airflow.sdk.definitions.connection import Connection
     from airflow.sdk.types import RuntimeTaskInstanceProtocol as RuntimeTI
 
-__all__ = ["ActivitySubprocess", "WatchedSubprocess", "supervise", 
"supervise_task"]
+
+class ResponseSent(Enum):
+    ALREADY_SENT = auto()
+
+
+_RequestProcess = TypeVar("_RequestProcess")
+_RequestMessage = TypeVar("_RequestMessage", bound=BaseModel)
+RequestResult: TypeAlias = tuple[BaseModel | None, dict[str, bool]]
+RequestHandler: TypeAlias = Callable[
+    [_RequestProcess, BaseModel, "FilteringBoundLogger", int], RequestResult | 
ResponseSent
+]
+
+
+class _ClientSubprocess(Protocol):
+    @property
+    def client(self) -> Client: ...
+
+    @property
+    def id(self) -> UUID: ...
+
+
+def _register_request_handler(
+    message_type: type[_RequestMessage],
+    handler: Callable[
+        [_RequestProcess, _RequestMessage, FilteringBoundLogger, int], 
RequestResult | ResponseSent
+    ],
+) -> tuple[type[BaseModel], RequestHandler[_RequestProcess]]:
+    def dispatch(
+        process: _RequestProcess, msg: BaseModel, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult | ResponseSent:
+        # Registration checks the message/handler pair before the 
heterogeneous registry erases its type.
+        return handler(process, cast("_RequestMessage", msg), log, req_id)
+
+    return message_type, dispatch
+
+
+def register_request_method(
+    message_type: type[_RequestMessage],
+    method: Callable[
+        [_RequestProcess, _RequestMessage, FilteringBoundLogger, int], 
RequestResult | ResponseSent
+    ],
+) -> tuple[type[BaseModel], RequestHandler[_RequestProcess]]:
+    def dispatch(
+        process: _RequestProcess, msg: _RequestMessage, log: 
FilteringBoundLogger, req_id: int
+    ) -> RequestResult | ResponseSent:
+        bound_method = cast(
+            "Callable[[_RequestMessage, FilteringBoundLogger, int], 
RequestResult | ResponseSent]",
+            getattr(process, method.__name__),
+        )
+        return bound_method(msg, log, req_id)
+
+    return _register_request_handler(message_type, dispatch)
+
+
+def _register_client_handler(
+    message_type: type[_RequestMessage],
+    handler: Callable[[Client, _RequestMessage], RequestResult],
+) -> tuple[type[BaseModel], RequestHandler[_ClientSubprocess]]:
+    def dispatch(
+        process: _ClientSubprocess, msg: _RequestMessage, log: 
FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        return handler(process.client, msg)
+
+    return _register_request_handler(message_type, dispatch)
+
+
+def _handle_previous_successful_dag_run_request(
+    process: _ClientSubprocess, msg: GetPrevSuccessfulDagRun, log: 
FilteringBoundLogger, req_id: int
+) -> RequestResult:
+    return handle_get_prev_successful_dag_run(process.client, process.id)
+
+
+def _handle_mask_secret_request(
+    process: object, msg: MaskSecret, log: FilteringBoundLogger, req_id: int
+) -> RequestResult:
+    handle_mask_secret(msg)
+    return None, {}
+
+
+__all__ = [
+    "ActivitySubprocess",
+    "RequestHandler",
+    "RequestResult",
+    "ResponseSent",
+    "WatchedSubprocess",
+    "register_request_method",
+    "supervise",
+    "supervise_task",
+]
 
 log: FilteringBoundLogger = structlog.get_logger(logger_name="supervisor")
 
@@ -673,6 +773,37 @@ class WatchedSubprocess:
     socket handling, process monitoring, and request handling.
     """
 
+    _request_handlers: ClassVar[dict[type[BaseModel], RequestHandler[Any]] | 
None] = None
+    _shared_request_handlers: ClassVar[dict[type[BaseModel], 
RequestHandler[_ClientSubprocess]]] = dict(
+        [
+            _register_client_handler(DeleteVariable, handle_delete_variable),
+            _register_client_handler(DeleteXCom, handle_delete_xcom),
+            _register_client_handler(GetConnection, handle_get_connection),
+            _register_client_handler(GetDagRunState, handle_get_dag_run_state),
+            _register_client_handler(GetDRCount, handle_get_dr_count),
+            _register_client_handler(GetPreviousDagRun, 
handle_get_previous_dag_run),
+            _register_client_handler(GetPreviousTI, handle_get_previous_ti),
+            _register_client_handler(GetTaskStates, handle_get_task_states),
+            _register_client_handler(GetTICount, handle_get_ti_count),
+            _register_client_handler(GetVariable, handle_get_variable),
+            _register_client_handler(GetVariableKeys, 
handle_get_variable_keys),
+            _register_client_handler(GetXCom, handle_get_xcom),
+            _register_client_handler(GetXComCount, handle_get_xcom_count),
+            _register_client_handler(GetXComSequenceItem, 
handle_get_xcom_sequence_item),
+            _register_client_handler(GetXComSequenceSlice, 
handle_get_xcom_sequence_slice),
+            _register_client_handler(PutVariable, handle_put_variable),
+            _register_client_handler(SetXCom, handle_set_xcom),
+            _register_request_handler(GetPrevSuccessfulDagRun, 
_handle_previous_successful_dag_run_request),
+            _register_request_handler(MaskSecret, _handle_mask_secret_request),
+        ]
+    )
+
+    @classmethod
+    def _get_shared_request_handlers(
+        cls, *message_types: type[BaseModel]
+    ) -> dict[type[BaseModel], RequestHandler[_ClientSubprocess]]:
+        return {message_type: cls._shared_request_handlers[message_type] for 
message_type in message_types}
+
     id: UUID
 
     pid: int
@@ -1050,7 +1181,27 @@ class WatchedSubprocess:
                     otel_context.detach(token)
 
     def _handle_request(self, msg, log: FilteringBoundLogger, req_id: int) -> 
None:
-        raise NotImplementedError()
+        if self._request_handlers is None:
+            raise NotImplementedError(f"{type(self).__name__} must declare its 
request handlers")
+        handler = self._request_handlers.get(type(msg))
+        if handler is None:
+            self._reject_request(msg, log, req_id)
+            return
+        result = handler(self, msg, log, req_id)
+        if result is not ResponseSent.ALREADY_SENT:
+            resp, dump_opts = result
+            self.send_msg(resp, request_id=req_id, error=None, **dump_opts)
+
+    def _reject_request(self, msg, log: FilteringBoundLogger, req_id: int) -> 
None:
+        log.error("Unhandled request", msg=msg)
+        self.send_msg(
+            None,
+            request_id=req_id,
+            error=ErrorResponse(
+                error=ErrorType.API_SERVER_ERROR,
+                detail={"status_code": 400, "message": "Unhandled request"},
+            ),
+        )
 
     @staticmethod
     def _close_unused_sockets(*sockets):
@@ -1830,236 +1981,346 @@ class ActivitySubprocess(WatchedSubprocess):
 
         return TaskInstanceState.FAILED
 
-    def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, 
req_id: int):
+    def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, 
req_id: int) -> None:
         if isinstance(msg, MaskSecret):
             log.debug("Received message from task runner (body omitted)", 
msg=type(msg))
         else:
             log.debug("Received message from task runner", msg=msg)
-        resp: BaseModel | None = None
-        dump_opts: dict[str, bool] = {}
-        if isinstance(msg, TaskState):
-            # No direct API call here — the recovery path in
-            # `update_task_state_if_needed` will call `finish()` for
-            # non-direct states (FAILED, etc.) once the subprocess exits.
-            self._terminal_state = msg.state
-            self._task_end_time_monotonic = time.monotonic()
-            self._rendered_map_index = msg.rendered_map_index
-        elif isinstance(msg, SucceedTask):
-            self._task_end_time_monotonic = time.monotonic()
-            self._rendered_map_index = msg.rendered_map_index
-            self._send_terminal_state_msg(msg)
-        elif isinstance(msg, RetryTask):
-            self._task_end_time_monotonic = time.monotonic()
-            self._rendered_map_index = msg.rendered_map_index
-            self._send_terminal_state_msg(msg)
-        elif isinstance(msg, GetConnection):
-            resp, dump_opts = handle_get_connection(self.client, msg)
-        elif isinstance(msg, GetVariable):
-            resp, dump_opts = handle_get_variable(self.client, msg)
-        elif isinstance(msg, GetVariableKeys):
-            resp, dump_opts = handle_get_variable_keys(self.client, msg)
-        elif isinstance(msg, GetXCom):
-            resp, dump_opts = handle_get_xcom(self.client, msg)
-        elif isinstance(msg, GetXComSequenceItem):
-            resp, dump_opts = handle_get_xcom_sequence_item(self.client, msg)
-        elif isinstance(msg, GetXComSequenceSlice):
-            resp, dump_opts = handle_get_xcom_sequence_slice(self.client, msg)
-        elif isinstance(msg, DeferTask):
-            self._rendered_map_index = msg.rendered_map_index
-            self._send_terminal_state_msg(msg)
-        elif isinstance(msg, AwaitInputTask):
-            self._rendered_map_index = msg.rendered_map_index
-            self._send_terminal_state_msg(msg)
-        elif isinstance(msg, RescheduleTask):
-            self._send_terminal_state_msg(msg)
-        elif isinstance(msg, SkipDownstreamTasks):
-            self.client.task_instances.skip_downstream_tasks(self.id, msg)
-        elif isinstance(msg, SetXCom):
-            resp, dump_opts = handle_set_xcom(self.client, msg)
-        elif isinstance(msg, DeleteXCom):
-            resp, dump_opts = handle_delete_xcom(self.client, msg)
-        elif isinstance(msg, PutVariable):
-            resp, dump_opts = handle_put_variable(self.client, msg)
-        elif isinstance(msg, SetRenderedFields):
-            try:
-                self.client.task_instances.set_rtif(self.id, 
msg.rendered_fields)
-            except ServerResponseError as e:
-                # On retry/clear the server replaces the TI id (archiving the 
old one), so a late RTIF
-                # overwrite from finalize() lands on an id that no longer 
exists. Supervisor kills such
-                # a worker when handling 410 heartbeat response. We only need 
to skip this stale overwrite here.
-                if e.response.status_code != HTTPStatus.GONE:
-                    raise
-                log.debug("Skipping RTIF overwrite; task instance archived on 
retry/clear", ti_id=self.id)
-        elif isinstance(msg, SetRenderedMapIndex):
-            self.client.task_instances.set_rendered_map_index(self.id, 
msg.rendered_map_index)
-        elif isinstance(msg, GetAssetByName):
-            asset_resp = self.client.assets.get(name=msg.name)
-            if isinstance(asset_resp, AssetResponse):
-                asset_result = AssetResult.from_asset_response(asset_resp)
-                resp = asset_result
-                dump_opts = {"exclude_unset": True}
-            else:
-                resp = asset_resp
-        elif isinstance(msg, GetAssetByUri):
-            asset_resp = self.client.assets.get(uri=msg.uri)
-            if isinstance(asset_resp, AssetResponse):
-                asset_result = AssetResult.from_asset_response(asset_resp)
-                resp = asset_result
-                dump_opts = {"exclude_unset": True}
-            else:
-                resp = asset_resp
-        elif isinstance(msg, GetAssetsByAlias):
-            resp = self.client.assets.get_by_alias(alias_name=msg.alias_name)
-        elif isinstance(msg, GetAssetEventByAsset):
-            asset_event_resp = self.client.asset_events.get(
-                uri=msg.uri,
-                name=msg.name,
-                after=msg.after,
-                before=msg.before,
-                ascending=msg.ascending,
-                limit=msg.limit,
-                partition_key=msg.partition_key,
-                partition_key_regexp_pattern=msg.partition_key_regexp_pattern,
-                extra=msg.extra,
-            )
-            asset_event_result = 
AssetEventsResult.from_asset_events_response(asset_event_resp)
-            resp = asset_event_result
-            dump_opts = {"exclude_unset": True}
-        elif isinstance(msg, GetAssetEventByAssetAlias):
-            asset_event_resp = self.client.asset_events.get(
-                alias_name=msg.alias_name,
-                after=msg.after,
-                before=msg.before,
-                ascending=msg.ascending,
-                limit=msg.limit,
-                partition_key=msg.partition_key,
-                partition_key_regexp_pattern=msg.partition_key_regexp_pattern,
-                extra=msg.extra,
-            )
-            asset_event_result = 
AssetEventsResult.from_asset_events_response(asset_event_resp)
-            resp = asset_event_result
-            dump_opts = {"exclude_unset": True}
-        elif isinstance(msg, GetPrevSuccessfulDagRun):
-            resp, dump_opts = handle_get_prev_successful_dag_run(self.client, 
self.id)
-        elif isinstance(msg, GetXComCount):
-            resp, dump_opts = handle_get_xcom_count(self.client, msg)
-        elif isinstance(msg, TriggerDagRun):
-            resp = self.client.dag_runs.trigger(
-                msg.dag_id, msg.run_id, msg.conf, msg.logical_date, 
msg.run_after, msg.reset_dag_run, msg.note
-            )
-        elif isinstance(msg, GetDagRun):
-            dr_resp = self.client.dag_runs.get_detail(msg.dag_id, msg.run_id)
-            resp = DagRunResult.from_api_response(dr_resp)
-        elif isinstance(msg, GetTaskRescheduleStartDate):
-            resp = 
self.client.task_instances.get_reschedule_start_date(msg.ti_id, msg.try_number)
-        elif isinstance(msg, GetTICount):
-            resp, dump_opts = handle_get_ti_count(self.client, msg)
-        elif isinstance(msg, GetTaskStates):
-            resp, dump_opts = handle_get_task_states(self.client, msg)
-        elif isinstance(msg, GetTaskBreadcrumbs):
-            api_resp = 
self.client.task_instances.get_task_breakcrumbs(dag_id=msg.dag_id, 
run_id=msg.run_id)
-            resp = TaskBreadcrumbsResult.from_api_response(api_resp)
-        elif isinstance(msg, GetDRCount):
-            resp, dump_opts = handle_get_dr_count(self.client, msg)
-        elif isinstance(msg, GetDagRunState):
-            resp, dump_opts = handle_get_dag_run_state(self.client, msg)
-        elif isinstance(msg, GetPreviousDagRun):
-            resp, dump_opts = handle_get_previous_dag_run(self.client, msg)
-        elif isinstance(msg, GetPreviousTI):
-            resp, dump_opts = handle_get_previous_ti(self.client, msg)
-        elif isinstance(msg, DeleteVariable):
-            resp, dump_opts = handle_delete_variable(self.client, msg)
-        elif isinstance(msg, ValidateInletsAndOutlets):
-            inactive_assets_resp = 
self.client.task_instances.validate_inlets_and_outlets(msg.ti_id)
-            resp = 
InactiveAssetsResult.from_inactive_assets_response(inactive_assets_resp)
-            dump_opts = {"exclude_unset": True}
-        elif isinstance(msg, ResendLoggingFD):
-            # We need special handling here!
-            if send_fds is not None:
-                self._send_new_log_fd(req_id)
-                # Since we've sent the message, return. Nothing else in this 
ifelse/switch should return directly
-                return
-        elif isinstance(msg, CreateHITLDetailPayload):
-            hitl_detail_request = self.client.hitl.add_response(
-                ti_id=msg.ti_id,
-                options=msg.options,
-                subject=msg.subject,
-                body=msg.body,
-                defaults=msg.defaults,
-                params=msg.params,
-                multiple=msg.multiple,
-                assigned_users=msg.assigned_users,
-            )
-            resp = 
HITLDetailRequestResult.from_api_response(hitl_detail_request)
-            dump_opts = {"exclude_unset": True}
-        elif isinstance(msg, MaskSecret):
-            handle_mask_secret(msg)
-        elif isinstance(msg, GetDag):
-            dag = self.client.dags.get(
-                dag_id=msg.dag_id,
-            )
-            resp = DagResult.from_api_response(dag)
-        elif isinstance(msg, GetTaskStateStore):
-            task_store = self.client.task_state_store.get(msg.ti_id, msg.key)
-            resp = (
-                task_store
-                if isinstance(task_store, ErrorResponse)
-                else 
TaskStateStoreResult.from_task_state_store_response(task_store)
-            )
-        elif isinstance(msg, SetTaskStateStore):
-            self.client.task_state_store.set(msg.ti_id, msg.key, msg.value, 
expires_at=msg.expires_at)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, DeleteTaskStateStore):
-            self.client.task_state_store.delete(msg.ti_id, msg.key)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, ClearTaskStateStore):
-            self.client.task_state_store.clear(msg.ti_id)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, GetAssetStateStoreByName):
-            asset_store = self.client.asset_state_store.get(msg.key, 
name=msg.name)
-            resp = (
-                asset_store
-                if isinstance(asset_store, ErrorResponse)
-                else 
AssetStateStoreResult.from_asset_state_store_response(asset_store)
-            )
-        elif isinstance(msg, GetAssetStateStoreByUri):
-            asset_store = self.client.asset_state_store.get(msg.key, 
uri=msg.uri)
-            resp = (
-                asset_store
-                if isinstance(asset_store, ErrorResponse)
-                else 
AssetStateStoreResult.from_asset_state_store_response(asset_store)
-            )
-        elif isinstance(msg, SetAssetStateStoreByName):
-            self.client.asset_state_store.set(msg.key, msg.value, 
name=msg.name)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, SetAssetStateStoreByUri):
-            self.client.asset_state_store.set(msg.key, msg.value, uri=msg.uri)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, DeleteAssetStateStoreByName):
-            self.client.asset_state_store.delete(msg.key, name=msg.name)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, DeleteAssetStateStoreByUri):
-            self.client.asset_state_store.delete(msg.key, uri=msg.uri)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, ClearAssetStateStoreByName):
-            self.client.asset_state_store.clear(name=msg.name)
-            resp = OKResponse(ok=True)
-        elif isinstance(msg, ClearAssetStateStoreByUri):
-            self.client.asset_state_store.clear(uri=msg.uri)
-            resp = OKResponse(ok=True)
-        else:
-            log.error("Unhandled request", msg=msg)
-            self.send_msg(
-                None,
-                request_id=req_id,
-                error=ErrorResponse(
-                    error=ErrorType.API_SERVER_ERROR,
-                    detail={"status_code": 400, "message": "Unhandled 
request"},
-                ),
-            )
-            return
+        super()._handle_request(msg, log, req_id)
+
+    def _handle_task_state(self, msg: TaskState, log: FilteringBoundLogger, 
req_id: int) -> RequestResult:
+        # No direct API call here — the recovery path in
+        # `update_task_state_if_needed` will call `finish()` for
+        # non-direct states (FAILED, etc.) once the subprocess exits.
+        self._terminal_state = msg.state
+        self._task_end_time_monotonic = time.monotonic()
+        self._rendered_map_index = msg.rendered_map_index
+        return None, {}
+
+    def _handle_finished_task(
+        self, msg: SucceedTask | RetryTask, log: FilteringBoundLogger, req_id: 
int
+    ) -> RequestResult:
+        self._task_end_time_monotonic = time.monotonic()
+        self._rendered_map_index = msg.rendered_map_index
+        self._send_terminal_state_msg(msg)
+        return None, {}
+
+    def _handle_suspended_task(
+        self, msg: DeferTask | AwaitInputTask, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        self._rendered_map_index = msg.rendered_map_index
+        self._send_terminal_state_msg(msg)
+        return None, {}
+
+    def _handle_reschedule_task(
+        self, msg: RescheduleTask, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        self._send_terminal_state_msg(msg)
+        return None, {}
+
+    def _handle_skip_downstream_tasks(
+        self, msg: SkipDownstreamTasks, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        self.client.task_instances.skip_downstream_tasks(self.id, msg)
+        return None, {}
+
+    def _handle_set_rendered_fields(
+        self, msg: SetRenderedFields, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        try:
+            self.client.task_instances.set_rtif(self.id, msg.rendered_fields)
+        except ServerResponseError as e:
+            # On retry/clear the server replaces the TI id (archiving the old 
one), so a late RTIF
+            # overwrite from finalize() lands on an id that no longer exists. 
Supervisor kills such
+            # a worker when handling 410 heartbeat response. We only need to 
skip this stale overwrite here.
+            if e.response.status_code != HTTPStatus.GONE:
+                raise
+            log.debug("Skipping RTIF overwrite; task instance archived on 
retry/clear", ti_id=self.id)
+        return None, {}
+
+    def _handle_set_rendered_map_index(
+        self, msg: SetRenderedMapIndex, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        self.client.task_instances.set_rendered_map_index(self.id, 
msg.rendered_map_index)
+        return None, {}
+
+    def _handle_get_asset_by_name(
+        self, msg: GetAssetByName, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        asset_resp = self.client.assets.get(name=msg.name)
+        if isinstance(asset_resp, AssetResponse):
+            return AssetResult.from_asset_response(asset_resp), 
{"exclude_unset": True}
+        return asset_resp, {}
+
+    def _handle_get_asset_by_uri(
+        self, msg: GetAssetByUri, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        asset_resp = self.client.assets.get(uri=msg.uri)
+        if isinstance(asset_resp, AssetResponse):
+            return AssetResult.from_asset_response(asset_resp), 
{"exclude_unset": True}
+        return asset_resp, {}
+
+    def _handle_get_assets_by_alias(
+        self, msg: GetAssetsByAlias, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        resp = self.client.assets.get_by_alias(alias_name=msg.alias_name)
+        return resp, {}
+
+    def _handle_get_asset_event_by_asset(
+        self, msg: GetAssetEventByAsset, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        asset_event_resp = self.client.asset_events.get(
+            uri=msg.uri,
+            name=msg.name,
+            after=msg.after,
+            before=msg.before,
+            ascending=msg.ascending,
+            limit=msg.limit,
+            partition_key=msg.partition_key,
+            partition_key_regexp_pattern=msg.partition_key_regexp_pattern,
+            extra=msg.extra,
+        )
+        return AssetEventsResult.from_asset_events_response(asset_event_resp), 
{"exclude_unset": True}
+
+    def _handle_get_asset_event_by_asset_alias(
+        self, msg: GetAssetEventByAssetAlias, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        asset_event_resp = self.client.asset_events.get(
+            alias_name=msg.alias_name,
+            after=msg.after,
+            before=msg.before,
+            ascending=msg.ascending,
+            limit=msg.limit,
+            partition_key=msg.partition_key,
+            partition_key_regexp_pattern=msg.partition_key_regexp_pattern,
+            extra=msg.extra,
+        )
+        return AssetEventsResult.from_asset_events_response(asset_event_resp), 
{"exclude_unset": True}
 
-        self.send_msg(resp, request_id=req_id, error=None, **dump_opts)
+    def _handle_trigger_dag_run(
+        self, msg: TriggerDagRun, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        resp = self.client.dag_runs.trigger(
+            msg.dag_id, msg.run_id, msg.conf, msg.logical_date, msg.run_after, 
msg.reset_dag_run, msg.note
+        )
+        return resp, {}
+
+    def _handle_get_dag_run(self, msg: GetDagRun, log: FilteringBoundLogger, 
req_id: int) -> RequestResult:
+        dr_resp = self.client.dag_runs.get_detail(msg.dag_id, msg.run_id)
+        resp = DagRunResult.from_api_response(dr_resp)
+        return resp, {}
+
+    def _handle_get_task_reschedule_start_date(
+        self, msg: GetTaskRescheduleStartDate, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        resp = self.client.task_instances.get_reschedule_start_date(msg.ti_id, 
msg.try_number)
+        return resp, {}
+
+    def _handle_get_task_breadcrumbs(
+        self, msg: GetTaskBreadcrumbs, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        api_resp = 
self.client.task_instances.get_task_breakcrumbs(dag_id=msg.dag_id, 
run_id=msg.run_id)
+        resp = TaskBreadcrumbsResult.from_api_response(api_resp)
+        return resp, {}
+
+    def _handle_validate_inlets_and_outlets(
+        self, msg: ValidateInletsAndOutlets, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        inactive_assets_resp = 
self.client.task_instances.validate_inlets_and_outlets(msg.ti_id)
+        return 
InactiveAssetsResult.from_inactive_assets_response(inactive_assets_resp), {
+            "exclude_unset": True
+        }
+
+    def _handle_resend_logging_fd(
+        self, msg: ResendLoggingFD, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult | ResponseSent:
+        if send_fds is not None:
+            self._send_new_log_fd(req_id)
+            return ResponseSent.ALREADY_SENT
+        # Preserve the empty acknowledgment for this no-op when descriptor 
passing is unavailable.
+        return None, {}
+
+    def _handle_create_hitl_detail_payload(
+        self, msg: CreateHITLDetailPayload, log: FilteringBoundLogger, req_id: 
int
+    ) -> RequestResult:
+        hitl_detail_request = self.client.hitl.add_response(
+            ti_id=msg.ti_id,
+            options=msg.options,
+            subject=msg.subject,
+            body=msg.body,
+            defaults=msg.defaults,
+            params=msg.params,
+            multiple=msg.multiple,
+            assigned_users=msg.assigned_users,
+        )
+        return HITLDetailRequestResult.from_api_response(hitl_detail_request), 
{"exclude_unset": True}
+
+    def _handle_get_dag(self, msg: GetDag, log: FilteringBoundLogger, req_id: 
int) -> RequestResult:
+        dag = self.client.dags.get(
+            dag_id=msg.dag_id,
+        )
+        resp = DagResult.from_api_response(dag)
+        return resp, {}
+
+    def _handle_get_task_state_store(
+        self, msg: GetTaskStateStore, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        task_store = self.client.task_state_store.get(msg.ti_id, msg.key)
+        resp = (
+            task_store
+            if isinstance(task_store, ErrorResponse)
+            else 
TaskStateStoreResult.from_task_state_store_response(task_store)
+        )
+        return resp, {}
+
+    def _handle_set_task_state_store(
+        self, msg: SetTaskStateStore, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        self.client.task_state_store.set(msg.ti_id, msg.key, msg.value, 
expires_at=msg.expires_at)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_delete_task_state_store(
+        self, msg: DeleteTaskStateStore, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        self.client.task_state_store.delete(msg.ti_id, msg.key)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_clear_task_state_store(
+        self, msg: ClearTaskStateStore, log: FilteringBoundLogger, req_id: int
+    ) -> RequestResult:
+        self.client.task_state_store.clear(msg.ti_id)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_get_asset_state_store_by_name(
+        self, msg: GetAssetStateStoreByName, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        asset_store = self.client.asset_state_store.get(msg.key, name=msg.name)
+        resp = (
+            asset_store
+            if isinstance(asset_store, ErrorResponse)
+            else 
AssetStateStoreResult.from_asset_state_store_response(asset_store)
+        )
+        return resp, {}
+
+    def _handle_get_asset_state_store_by_uri(
+        self, msg: GetAssetStateStoreByUri, log: FilteringBoundLogger, req_id: 
int
+    ) -> RequestResult:
+        asset_store = self.client.asset_state_store.get(msg.key, uri=msg.uri)
+        resp = (
+            asset_store
+            if isinstance(asset_store, ErrorResponse)
+            else 
AssetStateStoreResult.from_asset_state_store_response(asset_store)
+        )
+        return resp, {}
+
+    def _handle_set_asset_state_store_by_name(
+        self, msg: SetAssetStateStoreByName, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        self.client.asset_state_store.set(msg.key, msg.value, name=msg.name)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_set_asset_state_store_by_uri(
+        self, msg: SetAssetStateStoreByUri, log: FilteringBoundLogger, req_id: 
int
+    ) -> RequestResult:
+        self.client.asset_state_store.set(msg.key, msg.value, uri=msg.uri)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_delete_asset_state_store_by_name(
+        self, msg: DeleteAssetStateStoreByName, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        self.client.asset_state_store.delete(msg.key, name=msg.name)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_delete_asset_state_store_by_uri(
+        self, msg: DeleteAssetStateStoreByUri, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        self.client.asset_state_store.delete(msg.key, uri=msg.uri)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_clear_asset_state_store_by_name(
+        self, msg: ClearAssetStateStoreByName, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        self.client.asset_state_store.clear(name=msg.name)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    def _handle_clear_asset_state_store_by_uri(
+        self, msg: ClearAssetStateStoreByUri, log: FilteringBoundLogger, 
req_id: int
+    ) -> RequestResult:
+        self.client.asset_state_store.clear(uri=msg.uri)
+        resp = OKResponse(ok=True)
+        return resp, {}
+
+    _request_handlers: ClassVar[dict[type[BaseModel], 
RequestHandler[ActivitySubprocess]]] = {
+        **WatchedSubprocess._get_shared_request_handlers(
+            DeleteVariable,
+            DeleteXCom,
+            GetConnection,
+            GetDRCount,
+            GetDagRunState,
+            GetPrevSuccessfulDagRun,
+            GetPreviousDagRun,
+            GetPreviousTI,
+            GetTICount,
+            GetTaskStates,
+            GetVariable,
+            GetVariableKeys,
+            GetXCom,
+            GetXComCount,
+            GetXComSequenceItem,
+            GetXComSequenceSlice,
+            MaskSecret,
+            PutVariable,
+            SetXCom,
+        ),
+        **dict(
+            [
+                register_request_method(AwaitInputTask, 
_handle_suspended_task),
+                register_request_method(ClearAssetStateStoreByName, 
_handle_clear_asset_state_store_by_name),
+                register_request_method(ClearAssetStateStoreByUri, 
_handle_clear_asset_state_store_by_uri),
+                register_request_method(ClearTaskStateStore, 
_handle_clear_task_state_store),
+                register_request_method(CreateHITLDetailPayload, 
_handle_create_hitl_detail_payload),
+                register_request_method(DeferTask, _handle_suspended_task),
+                register_request_method(
+                    DeleteAssetStateStoreByName, 
_handle_delete_asset_state_store_by_name
+                ),
+                register_request_method(DeleteAssetStateStoreByUri, 
_handle_delete_asset_state_store_by_uri),
+                register_request_method(DeleteTaskStateStore, 
_handle_delete_task_state_store),
+                register_request_method(GetAssetByName, 
_handle_get_asset_by_name),
+                register_request_method(GetAssetByUri, 
_handle_get_asset_by_uri),
+                register_request_method(GetAssetEventByAsset, 
_handle_get_asset_event_by_asset),
+                register_request_method(GetAssetEventByAssetAlias, 
_handle_get_asset_event_by_asset_alias),
+                register_request_method(GetAssetsByAlias, 
_handle_get_assets_by_alias),
+                register_request_method(GetAssetStateStoreByName, 
_handle_get_asset_state_store_by_name),
+                register_request_method(GetAssetStateStoreByUri, 
_handle_get_asset_state_store_by_uri),
+                register_request_method(GetDag, _handle_get_dag),
+                register_request_method(GetDagRun, _handle_get_dag_run),
+                register_request_method(GetTaskBreadcrumbs, 
_handle_get_task_breadcrumbs),
+                register_request_method(GetTaskRescheduleStartDate, 
_handle_get_task_reschedule_start_date),
+                register_request_method(GetTaskStateStore, 
_handle_get_task_state_store),
+                register_request_method(RescheduleTask, 
_handle_reschedule_task),
+                register_request_method(ResendLoggingFD, 
_handle_resend_logging_fd),
+                register_request_method(RetryTask, _handle_finished_task),
+                register_request_method(SetAssetStateStoreByName, 
_handle_set_asset_state_store_by_name),
+                register_request_method(SetAssetStateStoreByUri, 
_handle_set_asset_state_store_by_uri),
+                register_request_method(SetRenderedFields, 
_handle_set_rendered_fields),
+                register_request_method(SetRenderedMapIndex, 
_handle_set_rendered_map_index),
+                register_request_method(SetTaskStateStore, 
_handle_set_task_state_store),
+                register_request_method(SkipDownstreamTasks, 
_handle_skip_downstream_tasks),
+                register_request_method(SucceedTask, _handle_finished_task),
+                register_request_method(TaskState, _handle_task_state),
+                register_request_method(TriggerDagRun, 
_handle_trigger_dag_run),
+                register_request_method(ValidateInletsAndOutlets, 
_handle_validate_inlets_and_outlets),
+            ]
+        ),
+    }
 
     def _send_new_log_fd(self, req_id: int) -> None:
         if send_fds is None:
diff --git 
a/task-sdk/tests/task_sdk/execution_time/_request_registration_types.py 
b/task-sdk/tests/task_sdk/execution_time/_request_registration_types.py
new file mode 100644
index 00000000000..ea7935ecb66
--- /dev/null
+++ b/task-sdk/tests/task_sdk/execution_time/_request_registration_types.py
@@ -0,0 +1,41 @@
+# 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.
+# mypy: warn-unused-ignores
+"""Unused-ignore errors ensure these invalid registrations stay rejected by 
mypy-task-sdk."""
+
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
+if TYPE_CHECKING:
+    from pydantic import BaseModel
+
+    from airflow.sdk.execution_time.comms import GetConnection, GetVariable
+    from airflow.sdk.execution_time.request_handlers import 
handle_get_connection
+    from airflow.sdk.execution_time.supervisor import (
+        ActivitySubprocess,
+        RequestHandler,
+        WatchedSubprocess,
+        _register_client_handler,
+        register_request_method,
+    )
+
+    _register_client_handler(GetVariable, handle_get_connection)  # type: 
ignore[arg-type]
+    register_request_method(GetConnection, ActivitySubprocess._handle_get_dag) 
 # type: ignore[arg-type]
+    clientless_handlers: dict[type[BaseModel], 
RequestHandler[WatchedSubprocess]] = (
+        WatchedSubprocess._get_shared_request_handlers(GetConnection)  # type: 
ignore[assignment]
+    )
diff --git a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py 
b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
index 807fd0475d1..e94b901179d 100644
--- a/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
+++ b/task-sdk/tests/task_sdk/execution_time/test_supervisor.py
@@ -35,7 +35,7 @@ from operator import attrgetter
 from random import randint
 from textwrap import dedent
 from time import sleep
-from typing import TYPE_CHECKING, Any
+from typing import TYPE_CHECKING, Any, get_args
 from unittest import mock
 from unittest.mock import MagicMock, patch
 
@@ -2643,7 +2643,7 @@ REQUEST_TEST_CASES = [
         test_id="validate_inlets_and_outlets",
     ),
     RequestTestCase(
-        message=GetPrevSuccessfulDagRun(ti_id=TI_ID),
+        message=GetPrevSuccessfulDagRun(ti_id=uuid7()),
         expected_body={
             "data_interval_start": timezone.parse("2025-01-10T12:00:00Z"),
             "data_interval_end": timezone.parse("2025-01-10T14:00:00Z"),
@@ -3330,20 +3330,143 @@ REQUEST_TEST_CASES = [
 
 
 class TestHandleRequest:
+    class _OverrideActivitySubprocess(ActivitySubprocess):
+        def _handle_set_rendered_map_index(
+            self, msg: SetRenderedMapIndex, log: FilteringBoundLogger, req_id: 
int
+        ) -> supervisor.RequestResult:
+            return OKResponse(ok=True), {}
+
+    @patch.object(
+        ActivitySubprocess, "_handle_set_rendered_map_index", autospec=True, 
return_value=(None, {})
+    )
+    @patch.object(ActivitySubprocess, "send_msg", autospec=True)
+    def test_dispatch_resolves_patched_method(self, send_msg, handler, 
watched_subprocess):
+        process, _ = watched_subprocess
+        msg = SetRenderedMapIndex(rendered_map_index="label")
+        log = structlog.get_logger()
+
+        process._handle_request(msg, log, req_id=42)
+
+        handler.assert_called_once_with(process, msg, log, 42)
+        send_msg.assert_called_once_with(process, None, request_id=42, 
error=None)
+        assert not process.client.mock_calls
+
+    @patch.object(ActivitySubprocess, "send_msg", autospec=True)
+    @pytest.mark.parametrize("watched_subprocess", 
[_OverrideActivitySubprocess], indirect=True)
+    def test_dispatch_resolves_subclass_override(self, send_msg, 
watched_subprocess):
+        process, _ = watched_subprocess
+
+        process._handle_request(
+            SetRenderedMapIndex(rendered_map_index="label"), 
structlog.get_logger(), req_id=42
+        )
+
+        send_msg.assert_called_once_with(process, OKResponse(ok=True), 
request_id=42, error=None)
+        assert not process.client.mock_calls
+
+    @patch.object(ActivitySubprocess, "_handle_set_rendered_map_index", 
autospec=True, return_value=None)
+    def test_missing_handler_result_sends_error(self, handler, 
watched_subprocess, mocker):
+        process, read_socket = watched_subprocess
+        generator = 
process.handle_requests(log=mocker.Mock(spec=FilteringBoundLogger))
+        next(generator)
+
+        msg = SetRenderedMapIndex(rendered_map_index="label")
+        generator.send(_RequestFrame(id=42, body=msg.model_dump()))
+
+        read_socket.settimeout(0.1)
+        frame_len = int.from_bytes(read_socket.recv(4), "big")
+        frame = 
msgspec.msgpack.Decoder(_ResponseFrame).decode(read_socket.recv(frame_len))
+        assert frame.id == 42
+        assert frame.error is not None
+        assert frame.error["error"] == ErrorType.API_SERVER_ERROR.value
+        assert frame.error["detail"]["exception_type"] == "TypeError"
+        handler.assert_called_once()
+
+    @patch.object(ActivitySubprocess, "send_msg", autospec=True)
+    @pytest.mark.parametrize("registry", [None, {}], ids=["undeclared", 
"empty"])
+    def test_undeclared_registry_is_distinct_from_empty_registry(
+        self, send_msg, watched_subprocess, monkeypatch, registry
+    ):
+        process, _ = watched_subprocess
+        monkeypatch.setattr(ActivitySubprocess, "_request_handlers", registry)
+        msg = GetVariable(key="key")
+
+        if registry is None:
+            with pytest.raises(NotImplementedError, match="must declare its 
request handlers"):
+                process._handle_request(msg, structlog.get_logger(), req_id=42)
+            send_msg.assert_not_called()
+        else:
+            process._handle_request(msg, structlog.get_logger(), req_id=42)
+            send_msg.assert_called_once_with(
+                process,
+                None,
+                request_id=42,
+                error=ErrorResponse(
+                    error=ErrorType.API_SERVER_ERROR,
+                    detail={"status_code": 400, "message": "Unhandled 
request"},
+                ),
+            )
+
+    @patch.object(ActivitySubprocess, "send_msg", autospec=True)
+    @pytest.mark.parametrize(
+        "message",
+        [
+            GetHITLDetailResponse(ti_id=TI_ID),
+            UpdateHITLDetail(ti_id=TI_ID, chosen_options=["approved"]),
+        ],
+    )
+    def test_rejects_unsupported_task_messages(self, send_msg, 
watched_subprocess, message):
+        process, _ = watched_subprocess
+        process._handle_request(message, structlog.get_logger(), req_id=42)
+
+        send_msg.assert_called_once_with(
+            process,
+            None,
+            request_id=42,
+            error=ErrorResponse(
+                error=ErrorType.API_SERVER_ERROR,
+                detail={"status_code": 400, "message": "Unhandled request"},
+            ),
+        )
+        assert not process.client.mock_calls
+
+    @patch.object(ActivitySubprocess, "_send_new_log_fd", autospec=True)
+    @patch.object(ActivitySubprocess, "send_msg", autospec=True)
+    @pytest.mark.parametrize("fd_supported", [True, False])
+    def test_resend_logging_fd_sends_one_response(
+        self, send_msg, send_new_log_fd, watched_subprocess, monkeypatch, 
fd_supported
+    ):
+        process, _ = watched_subprocess
+        monkeypatch.setattr(supervisor, "send_fds", object() if fd_supported 
else None)
+
+        process._handle_request(ResendLoggingFD(), structlog.get_logger(), 
req_id=42)
+
+        if fd_supported:
+            send_new_log_fd.assert_called_once_with(process, 42)
+            send_msg.assert_not_called()
+        else:
+            send_new_log_fd.assert_not_called()
+            send_msg.assert_called_once_with(process, None, request_id=42, 
error=None)
+
     @pytest.fixture
-    def watched_subprocess(self, mocker):
+    def watched_subprocess(self, mocker, request):
         read_end, write_end = socket.socketpair()
+        process_type = getattr(request, "param", ActivitySubprocess)
 
-        subprocess = ActivitySubprocess(
-            process_log=mocker.MagicMock(),
+        subprocess = process_type(
+            process_log=mocker.MagicMock(spec=FilteringBoundLogger),
             id=TI_ID,
             pid=12345,
             stdin=write_end,
-            client=mocker.Mock(),
-            process=mocker.Mock(),
+            client=mocker.Mock(spec=sdk_client.Client),
+            process=mocker.Mock(spec=psutil.Process),
         )
 
-        return subprocess, read_end
+        try:
+            yield subprocess, read_end
+        finally:
+            subprocess.selector.close()
+            read_end.close()
+            write_end.close()
 
     @patch("airflow.sdk.execution_time.request_handlers.mask_secret")
     @pytest.mark.parametrize("test_case", REQUEST_TEST_CASES, ids=lambda tc: 
tc.test_id)
@@ -3416,31 +3539,10 @@ class TestHandleRequest:
             decoder = CommsDecoder(socket=None).body_decoder  # type: 
ignore[var-annotated, arg-type]
             assert decoder.validate_python(frame.body) == client_mock.response
 
-    def test_all_to_supervisor_messages_are_covered(self):
-        """Ensure all ToSupervisor message types have test coverage."""
-
-        # Extract the individual message types from the Union
-        union_type = ToSupervisor.__args__[0]
-        supervisor_message_types = set(union_type.__args__)
-
-        # Get all message types covered in our test cases
-        tested_message_types = {type(test_case.message) for test_case in 
REQUEST_TEST_CASES}
-
-        # Message types which are excluded for a good reason
-        excluded_message_types = {
-            GetHITLDetailResponse,  # Only used in Triggerer, not needed in 
worker
-            UpdateHITLDetail,  # Only used in Triggerer, not needed in worker
-        }
-
-        untested_types = supervisor_message_types - tested_message_types - 
excluded_message_types
-
-        # Assert all types are covered
-        assert not untested_types, (
-            f"Missing test coverage for 
{len(untested_types)}/{len(supervisor_message_types)} "
-            f"ToSupervisor message types:\n"
-            + "\n".join(f"  - {t.__name__}" for t in sorted(untested_types, 
key=lambda x: x.__name__))
-            + "\n\nPlease add test cases to REQUEST_TEST_CASES."
-        )
+    def test_registered_message_types(self):
+        expected = set(get_args(get_args(ToSupervisor)[0])) - 
{GetHITLDetailResponse, UpdateHITLDetail}
+        assert set(ActivitySubprocess._request_handlers) == expected
+        assert {type(case.message) for case in REQUEST_TEST_CASES} == expected
 
     def test_handle_requests_api_server_error(self, watched_subprocess, 
mocker):
         """Test that API server errors are properly handled and sent back to 
the task."""

Reply via email to