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

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 

18 

19import uuid 

20from collections.abc import Iterable 

21from datetime import timedelta 

22from enum import Enum 

23from typing import Annotated, Any, Literal 

24 

25from pydantic import ( 

26 AwareDatetime, 

27 Field, 

28 JsonValue, 

29 Tag, 

30 TypeAdapter, 

31 WithJsonSchema, 

32 model_validator, 

33) 

34 

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 

47 

48AwareDatetimeAdapter = TypeAdapter(AwareDatetime) 

49 

50 

51class TIEnterRunningPayload(StrictBaseModel): 

52 """Schema for updating TaskInstance to 'RUNNING' state with minimal required fields.""" 

53 

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

67 

68 

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.""" 

72 

73 FAILED = TerminalTIState.FAILED 

74 SKIPPED = TerminalTIState.SKIPPED 

75 REMOVED = TerminalTIState.REMOVED 

76 UPSTREAM_FAILED = TerminalTIState.UPSTREAM_FAILED 

77 

78 

79class TITerminalStatePayload(StrictBaseModel): 

80 """Schema for updating TaskInstance to a terminal state except SUCCESS state.""" 

81 

82 state: TerminalStateNonSuccess 

83 

84 end_date: UtcDateTime 

85 """When the task completed executing""" 

86 rendered_map_index: str | None = None 

87 

88 

89class TISuccessStatePayload(StrictBaseModel): 

90 """Schema for updating TaskInstance to success state.""" 

91 

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 ] 

103 

104 end_date: UtcDateTime 

105 """When the task completed executing""" 

106 

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 

110 

111 

112class TITargetStatePayload(StrictBaseModel): 

113 """Schema for updating TaskInstance to a target state, excluding terminal and running states.""" 

114 

115 state: IntermediateTIState 

116 

117 

118class TIDeferredStatePayload(StrictBaseModel): 

119 """Schema for updating TaskInstance to a deferred state.""" 

120 

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. 

136 

137 Both forms will be passed along to the trigger, the server will not handle either. 

138 """ 

139 

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. 

147 

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 

151 

152 

153class TIRescheduleStatePayload(StrictBaseModel): 

154 """Schema for updating TaskInstance to a up_for_reschedule state.""" 

155 

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 

169 

170 

171class TIAwaitingInputStatePayload(StrictBaseModel): 

172 """Schema for parking a TaskInstance in an awaiting_input state (Human-in-the-loop, no trigger).""" 

173 

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. 

192 

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 

196 

197 

198class TIRetryStatePayload(StrictBaseModel): 

199 """Schema for updating TaskInstance to up_for_retry.""" 

200 

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 

216 

217 

218class TISkippedDownstreamTasksStatePayload(StrictBaseModel): 

219 """Schema for updating downstream tasks to a skipped state.""" 

220 

221 tasks: list[str | tuple[str, int]] 

222 

223 

224def ti_state_discriminator(v: dict[str, str] | StrictBaseModel) -> str: 

225 """ 

226 Determine the discriminator key for TaskInstance state transitions. 

227 

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) 

237 

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_" 

251 

252 

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] 

265 

266 

267class TIHeartbeatInfo(StrictBaseModel): 

268 """Schema for TaskInstance heartbeat endpoint.""" 

269 

270 hostname: str 

271 pid: int 

272 

273 

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.""" 

278 

279 id: uuid.UUID 

280 

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" 

293 

294 

295class AssetReferenceAssetEventDagRun(StrictBaseModel): 

296 """Schema for AssetModel used in AssetEventDagRunReference.""" 

297 

298 name: str 

299 uri: str 

300 extra: dict[str, JsonValue] 

301 

302 

303class AssetAliasReferenceAssetEventDagRun(StrictBaseModel): 

304 """Schema for AssetAliasModel used in AssetEventDagRunReference.""" 

305 

306 name: str 

307 

308 

309class AssetEventDagRunReference(StrictBaseModel): 

310 """Schema for AssetEvent model used in DagRun.""" 

311 

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 

321 

322 

323class DagRun(StrictBaseModel): 

324 """Schema for DagRun model with minimal required fields needed for Runtime.""" 

325 

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 

331 

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 

348 

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. 

354 

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 

361 

362 if isinstance(data, dict): 

363 return data 

364 

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 

371 

372 values = {} 

373 

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 

382 

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"] = [] 

385 

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 

392 

393 return values 

394 

395 

396class TIRunContext(BaseModel): 

397 """Response schema for TaskInstance run context.""" 

398 

399 dag_run: DagRun 

400 """DAG run information for the task instance.""" 

401 

402 task_reschedule_count: int = 0 

403 """How many times the task has been rescheduled.""" 

404 

405 max_tries: int 

406 """Maximum number of tries for the task instance (from DB).""" 

407 

408 variables: Annotated[list[VariableResponse], Field(default_factory=list)] 

409 """Variables that can be accessed by the task instance.""" 

410 

411 connections: Annotated[list[ConnectionResponse], Field(default_factory=list)] 

412 """Connections that can be accessed by the task instance.""" 

413 

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``. 

419 

420 Can either be a "decorated" dict, or a string encrypted with the shared Fernet key. 

421 """ 

422 

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.""" 

425 

426 should_retry: bool = False 

427 """If the ti encounters an error, whether it should enter retry or failed state.""" 

428 

429 start_date: UtcDateTime | None = None 

430 """ 

431 The original start date of the task instance. 

432 

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

437 

438 

439class PrevSuccessfulDagRunResponse(BaseModel): 

440 """Schema for response with previous successful DagRun information for Task Template Context.""" 

441 

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 

446 

447 

448class PreviousTIResponse(BaseModel): 

449 """Schema for response with previous TaskInstance information.""" 

450 

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 

461 

462 

463class TaskStatesResponse(BaseModel): 

464 """Response for task states with run_id, task and state.""" 

465 

466 task_states: dict[str, Any] 

467 

468 

469class TaskBreadcrumbsResponse(BaseModel): 

470 """Response for task breadcrumbs.""" 

471 

472 breadcrumbs: Iterable[dict[str, Any]] 

473 

474 

475class InactiveAssetsResponse(BaseModel): 

476 """Response for inactive assets.""" 

477 

478 inactive_assets: Annotated[list[AssetProfile], Field(default_factory=list)]