From 5531d973f694849c2ccfa3d076a475e154a391ec Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 14:55:29 -0700 Subject: [PATCH 01/42] refactor: add handler registry and decorator for request messages Add the registry for request_handlers.py so we can refactor _handle_request in handlers --- .../sdk/execution_time/request_handlers.py | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) 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 959be43fe93cf..32eb15e43b066 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -27,6 +27,7 @@ from __future__ import annotations +from collections.abc import Callable from typing import TYPE_CHECKING from uuid import UUID @@ -75,6 +76,23 @@ from airflow.sdk.api.client import Client +_HANDLER_REGISTRY: dict[type, Callable] = {} + + +def handles(msg_type): + def decorator(fn): + _HANDLER_REGISTRY[msg_type] = fn + return fn + + return decorator + + +def get_handler(msg_type): + handler = _HANDLER_REGISTRY.get(msg_type) + if handler is None: + raise TypeError(f"No handler registered for {msg_type.__name__}") + return handler + def handle_get_connection(client: Client, msg: GetConnection) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch a connection and mask its sensitive fields.""" From f8d1de183497d3af2bf3cf54a225ab1214778ee6 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 15:20:56 -0700 Subject: [PATCH 02/42] refactor: register all shared handlers to all handlers with handles decorator --- .../sdk/execution_time/request_handlers.py | 25 ++++++++++++++++--- 1 file changed, 22 insertions(+), 3 deletions(-) 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 32eb15e43b066..78b3f232c93ff 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -29,7 +29,6 @@ from collections.abc import Callable from typing import TYPE_CHECKING -from uuid import UUID from airflow.sdk.api.datamodels._generated import ( ConnectionResponse, @@ -50,6 +49,7 @@ GetDRCount, GetPreviousDagRun, GetPreviousTI, + GetPrevSuccessfulDagRun, GetTaskStates, GetTICount, GetVariable, @@ -94,6 +94,7 @@ def get_handler(msg_type): return handler +@handles(GetConnection) def handle_get_connection(client: Client, msg: GetConnection) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch a connection and mask its sensitive fields.""" conn = client.connections.get(msg.conn_id) @@ -106,6 +107,7 @@ def handle_get_connection(client: Client, msg: GetConnection) -> tuple[BaseModel return conn, {} +@handles(GetVariable) def handle_get_variable(client: Client, msg: GetVariable) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch a variable and mask its value.""" var = client.variables.get(msg.key) @@ -116,6 +118,7 @@ def handle_get_variable(client: Client, msg: GetVariable) -> tuple[BaseModel | N return var, {} +@handles(GetVariableKeys) def handle_get_variable_keys( client: Client, msg: GetVariableKeys ) -> tuple[BaseModel | None, dict[str, bool]]: @@ -127,23 +130,27 @@ def handle_get_variable_keys( ) +@handles(MaskSecret) def handle_mask_secret(msg: MaskSecret) -> None: """Register a value with the secrets masker.""" mask_secret(msg.value, msg.name) +@handles(PutVariable) def handle_put_variable(client: Client, msg: PutVariable) -> tuple[BaseModel | None, dict[str, bool]]: """Store a variable value.""" client.variables.set(msg.key, msg.value, msg.description) return None, {} +@handles(DeleteVariable) def handle_delete_variable(client: Client, msg: DeleteVariable) -> tuple[BaseModel | None, dict[str, bool]]: """Delete a variable value.""" resp = client.variables.delete(msg.key) return resp, {} +@handles(GetTICount) def handle_get_ti_count(client: Client, msg: GetTICount) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch task instance counts.""" resp = client.task_instances.get_count( @@ -158,6 +165,7 @@ def handle_get_ti_count(client: Client, msg: GetTICount) -> tuple[BaseModel | No return resp, {} +@handles(GetTaskStates) def handle_get_task_states(client: Client, msg: GetTaskStates) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch task states and normalize them for supervisor response handling.""" task_states_map = client.task_instances.get_task_states( @@ -173,6 +181,7 @@ def handle_get_task_states(client: Client, msg: GetTaskStates) -> tuple[BaseMode return task_states_map, {} +@handles(GetPreviousTI) def handle_get_previous_ti(client: Client, msg: GetPreviousTI) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch the previous task instance.""" resp = client.task_instances.get_previous( @@ -185,6 +194,7 @@ def handle_get_previous_ti(client: Client, msg: GetPreviousTI) -> tuple[BaseMode return resp, {} +@handles(SetXCom) def handle_set_xcom(client: Client, msg: SetXCom) -> tuple[BaseModel | None, dict[str, bool]]: """Store an XCom value.""" client.xcoms.set( @@ -200,12 +210,14 @@ def handle_set_xcom(client: Client, msg: SetXCom) -> tuple[BaseModel | None, dic return None, {} +@handles(DeleteXCom) def handle_delete_xcom(client: Client, msg: DeleteXCom) -> tuple[BaseModel | None, dict[str, bool]]: """Delete an XCom value.""" client.xcoms.delete(msg.dag_id, msg.run_id, msg.task_id, msg.key, msg.map_index) return None, {} +@handles(GetDRCount) def handle_get_dr_count(client: Client, msg: GetDRCount) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch dag run counts.""" resp = client.dag_runs.get_count( @@ -217,6 +229,7 @@ def handle_get_dr_count(client: Client, msg: GetDRCount) -> tuple[BaseModel | No return resp, {} +@handles(GetDagRunState) def handle_get_dag_run_state(client: Client, msg: GetDagRunState) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch dag run state.""" dr_resp = client.dag_runs.get_state(msg.dag_id, msg.run_id) @@ -225,6 +238,7 @@ def handle_get_dag_run_state(client: Client, msg: GetDagRunState) -> tuple[BaseM return dr_resp, {} +@handles(GetPreviousDagRun) def handle_get_previous_dag_run( client: Client, msg: GetPreviousDagRun ) -> tuple[BaseModel | None, dict[str, bool]]: @@ -237,21 +251,24 @@ def handle_get_previous_dag_run( return resp, {} +@handles(GetPrevSuccessfulDagRun) def handle_get_prev_successful_dag_run( - client: Client, subprocess_id: UUID + client: Client, msg: GetPrevSuccessfulDagRun ) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch the previous successful dag run using the caller's current id.""" - dagrun_resp = client.task_instances.get_previous_successful_dagrun(subprocess_id) + dagrun_resp = client.task_instances.get_previous_successful_dagrun(msg.ti_id) dagrun_result = PrevSuccessfulDagRunResult.from_dagrun_response(dagrun_resp) return dagrun_result, {"exclude_unset": True} +@handles(GetXComCount) def handle_get_xcom_count(client: Client, msg: GetXComCount) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch XCom count metadata.""" resp = client.xcoms.head(msg.dag_id, msg.run_id, msg.task_id, msg.key) return resp, {} +@handles(GetXComSequenceItem) def handle_get_xcom_sequence_item( client: Client, msg: GetXComSequenceItem ) -> tuple[BaseModel | None, dict[str, bool]]: @@ -262,6 +279,7 @@ def handle_get_xcom_sequence_item( return xcom, {} +@handles(GetXComSequenceSlice) def handle_get_xcom_sequence_slice( client: Client, msg: GetXComSequenceSlice ) -> tuple[BaseModel | None, dict[str, bool]]: @@ -281,6 +299,7 @@ def handle_get_xcom_sequence_slice( return xcoms, {} +@handles(GetXCom) def handle_get_xcom(client: Client, msg: GetXCom) -> tuple[BaseModel | None, dict[str, bool]]: """Fetch an XCom and normalize it for supervisor response handling.""" xcom = client.xcoms.get( From 3f22f84fee02d2d3fd1bd03a13e460213ebaba5c Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 15:26:42 -0700 Subject: [PATCH 03/42] refactor: resolve handle_mask_secret to return none with empty set --- task-sdk/src/airflow/sdk/execution_time/request_handlers.py | 1 + 1 file changed, 1 insertion(+) 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 78b3f232c93ff..0f77b0341feb0 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -134,6 +134,7 @@ def handle_get_variable_keys( def handle_mask_secret(msg: MaskSecret) -> None: """Register a value with the secrets masker.""" mask_secret(msg.value, msg.name) + return (None, {}) @handles(PutVariable) From 89c3ed1da0c1288fe4cc68684433b6384fb15af0 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 15:46:15 -0700 Subject: [PATCH 04/42] fix: resolve return type error on handle_mask_secret --- task-sdk/src/airflow/sdk/execution_time/request_handlers.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 0f77b0341feb0..5079d1ca5b3ad 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -131,7 +131,7 @@ def handle_get_variable_keys( @handles(MaskSecret) -def handle_mask_secret(msg: MaskSecret) -> None: +def handle_mask_secret(msg: MaskSecret) -> tuple[BaseModel | None, dict[str, bool]]: """Register a value with the secrets masker.""" mask_secret(msg.value, msg.name) return (None, {}) From 66983218830fd33773f0b1a3729d1ad041198020 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 16:00:39 -0700 Subject: [PATCH 05/42] refactor: move client up to WatchedSubProcess base class and replace NotImplementedError with registry dispach --- airflow-core/src/airflow/dag_processing/processor.py | 1 - .../src/airflow/sdk/execution_time/callback_supervisor.py | 2 -- task-sdk/src/airflow/sdk/execution_time/supervisor.py | 7 +++++-- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index f7b3affd00830..a7215a6cbcd1e 100644 --- a/airflow-core/src/airflow/dag_processing/processor.py +++ b/airflow-core/src/airflow/dag_processing/processor.py @@ -551,7 +551,6 @@ class DagFileProcessorProcess(WatchedSubprocess): decoder: ClassVar[TypeAdapter[ToManager]] = TypeAdapter[ToManager](ToManager) had_callbacks: bool = False # Track if this process was started with callbacks to prevent stale DAG detection false positives - client: Client """The HTTP client to use for communication with the API server.""" bundle_name: str diff --git a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py index 7679c2328f963..117494d9f3ed3 100644 --- a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py @@ -164,8 +164,6 @@ class CallbackSubprocess(WatchedSubprocess): ``Connection.get()`` and ``Variable.get()`` via the supervisor's API client. """ - client: Client # The HTTP client to use for communication with the API server. - decoder: ClassVar[TypeAdapter[CallbackToSupervisor]] = TypeAdapter(CallbackToSupervisor) @classmethod diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 0005c36bc7e9b..a541f21466312 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -132,6 +132,7 @@ ) from airflow.sdk.execution_time.coordinator import get_coordinator_manager from airflow.sdk.execution_time.request_handlers import ( + get_handler, handle_delete_variable, handle_delete_xcom, handle_get_connection, @@ -635,6 +636,8 @@ class WatchedSubprocess: No migration is attempted if this is set to *None* (default). """ + client: Client + _exit_code: int | None = attrs.field(default=None, init=False) _process_exit_monotonic: float | None = attrs.field(default=None, init=False) _open_sockets: weakref.WeakKeyDictionary[socket, str] = attrs.field( @@ -948,7 +951,8 @@ def handle_requests(self, log: FilteringBoundLogger) -> Generator[None, _Request otel_context.detach(token) def _handle_request(self, msg, log: FilteringBoundLogger, req_id: int) -> None: - raise NotImplementedError() + resp, dump_opts = get_handler(type(msg))(self.client, msg) + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) @staticmethod def _close_unused_sockets(*sockets): @@ -1275,7 +1279,6 @@ def _remote_logging_conn(client: Client): @attrs.define(kw_only=True) class ActivitySubprocess(WatchedSubprocess): - client: Client """The HTTP client to use for communication with the API server.""" _terminal_state: str | None = attrs.field(default=None, init=False) From 4de5af16589afb4d89da22be1cadd731b2bbfd18 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 19:32:13 -0700 Subject: [PATCH 06/42] refactor: remove redundant elif chains --- .../src/airflow/dag_processing/processor.py | 84 +------ .../src/airflow/jobs/triggerer_job_runner.py | 68 +----- .../sdk/execution_time/callback_supervisor.py | 35 +-- .../airflow/sdk/execution_time/supervisor.py | 221 +++++++----------- 4 files changed, 102 insertions(+), 306 deletions(-) diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index a7215a6cbcd1e..0c804c9a0b1f4 100644 --- a/airflow-core/src/airflow/dag_processing/processor.py +++ b/airflow-core/src/airflow/dag_processing/processor.py @@ -69,24 +69,8 @@ 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.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 @@ -623,75 +607,11 @@ def _create_log_forwarder( ) 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"}, - ), - ) + self.send_msg(None, request_id=req_id, error=None) return - - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + super()._handle_request(msg, log, req_id) @property def is_ready(self) -> bool: diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index e8e1f66e7875c..b63194034f35a 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -33,6 +33,7 @@ from socket import socket from traceback import format_exception from typing import TYPE_CHECKING, Annotated, Any, BinaryIO, ClassVar, Literal, TextIO, TypedDict +from urllib import response from uuid import uuid4 import anyio @@ -89,22 +90,6 @@ _new_encoder, _RequestFrame, ) -from airflow.sdk.execution_time.request_handlers import ( - handle_delete_variable, - handle_delete_xcom, - handle_get_connection, - handle_get_dag_run_state, - handle_get_dr_count, - handle_get_previous_ti, - handle_get_task_states, - handle_get_ti_count, - handle_get_variable, - handle_get_variable_keys, - handle_get_xcom, - handle_mask_secret, - handle_put_variable, - handle_set_xcom, -) from airflow.sdk.execution_time.supervisor import WatchedSubprocess, make_buffered_socket_reader from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance from airflow.serialization.serialized_objects import DagSerialization @@ -515,7 +500,6 @@ def make_client(self) -> Client: def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, req_id: int) -> None: resp: BaseModel | None = None - dump_opts: dict[str, bool] = {} self._last_runner_comms = time.monotonic() if isinstance(msg, messages.TriggerStateChanges): @@ -536,11 +520,6 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r # handle leaks for every failed upload. factory.close() - response = messages.TriggerStateSync( - to_create=[], - to_cancel=self.cancelling_triggers, - ) - # Pull out of these dequeues in a thread-safe manner while self.creating_triggers: workload = self.creating_triggers.popleft() @@ -548,50 +527,23 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r self.running_triggers.update(m.id for m in response.to_create) resp = response - elif isinstance(msg, GetConnection): - resp, dump_opts = handle_get_connection(self.client, msg) - elif isinstance(msg, DeleteVariable): - resp, dump_opts = handle_delete_variable(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, PutVariable): - resp, dump_opts = handle_put_variable(self.client, msg) - elif isinstance(msg, DeleteXCom): - resp, dump_opts = handle_delete_xcom(self.client, msg) - elif isinstance(msg, GetXCom): - resp, dump_opts = handle_get_xcom(self.client, msg) - elif isinstance(msg, SetXCom): - resp, dump_opts = handle_set_xcom(self.client, msg) - 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, 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) - elif isinstance(msg, UpdateHITLDetail): + self.send_msg(response, request_id=req_id, error=None) + return + if isinstance(msg, UpdateHITLDetail): api_resp = self.client.hitl.update_response( ti_id=msg.ti_id, chosen_options=msg.chosen_options, params_input=msg.params_input, ) resp = HITLDetailResponseResult.from_api_response(response=api_resp) - elif isinstance(msg, GetHITLDetailResponse): + self.send_msg(resp, request_id=req_id, error=None) + return + if isinstance(msg, GetHITLDetailResponse): api_resp = self.client.hitl.get_detail_response(ti_id=msg.ti_id) resp = HITLDetailResponseResult.from_api_response(response=api_resp) - elif isinstance(msg, MaskSecret): - handle_mask_secret(msg) - else: - raise ValueError(f"Unknown message type {type(msg)}") - - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + self.send_msg(resp, request_id=req_id, error=None) + return + super()._handle_request(msg, log, req_id) def run(self) -> None: """Run synchronously and handle all database reads/writes.""" diff --git a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py index 117494d9f3ed3..87e1b8860f7b2 100644 --- a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py @@ -30,20 +30,12 @@ from pydantic import Field, TypeAdapter from airflow.sdk._shared.module_loading import accepts_context, accepts_keyword_args -from airflow.sdk.exceptions import ErrorType from airflow.sdk.execution_time.comms import ( - ErrorResponse, GetConnection, GetVariable, GetVariableKeys, MaskSecret, ) -from airflow.sdk.execution_time.request_handlers import ( - handle_get_connection, - handle_get_variable, - handle_get_variable_keys, - handle_mask_secret, -) from airflow.sdk.execution_time.supervisor import ( MIN_HEARTBEAT_INTERVAL, SOCKET_CLEANUP_TIMEOUT, @@ -53,7 +45,6 @@ ) if TYPE_CHECKING: - from pydantic import BaseModel from structlog.typing import FilteringBoundLogger from typing_extensions import Self @@ -302,31 +293,7 @@ def _handle_request(self, msg: CallbackToSupervisor, log: FilteringBoundLogger, log.debug("Received request from callback (body omitted)", msg=type(msg)) else: log.debug("Received request from callback", msg=msg) - - resp: BaseModel | None = None - dump_opts: dict[str, bool] = {} - - if 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, MaskSecret): - handle_mask_secret(msg) - else: - log.warning("Unhandled request from callback subprocess", 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 - - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + super()._handle_request(msg, log, req_id) def _configure_logging(log_path: str) -> tuple[FilteringBoundLogger, BinaryIO]: diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index a541f21466312..74a541bbc2fc0 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -74,8 +74,6 @@ DeleteAssetStoreByName, DeleteAssetStoreByUri, DeleteTaskStore, - DeleteVariable, - DeleteXCom, ErrorResponse, GetAssetByName, GetAssetByUri, @@ -84,30 +82,14 @@ GetAssetsByAlias, GetAssetStoreByName, GetAssetStoreByUri, - GetConnection, GetDag, GetDagRun, - GetDagRunState, - GetDRCount, - GetPreviousDagRun, - GetPreviousTI, - GetPrevSuccessfulDagRun, GetTaskBreadcrumbs, GetTaskRescheduleStartDate, - GetTaskStates, GetTaskStore, - GetTICount, - GetVariable, - GetVariableKeys, - GetXCom, - GetXComCount, - GetXComSequenceItem, - GetXComSequenceSlice, HITLDetailRequestResult, InactiveAssetsResult, - MaskSecret, OKResponse, - PutVariable, RescheduleTask, ResendLoggingFD, RetryTask, @@ -117,7 +99,6 @@ SetRenderedFields, SetRenderedMapIndex, SetTaskStore, - SetXCom, SkipDownstreamTasks, StartupDetails, SucceedTask, @@ -133,25 +114,6 @@ from airflow.sdk.execution_time.coordinator import get_coordinator_manager from airflow.sdk.execution_time.request_handlers import ( get_handler, - handle_delete_variable, - handle_delete_xcom, - handle_get_connection, - handle_get_dag_run_state, - handle_get_dr_count, - 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, - 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, - handle_set_xcom, ) from airflow.sdk.execution_time.schema import get_schema_version_migrator, resolve_body_class @@ -1674,10 +1636,6 @@ def final_state(self): return TaskInstanceState.FAILED def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: int): - 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): @@ -1687,44 +1645,34 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: self._terminal_state = msg.state self._task_end_time_monotonic = time.monotonic() self._rendered_map_index = msg.rendered_map_index - elif isinstance(msg, SucceedTask): + return + if 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): + return + if 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): + return + if isinstance(msg, DeferTask): self._rendered_map_index = msg.rendered_map_index self._send_terminal_state_msg(msg) - elif isinstance(msg, RescheduleTask): + return + if isinstance(msg, RescheduleTask): self._send_terminal_state_msg(msg) - elif isinstance(msg, SkipDownstreamTasks): + return + if 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): + return + if isinstance(msg, SetRenderedFields): self.client.task_instances.set_rtif(self.id, msg.rendered_fields) - elif isinstance(msg, SetRenderedMapIndex): + return + if isinstance(msg, SetRenderedMapIndex): self.client.task_instances.set_rendered_map_index(self.id, msg.rendered_map_index) - elif isinstance(msg, GetAssetByName): + return + if 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) @@ -1732,7 +1680,9 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: dump_opts = {"exclude_unset": True} else: resp = asset_resp - elif isinstance(msg, GetAssetByUri): + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + return + if 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) @@ -1740,9 +1690,13 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: dump_opts = {"exclude_unset": True} else: resp = asset_resp - elif isinstance(msg, GetAssetsByAlias): + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + return + if isinstance(msg, GetAssetsByAlias): resp = self.client.assets.get_by_alias(alias_name=msg.alias_name) - elif isinstance(msg, GetAssetEventByAsset): + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + return + if isinstance(msg, GetAssetEventByAsset): asset_event_resp = self.client.asset_events.get( uri=msg.uri, name=msg.name, @@ -1754,7 +1708,9 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) resp = asset_event_result dump_opts = {"exclude_unset": True} - elif isinstance(msg, GetAssetEventByAssetAlias): + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + return + if isinstance(msg, GetAssetEventByAssetAlias): asset_event_resp = self.client.asset_events.get( alias_name=msg.alias_name, after=msg.after, @@ -1765,47 +1721,42 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: 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): + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + return + if 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): + self.send_msg(resp, request_id=req_id, error=None) + return + if 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): + self.send_msg(resp, request_id=req_id, error=None) + return + if 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): + self.send_msg(resp, request_id=req_id, error=None) + return + if 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): + self.send_msg(resp, request_id=req_id, error=None) + return + if 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): + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + return + if 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): + return + if isinstance(msg, CreateHITLDetailPayload): hitl_detail_request = self.client.hitl.add_response( ti_id=msg.ti_id, options=msg.options, @@ -1818,74 +1769,80 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: ) 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): + self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + return + if isinstance(msg, GetDag): dag = self.client.dags.get( dag_id=msg.dag_id, ) resp = DagResult.from_api_response(dag) - elif isinstance(msg, GetTaskStore): + self.send_msg(resp, request_id=req_id, error=None) + return + if isinstance(msg, GetTaskStore): task_store = self.client.task_store.get(msg.ti_id, msg.key) resp = ( task_store if isinstance(task_store, ErrorResponse) else TaskStoreResult.from_task_store_response(task_store) ) - elif isinstance(msg, SetTaskStore): + self.send_msg(resp, request_id=req_id, error=None) + return + if isinstance(msg, SetTaskStore): self.client.task_store.set(msg.ti_id, msg.key, msg.value, expires_at=msg.expires_at) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteTaskStore): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, DeleteTaskStore): self.client.task_store.delete(msg.ti_id, msg.key) - resp = OKResponse(ok=True) - elif isinstance(msg, ClearTaskStore): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, ClearTaskStore): self.client.task_store.clear(msg.ti_id, all_map_indices=msg.all_map_indices) - resp = OKResponse(ok=True) - elif isinstance(msg, GetAssetStoreByName): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, GetAssetStoreByName): asset_store = self.client.asset_store.get(msg.key, name=msg.name) resp = ( asset_store if isinstance(asset_store, ErrorResponse) else AssetStoreResult.from_asset_store_response(asset_store) ) - elif isinstance(msg, GetAssetStoreByUri): + self.send_msg(resp, request_id=req_id, error=None) + return + if isinstance(msg, GetAssetStoreByUri): asset_store = self.client.asset_store.get(msg.key, uri=msg.uri) resp = ( asset_store if isinstance(asset_store, ErrorResponse) else AssetStoreResult.from_asset_store_response(asset_store) ) - elif isinstance(msg, SetAssetStoreByName): + self.send_msg(resp, request_id=req_id, error=None) + return + if isinstance(msg, SetAssetStoreByName): self.client.asset_store.set(msg.key, msg.value, name=msg.name) - resp = OKResponse(ok=True) - elif isinstance(msg, SetAssetStoreByUri): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, SetAssetStoreByUri): self.client.asset_store.set(msg.key, msg.value, uri=msg.uri) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteAssetStoreByName): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, DeleteAssetStoreByName): self.client.asset_store.delete(msg.key, name=msg.name) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteAssetStoreByUri): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, DeleteAssetStoreByUri): self.client.asset_store.delete(msg.key, uri=msg.uri) - resp = OKResponse(ok=True) - elif isinstance(msg, ClearAssetStoreByName): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, ClearAssetStoreByName): self.client.asset_store.clear(name=msg.name) - resp = OKResponse(ok=True) - elif isinstance(msg, ClearAssetStoreByUri): + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) + return + if isinstance(msg, ClearAssetStoreByUri): self.client.asset_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"}, - ), - ) + self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) return - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) + super()._handle_request(msg, log, req_id) def _send_new_log_fd(self, req_id: int) -> None: if send_fds is None: From 302ef8c64c47b0918d129b44f5f4b01ce98ba7b6 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 22:03:07 -0700 Subject: [PATCH 07/42] refactor: fix _RecordingSupervisor construction in schema integration tests client is now a required field in WatchedSubProcess so we now have to pass a mock to satisfy the constructor --- .../tests/task_sdk/execution_time/schema/test_integration.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/task-sdk/tests/task_sdk/execution_time/schema/test_integration.py b/task-sdk/tests/task_sdk/execution_time/schema/test_integration.py index d609e102f9af0..b432c7cec4471 100644 --- a/task-sdk/tests/task_sdk/execution_time/schema/test_integration.py +++ b/task-sdk/tests/task_sdk/execution_time/schema/test_integration.py @@ -116,6 +116,7 @@ def _new_supervisor(pinned_version: str) -> _RecordingSupervisor: stdin=MagicMock(), process=MagicMock(spec=psutil.Process), process_log=structlog.get_logger(), + client=MagicMock(), ) # In the reimplementation the field is ``_subprocess_schema_version``, # not ``lang_sdk_msg_schema_version``. @@ -312,6 +313,7 @@ def test_no_migration_when_subprocess_schema_version_unset(monkeypatch): stdin=MagicMock(), process=MagicMock(spec=psutil.Process), process_log=structlog.get_logger(), + client=MagicMock(), ) # ``_subprocess_schema_version`` is ``None`` by default; no version # negotiation has happened. From 115000e4fa22d169e37791260845127b9451c272 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 22:06:53 -0700 Subject: [PATCH 08/42] refactor: add client to input filed in MaskSecret handle_mask_secret does not use the client field but it requires the client:Client because it needs to satisfy the uniform handler contract stipulated by the registry dispatcher --- task-sdk/src/airflow/sdk/execution_time/request_handlers.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 5079d1ca5b3ad..69a3fdde616cc 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -131,7 +131,7 @@ def handle_get_variable_keys( @handles(MaskSecret) -def handle_mask_secret(msg: MaskSecret) -> tuple[BaseModel | None, dict[str, bool]]: +def handle_mask_secret(client: Client, msg: MaskSecret) -> tuple[BaseModel | None, dict[str, bool]]: """Register a value with the secrets masker.""" mask_secret(msg.value, msg.name) return (None, {}) From 3d7acdd1283eab6f44bc1072ddaae9f7a2949de3 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 22:12:22 -0700 Subject: [PATCH 09/42] refactor: added empty returns to handle_request to follow supervisor return contract --- task-sdk/src/airflow/sdk/execution_time/supervisor.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 74a541bbc2fc0..e3ba968a3f5fc 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -1645,32 +1645,41 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: self._terminal_state = msg.state self._task_end_time_monotonic = time.monotonic() self._rendered_map_index = msg.rendered_map_index + self.send_msg(None, request_id=req_id, error=None) return + if isinstance(msg, SucceedTask): self._task_end_time_monotonic = time.monotonic() self._rendered_map_index = msg.rendered_map_index self._send_terminal_state_msg(msg) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, RetryTask): self._task_end_time_monotonic = time.monotonic() self._rendered_map_index = msg.rendered_map_index self._send_terminal_state_msg(msg) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, DeferTask): self._rendered_map_index = msg.rendered_map_index self._send_terminal_state_msg(msg) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, RescheduleTask): self._send_terminal_state_msg(msg) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, SkipDownstreamTasks): self.client.task_instances.skip_downstream_tasks(self.id, msg) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, SetRenderedFields): self.client.task_instances.set_rtif(self.id, msg.rendered_fields) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, SetRenderedMapIndex): self.client.task_instances.set_rendered_map_index(self.id, msg.rendered_map_index) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, GetAssetByName): asset_resp = self.client.assets.get(name=msg.name) From 623f64eabae572fc849b47aa65fdd95a964333eb Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 5 Jun 2026 22:58:42 -0700 Subject: [PATCH 10/42] refactor: add empty message send on handle_request to meet message standard --- task-sdk/src/airflow/sdk/execution_time/supervisor.py | 1 + 1 file changed, 1 insertion(+) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index f3a53150a4229..26a8a0a41a40c 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -1675,6 +1675,7 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: if isinstance(msg, AwaitInputTask): self._rendered_map_index = msg.rendered_map_index self._send_terminal_state_msg(msg) + self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, RescheduleTask): self._send_terminal_state_msg(msg) From 1a9ca05a50ca1e68795796a27affda24a6bfaceb Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Sat, 6 Jun 2026 20:41:28 -0700 Subject: [PATCH 11/42] refactor: move isinstance cases to handler registry move GetAssetByName GetAssetByUri GetAssetsByAlias GetAssetEventByAsset GetAssetEventByAssetAlias TriggerDagRun GetDagRun GetDag GetTaskRescheduleStartDate GetTaskBreadcrumbs ValidateInletsAndOutlets CreateHITLDetailPayload GetTaskStore SetTaskStore DeleteTaskStore ClearTaskStore GetAssetStoreByName GetAssetStoreByUri SetAssetStoreByName SetAssetStoreByUri DeleteAssetStoreByName DeleteAssetStoreByUri ClearAssetStoreByName ClearAssetStoreByUri to decorator registry --- .../sdk/execution_time/request_handlers.py | 275 ++++++++++++++++++ 1 file changed, 275 insertions(+) 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 69a3fdde616cc..93edf228d7030 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -31,6 +31,7 @@ from typing import TYPE_CHECKING from airflow.sdk.api.datamodels._generated import ( + AssetResponse, ConnectionResponse, DagRunStateResponse, TaskStatesResponse, @@ -40,17 +41,42 @@ XComSequenceSliceResponse, ) from airflow.sdk.execution_time.comms import ( + AssetEventsResult, + AssetResult, + AssetStoreResult, + ClearAssetStoreByName, + ClearAssetStoreByUri, + ClearTaskStore, ConnectionResult, + CreateHITLDetailPayload, + DagResult, + DagRunResult, DagRunStateResult, + DeleteAssetStoreByName, + DeleteAssetStoreByUri, + DeleteTaskStore, DeleteVariable, DeleteXCom, + ErrorResponse, + GetAssetByName, + GetAssetByUri, + GetAssetEventByAsset, + GetAssetEventByAssetAlias, + GetAssetsByAlias, + GetAssetStoreByName, + GetAssetStoreByUri, GetConnection, + GetDag, + GetDagRun, GetDagRunState, GetDRCount, GetPreviousDagRun, GetPreviousTI, GetPrevSuccessfulDagRun, + GetTaskBreadcrumbs, + GetTaskRescheduleStartDate, GetTaskStates, + GetTaskStore, GetTICount, GetVariable, GetVariableKeys, @@ -58,11 +84,21 @@ GetXComCount, GetXComSequenceItem, GetXComSequenceSlice, + HITLDetailRequestResult, + InactiveAssetsResult, MaskSecret, + OKResponse, PrevSuccessfulDagRunResult, PutVariable, + SetAssetStoreByName, + SetAssetStoreByUri, + SetTaskStore, SetXCom, + TaskBreadcrumbsResult, TaskStatesResult, + TaskStoreResult, + TriggerDagRun, + ValidateInletsAndOutlets, VariableKeysResult, VariableResult, XComResult, @@ -310,3 +346,242 @@ def handle_get_xcom(client: Client, msg: GetXCom) -> tuple[BaseModel | None, dic xcom_result = XComResult.from_xcom_response(xcom) return xcom_result, {"exclude_unset": True} return xcom, {} + + +@handles(GetAssetByName) +def handle_get_asset_by_name(client: Client, msg: GetAssetByName) -> tuple[BaseModel | None, dict[str, bool]]: + asset_resp = client.assets.get(name=msg.name) + if isinstance(asset_resp, AssetResponse): + asset_result = AssetResult.from_asset_response(asset_resp) + return asset_result, {"exclude_unset": True} + return asset_resp, {} + + +@handles(GetAssetByUri) +def handle_get_asset_by_uri(client: Client, msg: GetAssetByUri) -> tuple[BaseModel | None, dict[str, bool]]: + asset_resp = client.assets.get(uri=msg.uri) + if isinstance(asset_resp, AssetResponse): + asset_result = AssetResult.from_asset_response(asset_resp) + return asset_result, {"exclude_unset": True} + return asset_resp, {} + + +@handles(GetAssetsByAlias) +def handle_get_assets_by_alias( + client: Client, msg: GetAssetsByAlias +) -> tuple[BaseModel | None, dict[str, bool]]: + asset_resp = client.assets.get_by_alias(alias_name=msg.alias_name) + if isinstance(asset_resp, AssetResponse): + asset_result = AssetResult.from_asset_response(asset_resp) + return asset_result, {"exclude_unset": True} + return asset_resp, {} + + +@handles(GetAssetEventByAsset) +def handle_get_asset_event_by_asset( + client: Client, msg: GetAssetEventByAsset +) -> tuple[BaseModel | None, dict[str, bool]]: + asset_event_resp = client.asset_events.get( + uri=msg.uri, + name=msg.name, + after=msg.after, + before=msg.before, + ascending=msg.ascending, + limit=msg.limit, + ) + asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) + return asset_event_result, {"exclude_unset": True} + + +@handles(GetAssetEventByAssetAlias) +def handle_get_asset_event_by_asset_alias( + client: Client, msg: GetAssetEventByAssetAlias +) -> tuple[BaseModel | None, dict[str, bool]]: + asset_event_resp = client.asset_events.get( + alias_name=msg.alias_name, + after=msg.after, + before=msg.before, + ascending=msg.ascending, + limit=msg.limit, + ) + asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) + return asset_event_result, {"exclude_unset": True} + + +@handles(TriggerDagRun) +def handle_trigger_dag_run(client: Client, msg: TriggerDagRun) -> tuple[BaseModel | None, dict[str, bool]]: + resp = client.dag_runs.trigger( + msg.dag_id, + msg.run_id, + msg.conf, + msg.logical_date, + msg.run_after, + bool(msg.reset_dag_run), + msg.note, + ) + return resp, {} + + +@handles(GetDagRun) +def handle_get_dag_run(client: Client, msg: GetDagRun) -> tuple[BaseModel | None, dict[str, bool]]: + dr_resp = client.dag_runs.get_detail(msg.dag_id, msg.run_id) + resp = DagRunResult.from_api_response(dr_resp) + return resp, {} + + +@handles(GetDag) +def handle_get_dag(client: Client, msg: GetDag) -> tuple[BaseModel | None, dict[str, bool]]: + dag = client.dags.get( + dag_id=msg.dag_id, + ) + resp = DagResult.from_api_response(dag) + return resp, {} + + +@handles(GetTaskRescheduleStartDate) +def handle_get_task_reschedule_start_date( + client: Client, msg: GetTaskRescheduleStartDate +) -> tuple[BaseModel | None, dict[str, bool]]: + resp = client.task_instances.get_reschedule_start_date(msg.ti_id, msg.try_number) + return resp, {} + + +@handles(GetTaskBreadcrumbs) +def handle_get_task_breadcrumbs( + client: Client, msg: GetTaskBreadcrumbs +) -> tuple[BaseModel | None, dict[str, bool]]: + api_resp = 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, {} + + +@handles(ValidateInletsAndOutlets) +def handle_validate_inlets_and_outlets( + client: Client, msg: ValidateInletsAndOutlets +) -> tuple[BaseModel | None, dict[str, bool]]: + inactive_assets_resp = client.task_instances.validate_inlets_and_outlets(msg.ti_id) + resp = InactiveAssetsResult.from_inactive_assets_response(inactive_assets_resp) + return resp, {"exclude_unset": True} + + +@handles(CreateHITLDetailPayload) +def handle_create_hitl_detail_payload( + client: Client, msg: CreateHITLDetailPayload +) -> tuple[BaseModel | None, dict[str, bool]]: + hitl_detail_request = 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) + return resp, {"exclude_unset": True} + + +@handles(GetTaskStore) +def handle_get_task_store(client: Client, msg: GetTaskStore) -> tuple[BaseModel | None, dict[str, bool]]: + task_store = client.task_store.get(msg.ti_id, msg.key) + resp = ( + task_store + if isinstance(task_store, ErrorResponse) + else TaskStoreResult.from_task_store_response(task_store) + ) + return resp, {} + + +@handles(SetTaskStore) +def handle_set_task_store(client: Client, msg: SetTaskStore) -> tuple[BaseModel | None, dict[str, bool]]: + client.task_store.set(msg.ti_id, msg.key, msg.value, expires_at=msg.expires_at) + return OKResponse(ok=True), {} + + +@handles(DeleteTaskStore) +def handle_delete_task_store( + client: Client, msg: DeleteTaskStore +) -> tuple[BaseModel | None, dict[str, bool]]: + client.task_store.delete(msg.ti_id, msg.key) + return OKResponse(ok=True), {} + + +@handles(ClearTaskStore) +def handle_clear_task_store(client: Client, msg: ClearTaskStore) -> tuple[BaseModel | None, dict[str, bool]]: + client.task_store.clear(msg.ti_id, all_map_indices=msg.all_map_indices) + return OKResponse(ok=True), {} + + +@handles(GetAssetStoreByName) +def handle_get_asset_store_by_name( + client: Client, msg: GetAssetStoreByName +) -> tuple[BaseModel | None, dict[str, bool]]: + asset_store = client.asset_store.get(msg.key, name=msg.name) + resp = ( + asset_store + if isinstance(asset_store, ErrorResponse) + else AssetStoreResult.from_asset_store_response(asset_store) + ) + return resp, {} + + +@handles(GetAssetStoreByUri) +def handle_get_asset_store_by_uri( + client: Client, msg: GetAssetStoreByUri +) -> tuple[BaseModel | None, dict[str, bool]]: + asset_store = client.asset_store.get(msg.key, uri=msg.uri) + resp = ( + asset_store + if isinstance(asset_store, ErrorResponse) + else AssetStoreResult.from_asset_store_response(asset_store) + ) + return resp, {} + + +@handles(SetAssetStoreByName) +def handle_set_asset_store_by_name( + client: Client, msg: SetAssetStoreByName +) -> tuple[BaseModel | None, dict[str, bool]]: + client.asset_store.set(msg.key, msg.value, name=msg.name) + return OKResponse(ok=True), {} + + +@handles(SetAssetStoreByUri) +def handle_set_asset_store_by_uri( + client: Client, msg: SetAssetStoreByUri +) -> tuple[BaseModel | None, dict[str, bool]]: + client.asset_store.set(msg.key, msg.value, uri=msg.uri) + return OKResponse(ok=True), {} + + +@handles(DeleteAssetStoreByName) +def handle_delete_asset_store_by_name( + client: Client, msg: DeleteAssetStoreByName +) -> tuple[BaseModel | None, dict[str, bool]]: + client.asset_store.delete(msg.key, name=msg.name) + return OKResponse(ok=True), {} + + +@handles(DeleteAssetStoreByUri) +def handle_delete_asset_store_by_uri( + client: Client, msg: DeleteAssetStoreByUri +) -> tuple[BaseModel | None, dict[str, bool]]: + client.asset_store.delete(msg.key, uri=msg.uri) + return OKResponse(ok=True), {} + + +@handles(ClearAssetStoreByName) +def handle_clear_asset_store_by_name( + client: Client, msg: ClearAssetStoreByName +) -> tuple[BaseModel | None, dict[str, bool]]: + client.asset_store.clear(name=msg.name) + return OKResponse(ok=True), {} + + +@handles(ClearAssetStoreByUri) +def handle_clear_asset_store_by_uri( + client: Client, msg: ClearAssetStoreByUri +) -> tuple[BaseModel | None, dict[str, bool]]: + client.asset_store.clear(uri=msg.uri) + return OKResponse(ok=True), {} From 196439d214b069154b471d3132dd5f92ff4976cf Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Sat, 6 Jun 2026 20:44:24 -0700 Subject: [PATCH 12/42] refactor: remove redundant isinstance cases remove GetAssetByName GetAssetByUri GetAssetsByAlias GetAssetEventByAsset GetAssetEventByAssetAlias TriggerDagRun GetDagRun GetDag GetTaskRescheduleStartDate GetTaskBreadcrumbs ValidateInletsAndOutlets CreateHITLDetailPayload GetTaskStore SetTaskStore DeleteTaskStore ClearTaskStore GetAssetStoreByName GetAssetStoreByUri SetAssetStoreByName SetAssetStoreByUri DeleteAssetStoreByName DeleteAssetStoreByUri ClearAssetStoreByName ClearAssetStoreByUri from isinstance chain --- .../airflow/sdk/execution_time/supervisor.py | 199 ------------------ 1 file changed, 199 deletions(-) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 26a8a0a41a40c..008b1765269cb 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -51,7 +51,6 @@ from airflow.sdk._shared.logging.structlog import reconfigure_logger from airflow.sdk.api.client import Client, ServerResponseError from airflow.sdk.api.datamodels._generated import ( - AssetResponse, ConnectionResponse, TaskInstance, TaskInstanceState, @@ -60,55 +59,21 @@ from airflow.sdk.exceptions import ErrorType from airflow.sdk.execution_time import comms from airflow.sdk.execution_time.comms import ( - AssetEventsResult, - AssetResult, - AssetStoreResult, AwaitInputTask, - ClearAssetStoreByName, - ClearAssetStoreByUri, - ClearTaskStore, ConnectionResult, - CreateHITLDetailPayload, - DagResult, - DagRunResult, DeferTask, - DeleteAssetStoreByName, - DeleteAssetStoreByUri, - DeleteTaskStore, ErrorResponse, - GetAssetByName, - GetAssetByUri, - GetAssetEventByAsset, - GetAssetEventByAssetAlias, - GetAssetsByAlias, - GetAssetStoreByName, - GetAssetStoreByUri, - GetDag, - GetDagRun, - GetTaskBreadcrumbs, - GetTaskRescheduleStartDate, - GetTaskStore, - HITLDetailRequestResult, - InactiveAssetsResult, - OKResponse, RescheduleTask, ResendLoggingFD, RetryTask, SentFDs, - SetAssetStoreByName, - SetAssetStoreByUri, SetRenderedFields, SetRenderedMapIndex, - SetTaskStore, SkipDownstreamTasks, StartupDetails, SucceedTask, - TaskBreadcrumbsResult, TaskState, - TaskStoreResult, ToSupervisor, - TriggerDagRun, - ValidateInletsAndOutlets, _RequestFrame, _ResponseFrame, ) @@ -1643,8 +1608,6 @@ def final_state(self): return TaskInstanceState.FAILED def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: int): - 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 @@ -1693,83 +1656,6 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: self.client.task_instances.set_rendered_map_index(self.id, msg.rendered_map_index) self.send_msg(None, request_id=req_id, error=None) return - if 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 - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) - return - if 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 - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) - return - if isinstance(msg, GetAssetsByAlias): - resp = self.client.assets.get_by_alias(alias_name=msg.alias_name) - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) - return - if 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, - ) - asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) - resp = asset_event_result - dump_opts = {"exclude_unset": True} - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) - return - if 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, - ) - asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) - resp = asset_event_result - dump_opts = {"exclude_unset": True} - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) - return - if 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 - ) - self.send_msg(resp, request_id=req_id, error=None) - return - if isinstance(msg, GetDagRun): - dr_resp = self.client.dag_runs.get_detail(msg.dag_id, msg.run_id) - resp = DagRunResult.from_api_response(dr_resp) - self.send_msg(resp, request_id=req_id, error=None) - return - if isinstance(msg, GetTaskRescheduleStartDate): - resp = self.client.task_instances.get_reschedule_start_date(msg.ti_id, msg.try_number) - self.send_msg(resp, request_id=req_id, error=None) - return - if 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) - self.send_msg(resp, request_id=req_id, error=None) - return - if 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} - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) - return if isinstance(msg, ResendLoggingFD): # We need special handling here! if send_fds is not None: @@ -1777,91 +1663,6 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: # Since we've sent the message, return. Nothing else in this ifelse/switch should return directly return return - if 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} - self.send_msg(resp, request_id=req_id, error=None, **dump_opts) - return - if isinstance(msg, GetDag): - dag = self.client.dags.get( - dag_id=msg.dag_id, - ) - resp = DagResult.from_api_response(dag) - self.send_msg(resp, request_id=req_id, error=None) - return - if isinstance(msg, GetTaskStore): - task_store = self.client.task_store.get(msg.ti_id, msg.key) - resp = ( - task_store - if isinstance(task_store, ErrorResponse) - else TaskStoreResult.from_task_store_response(task_store) - ) - self.send_msg(resp, request_id=req_id, error=None) - return - if isinstance(msg, SetTaskStore): - self.client.task_store.set(msg.ti_id, msg.key, msg.value, expires_at=msg.expires_at) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, DeleteTaskStore): - self.client.task_store.delete(msg.ti_id, msg.key) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, ClearTaskStore): - self.client.task_store.clear(msg.ti_id, all_map_indices=msg.all_map_indices) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, GetAssetStoreByName): - asset_store = self.client.asset_store.get(msg.key, name=msg.name) - resp = ( - asset_store - if isinstance(asset_store, ErrorResponse) - else AssetStoreResult.from_asset_store_response(asset_store) - ) - self.send_msg(resp, request_id=req_id, error=None) - return - if isinstance(msg, GetAssetStoreByUri): - asset_store = self.client.asset_store.get(msg.key, uri=msg.uri) - resp = ( - asset_store - if isinstance(asset_store, ErrorResponse) - else AssetStoreResult.from_asset_store_response(asset_store) - ) - self.send_msg(resp, request_id=req_id, error=None) - return - if isinstance(msg, SetAssetStoreByName): - self.client.asset_store.set(msg.key, msg.value, name=msg.name) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, SetAssetStoreByUri): - self.client.asset_store.set(msg.key, msg.value, uri=msg.uri) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, DeleteAssetStoreByName): - self.client.asset_store.delete(msg.key, name=msg.name) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, DeleteAssetStoreByUri): - self.client.asset_store.delete(msg.key, uri=msg.uri) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, ClearAssetStoreByName): - self.client.asset_store.clear(name=msg.name) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return - if isinstance(msg, ClearAssetStoreByUri): - self.client.asset_store.clear(uri=msg.uri) - self.send_msg(OKResponse(ok=True), request_id=req_id, error=None) - return super()._handle_request(msg, log, req_id) From 3a047ab82065dd5e0e693ab8ec5ff0d3aa8dd7f9 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Sun, 7 Jun 2026 15:43:59 -0700 Subject: [PATCH 13/42] refactor: updated test case to handle unified handler --- .../tests/unit/dag_processing/test_processor.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/airflow-core/tests/unit/dag_processing/test_processor.py b/airflow-core/tests/unit/dag_processing/test_processor.py index c1d5c76712ecd..98bcba5d38218 100644 --- a/airflow-core/tests/unit/dag_processing/test_processor.py +++ b/airflow-core/tests/unit/dag_processing/test_processor.py @@ -2086,9 +2086,10 @@ def test_handle_request_get_connection_masks_password_and_extra(self, proc): password="super-secret-password", extra='{"api_key":"super-secret-extra"}', ) + mock_masker = MagicMock() with ( - patch("airflow.dag_processing.processor.mask_secret") as mock_mask_secret, + patch("airflow.sdk._shared.secrets_masker._secrets_masker", return_value=mock_masker), patch.object(DagFileProcessorProcess, "send_msg", autospec=True) as mock_send_msg, ): proc._handle_request( @@ -2098,9 +2099,9 @@ def test_handle_request_get_connection_masks_password_and_extra(self, proc): ) proc.client.connections.get.assert_called_once_with("test_conn") - mock_mask_secret.assert_any_call("super-secret-password") - mock_mask_secret.assert_any_call('{"api_key":"super-secret-extra"}') - assert mock_mask_secret.call_count == 2 + mock_masker.add_mask.assert_any_call("super-secret-password", None) + mock_masker.add_mask.assert_any_call('{"api_key":"super-secret-extra"}', None) + assert mock_masker.add_mask.call_count == 2 mock_send_msg.assert_called_once() _, args, kwargs = mock_send_msg.mock_calls[0] @@ -2123,9 +2124,10 @@ def test_handle_request_get_variable_masks_value_with_key(self, proc): key="test_key", value="super-secret-value", ) + mock_masker = MagicMock() with ( - patch("airflow.dag_processing.processor.mask_secret") as mock_mask_secret, + patch("airflow.sdk._shared.secrets_masker._secrets_masker", return_value=mock_masker), patch.object(DagFileProcessorProcess, "send_msg", autospec=True) as mock_send_msg, ): proc._handle_request( @@ -2135,7 +2137,7 @@ def test_handle_request_get_variable_masks_value_with_key(self, proc): ) proc.client.variables.get.assert_called_once_with("test_key") - mock_mask_secret.assert_called_once_with("super-secret-value", "test_key") + mock_masker.add_mask.assert_called_once_with("super-secret-value", "test_key") mock_send_msg.assert_called_once() _, args, kwargs = mock_send_msg.mock_calls[0] From d3bac47df7a1a00da997ee0a2f46034e0aadf280 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Sun, 7 Jun 2026 16:17:29 -0700 Subject: [PATCH 14/42] fix: resolve None client caused by attrs/cached_property conflict Replace cached_property with a property backed by _client. attrs was silently resolving the client to None at runtime, causing NoneType failures when trigger code requested connections or dag run state. The setter allows subclasses like TriggerRunnerSupervisor to inject or override the client during construction. --- .../src/airflow/jobs/triggerer_job_runner.py | 31 ++++++++++--------- 1 file changed, 16 insertions(+), 15 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index b63194034f35a..1e5b8e529c35d 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -17,7 +17,6 @@ from __future__ import annotations import asyncio -import functools import logging import math import os @@ -33,7 +32,6 @@ from socket import socket from traceback import format_exception from typing import TYPE_CHECKING, Annotated, Any, BinaryIO, ClassVar, Literal, TextIO, TypedDict -from urllib import response from uuid import uuid4 import anyio @@ -476,9 +474,15 @@ def start( # type: ignore[override] proc.send_msg(msg, request_id=0) return proc - @functools.cached_property + @property def client(self) -> Client: - return self.make_client() + if self._client is None: + self._client = self.make_client() + return self._client + + @client.setter + def client(self, value: Client | None) -> None: + self._client = value def make_client(self) -> Client: """ @@ -498,8 +502,6 @@ def make_client(self) -> Client: return client def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, req_id: int) -> None: - - resp: BaseModel | None = None self._last_runner_comms = time.monotonic() if isinstance(msg, messages.TriggerStateChanges): @@ -516,19 +518,16 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r except Exception: log.exception("Failed to upload trigger logs to remote", trigger_id=id) finally: - # Close the FD explicitly even if upload raised, otherwise the file - # handle leaks for every failed upload. factory.close() - - # Pull out of these dequeues in a thread-safe manner + sync = messages.TriggerStateSync(to_create=[], to_cancel=set()) + sync.to_cancel = self.cancelling_triggers.copy() while self.creating_triggers: workload = self.creating_triggers.popleft() - response.to_create.append(workload) - self.running_triggers.update(m.id for m in response.to_create) - resp = response - - self.send_msg(response, request_id=req_id, error=None) + sync.to_create.append(workload) + self.running_triggers.update(m.id for m in sync.to_create) + self.send_msg(sync, request_id=req_id, error=None) return + if isinstance(msg, UpdateHITLDetail): api_resp = self.client.hitl.update_response( ti_id=msg.ti_id, @@ -538,11 +537,13 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r resp = HITLDetailResponseResult.from_api_response(response=api_resp) self.send_msg(resp, request_id=req_id, error=None) return + if isinstance(msg, GetHITLDetailResponse): api_resp = self.client.hitl.get_detail_response(ti_id=msg.ti_id) resp = HITLDetailResponseResult.from_api_response(response=api_resp) self.send_msg(resp, request_id=req_id, error=None) return + super()._handle_request(msg, log, req_id) def run(self) -> None: From 10e8db6d8369a257e53cc34726e3b43733fe4a23 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Sun, 7 Jun 2026 16:19:52 -0700 Subject: [PATCH 15/42] fix: rename client field to _client to allow property override attrs was preventing a @property named 'client' from being defined because it owned a field with the same name. Renaming the field to _client with alias='client' keeps the constructor interface unchanged while freeing up the 'client' name for the lazy-init property. --- task-sdk/src/airflow/sdk/execution_time/supervisor.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 008b1765269cb..732bc71a96b03 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -565,7 +565,7 @@ class WatchedSubprocess: No migration is attempted if this is set to *None* (default). """ - client: Client + _client: Client | None = attrs.field(default=None, alias="client", repr=False) _exit_code: int | None = attrs.field(default=None, init=False) _process_exit_monotonic: float | None = attrs.field(default=None, init=False) @@ -585,6 +585,14 @@ class WatchedSubprocess: start_time: float = attrs.field(factory=time.monotonic) """The start time of the child process.""" + @property + def client(self) -> Client | None: + return self._client + + @client.setter + def client(self, value: Client | None) -> None: + self._client = value + @classmethod def start( cls, From 935527df4a3b09cd7cfd1dc6a2249693eee7521e Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Mon, 8 Jun 2026 15:15:32 -0700 Subject: [PATCH 16/42] fix: resolve lazy initialization for client in tests test cases expect client when running test_triggerer_job.py but don't receive it due to lazy initialization being overwritten by attrs setting it to none. Change tests and client initialization to receive client on build time instead of as needed. --- .../src/airflow/jobs/triggerer_job_runner.py | 56 +++++++++---------- .../tests/unit/jobs/test_triggerer_job.py | 33 ++++++----- .../airflow/sdk/execution_time/supervisor.py | 12 +--- 3 files changed, 47 insertions(+), 54 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 1e5b8e529c35d..f840d6437809f 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -468,38 +468,36 @@ def start( # type: ignore[override] **kwargs, ): proc_id = job.id if job is not None else uuid4() - proc = super().start(id=proc_id, job=job, target=cls.run_in_process, logger=logger, **kwargs) - msg = messages.StartTriggerer() - proc.send_msg(msg, request_id=0) - return proc - - @property - def client(self) -> Client: - if self._client is None: - self._client = self.make_client() - return self._client - - @client.setter - def client(self, value: Client | None) -> None: - self._client = value - - def make_client(self) -> Client: - """ - Build the API client used to talk to the API server. - - Subclasses may override this to substitute a different transport — e.g. a - real HTTP client pointing at a remote API server — instead of the default - in-process one. The returned client must have ``base_url`` set; downstream - request handling (``self.client.variables``, ``.xcoms``, etc.) reads it - when issuing requests. - """ from airflow.sdk.api.client import Client - - client = Client(base_url=None, token="", dry_run=True, transport=in_process_api_server().transport) - # Mypy is wrong -- the setter accepts a string on the property setter! `URLType = URL | str` + api = in_process_api_server() + client = Client(base_url=None, token="", dry_run=True, transport=api.transport) client.base_url = "http://in-process.invalid./" - return client + + proc = super().start(id=proc_id, job=job, client=client,target=cls.run_in_process, logger=logger, **kwargs) + proc.send_msg(messages.StartTriggerer(), request_id=0) + return proc + + # @functools.cached_property + # def client(self) -> Client: + # return self.make_client() + + # def make_client(self) -> Client: + # """ + # Build the API client used to talk to the API server. + + # Subclasses may override this to substitute a different transport — e.g. a + # real HTTP client pointing at a remote API server — instead of the default + # in-process one. The returned client must have ``base_url`` set; downstream + # request handling (``self.client.variables``, ``.xcoms``, etc.) reads it + # when issuing requests. + # """ + # from airflow.sdk.api.client import Client + + # client = Client(base_url=None, token="", dry_run=True, transport=in_process_api_server().transport) + # # Mypy is wrong -- the setter accepts a string on the property setter! `URLType = URL | str` + # client.base_url = "http://in-process.invalid./" + # return client def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, req_id: int) -> None: self._last_runner_comms = time.monotonic() diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 4984c0864a816..42c17fd130216 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -34,6 +34,7 @@ from unittest import mock from unittest.mock import ANY, AsyncMock, MagicMock, patch +from airflow.sdk.api.client import Client import greenback import msgspec import pendulum @@ -232,6 +233,7 @@ def builder(job=None): stdin=mock_stdin, process=process, capacity=10, + client = mocker.Mock(spec=Client), ) # Mock the selector mock_selector = mocker.Mock(spec=selectors.DefaultSelector) @@ -264,6 +266,7 @@ def test_supervisor_stores_team_name(supervisor_builder, mocker, session): process=process, capacity=10, team_name="team_x", + client = mocker.Mock(spec=Client), ) assert proc.team_name == "team_x" @@ -276,6 +279,7 @@ def test_supervisor_stores_team_name(supervisor_builder, mocker, session): process=process, capacity=10, team_name=None, + client = mocker.Mock(spec=Client), ) assert proc_global.team_name is None @@ -311,22 +315,22 @@ def fake_run_once(self): assert events == ["enter", "tick-1", "tick-2", "tick-3", "exit"] -def test_client_delegates_to_make_client_and_caches_result(supervisor_builder, mocker): - """``supervisor.client`` delegates to ``make_client`` (the subclass-override hook) - and caches the result across accesses.""" - supervisor = supervisor_builder() - make_client = mocker.patch.object( - TriggerRunnerSupervisor, - "make_client", - autospec=True, - return_value=mocker.sentinel.client, - ) +# def test_client_delegates_to_make_client_and_caches_result(supervisor_builder, mocker): +# """``supervisor.client`` delegates to ``make_client`` (the subclass-override hook) +# and caches the result across accesses.""" +# supervisor = supervisor_builder() +# make_client = mocker.patch.object( +# TriggerRunnerSupervisor, +# "make_client", +# autospec=True, +# return_value=mocker.sentinel.client, +# ) - first = supervisor.client - second = supervisor.client +# first = supervisor.client +# second = supervisor.client - assert first is second is mocker.sentinel.client # cached — same object - make_client.assert_called_once_with(supervisor) +# assert first is second is mocker.sentinel.client # cached — same object +# make_client.assert_called_once_with(supervisor) def test_run_context_exits_when_subprocess_dies(supervisor_builder, mocker): @@ -372,6 +376,7 @@ def jobless_supervisor(mocker): stdin=mock_stdin, process=process, capacity=10, + client=mocker.Mock(spec=Client), ) mock_selector = mocker.Mock(spec=selectors.DefaultSelector) mock_selector.select.return_value = [] diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 732bc71a96b03..e47f52839a0a8 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -565,8 +565,6 @@ class WatchedSubprocess: No migration is attempted if this is set to *None* (default). """ - _client: Client | None = attrs.field(default=None, alias="client", repr=False) - _exit_code: int | None = attrs.field(default=None, init=False) _process_exit_monotonic: float | None = attrs.field(default=None, init=False) _open_sockets: weakref.WeakKeyDictionary[socket, str] = attrs.field( @@ -585,13 +583,7 @@ class WatchedSubprocess: start_time: float = attrs.field(factory=time.monotonic) """The start time of the child process.""" - @property - def client(self) -> Client | None: - return self._client - - @client.setter - def client(self, value: Client | None) -> None: - self._client = value + client: Client = attrs.field(repr=False) @classmethod def start( @@ -1216,8 +1208,6 @@ def _remote_logging_conn(client: Client): @attrs.define(kw_only=True) class ActivitySubprocess(WatchedSubprocess): - """The HTTP client to use for communication with the API server.""" - _terminal_state: str | None = attrs.field(default=None, init=False) _final_state: str | None = attrs.field(default=None, init=False) # The terminal-state message currently being processed by `_handle_request`, From 650b8bf364c996154a0f3ae4e72db440a2d38411 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Mon, 8 Jun 2026 15:44:31 -0700 Subject: [PATCH 17/42] refactor: removed redundant commented out code --- .../tests/unit/jobs/test_triggerer_job.py | 18 ------------------ 1 file changed, 18 deletions(-) diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 381db70cf0d3a..73f9b0d6ea4b9 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -315,24 +315,6 @@ def fake_run_once(self): assert events == ["enter", "tick-1", "tick-2", "tick-3", "exit"] -# def test_client_delegates_to_make_client_and_caches_result(supervisor_builder, mocker): -# """``supervisor.client`` delegates to ``make_client`` (the subclass-override hook) -# and caches the result across accesses.""" -# supervisor = supervisor_builder() -# make_client = mocker.patch.object( -# TriggerRunnerSupervisor, -# "make_client", -# autospec=True, -# return_value=mocker.sentinel.client, -# ) - -# first = supervisor.client -# second = supervisor.client - -# assert first is second is mocker.sentinel.client # cached — same object -# make_client.assert_called_once_with(supervisor) - - def test_run_context_exits_when_subprocess_dies(supervisor_builder, mocker): """Breaking out of the loop on a dead subprocess still unwinds run_context.""" from contextlib import contextmanager From 6fd216b60d9e600c30dfef2e47715742b2be4c80 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 9 Jun 2026 16:26:36 -0700 Subject: [PATCH 18/42] chore: run prek --- .../src/airflow/jobs/triggerer_job_runner.py | 27 +++---------------- .../tests/unit/jobs/test_triggerer_job.py | 8 +++--- generated/provider_dependencies.json | 3 +++ 3 files changed, 11 insertions(+), 27 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 7450e7040f3b4..879f5a5d22d48 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -104,7 +104,6 @@ from airflow.api_fastapi.execution_api.app import InProcessExecutionAPI from airflow.jobs.job import Job - from airflow.sdk.api.client import Client from airflow.sdk.definitions.context import Context from airflow.sdk.types import RuntimeTaskInstanceProtocol as RuntimeTI @@ -480,35 +479,17 @@ def start( # type: ignore[override] proc_id = job.id if job is not None else uuid4() from airflow.sdk.api.client import Client + api = in_process_api_server() client = Client(base_url=None, token="", dry_run=True, transport=api.transport) client.base_url = "http://in-process.invalid./" - proc = super().start(id=proc_id, job=job, client=client,target=cls.run_in_process, logger=logger, **kwargs) + proc = super().start( + id=proc_id, job=job, client=client, target=cls.run_in_process, logger=logger, **kwargs + ) proc.send_msg(messages.StartTriggerer(), request_id=0) return proc - # @functools.cached_property - # def client(self) -> Client: - # return self.make_client() - - # def make_client(self) -> Client: - # """ - # Build the API client used to talk to the API server. - - # Subclasses may override this to substitute a different transport — e.g. a - # real HTTP client pointing at a remote API server — instead of the default - # in-process one. The returned client must have ``base_url`` set; downstream - # request handling (``self.client.variables``, ``.xcoms``, etc.) reads it - # when issuing requests. - # """ - # from airflow.sdk.api.client import Client - - # client = Client(base_url=None, token="", dry_run=True, transport=in_process_api_server().transport) - # # Mypy is wrong -- the setter accepts a string on the property setter! `URLType = URL | str` - # client.base_url = "http://in-process.invalid./" - # return client - def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, req_id: int) -> None: self._last_runner_comms = time.monotonic() diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 73f9b0d6ea4b9..991088d0a0f66 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -34,7 +34,6 @@ from unittest import mock from unittest.mock import ANY, AsyncMock, MagicMock, patch -from airflow.sdk.api.client import Client import greenback import msgspec import pendulum @@ -74,6 +73,7 @@ from airflow.providers.standard.triggers.file import FileDeleteTrigger from airflow.providers.standard.triggers.temporal import DateTimeTrigger, TimeDeltaTrigger from airflow.sdk import DAG, BaseHook, BaseOperator +from airflow.sdk.api.client import Client from airflow.sdk.execution_time.comms import ToSupervisor, ToTask, _RequestFrame, _ResponseFrame from airflow.serialization.serialized_objects import LazyDeserializedDAG from airflow.triggers.base import BaseTrigger, TriggerEvent @@ -233,7 +233,7 @@ def builder(job=None): stdin=mock_stdin, process=process, capacity=10, - client = mocker.Mock(spec=Client), + client=mocker.Mock(spec=Client), ) # Mock the selector mock_selector = mocker.Mock(spec=selectors.DefaultSelector) @@ -266,7 +266,7 @@ def test_supervisor_stores_team_name(supervisor_builder, mocker, session): process=process, capacity=10, team_name="team_x", - client = mocker.Mock(spec=Client), + client=mocker.Mock(spec=Client), ) assert proc.team_name == "team_x" @@ -279,7 +279,7 @@ def test_supervisor_stores_team_name(supervisor_builder, mocker, session): process=process, capacity=10, team_name=None, - client = mocker.Mock(spec=Client), + client=mocker.Mock(spec=Client), ) assert proc_global.team_name is None diff --git a/generated/provider_dependencies.json b/generated/provider_dependencies.json index e97e830bbd0a9..296bba3c24ab3 100644 --- a/generated/provider_dependencies.json +++ b/generated/provider_dependencies.json @@ -1013,6 +1013,7 @@ "http", "microsoft.azure", "microsoft.mssql", + "mongo", "mysql", "openlineage", "oracle", @@ -1253,6 +1254,8 @@ "amazon", "common.compat", "common.messaging", + "google", + "openlineage", "oracle", "sftp" ], From 4c5751a2aec7f0cde240ac910cf4bc2691d99f47 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 9 Jun 2026 16:30:43 -0700 Subject: [PATCH 19/42] add comment to handle_request for explicit close of factory --- airflow-core/src/airflow/jobs/triggerer_job_runner.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 879f5a5d22d48..616537f930c41 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -507,6 +507,8 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r except Exception: log.exception("Failed to upload trigger logs to remote", trigger_id=id) finally: + # Close the FD explicitly even if upload raised, otherwise the file + # handle leaks for every failed upload. factory.close() sync = messages.TriggerStateSync(to_create=[], to_cancel=set()) sync.to_cancel = self.cancelling_triggers.copy() From 51d7ef5ba3bb8571b8e07e301076aed013160a33 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 9 Jun 2026 17:21:31 -0700 Subject: [PATCH 20/42] fix: revert test to original, update mask_secret patch path --- .../tests/unit/dag_processing/test_processor.py | 14 ++++++-------- generated/provider_dependencies.json.sha256sum | 2 +- 2 files changed, 7 insertions(+), 9 deletions(-) diff --git a/airflow-core/tests/unit/dag_processing/test_processor.py b/airflow-core/tests/unit/dag_processing/test_processor.py index 98bcba5d38218..f02cdc0aa80c8 100644 --- a/airflow-core/tests/unit/dag_processing/test_processor.py +++ b/airflow-core/tests/unit/dag_processing/test_processor.py @@ -2086,10 +2086,9 @@ def test_handle_request_get_connection_masks_password_and_extra(self, proc): password="super-secret-password", extra='{"api_key":"super-secret-extra"}', ) - mock_masker = MagicMock() with ( - patch("airflow.sdk._shared.secrets_masker._secrets_masker", return_value=mock_masker), + patch("airflow.sdk.execution_time.request_handlers.mask_secret") as mock_mask_secret, patch.object(DagFileProcessorProcess, "send_msg", autospec=True) as mock_send_msg, ): proc._handle_request( @@ -2099,9 +2098,9 @@ def test_handle_request_get_connection_masks_password_and_extra(self, proc): ) proc.client.connections.get.assert_called_once_with("test_conn") - mock_masker.add_mask.assert_any_call("super-secret-password", None) - mock_masker.add_mask.assert_any_call('{"api_key":"super-secret-extra"}', None) - assert mock_masker.add_mask.call_count == 2 + mock_mask_secret.assert_any_call("super-secret-password") + mock_mask_secret.assert_any_call('{"api_key":"super-secret-extra"}') + assert mock_mask_secret.call_count == 2 mock_send_msg.assert_called_once() _, args, kwargs = mock_send_msg.mock_calls[0] @@ -2124,10 +2123,9 @@ def test_handle_request_get_variable_masks_value_with_key(self, proc): key="test_key", value="super-secret-value", ) - mock_masker = MagicMock() with ( - patch("airflow.sdk._shared.secrets_masker._secrets_masker", return_value=mock_masker), + patch("airflow.sdk.execution_time.request_handlers.mask_secret") as mock_mask_secret, patch.object(DagFileProcessorProcess, "send_msg", autospec=True) as mock_send_msg, ): proc._handle_request( @@ -2137,7 +2135,7 @@ def test_handle_request_get_variable_masks_value_with_key(self, proc): ) proc.client.variables.get.assert_called_once_with("test_key") - mock_masker.add_mask.assert_called_once_with("super-secret-value", "test_key") + mock_mask_secret.assert_called_once_with("super-secret-value", "test_key") mock_send_msg.assert_called_once() _, args, kwargs = mock_send_msg.mock_calls[0] diff --git a/generated/provider_dependencies.json.sha256sum b/generated/provider_dependencies.json.sha256sum index 49ab0db88c08e..659660a6a7bf5 100644 --- a/generated/provider_dependencies.json.sha256sum +++ b/generated/provider_dependencies.json.sha256sum @@ -1 +1 @@ -f042436099826662d45d5f59c100a363d5e12facd51a7c8b850ccbce08d8c4ee +c2a0259b8dbc5d60fdf336b2dbc8ee4860bc94313f620fb70b1d21e8a612072b From b303d9e2fad4a46cc399f777d0fd04c22405565a Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 9 Jun 2026 17:32:52 -0700 Subject: [PATCH 21/42] add: add debug logging for _handle_request --- task-sdk/src/airflow/sdk/execution_time/supervisor.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 26f3f3ff939f2..4c21f75e13bc2 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -63,6 +63,7 @@ ConnectionResult, DeferTask, ErrorResponse, + MaskSecret, RescheduleTask, ResendLoggingFD, RetryTask, @@ -1606,6 +1607,11 @@ def final_state(self): return TaskInstanceState.FAILED def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: int): + 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) + if isinstance(msg, TaskState): # No direct API call here — the recovery path in # `update_task_state_if_needed` will call `finish()` for From 83b5e0308e3fe7ab5e77f148afef82f8998a29c6 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 9 Jun 2026 17:41:31 -0700 Subject: [PATCH 22/42] refactor: removed unnecessary isinstance in getAssetsByAlias --- task-sdk/src/airflow/sdk/execution_time/request_handlers.py | 4 ---- 1 file changed, 4 deletions(-) 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 93edf228d7030..70429b10d7b39 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -371,12 +371,8 @@ def handle_get_assets_by_alias( client: Client, msg: GetAssetsByAlias ) -> tuple[BaseModel | None, dict[str, bool]]: asset_resp = client.assets.get_by_alias(alias_name=msg.alias_name) - if isinstance(asset_resp, AssetResponse): - asset_result = AssetResult.from_asset_response(asset_resp) - return asset_result, {"exclude_unset": True} return asset_resp, {} - @handles(GetAssetEventByAsset) def handle_get_asset_event_by_asset( client: Client, msg: GetAssetEventByAsset From 345af5cfd0a2bd51616e90320902a6faddeddd58 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 9 Jun 2026 18:11:35 -0700 Subject: [PATCH 23/42] refactor: add make_client as classmethod for subclasses to override --- .../src/airflow/jobs/triggerer_job_runner.py | 27 ++++++++++++++----- 1 file changed, 20 insertions(+), 7 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 616537f930c41..26a609d283d19 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -56,6 +56,7 @@ from airflow.models.dagbag import DBDagBag from airflow.models.trigger import Trigger from airflow.observability.metrics import stats_utils +from airflow.sdk.api.client import Client from airflow.sdk.api.datamodels._generated import HITLDetailResponse from airflow.sdk.execution_time.comms import ( CommsDecoder, @@ -478,18 +479,30 @@ def start( # type: ignore[override] ): proc_id = job.id if job is not None else uuid4() - from airflow.sdk.api.client import Client - - api = in_process_api_server() - client = Client(base_url=None, token="", dry_run=True, transport=api.transport) - client.base_url = "http://in-process.invalid./" - proc = super().start( - id=proc_id, job=job, client=client, target=cls.run_in_process, logger=logger, **kwargs + id=proc_id, job=job, client=cls.make_client(), target=cls.run_in_process, logger=logger, **kwargs ) proc.send_msg(messages.StartTriggerer(), request_id=0) return proc + @classmethod + def make_client(cls) -> Client: + """ + Build the API client used to talk to the API server. + + Subclasses may override this to substitute a different transport — e.g. a + real HTTP client pointing at a remote API server — instead of the default + in-process one. The returned client must have ``base_url`` set; downstream + request handling (``self.client.variables``, ``.xcoms``, etc.) reads it + when issuing requests. + """ + from airflow.sdk.api.client import Client + + client = Client(base_url=None, token="", dry_run=True, transport=in_process_api_server().transport) + # Mypy is wrong -- the setter accepts a string on the property setter! `URLType = URL | str` + client.base_url = "http://in-process.invalid./" + return client + def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, req_id: int) -> None: self._last_runner_comms = time.monotonic() From da6d4745d316e7722b435b439c70abc251739f13 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 9 Jun 2026 18:13:19 -0700 Subject: [PATCH 24/42] chore: run prek --- task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py | 1 - task-sdk/src/airflow/sdk/execution_time/request_handlers.py | 1 + task-sdk/src/airflow/sdk/execution_time/supervisor.py | 2 +- 3 files changed, 2 insertions(+), 2 deletions(-) diff --git a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py index d0e553341a47a..9889bb7d8353a 100644 --- a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py @@ -33,7 +33,6 @@ from pydantic import Field, TypeAdapter from airflow.sdk._shared.module_loading import UNUSUAL_MODULE_PREFIX, accepts_context, accepts_keyword_args -from airflow.sdk.exceptions import ErrorType from airflow.sdk.execution_time.comms import ( GetConnection, GetVariable, 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 70429b10d7b39..8c334c3cb41e7 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -373,6 +373,7 @@ def handle_get_assets_by_alias( asset_resp = client.assets.get_by_alias(alias_name=msg.alias_name) return asset_resp, {} + @handles(GetAssetEventByAsset) def handle_get_asset_event_by_asset( client: Client, msg: GetAssetEventByAsset diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 4c21f75e13bc2..e9847f18b5949 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -1611,7 +1611,7 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: log.debug("Received message from task runner (body omitted)", msg=type(msg)) else: log.debug("Received message from task runner", msg=msg) - + if isinstance(msg, TaskState): # No direct API call here — the recovery path in # `update_task_state_if_needed` will call `finish()` for From 0f68dda811b431fef9d9b32a1fe58a126420c1c8 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 10 Jun 2026 13:12:29 -0700 Subject: [PATCH 25/42] chore: prek run --- airflow-core/src/airflow/jobs/triggerer_job_runner.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 26a609d283d19..39ec846a8de85 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -56,7 +56,6 @@ from airflow.models.dagbag import DBDagBag from airflow.models.trigger import Trigger from airflow.observability.metrics import stats_utils -from airflow.sdk.api.client import Client from airflow.sdk.api.datamodels._generated import HITLDetailResponse from airflow.sdk.execution_time.comms import ( CommsDecoder, @@ -105,6 +104,7 @@ from airflow.api_fastapi.execution_api.app import InProcessExecutionAPI from airflow.jobs.job import Job + from airflow.sdk.api.client import Client from airflow.sdk.definitions.context import Context from airflow.sdk.types import RuntimeTaskInstanceProtocol as RuntimeTI From 30db227c8b34ca25e7955a845c96ada0eeabc1b3 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 10 Jun 2026 17:17:26 -0700 Subject: [PATCH 26/42] refactor: remove redundant client comment --- airflow-core/src/airflow/dag_processing/processor.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index 0c804c9a0b1f4..ee16bc77940ec 100644 --- a/airflow-core/src/airflow/dag_processing/processor.py +++ b/airflow-core/src/airflow/dag_processing/processor.py @@ -535,8 +535,6 @@ class DagFileProcessorProcess(WatchedSubprocess): decoder: ClassVar[TypeAdapter[ToManager]] = TypeAdapter[ToManager](ToManager) had_callbacks: bool = False # Track if this process was started with callbacks to prevent stale DAG detection false positives - """The HTTP client to use for communication with the API server.""" - bundle_name: str dag_file_rel_path: str From bf958d3f72a2e17c55d3c0035cbad05a8bd238fe Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 17 Jun 2026 14:03:43 -0700 Subject: [PATCH 27/42] refactor: rename comms functions to follow comms.py definitions --- .../sdk/execution_time/request_handlers.py | 142 +++++++++--------- 1 file changed, 74 insertions(+), 68 deletions(-) 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 8c334c3cb41e7..54cca5ea2d910 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -43,18 +43,18 @@ from airflow.sdk.execution_time.comms import ( AssetEventsResult, AssetResult, - AssetStoreResult, - ClearAssetStoreByName, - ClearAssetStoreByUri, - ClearTaskStore, + AssetStateStoreResult, + ClearAssetStateStoreByName, + ClearAssetStateStoreByUri, + ClearTaskStateStore, ConnectionResult, CreateHITLDetailPayload, DagResult, DagRunResult, DagRunStateResult, - DeleteAssetStoreByName, - DeleteAssetStoreByUri, - DeleteTaskStore, + DeleteAssetStateStoreByName, + DeleteAssetStateStoreByUri, + DeleteTaskStateStore, DeleteVariable, DeleteXCom, ErrorResponse, @@ -63,8 +63,8 @@ GetAssetEventByAsset, GetAssetEventByAssetAlias, GetAssetsByAlias, - GetAssetStoreByName, - GetAssetStoreByUri, + GetAssetStateStoreByName, + GetAssetStateStoreByUri, GetConnection, GetDag, GetDagRun, @@ -76,7 +76,7 @@ GetTaskBreadcrumbs, GetTaskRescheduleStartDate, GetTaskStates, - GetTaskStore, + GetTaskStateStore, GetTICount, GetVariable, GetVariableKeys, @@ -90,13 +90,13 @@ OKResponse, PrevSuccessfulDagRunResult, PutVariable, - SetAssetStoreByName, - SetAssetStoreByUri, - SetTaskStore, + SetAssetStateStoreByName, + SetAssetStateStoreByUri, + SetTaskStateStore, SetXCom, TaskBreadcrumbsResult, TaskStatesResult, - TaskStoreResult, + TaskStateStoreResult, TriggerDagRun, ValidateInletsAndOutlets, VariableKeysResult, @@ -479,106 +479,112 @@ def handle_create_hitl_detail_payload( return resp, {"exclude_unset": True} -@handles(GetTaskStore) -def handle_get_task_store(client: Client, msg: GetTaskStore) -> tuple[BaseModel | None, dict[str, bool]]: - task_store = client.task_store.get(msg.ti_id, msg.key) +@handles(GetTaskStateStore) +def handle_get_task_state_store( + client: Client, msg: GetTaskStateStore +) -> tuple[BaseModel | None, dict[str, bool]]: + task_state_store = client.task_state_store.get(msg.ti_id, msg.key) resp = ( - task_store - if isinstance(task_store, ErrorResponse) - else TaskStoreResult.from_task_store_response(task_store) + task_state_store + if isinstance(task_state_store, ErrorResponse) + else TaskStateStoreResult.from_task_state_store_response(task_state_store) ) return resp, {} -@handles(SetTaskStore) -def handle_set_task_store(client: Client, msg: SetTaskStore) -> tuple[BaseModel | None, dict[str, bool]]: - client.task_store.set(msg.ti_id, msg.key, msg.value, expires_at=msg.expires_at) +@handles(SetTaskStateStore) +def handle_set_task_state_store( + client: Client, msg: SetTaskStateStore +) -> tuple[BaseModel | None, dict[str, bool]]: + client.task_state_store.set(msg.ti_id, msg.key, msg.value, expires_at=msg.expires_at) return OKResponse(ok=True), {} -@handles(DeleteTaskStore) -def handle_delete_task_store( - client: Client, msg: DeleteTaskStore +@handles(DeleteTaskStateStore) +def handle_delete_task_state_store( + client: Client, msg: DeleteTaskStateStore ) -> tuple[BaseModel | None, dict[str, bool]]: - client.task_store.delete(msg.ti_id, msg.key) + client.task_state_store.delete(msg.ti_id, msg.key) return OKResponse(ok=True), {} -@handles(ClearTaskStore) -def handle_clear_task_store(client: Client, msg: ClearTaskStore) -> tuple[BaseModel | None, dict[str, bool]]: - client.task_store.clear(msg.ti_id, all_map_indices=msg.all_map_indices) +@handles(ClearTaskStateStore) +def handle_clear_task_state_store( + client: Client, msg: ClearTaskStateStore +) -> tuple[BaseModel | None, dict[str, bool]]: + client.task_state_store.clear(msg.ti_id, all_map_indices=msg.all_map_indices) return OKResponse(ok=True), {} -@handles(GetAssetStoreByName) -def handle_get_asset_store_by_name( - client: Client, msg: GetAssetStoreByName +@handles(GetAssetStateStoreByName) +def handle_get_asset_state_store_by_name( + client: Client, msg: GetAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: - asset_store = client.asset_store.get(msg.key, name=msg.name) + asset_state_store = client.asset_state_store.get(msg.key, name=msg.name) resp = ( - asset_store - if isinstance(asset_store, ErrorResponse) - else AssetStoreResult.from_asset_store_response(asset_store) + asset_state_store + if isinstance(asset_state_store, ErrorResponse) + else AssetStateStoreResult.from_asset_state_store_response(asset_state_store) ) return resp, {} -@handles(GetAssetStoreByUri) -def handle_get_asset_store_by_uri( - client: Client, msg: GetAssetStoreByUri +@handles(GetAssetStateStoreByUri) +def handle_get_asset_state_store_by_uri( + client: Client, msg: GetAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: - asset_store = client.asset_store.get(msg.key, uri=msg.uri) + asset_state_store = client.asset_state_store.get(msg.key, uri=msg.uri) resp = ( - asset_store - if isinstance(asset_store, ErrorResponse) - else AssetStoreResult.from_asset_store_response(asset_store) + asset_state_store + if isinstance(asset_state_store, ErrorResponse) + else AssetStateStoreResult.from_asset_state_store_response(asset_state_store) ) return resp, {} -@handles(SetAssetStoreByName) -def handle_set_asset_store_by_name( - client: Client, msg: SetAssetStoreByName +@handles(SetAssetStateStoreByName) +def handle_set_asset_state_store_by_name( + client: Client, msg: SetAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_store.set(msg.key, msg.value, name=msg.name) + client.asset_state_store.set(msg.key, msg.value, name=msg.name) return OKResponse(ok=True), {} -@handles(SetAssetStoreByUri) -def handle_set_asset_store_by_uri( - client: Client, msg: SetAssetStoreByUri +@handles(SetAssetStateStoreByUri) +def handle_set_asset_state_store_by_uri( + client: Client, msg: SetAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_store.set(msg.key, msg.value, uri=msg.uri) + client.asset_state_store.set(msg.key, msg.value, uri=msg.uri) return OKResponse(ok=True), {} -@handles(DeleteAssetStoreByName) -def handle_delete_asset_store_by_name( - client: Client, msg: DeleteAssetStoreByName +@handles(DeleteAssetStateStoreByName) +def handle_delete_asset_state_store_by_name( + client: Client, msg: DeleteAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_store.delete(msg.key, name=msg.name) + client.asset_state_store.delete(msg.key, name=msg.name) return OKResponse(ok=True), {} -@handles(DeleteAssetStoreByUri) -def handle_delete_asset_store_by_uri( - client: Client, msg: DeleteAssetStoreByUri +@handles(DeleteAssetStateStoreByUri) +def handle_delete_asset_state_store_by_uri( + client: Client, msg: DeleteAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_store.delete(msg.key, uri=msg.uri) + client.asset_state_store.delete(msg.key, uri=msg.uri) return OKResponse(ok=True), {} -@handles(ClearAssetStoreByName) -def handle_clear_asset_store_by_name( - client: Client, msg: ClearAssetStoreByName +@handles(ClearAssetStateStoreByName) +def handle_clear_asset_state_store_by_name( + client: Client, msg: ClearAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_store.clear(name=msg.name) + client.asset_state_store.clear(name=msg.name) return OKResponse(ok=True), {} -@handles(ClearAssetStoreByUri) -def handle_clear_asset_store_by_uri( - client: Client, msg: ClearAssetStoreByUri +@handles(ClearAssetStateStoreByUri) +def handle_clear_asset_state_store_by_uri( + client: Client, msg: ClearAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_store.clear(uri=msg.uri) + client.asset_state_store.clear(uri=msg.uri) return OKResponse(ok=True), {} From 5eeb88fd511115e4c8eb12c5914ddd9f44267856 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 17 Jun 2026 14:04:19 -0700 Subject: [PATCH 28/42] chore: lint check --- airflow-core/src/airflow/jobs/triggerer_job_runner.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 721442a6d8015..eff9537d414d7 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -539,7 +539,7 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r # Close the FD explicitly even if upload raised, otherwise the file # handle leaks for every failed upload. factory.close() - + # Drain the persist confirmations accumulated since the last sync. events_persisted: list[int] = [] while self.persisted_event_seqs: @@ -556,7 +556,7 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r workload = self.creating_triggers.popleft() sync.to_create.append(workload) self.running_triggers.update(m.id for m in sync.to_create) - resp=sync + resp = sync return if isinstance(msg, UpdateHITLDetail): From df344c2b2367b21d735141531acce59bd60bed04 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 17 Jun 2026 14:58:03 -0700 Subject: [PATCH 29/42] refactor: add frozen set to resolve toSupervisor type alias problem toSupervisor is a type alias and not a collection so in is supported. To resolve create a helper function to make a frozen set and using get_args to unpack it. Add ovveride as well so the uniion set can be checked for specific input before running --- .../src/airflow/sdk/execution_time/supervisor.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 8eb2ce99f5cfd..557a8facb0640 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -37,7 +37,7 @@ from datetime import datetime, timezone 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, TextIO, cast, get_args from urllib.parse import urlparse from uuid import UUID @@ -68,17 +68,12 @@ ResendLoggingFD, RetryTask, SentFDs, - SetAssetStateStoreByName, - SetAssetStateStoreByUri, SetRenderedFields, SetRenderedMapIndex, - SetTaskStateStore, - SetXCom, SkipDownstreamTasks, StartupDetails, SucceedTask, TaskState, - TaskStateStoreResult, ToSupervisor, _RequestFrame, _ResponseFrame, @@ -545,6 +540,13 @@ class WatchedSubprocess: socket handling, process monitoring, and request handling. """ + _msg_union: ClassVar[type] # subclasses override + + @classmethod + @functools.cache + def _allowed_msg_types(cls) -> frozenset[type]: + return frozenset(get_args(get_args(cls._msg_union)[0])) + id: UUID pid: int From 51a28fd67e52938d7bc549401cac50e92a340f7f Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 17 Jun 2026 14:58:12 -0700 Subject: [PATCH 30/42] refactor: add frozen set to resolve toSupervisor type alias problem toSupervisor is a type alias and not a collection so in is supported. To resolve create a helper function to make a frozen set and using get_args to unpack it. Add ovveride as well so the uniion set can be checked for specific input before running --- task-sdk/src/airflow/sdk/execution_time/supervisor.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 557a8facb0640..e3866da242fee 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -887,6 +887,10 @@ def handle_requests(self, log: FilteringBoundLogger) -> Generator[None, _Request otel_context.detach(token) def _handle_request(self, msg, log: FilteringBoundLogger, req_id: int) -> None: + if type(msg) not in self._allowed_msg_types(): + log.error("Unhandled request", msg=msg) + self.send_msg(None, request_id=req_id, error=ErrorResponse(error=ErrorType.UNKNOWN_REQUEST, detail={"type": type(msg).__name__})) + return resp, dump_opts = get_handler(type(msg))(self.client, msg) self.send_msg(resp, request_id=req_id, error=None, **dump_opts) From 9e240c619476f376c9b27b7d36b07365d0ed6019 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 17 Jun 2026 17:58:44 -0700 Subject: [PATCH 31/42] refactor: add unioned type checking create a inherited classmethod in watchedSubProcess that is overrode by subprocesses as a frozen set so that processes can check that they only use valid handlers from request_handlers --- airflow-core/src/airflow/dag_processing/processor.py | 3 ++- airflow-core/src/airflow/jobs/triggerer_job_runner.py | 1 + .../src/airflow/sdk/execution_time/callback_supervisor.py | 3 ++- task-sdk/src/airflow/sdk/execution_time/supervisor.py | 8 ++++++-- 4 files changed, 11 insertions(+), 4 deletions(-) diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index ee16bc77940ec..1246429c78089 100644 --- a/airflow-core/src/airflow/dag_processing/processor.py +++ b/airflow-core/src/airflow/dag_processing/processor.py @@ -23,7 +23,7 @@ import traceback from collections.abc import Callable, Sequence from pathlib import Path -from typing import TYPE_CHECKING, Annotated, BinaryIO, ClassVar, Literal +from typing import TYPE_CHECKING, Annotated, Any, BinaryIO, ClassVar, Literal import attrs from pydantic import BaseModel, Field, TypeAdapter @@ -530,6 +530,7 @@ class DagFileProcessorProcess(WatchedSubprocess): in core Airflow. """ + _msg_union: ClassVar[Any] = ToManager logger_filehandle: BinaryIO parsing_result: DagFileParsingResult | None = None decoder: ClassVar[TypeAdapter[ToManager]] = TypeAdapter[ToManager](ToManager) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index eff9537d414d7..54c715e498650 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -437,6 +437,7 @@ class TriggerRunnerSupervisor(WatchedSubprocess): rather than silently zombieing the supervisor. """ + _msg_union: ClassVar[Any] = ToTriggerSupervisor job: Job | None = None capacity: int queues: set[str] | None = None diff --git a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py index 9889bb7d8353a..7527c427690db 100644 --- a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py @@ -25,7 +25,7 @@ from importlib import import_module from importlib.util import module_from_spec, spec_from_file_location from pathlib import Path -from typing import TYPE_CHECKING, Annotated, BinaryIO, ClassVar, Protocol +from typing import TYPE_CHECKING, Annotated, Any, BinaryIO, ClassVar, Protocol from uuid import UUID import attrs @@ -178,6 +178,7 @@ class CallbackSubprocess(WatchedSubprocess): ``Connection.get()`` and ``Variable.get()`` via the supervisor's API client. """ + _msg_union: ClassVar[Any] = CallbackToSupervisor decoder: ClassVar[TypeAdapter[CallbackToSupervisor]] = TypeAdapter(CallbackToSupervisor) @classmethod diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index e3866da242fee..0e7daccde578c 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -540,7 +540,7 @@ class WatchedSubprocess: socket handling, process monitoring, and request handling. """ - _msg_union: ClassVar[type] # subclasses override + _msg_union: ClassVar[Any] = ToSupervisor # subclasses override @classmethod @functools.cache @@ -889,7 +889,11 @@ def handle_requests(self, log: FilteringBoundLogger) -> Generator[None, _Request def _handle_request(self, msg, log: FilteringBoundLogger, req_id: int) -> None: if type(msg) not in self._allowed_msg_types(): log.error("Unhandled request", msg=msg) - self.send_msg(None, request_id=req_id, error=ErrorResponse(error=ErrorType.UNKNOWN_REQUEST, detail={"type": type(msg).__name__})) + self.send_msg( + None, + request_id=req_id, + error=ErrorResponse(error=ErrorType.API_SERVER_ERROR, detail={"type": type(msg).__name__}), + ) return resp, dump_opts = get_handler(type(msg))(self.client, msg) self.send_msg(resp, request_id=req_id, error=None, **dump_opts) From 865e0f2f6b7d35cefb1ab3a1b6af43d5f4722c78 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 17 Jun 2026 18:25:56 -0700 Subject: [PATCH 32/42] refactor: add self.send_msg to triggerstatechange --- airflow-core/src/airflow/jobs/triggerer_job_runner.py | 1 + 1 file changed, 1 insertion(+) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 54c715e498650..ef79c1503f194 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -558,6 +558,7 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r sync.to_create.append(workload) self.running_triggers.update(m.id for m in sync.to_create) resp = sync + self.send_msg(resp, request_id=req_id, error=None) return if isinstance(msg, UpdateHITLDetail): From 5502655246d9633119303071bbd83fce7d91739d Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Wed, 17 Jun 2026 19:55:47 -0700 Subject: [PATCH 33/42] chore: remove redundant variable assignment --- airflow-core/src/airflow/jobs/triggerer_job_runner.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index ef79c1503f194..38c435950353b 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -557,8 +557,7 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r workload = self.creating_triggers.popleft() sync.to_create.append(workload) self.running_triggers.update(m.id for m in sync.to_create) - resp = sync - self.send_msg(resp, request_id=req_id, error=None) + self.send_msg(sync, request_id=req_id, error=None) return if isinstance(msg, UpdateHITLDetail): From 253b8a8b2ba21ee04ab69ad8b5907a873748515b Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Mon, 22 Jun 2026 17:43:38 -0700 Subject: [PATCH 34/42] resolve broken imports --- airflow-core/tests/unit/jobs/test_triggerer_job.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 86c48ab76827e..988dcec133e07 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -77,8 +77,16 @@ from airflow.providers.standard.triggers.temporal import DateTimeTrigger, TimeDeltaTrigger from airflow.sdk import DAG, BaseHook, BaseOperator from airflow.sdk.api.client import Client -from airflow.sdk.execution_time.comms import ToSupervisor, ToTask, _RequestFrame, _ResponseFrame -from airflow.sdk.execution_time.context import AssetStateStoreAccessors +from airflow.sdk.execution_time.comms import ( + AssetStateStoreResult, + GetAssetStateStoreByName, + SetAssetStateStoreByName, + ToSupervisor, + ToTask, + _RequestFrame, + _ResponseFrame, +) +from airflow.sdk.execution_time.context import Asset, AssetStateStoreAccessors from airflow.serialization.serialized_objects import LazyDeserializedDAG from airflow.triggers.base import BaseEventTrigger, BaseTrigger, TriggerEvent from airflow.triggers.shared_stream import SharedStreamProducer From 3173a1a6540115f96877a4af909821644a3a92c2 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 23 Jun 2026 14:40:37 -0700 Subject: [PATCH 35/42] fix task-sdk failing test case --- task-sdk/src/airflow/sdk/execution_time/request_handlers.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 54cca5ea2d910..c2471b11d8afe 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -512,7 +512,7 @@ def handle_delete_task_state_store( def handle_clear_task_state_store( client: Client, msg: ClearTaskStateStore ) -> tuple[BaseModel | None, dict[str, bool]]: - client.task_state_store.clear(msg.ti_id, all_map_indices=msg.all_map_indices) + client.task_state_store.clear(msg.ti_id) return OKResponse(ok=True), {} From e1066c50cbddddf4c55046283c872e030b72b677 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Thu, 25 Jun 2026 16:31:22 -0700 Subject: [PATCH 36/42] resolve test failures due to misaligned test contracts with request handlers and supervisor tests --- .../sdk/execution_time/request_handlers.py | 27 +++++++++---------- .../execution_time/test_request_handlers.py | 27 ++++++++++++++----- 2 files changed, 33 insertions(+), 21 deletions(-) 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 fd9d7522779f0..f760e062b6321 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -31,7 +31,6 @@ from typing import TYPE_CHECKING from airflow.sdk.api.datamodels._generated import ( - AssetStateStoreResponse, ConnectionResponse, DagRunStateResponse, TaskStatesResponse, @@ -42,7 +41,9 @@ ) from airflow.sdk.execution_time.comms import ( AssetEventsResult, + AssetResponse, AssetResult, + AssetStateStoreResponse, AssetStateStoreResult, ClearAssetStateStoreByName, ClearAssetStateStoreByUri, @@ -520,26 +521,22 @@ def handle_clear_task_state_store( def handle_get_asset_state_store_by_name( client: Client, msg: GetAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: - asset_state_store = client.asset_state_store.get(msg.key, name=msg.name) - resp = ( - asset_state_store - if isinstance(asset_state_store, ErrorResponse) - else AssetStateStoreResult.from_asset_state_store_response(asset_state_store) - ) - return resp, {} + asset_state = client.asset_state_store.get(msg.key, name=msg.name) + + if isinstance(asset_state, AssetStateStoreResponse): + return AssetStateStoreResult.from_asset_state_store_response(asset_state), {} + return asset_state, {} @handles(GetAssetStateStoreByUri) def handle_get_asset_state_store_by_uri( client: Client, msg: GetAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: - asset_state_store = client.asset_state_store.get(msg.key, uri=msg.uri) - resp = ( - asset_state_store - if isinstance(asset_state_store, ErrorResponse) - else AssetStateStoreResult.from_asset_state_store_response(asset_state_store) - ) - return resp, {} + asset_state = client.asset_state_store.get(msg.key, uri=msg.uri) + + if isinstance(asset_state, AssetStateStoreResponse): + return AssetStateStoreResult.from_asset_state_store_response(asset_state), {} + return asset_state, {} @handles(SetAssetStateStoreByName) diff --git a/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py b/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py index c3999f9baca9a..67647db6c3ce5 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py +++ b/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py @@ -32,6 +32,7 @@ ErrorResponse, GetAssetStateStoreByName, GetAssetStateStoreByUri, + OKResponse, SetAssetStateStoreByName, SetAssetStateStoreByUri, ) @@ -59,7 +60,7 @@ def test_get_asset_state_store_by_name_wraps_response_as_result(client): client, GetAssetStateStoreByName(name="asset_a", key="watermark") ) - client.asset_state_store.get.assert_called_once_with(key="watermark", name="asset_a") + client.asset_state_store.get.assert_called_once_with("watermark", name="asset_a") assert result == AssetStateStoreResult(value="2026-01-01") assert dump_opts == {} @@ -83,7 +84,7 @@ def test_get_asset_state_store_by_uri_wraps_response_as_result(client): client, GetAssetStateStoreByUri(uri="s3://bucket/a", key="watermark") ) - client.asset_state_store.get.assert_called_once_with(key="watermark", uri="s3://bucket/a") + client.asset_state_store.get.assert_called_once_with("watermark", uri="s3://bucket/a") assert result == AssetStateStoreResult(value="2026-01-01") assert dump_opts == {} @@ -101,7 +102,7 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): @pytest.mark.parametrize( - ("handler", "msg", "call_kwargs", "method"), + ("handler", "msg", "call_kwargs", "method", "expected_args", "expected_kwargs"), [ ( handle_set_asset_state_store_by_name, @@ -112,6 +113,8 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): "value": "2026-01-01", }, "set", + ("watermark", "2026-01-01"), + {"name": "asset_a"}, ), ( handle_set_asset_state_store_by_uri, @@ -122,6 +125,8 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): "value": "2026-01-01", }, "set", + ("watermark", "2026-01-01"), + {"uri": "s3://bucket/a"}, ), ( handle_delete_asset_state_store_by_name, @@ -131,6 +136,8 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): "key": "watermark", }, "delete", + ("watermark",), + {"name": "asset_a"}, ), ( handle_delete_asset_state_store_by_uri, @@ -140,6 +147,8 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): "key": "watermark", }, "delete", + ("watermark",), + {"uri": "s3://bucket/a"}, ), ( handle_clear_asset_state_store_by_name, @@ -148,6 +157,8 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): "name": "asset_a", }, "clear", + (), + {"name": "asset_a"}, ), ( handle_clear_asset_state_store_by_uri, @@ -156,12 +167,16 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): "uri": "s3://bucket/a", }, "clear", + (), + {"uri": "s3://bucket/a"}, ), ], ) -def test_asset_store_delegates_to_client(client, handler, msg, call_kwargs, method): +def test_asset_store_delegates_to_client( + client, handler, msg, call_kwargs, method, expected_args, expected_kwargs +): result, dump_opts = handler(client, msg(**call_kwargs)) - getattr(client.asset_state_store, method).assert_called_once_with(**call_kwargs) - assert result is None + getattr(client.asset_state_store, method).assert_called_once_with(*expected_args, **expected_kwargs) + assert result == OKResponse(ok=True) assert dump_opts == {} From c62a4b4e7df0c4d87e5aa21548a0338c8bc12f42 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Thu, 25 Jun 2026 17:27:00 -0700 Subject: [PATCH 37/42] resolve undefined names in test cases --- airflow-core/tests/unit/jobs/test_triggerer_job.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index ddb71ee83b653..625bcbd48c0ed 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -78,12 +78,14 @@ from airflow.sdk import DAG, BaseHook, BaseOperator from airflow.sdk.api.client import Client from airflow.sdk.execution_time.comms import ( + AssetStateStoreResponse, AssetStateStoreResult, ClearAssetStateStoreByName, ClearAssetStateStoreByUri, DeleteAssetStateStoreByName, DeleteAssetStateStoreByUri, ErrorResponse, + ErrorType, GetAssetStateStoreByName, GetAssetStateStoreByUri, OKResponse, From 3d23ebf4e4be5a04ff3bbb85fc69c4db73d78ce0 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Thu, 25 Jun 2026 18:51:27 -0700 Subject: [PATCH 38/42] align the input output contract in the test cases for request handlers --- airflow-core/tests/unit/jobs/test_triggerer_job.py | 14 ++++++-------- 1 file changed, 6 insertions(+), 8 deletions(-) diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 625bcbd48c0ed..b89ca3ba2dde1 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -854,7 +854,7 @@ def test_get_by_name_wraps_response_and_replies(self, supervisor): self._handle(supervisor, GetAssetStateStoreByName(name="asset_a", key="watermark")) - supervisor.client.asset_state_store.get.assert_called_once_with(key="watermark", name="asset_a") + supervisor.client.asset_state_store.get.assert_called_once_with("watermark", name="asset_a") supervisor.send_msg.assert_called_once_with( AssetStateStoreResult(value="2026-01-01"), request_id=7, error=None ) @@ -864,7 +864,7 @@ def test_get_by_uri_wraps_response_and_replies(self, supervisor): self._handle(supervisor, GetAssetStateStoreByUri(uri="s3://bucket/a", key="watermark")) - supervisor.client.asset_state_store.get.assert_called_once_with(key="watermark", uri="s3://bucket/a") + supervisor.client.asset_state_store.get.assert_called_once_with("watermark", uri="s3://bucket/a") supervisor.send_msg.assert_called_once_with( AssetStateStoreResult(value="2026-01-01"), request_id=7, error=None ) @@ -883,7 +883,7 @@ def test_set_by_name(self, supervisor): ) supervisor.client.asset_state_store.set.assert_called_once_with( - key="watermark", value="2026-01-01", name="asset_a" + "watermark", "2026-01-01", name="asset_a" ) supervisor.send_msg.assert_called_once_with(OKResponse(ok=True), request_id=7, error=None) @@ -893,22 +893,20 @@ def test_set_by_uri(self, supervisor): ) supervisor.client.asset_state_store.set.assert_called_once_with( - key="watermark", value="2026-01-01", uri="s3://bucket/a" + "watermark", "2026-01-01", uri="s3://bucket/a" ) supervisor.send_msg.assert_called_once_with(OKResponse(ok=True), request_id=7, error=None) def test_delete_by_name(self, supervisor): self._handle(supervisor, DeleteAssetStateStoreByName(name="asset_a", key="watermark")) - supervisor.client.asset_state_store.delete.assert_called_once_with(key="watermark", name="asset_a") + supervisor.client.asset_state_store.delete.assert_called_once_with("watermark", name="asset_a") supervisor.send_msg.assert_called_once_with(OKResponse(ok=True), request_id=7, error=None) def test_delete_by_uri(self, supervisor): self._handle(supervisor, DeleteAssetStateStoreByUri(uri="s3://bucket/a", key="watermark")) - supervisor.client.asset_state_store.delete.assert_called_once_with( - key="watermark", uri="s3://bucket/a" - ) + supervisor.client.asset_state_store.delete.assert_called_once_with("watermark", uri="s3://bucket/a") supervisor.send_msg.assert_called_once_with(OKResponse(ok=True), request_id=7, error=None) def test_clear_by_name(self, supervisor): From e32ff462b80ff19e011c8dee2cc49660b87bb71a Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Sat, 4 Jul 2026 13:10:13 -0700 Subject: [PATCH 39/42] fix breaking merge conflicts with triggerer and broken imports --- airflow-core/src/airflow/jobs/triggerer_job_runner.py | 1 + airflow-core/tests/unit/jobs/test_triggerer_job.py | 2 -- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 17130e5ed76f0..77805b9707234 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -522,6 +522,7 @@ def start( # type: ignore[override] proc = super().start( id=proc_id, job=job, + client=cls.make_client(), target=cls.run_in_process, logger=logger, use_exec=supervisor._should_use_exec(), diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index 35389ec043849..22a25d5fa233e 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -77,8 +77,6 @@ from airflow.providers.standard.triggers.temporal import DateTimeTrigger, TimeDeltaTrigger from airflow.sdk import DAG, BaseHook, BaseOperator from airflow.sdk.api.client import Client -from airflow.sdk.api.datamodels._generated import AssetStateStoreResponse -from airflow.sdk.exceptions import ErrorType from airflow.sdk.execution_time import supervisor from airflow.sdk.execution_time.comms import ( AssetStateStoreResponse, From d2906da1f1db0717a5b4ac2f73426adcf097be5d Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 10 Jul 2026 22:06:31 -0700 Subject: [PATCH 40/42] resolve failing test case by realigning output contracts --- .../src/airflow/sdk/execution_time/request_handlers.py | 3 +++ task-sdk/src/airflow/sdk/execution_time/supervisor.py | 10 +++++++++- 2 files changed, 12 insertions(+), 1 deletion(-) 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 f760e062b6321..0cbed8cd8d94c 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -30,6 +30,7 @@ from collections.abc import Callable from typing import TYPE_CHECKING +from airflow.sdk import timezone from airflow.sdk.api.datamodels._generated import ( ConnectionResponse, DagRunStateResponse, @@ -386,6 +387,7 @@ def handle_get_asset_event_by_asset( before=msg.before, ascending=msg.ascending, limit=msg.limit, + extra=msg.extra, ) asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) return asset_event_result, {"exclude_unset": True} @@ -401,6 +403,7 @@ def handle_get_asset_event_by_asset_alias( before=msg.before, ascending=msg.ascending, limit=msg.limit, + extra=msg.extra, ) asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) return asset_event_result, {"exclude_unset": True} diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index c31ffaca83636..44728a2f3627b 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -1677,7 +1677,15 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, SetRenderedFields): - self.client.task_instances.set_rtif(self.id, msg.rendered_fields) + 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) self.send_msg(None, request_id=req_id, error=None) return if isinstance(msg, SetRenderedMapIndex): From 518e8a8c3a58a35edf696d86b04dbabd603d814d Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Fri, 10 Jul 2026 22:07:43 -0700 Subject: [PATCH 41/42] remove unused import --- task-sdk/src/airflow/sdk/execution_time/request_handlers.py | 1 - 1 file changed, 1 deletion(-) 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 0cbed8cd8d94c..efd1ff2a32473 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -30,7 +30,6 @@ from collections.abc import Callable from typing import TYPE_CHECKING -from airflow.sdk import timezone from airflow.sdk.api.datamodels._generated import ( ConnectionResponse, DagRunStateResponse, From adb7b3fae2740de11e7ee2aa2276279ef362c8b5 Mon Sep 17 00:00:00 2001 From: Jung-Hyun Andrew Kim Date: Tue, 21 Jul 2026 13:23:12 -0700 Subject: [PATCH 42/42] resolve failing test cases by updating output to include partition key --- .../src/airflow/sdk/execution_time/request_handlers.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) 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 efd1ff2a32473..6cef3a759ab28 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -384,8 +384,10 @@ def handle_get_asset_event_by_asset( name=msg.name, after=msg.after, before=msg.before, - ascending=msg.ascending, limit=msg.limit, + ascending=msg.ascending, + 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) @@ -400,8 +402,10 @@ def handle_get_asset_event_by_asset_alias( alias_name=msg.alias_name, after=msg.after, before=msg.before, - ascending=msg.ascending, limit=msg.limit, + ascending=msg.ascending, + 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)