diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index ca7539131f60e..680c1973cfbc7 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 @@ -71,24 +71,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 @@ -562,14 +546,12 @@ class DagFileProcessorProcess(WatchedSubprocess, LoggingMixin): in core Airflow. """ + _msg_union: ClassVar[Any] = ToManager logger_filehandle: BinaryIO parsing_result: DagFileParsingResult | None = None 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 dag_file_rel_path: str @@ -651,75 +633,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 759735242ab1d..aae1b69518799 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 @@ -101,30 +100,6 @@ _RequestFrame, ) from airflow.sdk.execution_time.context import AssetStateStoreAccessors -from airflow.sdk.execution_time.request_handlers import ( - handle_clear_asset_state_store_by_name, - handle_clear_asset_state_store_by_uri, - handle_delete_asset_state_store_by_name, - handle_delete_asset_state_store_by_uri, - handle_delete_variable, - handle_delete_xcom, - handle_get_asset_state_store_by_name, - handle_get_asset_state_store_by_uri, - 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_asset_state_store_by_name, - handle_set_asset_state_store_by_uri, - 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 @@ -484,6 +459,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 @@ -547,6 +523,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(), @@ -558,11 +535,8 @@ def start( # type: ignore[override] proc.send_msg(msg, request_id=0) return proc - @functools.cached_property - def client(self) -> Client: - return self.make_client() - - def make_client(self) -> Client: + @classmethod + def make_client(cls) -> Client: """ Build the API client used to talk to the API server. @@ -580,9 +554,6 @@ def make_client(self) -> Client: return 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): @@ -608,85 +579,37 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r while self.persisted_event_seqs: events_persisted.append(self.persisted_event_seqs.popleft()) - response = messages.TriggerStateSync( + sync = messages.TriggerStateSync( to_create=[], - to_cancel=self.cancelling_triggers, + to_cancel=self.cancelling_triggers.copy(), events_persisted=events_persisted or None, ) # Pull out of these dequeues in a thread-safe manner 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 - - 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): + 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, 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) - elif isinstance(msg, ClearAssetStateStoreByName): - handle_clear_asset_state_store_by_name(self.client, msg) - resp = OKResponse(ok=True) - elif isinstance(msg, ClearAssetStateStoreByUri): - handle_clear_asset_state_store_by_uri(self.client, msg) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteAssetStateStoreByName): - handle_delete_asset_state_store_by_name(self.client, msg) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteAssetStateStoreByUri): - handle_delete_asset_state_store_by_uri(self.client, msg) - resp = OKResponse(ok=True) - elif isinstance(msg, GetAssetStateStoreByName): - resp, dump_opts = handle_get_asset_state_store_by_name(self.client, msg) - elif isinstance(msg, GetAssetStateStoreByUri): - resp, dump_opts = handle_get_asset_state_store_by_uri(self.client, msg) - elif isinstance(msg, SetAssetStateStoreByName): - handle_set_asset_state_store_by_name(self.client, msg) - resp = OKResponse(ok=True) - elif isinstance(msg, SetAssetStateStoreByUri): - handle_set_asset_state_store_by_uri(self.client, msg) - resp = OKResponse(ok=True) - else: - raise ValueError(f"Unknown message type {type(msg)}") + self.send_msg(resp, 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 run(self) -> None: """Run synchronously and handle all database reads/writes.""" diff --git a/airflow-core/tests/unit/dag_processing/test_processor.py b/airflow-core/tests/unit/dag_processing/test_processor.py index dd7f603b35bb0..fbee3c75505b2 100644 --- a/airflow-core/tests/unit/dag_processing/test_processor.py +++ b/airflow-core/tests/unit/dag_processing/test_processor.py @@ -2214,7 +2214,7 @@ def test_handle_request_get_connection_masks_password_and_extra(self, proc): ) with ( - patch("airflow.dag_processing.processor.mask_secret") as mock_mask_secret, + 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( @@ -2251,7 +2251,7 @@ def test_handle_request_get_variable_masks_value_with_key(self, proc): ) with ( - patch("airflow.dag_processing.processor.mask_secret") as mock_mask_secret, + 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( diff --git a/airflow-core/tests/unit/jobs/test_triggerer_job.py b/airflow-core/tests/unit/jobs/test_triggerer_job.py index fa4d68f5fff6a..21a402147d594 100644 --- a/airflow-core/tests/unit/jobs/test_triggerer_job.py +++ b/airflow-core/tests/unit/jobs/test_triggerer_job.py @@ -75,18 +75,18 @@ from airflow.providers.standard.operators.python import PythonOperator from airflow.providers.standard.triggers.file import FileDeleteTrigger from airflow.providers.standard.triggers.temporal import DateTimeTrigger, TimeDeltaTrigger -from airflow.sdk import DAG, Asset, BaseHook, BaseOperator +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, AssetStateStoreResult, ClearAssetStateStoreByName, ClearAssetStateStoreByUri, DeleteAssetStateStoreByName, DeleteAssetStateStoreByUri, ErrorResponse, + ErrorType, GetAssetStateStoreByName, GetAssetStateStoreByUri, OKResponse, @@ -97,7 +97,7 @@ _RequestFrame, _ResponseFrame, ) -from airflow.sdk.execution_time.context import AssetStateStoreAccessors +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 @@ -283,6 +283,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) @@ -315,6 +316,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" @@ -327,6 +329,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 @@ -362,24 +365,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 @@ -423,6 +408,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 = [] @@ -895,7 +881,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 ) @@ -905,7 +891,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 ) @@ -924,7 +910,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) @@ -934,22 +920,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): 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 9830771701293..7eb0b3eab64fd 100644 --- a/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/callback_supervisor.py @@ -33,20 +33,12 @@ 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 ( - 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, @@ -56,7 +48,6 @@ ) if TYPE_CHECKING: - from pydantic import BaseModel from structlog.typing import FilteringBoundLogger from typing_extensions import Self @@ -188,8 +179,7 @@ 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. - + _msg_union: ClassVar[Any] = CallbackToSupervisor decoder: ClassVar[TypeAdapter[CallbackToSupervisor]] = TypeAdapter(CallbackToSupervisor) @classmethod @@ -352,31 +342,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, client: Client) -> tuple[FilteringBoundLogger, BinaryIO]: 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 a31596b3d2b82..6cef3a759ab28 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -27,11 +27,10 @@ from __future__ import annotations +from collections.abc import Callable from typing import TYPE_CHECKING -from uuid import UUID from airflow.sdk.api.datamodels._generated import ( - AssetStateStoreResponse, ConnectionResponse, DagRunStateResponse, TaskStatesResponse, @@ -41,23 +40,44 @@ XComSequenceSliceResponse, ) from airflow.sdk.execution_time.comms import ( + AssetEventsResult, + AssetResponse, + AssetResult, + AssetStateStoreResponse, AssetStateStoreResult, ClearAssetStateStoreByName, ClearAssetStateStoreByUri, + ClearTaskStateStore, ConnectionResult, + CreateHITLDetailPayload, + DagResult, + DagRunResult, DagRunStateResult, DeleteAssetStateStoreByName, DeleteAssetStateStoreByUri, + DeleteTaskStateStore, DeleteVariable, DeleteXCom, + ErrorResponse, + GetAssetByName, + GetAssetByUri, + GetAssetEventByAsset, + GetAssetEventByAssetAlias, + GetAssetsByAlias, GetAssetStateStoreByName, GetAssetStateStoreByUri, GetConnection, + GetDag, + GetDagRun, GetDagRunState, GetDRCount, GetPreviousDagRun, GetPreviousTI, + GetPrevSuccessfulDagRun, + GetTaskBreadcrumbs, + GetTaskRescheduleStartDate, GetTaskStates, + GetTaskStateStore, GetTICount, GetVariable, GetVariableKeys, @@ -65,13 +85,21 @@ GetXComCount, GetXComSequenceItem, GetXComSequenceSlice, + HITLDetailRequestResult, + InactiveAssetsResult, MaskSecret, + OKResponse, PrevSuccessfulDagRunResult, PutVariable, SetAssetStateStoreByName, SetAssetStateStoreByUri, + SetTaskStateStore, SetXCom, + TaskBreadcrumbsResult, TaskStatesResult, + TaskStateStoreResult, + TriggerDagRun, + ValidateInletsAndOutlets, VariableKeysResult, VariableResult, XComResult, @@ -85,7 +113,25 @@ 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 + + +@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) @@ -98,6 +144,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) @@ -108,6 +155,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]]: @@ -119,23 +167,28 @@ def handle_get_variable_keys( ) -def handle_mask_secret(msg: MaskSecret) -> None: +@handles(MaskSecret) +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, {}) +@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( @@ -150,6 +203,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( @@ -165,6 +219,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( @@ -177,6 +232,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( @@ -192,12 +248,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( @@ -209,6 +267,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) @@ -217,6 +276,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]]: @@ -229,21 +289,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]]: @@ -254,6 +317,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]]: @@ -273,6 +337,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( @@ -284,85 +349,245 @@ def handle_get_xcom(client: Client, msg: GetXCom) -> tuple[BaseModel | None, dic return xcom, {} -def handle_get_asset_state_store_by_name( - client: Client, msg: GetAssetStateStoreByName +@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) + return asset_resp, {} + + +@handles(GetAssetEventByAsset) +def handle_get_asset_event_by_asset( + client: Client, msg: GetAssetEventByAsset ) -> tuple[BaseModel | None, dict[str, bool]]: - asset_state = client.asset_state_store.get( - key=msg.key, + asset_event_resp = client.asset_events.get( + uri=msg.uri, name=msg.name, + after=msg.after, + before=msg.before, + 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) + 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, + 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) + 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(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_state_store + if isinstance(task_state_store, ErrorResponse) + else TaskStateStoreResult.from_task_state_store_response(task_state_store) ) + return resp, {} + + +@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(DeleteTaskStateStore) +def handle_delete_task_state_store( + client: Client, msg: DeleteTaskStateStore +) -> tuple[BaseModel | None, dict[str, bool]]: + client.task_state_store.delete(msg.ti_id, msg.key) + return OKResponse(ok=True), {} + + +@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) + return OKResponse(ok=True), {} + + +@handles(GetAssetStateStoreByName) +def handle_get_asset_state_store_by_name( + client: Client, msg: GetAssetStateStoreByName +) -> tuple[BaseModel | None, dict[str, bool]]: + 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 = client.asset_state_store.get( - key=msg.key, - uri=msg.uri, - ) + 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) def handle_set_asset_state_store_by_name( client: Client, msg: SetAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_state_store.set( - key=msg.key, - value=msg.value, - name=msg.name, - ) - return None, {} + client.asset_state_store.set(msg.key, msg.value, name=msg.name) + return OKResponse(ok=True), {} +@handles(SetAssetStateStoreByUri) def handle_set_asset_state_store_by_uri( client: Client, msg: SetAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_state_store.set( - key=msg.key, - value=msg.value, - uri=msg.uri, - ) - return None, {} + client.asset_state_store.set(msg.key, msg.value, uri=msg.uri) + return OKResponse(ok=True), {} +@handles(DeleteAssetStateStoreByName) def handle_delete_asset_state_store_by_name( client: Client, msg: DeleteAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_state_store.delete( - key=msg.key, - name=msg.name, - ) - return None, {} + client.asset_state_store.delete(msg.key, name=msg.name) + return OKResponse(ok=True), {} +@handles(DeleteAssetStateStoreByUri) def handle_delete_asset_state_store_by_uri( client: Client, msg: DeleteAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: - client.asset_state_store.delete( - key=msg.key, - uri=msg.uri, - ) - return None, {} + client.asset_state_store.delete(msg.key, uri=msg.uri) + return OKResponse(ok=True), {} +@handles(ClearAssetStateStoreByName) def handle_clear_asset_state_store_by_name( client: Client, msg: ClearAssetStateStoreByName ) -> tuple[BaseModel | None, dict[str, bool]]: client.asset_state_store.clear(name=msg.name) - return None, {} + return OKResponse(ok=True), {} +@handles(ClearAssetStateStoreByUri) def handle_clear_asset_state_store_by_uri( client: Client, msg: ClearAssetStateStoreByUri ) -> tuple[BaseModel | None, dict[str, bool]]: client.asset_state_store.clear(uri=msg.uri) - return None, {} + return OKResponse(ok=True), {} diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 87311f02da7a1..44728a2f3627b 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -38,7 +38,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 @@ -52,7 +52,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, @@ -61,98 +60,28 @@ from airflow.sdk.exceptions import ErrorType from airflow.sdk.execution_time import comms from airflow.sdk.execution_time.comms import ( - AssetEventsResult, - AssetResult, - AssetStateStoreResult, AwaitInputTask, - ClearAssetStateStoreByName, - ClearAssetStateStoreByUri, - ClearTaskStateStore, ConnectionResult, - CreateHITLDetailPayload, - DagResult, - DagRunResult, DeferTask, - DeleteAssetStateStoreByName, - DeleteAssetStateStoreByUri, - DeleteTaskStateStore, - DeleteVariable, - DeleteXCom, ErrorResponse, - GetAssetByName, - GetAssetByUri, - GetAssetEventByAsset, - GetAssetEventByAssetAlias, - GetAssetsByAlias, - GetAssetStateStoreByName, - GetAssetStateStoreByUri, - GetConnection, - GetDag, - GetDagRun, - GetDagRunState, - GetDRCount, - GetPreviousDagRun, - GetPreviousTI, - GetPrevSuccessfulDagRun, - GetTaskBreadcrumbs, - GetTaskRescheduleStartDate, - GetTaskStates, - GetTaskStateStore, - GetTICount, - GetVariable, - GetVariableKeys, - GetXCom, - GetXComCount, - GetXComSequenceItem, - GetXComSequenceSlice, - HITLDetailRequestResult, - InactiveAssetsResult, MaskSecret, - OKResponse, - PutVariable, RescheduleTask, ResendLoggingFD, RetryTask, SentFDs, - SetAssetStateStoreByName, - SetAssetStateStoreByUri, SetRenderedFields, SetRenderedMapIndex, - SetTaskStateStore, - SetXCom, SkipDownstreamTasks, StartupDetails, SucceedTask, - TaskBreadcrumbsResult, TaskState, - TaskStateStoreResult, ToSupervisor, - TriggerDagRun, - ValidateInletsAndOutlets, _RequestFrame, _ResponseFrame, ) from airflow.sdk.execution_time.coordinator import get_coordinator_manager 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_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, + get_handler, ) from airflow.sdk.execution_time.schema import get_schema_version_migrator, resolve_body_class @@ -633,6 +562,13 @@ class WatchedSubprocess: socket handling, process monitoring, and request handling. """ + _msg_union: ClassVar[Any] = ToSupervisor # 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 @@ -676,6 +612,8 @@ class WatchedSubprocess: start_time: float = attrs.field(factory=time.monotonic) """The start time of the child process.""" + client: Client = attrs.field(repr=False) + @classmethod def start( cls, @@ -978,7 +916,16 @@ 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() + 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.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) @staticmethod def _close_unused_sockets(*sockets): @@ -1305,9 +1252,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) _final_state: str | None = attrs.field(default=None, init=False) # The terminal-state message currently being processed by `_handle_request`, @@ -1691,8 +1635,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) - 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 @@ -1700,43 +1643,40 @@ 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): + 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) - elif isinstance(msg, RetryTask): + 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) - elif isinstance(msg, GetConnection): - resp, dump_opts = handle_get_connection(self.client, msg) - elif isinstance(msg, GetVariable): - resp, dump_opts = handle_get_variable(self.client, msg) - elif isinstance(msg, GetVariableKeys): - resp, dump_opts = handle_get_variable_keys(self.client, msg) - elif isinstance(msg, GetXCom): - resp, dump_opts = handle_get_xcom(self.client, msg) - elif isinstance(msg, GetXComSequenceItem): - resp, dump_opts = handle_get_xcom_sequence_item(self.client, msg) - elif isinstance(msg, GetXComSequenceSlice): - resp, dump_opts = handle_get_xcom_sequence_slice(self.client, msg) - elif isinstance(msg, DeferTask): + self.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) - elif isinstance(msg, AwaitInputTask): + self.send_msg(None, request_id=req_id, error=None) + return + if isinstance(msg, AwaitInputTask): self._rendered_map_index = msg.rendered_map_index self._send_terminal_state_msg(msg) - elif isinstance(msg, RescheduleTask): + self.send_msg(None, request_id=req_id, error=None) + return + if isinstance(msg, RescheduleTask): self._send_terminal_state_msg(msg) - elif isinstance(msg, SkipDownstreamTasks): + 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) - 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): + self.send_msg(None, request_id=req_id, error=None) + return + if isinstance(msg, SetRenderedFields): try: self.client.task_instances.set_rtif(self.id, msg.rendered_fields) except ServerResponseError as e: @@ -1746,176 +1686,21 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: if e.response.status_code != HTTPStatus.GONE: raise log.debug("Skipping RTIF overwrite; task instance archived on retry/clear", ti_id=self.id) - elif isinstance(msg, SetRenderedMapIndex): + self.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) - elif isinstance(msg, GetAssetByName): - asset_resp = self.client.assets.get(name=msg.name) - if isinstance(asset_resp, AssetResponse): - asset_result = AssetResult.from_asset_response(asset_resp) - resp = asset_result - dump_opts = {"exclude_unset": True} - else: - resp = asset_resp - elif isinstance(msg, GetAssetByUri): - asset_resp = self.client.assets.get(uri=msg.uri) - if isinstance(asset_resp, AssetResponse): - asset_result = AssetResult.from_asset_response(asset_resp) - resp = asset_result - dump_opts = {"exclude_unset": True} - else: - resp = asset_resp - elif isinstance(msg, GetAssetsByAlias): - resp = self.client.assets.get_by_alias(alias_name=msg.alias_name) - elif isinstance(msg, GetAssetEventByAsset): - asset_event_resp = self.client.asset_events.get( - uri=msg.uri, - name=msg.name, - after=msg.after, - before=msg.before, - ascending=msg.ascending, - limit=msg.limit, - partition_key=msg.partition_key, - partition_key_regexp_pattern=msg.partition_key_regexp_pattern, - extra=msg.extra, - ) - asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) - resp = asset_event_result - dump_opts = {"exclude_unset": True} - elif isinstance(msg, GetAssetEventByAssetAlias): - asset_event_resp = self.client.asset_events.get( - alias_name=msg.alias_name, - after=msg.after, - before=msg.before, - ascending=msg.ascending, - limit=msg.limit, - partition_key=msg.partition_key, - partition_key_regexp_pattern=msg.partition_key_regexp_pattern, - extra=msg.extra, - ) - asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) - resp = asset_event_result - dump_opts = {"exclude_unset": True} - elif isinstance(msg, GetPrevSuccessfulDagRun): - resp, dump_opts = handle_get_prev_successful_dag_run(self.client, self.id) - elif isinstance(msg, GetXComCount): - resp, dump_opts = handle_get_xcom_count(self.client, msg) - elif isinstance(msg, TriggerDagRun): - resp = self.client.dag_runs.trigger( - msg.dag_id, msg.run_id, msg.conf, msg.logical_date, msg.run_after, msg.reset_dag_run, msg.note - ) - elif isinstance(msg, GetDagRun): - dr_resp = self.client.dag_runs.get_detail(msg.dag_id, msg.run_id) - resp = DagRunResult.from_api_response(dr_resp) - elif isinstance(msg, GetTaskRescheduleStartDate): - resp = self.client.task_instances.get_reschedule_start_date(msg.ti_id, msg.try_number) - elif isinstance(msg, GetTICount): - resp, dump_opts = handle_get_ti_count(self.client, msg) - elif isinstance(msg, GetTaskStates): - resp, dump_opts = handle_get_task_states(self.client, msg) - elif isinstance(msg, GetTaskBreadcrumbs): - api_resp = self.client.task_instances.get_task_breakcrumbs(dag_id=msg.dag_id, run_id=msg.run_id) - resp = TaskBreadcrumbsResult.from_api_response(api_resp) - elif isinstance(msg, GetDRCount): - resp, dump_opts = handle_get_dr_count(self.client, msg) - elif isinstance(msg, GetDagRunState): - resp, dump_opts = handle_get_dag_run_state(self.client, msg) - elif isinstance(msg, GetPreviousDagRun): - resp, dump_opts = handle_get_previous_dag_run(self.client, msg) - elif isinstance(msg, GetPreviousTI): - resp, dump_opts = handle_get_previous_ti(self.client, msg) - elif isinstance(msg, DeleteVariable): - resp, dump_opts = handle_delete_variable(self.client, msg) - elif isinstance(msg, ValidateInletsAndOutlets): - inactive_assets_resp = self.client.task_instances.validate_inlets_and_outlets(msg.ti_id) - resp = InactiveAssetsResult.from_inactive_assets_response(inactive_assets_resp) - dump_opts = {"exclude_unset": True} - elif isinstance(msg, ResendLoggingFD): + self.send_msg(None, request_id=req_id, error=None) + 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): - hitl_detail_request = self.client.hitl.add_response( - ti_id=msg.ti_id, - options=msg.options, - subject=msg.subject, - body=msg.body, - defaults=msg.defaults, - params=msg.params, - multiple=msg.multiple, - assigned_users=msg.assigned_users, - ) - resp = HITLDetailRequestResult.from_api_response(hitl_detail_request) - dump_opts = {"exclude_unset": True} - elif isinstance(msg, MaskSecret): - handle_mask_secret(msg) - elif isinstance(msg, GetDag): - dag = self.client.dags.get( - dag_id=msg.dag_id, - ) - resp = DagResult.from_api_response(dag) - elif isinstance(msg, GetTaskStateStore): - task_store = self.client.task_state_store.get(msg.ti_id, msg.key) - resp = ( - task_store - if isinstance(task_store, ErrorResponse) - else TaskStateStoreResult.from_task_state_store_response(task_store) - ) - elif isinstance(msg, SetTaskStateStore): - self.client.task_state_store.set(msg.ti_id, msg.key, msg.value, expires_at=msg.expires_at) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteTaskStateStore): - self.client.task_state_store.delete(msg.ti_id, msg.key) - resp = OKResponse(ok=True) - elif isinstance(msg, ClearTaskStateStore): - self.client.task_state_store.clear(msg.ti_id) - resp = OKResponse(ok=True) - elif isinstance(msg, GetAssetStateStoreByName): - asset_store = self.client.asset_state_store.get(msg.key, name=msg.name) - resp = ( - asset_store - if isinstance(asset_store, ErrorResponse) - else AssetStateStoreResult.from_asset_state_store_response(asset_store) - ) - elif isinstance(msg, GetAssetStateStoreByUri): - asset_store = self.client.asset_state_store.get(msg.key, uri=msg.uri) - resp = ( - asset_store - if isinstance(asset_store, ErrorResponse) - else AssetStateStoreResult.from_asset_state_store_response(asset_store) - ) - elif isinstance(msg, SetAssetStateStoreByName): - self.client.asset_state_store.set(msg.key, msg.value, name=msg.name) - resp = OKResponse(ok=True) - elif isinstance(msg, SetAssetStateStoreByUri): - self.client.asset_state_store.set(msg.key, msg.value, uri=msg.uri) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteAssetStateStoreByName): - self.client.asset_state_store.delete(msg.key, name=msg.name) - resp = OKResponse(ok=True) - elif isinstance(msg, DeleteAssetStateStoreByUri): - self.client.asset_state_store.delete(msg.key, uri=msg.uri) - resp = OKResponse(ok=True) - elif isinstance(msg, ClearAssetStateStoreByName): - self.client.asset_state_store.clear(name=msg.name) - resp = OKResponse(ok=True) - elif isinstance(msg, ClearAssetStateStoreByUri): - self.client.asset_state_store.clear(uri=msg.uri) - resp = OKResponse(ok=True) - else: - log.error("Unhandled request", msg=msg) - self.send_msg( - None, - request_id=req_id, - error=ErrorResponse( - error=ErrorType.API_SERVER_ERROR, - detail={"status_code": 400, "message": "Unhandled request"}, - ), - ) return - 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: 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. 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 == {}