Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/execution_api/datamodels/taskinstance.py: 83%
215 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
1# Licensed to the Apache Software Foundation (ASF) under one
2# or more contributor license agreements. See the NOTICE file
3# distributed with this work for additional information
4# regarding copyright ownership. The ASF licenses this file
5# to you under the Apache License, Version 2.0 (the
6# "License"); you may not use this file except in compliance
7# with the License. You may obtain a copy of the License at
8#
9# http://www.apache.org/licenses/LICENSE-2.0
10#
11# Unless required by applicable law or agreed to in writing,
12# software distributed under the License is distributed on an
13# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14# KIND, either express or implied. See the License for the
15# specific language governing permissions and limitations
16# under the License.
17from __future__ import annotations
19import uuid
20from collections.abc import Iterable
21from datetime import timedelta
22from enum import Enum
23from typing import Annotated, Any, Literal
25from pydantic import (
26 AwareDatetime,
27 Field,
28 JsonValue,
29 Tag,
30 TypeAdapter,
31 WithJsonSchema,
32 model_validator,
33)
35from airflow.api_fastapi.common.types import UtcDateTime
36from airflow.api_fastapi.core_api.base import BaseModel, StrictBaseModel
37from airflow.api_fastapi.execution_api.datamodels.asset import AssetProfile
38from airflow.api_fastapi.execution_api.datamodels.connection import ConnectionResponse
39from airflow.api_fastapi.execution_api.datamodels.variable import VariableResponse
40from airflow.utils.state import (
41 DagRunState,
42 IntermediateTIState,
43 TaskInstanceState as TIState,
44 TerminalTIState,
45)
46from airflow.utils.types import DagRunType
48AwareDatetimeAdapter = TypeAdapter(AwareDatetime)
51class TIEnterRunningPayload(StrictBaseModel):
52 """Schema for updating TaskInstance to 'RUNNING' state with minimal required fields."""
54 state: Annotated[
55 Literal[TIState.RUNNING],
56 # Specify a default in the schema, but not in code.
57 WithJsonSchema({"type": "string", "enum": [TIState.RUNNING], "default": TIState.RUNNING}),
58 ]
59 hostname: str
60 """Hostname where this task has started"""
61 unixname: str
62 """Local username of the process where this task has started"""
63 pid: int
64 """Process Identifier on `hostname`"""
65 start_date: UtcDateTime
66 """When the task started executing"""
69# Create an enum to give a nice name in the generated datamodels
70class TerminalStateNonSuccess(str, Enum):
71 """TaskInstance states that can be reported without extra information."""
73 FAILED = TerminalTIState.FAILED
74 SKIPPED = TerminalTIState.SKIPPED
75 REMOVED = TerminalTIState.REMOVED
76 UPSTREAM_FAILED = TerminalTIState.UPSTREAM_FAILED
79class TITerminalStatePayload(StrictBaseModel):
80 """Schema for updating TaskInstance to a terminal state except SUCCESS state."""
82 state: TerminalStateNonSuccess
84 end_date: UtcDateTime
85 """When the task completed executing"""
86 rendered_map_index: str | None = None
89class TISuccessStatePayload(StrictBaseModel):
90 """Schema for updating TaskInstance to success state."""
92 state: Annotated[
93 Literal[TerminalTIState.SUCCESS],
94 # Specify a default in the schema, but not in code, so Pydantic marks it as required.
95 WithJsonSchema(
96 {
97 "type": "string",
98 "enum": [TerminalTIState.SUCCESS],
99 "default": TerminalTIState.SUCCESS,
100 }
101 ),
102 ]
104 end_date: UtcDateTime
105 """When the task completed executing"""
107 task_outlets: Annotated[list[AssetProfile], Field(default_factory=list)]
108 outlet_events: Annotated[list[dict[str, Any]], Field(default_factory=list)]
109 rendered_map_index: str | None = None
112class TITargetStatePayload(StrictBaseModel):
113 """Schema for updating TaskInstance to a target state, excluding terminal and running states."""
115 state: IntermediateTIState
118class TIDeferredStatePayload(StrictBaseModel):
119 """Schema for updating TaskInstance to a deferred state."""
121 state: Annotated[
122 Literal[IntermediateTIState.DEFERRED],
123 # Specify a default in the schema, but not in code, so Pydantic marks it as required.
124 WithJsonSchema(
125 {
126 "type": "string",
127 "enum": [IntermediateTIState.DEFERRED],
128 "default": IntermediateTIState.DEFERRED,
129 }
130 ),
131 ]
132 classpath: str
133 trigger_kwargs: Annotated[dict[str, JsonValue] | str, Field(default_factory=dict)]
134 """
135 Kwargs to pass to the trigger constructor, either a plain dict or an encrypted string.
137 Both forms will be passed along to the trigger, the server will not handle either.
138 """
140 trigger_timeout: timedelta | None = None
141 queue: str | None = None
142 next_method: str
143 """The name of the method on the operator to call in the worker after the trigger has fired."""
144 next_kwargs: Annotated[dict[str, JsonValue], Field(default_factory=dict)]
145 """
146 Kwargs to pass to the above method, either a plain dict or an encrypted string.
148 Both forms will be passed along to the TaskSDK upon resume, the server will not handle either.
149 """
150 rendered_map_index: str | None = None
153class TIRescheduleStatePayload(StrictBaseModel):
154 """Schema for updating TaskInstance to a up_for_reschedule state."""
156 state: Annotated[
157 Literal[IntermediateTIState.UP_FOR_RESCHEDULE],
158 # Specify a default in the schema, but not in code, so Pydantic marks it as required.
159 WithJsonSchema(
160 {
161 "type": "string",
162 "enum": [IntermediateTIState.UP_FOR_RESCHEDULE],
163 "default": IntermediateTIState.UP_FOR_RESCHEDULE,
164 }
165 ),
166 ]
167 reschedule_date: UtcDateTime
168 end_date: UtcDateTime
171class TIAwaitingInputStatePayload(StrictBaseModel):
172 """Schema for parking a TaskInstance in an awaiting_input state (Human-in-the-loop, no trigger)."""
174 state: Annotated[
175 Literal[IntermediateTIState.AWAITING_INPUT],
176 # Specify a default in the schema, but not in code, so Pydantic marks it as required.
177 WithJsonSchema(
178 {
179 "type": "string",
180 "enum": [IntermediateTIState.AWAITING_INPUT],
181 "default": IntermediateTIState.AWAITING_INPUT,
182 }
183 ),
184 ]
185 timeout: timedelta | None = None
186 """Optional response deadline (relative); converted to an absolute datetime server-side."""
187 next_method: str
188 """The name of the method on the operator to call in the worker after input is received."""
189 next_kwargs: Annotated[dict[str, JsonValue], Field(default_factory=dict)]
190 """
191 Kwargs to pass to the above method, either a plain dict or an encrypted string.
193 Both forms will be passed along to the TaskSDK upon resume, the server will not handle either.
194 """
195 rendered_map_index: str | None = None
198class TIRetryStatePayload(StrictBaseModel):
199 """Schema for updating TaskInstance to up_for_retry."""
201 state: Annotated[
202 Literal[IntermediateTIState.UP_FOR_RETRY],
203 # Specify a default in the schema, but not in code, so Pydantic marks it as required.
204 WithJsonSchema(
205 {
206 "type": "string",
207 "enum": [IntermediateTIState.UP_FOR_RETRY],
208 "default": IntermediateTIState.UP_FOR_RETRY,
209 }
210 ),
211 ]
212 end_date: UtcDateTime
213 rendered_map_index: str | None = None
214 retry_delay_seconds: float | None = None
215 retry_reason: str | None = None
218class TISkippedDownstreamTasksStatePayload(StrictBaseModel):
219 """Schema for updating downstream tasks to a skipped state."""
221 tasks: list[str | tuple[str, int]]
224def ti_state_discriminator(v: dict[str, str] | StrictBaseModel) -> str:
225 """
226 Determine the discriminator key for TaskInstance state transitions.
228 This function serves as a discriminator for the TIStateUpdate union schema,
229 categorizing the payload based on the ``state`` attribute in the input data.
230 It returns a key that directs FastAPI to the appropriate subclass (schema)
231 based on the requested state.
232 """
233 if isinstance(v, dict):
234 state = v.get("state")
235 else:
236 state = getattr(v, "state", None)
238 if state == TIState.SUCCESS:
239 return "success"
240 if state in set(TerminalTIState):
241 return "_terminal_"
242 if state == TIState.DEFERRED:
243 return "deferred"
244 if state == TIState.UP_FOR_RESCHEDULE:
245 return "up_for_reschedule"
246 if state == TIState.AWAITING_INPUT:
247 return "awaiting_input"
248 if state == TIState.UP_FOR_RETRY:
249 return "up_for_retry"
250 return "_other_"
253# It is called "_terminal_" to avoid future conflicts if we added an actual state named "terminal"
254# and "_other_" is a catch-all for all other states that are not covered by the other schemas.
255TIStateUpdate = Annotated[
256 Annotated[TITerminalStatePayload, Tag("_terminal_")]
257 | Annotated[TISuccessStatePayload, Tag("success")]
258 | Annotated[TITargetStatePayload, Tag("_other_")]
259 | Annotated[TIDeferredStatePayload, Tag("deferred")]
260 | Annotated[TIRescheduleStatePayload, Tag("up_for_reschedule")]
261 | Annotated[TIAwaitingInputStatePayload, Tag("awaiting_input")]
262 | Annotated[TIRetryStatePayload, Tag("up_for_retry")],
263 Field(discriminator=ti_state_discriminator),
264]
267class TIHeartbeatInfo(StrictBaseModel):
268 """Schema for TaskInstance heartbeat endpoint."""
270 hostname: str
271 pid: int
274# This model is not used in the API, but it is included in generated OpenAPI schema
275# for use in the client SDKs.
276class TaskInstance(BaseModel):
277 """Schema for TaskInstance model with minimal required fields needed for Runtime."""
279 id: uuid.UUID
281 task_id: str
282 dag_id: str
283 run_id: str
284 try_number: int
285 dag_version_id: uuid.UUID
286 map_index: int = -1
287 hostname: str | None = None
288 context_carrier: dict | None = None
289 # The supervisor routes tasks to a coordinator by queue. The default keeps
290 # hand-built instances (tests, dry runs) valid; the executor workload
291 # always sends the real value.
292 queue: str = "default"
295class AssetReferenceAssetEventDagRun(StrictBaseModel):
296 """Schema for AssetModel used in AssetEventDagRunReference."""
298 name: str
299 uri: str
300 extra: dict[str, JsonValue]
303class AssetAliasReferenceAssetEventDagRun(StrictBaseModel):
304 """Schema for AssetAliasModel used in AssetEventDagRunReference."""
306 name: str
309class AssetEventDagRunReference(StrictBaseModel):
310 """Schema for AssetEvent model used in DagRun."""
312 asset: AssetReferenceAssetEventDagRun
313 extra: dict[str, JsonValue]
314 source_task_id: str | None
315 source_dag_id: str | None
316 source_run_id: str | None
317 source_map_index: int | None
318 source_aliases: list[AssetAliasReferenceAssetEventDagRun]
319 timestamp: UtcDateTime
320 partition_key: str | None = None
323class DagRun(StrictBaseModel):
324 """Schema for DagRun model with minimal required fields needed for Runtime."""
326 # TODO: `dag_id` and `run_id` are duplicated from TaskInstance
327 # See if we can avoid sending these fields from API server and instead
328 # use the TaskInstance data to get the DAG run information in the client (Task Execution Interface).
329 dag_id: str
330 run_id: str
332 logical_date: UtcDateTime | None
333 data_interval_start: UtcDateTime | None
334 data_interval_end: UtcDateTime | None
335 run_after: UtcDateTime
336 start_date: UtcDateTime | None
337 end_date: UtcDateTime | None
338 clear_number: int = 0
339 run_type: DagRunType
340 state: DagRunState
341 conf: dict[str, Any] | None = None
342 triggering_user_name: str | None = None
343 consumed_asset_events: list[AssetEventDagRunReference]
344 partition_key: str | None
345 partition_date: UtcDateTime | None = None
346 note: str | None = None
347 team_name: str | None = None
349 @model_validator(mode="before")
350 @classmethod
351 def safe_extract_from_orm(cls, data: Any) -> Any:
352 """
353 Safely extract data from SQLAlchemy DagRun instances.
355 Handles the 'note' association proxy and provides defaults for unloaded relationships
356 to prevent DetachedInstanceError when the instance is not bound to a session.
357 """
358 from sqlalchemy import inspect as sa_inspect
359 from sqlalchemy.exc import NoInspectionAvailable
360 from sqlalchemy.orm.state import InstanceState
362 if isinstance(data, dict):
363 return data
365 # Check if this is a SQLAlchemy model by looking for the inspection interface
366 try:
367 insp: InstanceState = sa_inspect(data)
368 except NoInspectionAvailable:
369 # Not a SQLAlchemy object, return as-is for Pydantic to handle
370 return data
372 values = {}
374 for field_name in cls.model_fields:
375 if field_name in insp.dict:
376 values[field_name] = insp.dict[field_name]
377 elif field_name == "state":
378 if "_state" in insp.dict: 378 ↛ 380line 378 didn't jump to line 380 because the condition on line 378 was always true
379 values["state"] = insp.dict["_state"]
380 elif not insp.detached and (state_val := data._state) is not None:
381 values["state"] = state_val
383 if "consumed_asset_events" not in values: 383 ↛ 384line 383 didn't jump to line 384 because the condition on line 383 was never true
384 values["consumed_asset_events"] = []
386 # Check if dag_run_note is already loaded (avoid lazy load on detached instance)
387 if "note" not in values: 387 ↛ 393line 387 didn't jump to line 393 because the condition on line 387 was always true
388 if "dag_run_note" in insp.dict: 388 ↛ 389line 388 didn't jump to line 389 because the condition on line 388 was never true
389 values["note"] = data.note
390 else:
391 values["note"] = None
393 return values
396class TIRunContext(BaseModel):
397 """Response schema for TaskInstance run context."""
399 dag_run: DagRun
400 """DAG run information for the task instance."""
402 task_reschedule_count: int = 0
403 """How many times the task has been rescheduled."""
405 max_tries: int
406 """Maximum number of tries for the task instance (from DB)."""
408 variables: Annotated[list[VariableResponse], Field(default_factory=list)]
409 """Variables that can be accessed by the task instance."""
411 connections: Annotated[list[ConnectionResponse], Field(default_factory=list)]
412 """Connections that can be accessed by the task instance."""
414 next_method: str | None = None
415 """Method to call. Set when task resumes from a trigger."""
416 next_kwargs: dict[str, Any] | str | None = None
417 """
418 Args to pass to ``next_method``.
420 Can either be a "decorated" dict, or a string encrypted with the shared Fernet key.
421 """
423 xcom_keys_to_clear: Annotated[list[str], Field(default_factory=list)]
424 """List of Xcom keys that need to be cleared and purged on by the worker."""
426 should_retry: bool = False
427 """If the ti encounters an error, whether it should enter retry or failed state."""
429 start_date: UtcDateTime | None = None
430 """
431 The original start date of the task instance.
433 When resuming from deferral, this is set to the task's original ``start_date`` so the
434 supervisor uses it instead of ``datetime.now()``. This ensures ``context["ti"].start_date``
435 always reflects when the task *first* started, not when it was rescheduled/resumed.
436 """
439class PrevSuccessfulDagRunResponse(BaseModel):
440 """Schema for response with previous successful DagRun information for Task Template Context."""
442 data_interval_start: UtcDateTime | None = None
443 data_interval_end: UtcDateTime | None = None
444 start_date: UtcDateTime | None = None
445 end_date: UtcDateTime | None = None
448class PreviousTIResponse(BaseModel):
449 """Schema for response with previous TaskInstance information."""
451 task_id: str
452 dag_id: str
453 run_id: str
454 logical_date: UtcDateTime | None = None
455 start_date: UtcDateTime | None = None
456 end_date: UtcDateTime | None = None
457 state: str | None = None
458 try_number: int
459 map_index: int | None = -1
460 duration: float | None = None
463class TaskStatesResponse(BaseModel):
464 """Response for task states with run_id, task and state."""
466 task_states: dict[str, Any]
469class TaskBreadcrumbsResponse(BaseModel):
470 """Response for task breadcrumbs."""
472 breadcrumbs: Iterable[dict[str, Any]]
475class InactiveAssetsResponse(BaseModel):
476 """Response for inactive assets."""
478 inactive_assets: Annotated[list[AssetProfile], Field(default_factory=list)]