Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/execution_api/routes/task_instances.py: 60%

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

17 

18from __future__ import annotations 

19 

20import contextlib 

21import itertools 

22import json 

23from collections import defaultdict 

24from collections.abc import Callable, Iterator, Sequence 

25from typing import TYPE_CHECKING, Annotated, Any, NoReturn, cast 

26from uuid import UUID 

27 

28import attrs 

29import structlog 

30from cadwyn import VersionedAPIRouter 

31from fastapi import Body, HTTPException, Query, Response, Security, status 

32from opentelemetry import trace 

33from opentelemetry.trace import StatusCode 

34from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator 

35from pydantic import JsonValue 

36from sqlalchemy import and_, func, or_, tuple_, update 

37from sqlalchemy.engine import CursorResult 

38from sqlalchemy.exc import DataError, NoResultFound, SQLAlchemyError 

39from sqlalchemy.orm import contains_eager, joinedload 

40from sqlalchemy.sql import select 

41from structlog.contextvars import bind_contextvars 

42 

43from airflow._shared.observability.traces import override_ids 

44from airflow._shared.state import TaskScope 

45from airflow._shared.timezones import timezone 

46from airflow.api_fastapi.auth.tokens import JWTGenerator 

47from airflow.api_fastapi.common.dagbag import DagBagDep, get_latest_version_of_dag 

48from airflow.api_fastapi.common.db.common import SessionDep 

49from airflow.api_fastapi.common.types import UtcDateTime 

50from airflow.api_fastapi.compat import HTTP_422_UNPROCESSABLE_CONTENT 

51from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc 

52from airflow.api_fastapi.execution_api.datamodels.taskinstance import ( 

53 InactiveAssetsResponse, 

54 PreviousTIResponse, 

55 PrevSuccessfulDagRunResponse, 

56 TaskBreadcrumbsResponse, 

57 TaskStatesResponse, 

58 TIAwaitingInputStatePayload, 

59 TIDeferredStatePayload, 

60 TIEnterRunningPayload, 

61 TIHeartbeatInfo, 

62 TIRescheduleStatePayload, 

63 TIRetryStatePayload, 

64 TIRunContext, 

65 TISkippedDownstreamTasksStatePayload, 

66 TIStateUpdate, 

67 TISuccessStatePayload, 

68 TITerminalStatePayload, 

69) 

70from airflow.api_fastapi.execution_api.datamodels.token import TIToken 

71from airflow.api_fastapi.execution_api.deps import DepContainer 

72from airflow.api_fastapi.execution_api.security import ( 

73 CurrentTIToken, 

74 ExecutionAPIRoute, 

75 get_team_name_for_ti, 

76 require_auth, 

77) 

78from airflow.configuration import conf 

79from airflow.exceptions import InvalidPartitionKeyError, TaskNotFound 

80from airflow.models.asset import AssetActive 

81from airflow.models.base import ID_LEN 

82from airflow.models.dag import DagModel 

83from airflow.models.dagrun import DagRun as DR 

84from airflow.models.hitl import HITLDetail 

85from airflow.models.log import Log 

86from airflow.models.taskinstance import TaskInstance as TI, _stop_remaining_tasks 

87from airflow.models.taskinstancehistory import TaskInstanceHistory as TIH 

88from airflow.models.taskreschedule import TaskReschedule 

89from airflow.models.trigger import Trigger, handle_event_submit 

90from airflow.models.xcom import XComModel 

91from airflow.serialization.definitions.assets import SerializedAsset, SerializedAssetUniqueKey 

92from airflow.state import get_state_backend 

93from airflow.triggers.base import TriggerEvent 

94from airflow.utils.sqlalchemy import get_dialect_name 

95from airflow.utils.state import DagRunState, TaskInstanceState, TerminalTIState 

96 

97if TYPE_CHECKING: 97 ↛ 98line 97 didn't jump to line 98 because the condition on line 97 was never true

98 from sqlalchemy.sql.dml import Update 

99 

100router = VersionedAPIRouter() 

101 

102ti_id_router = VersionedAPIRouter( 

103 route_class=ExecutionAPIRoute, 

104 dependencies=[ 

105 Security(require_auth, scopes=["ti:self"]), 

106 ], 

107) 

108 

109 

110log = structlog.get_logger(__name__) 

111tracer = trace.get_tracer(__name__) 

112 

113 

114@ti_id_router.patch( 

115 "/{task_instance_id}/run", 

116 status_code=status.HTTP_200_OK, 

117 dependencies=[Security(require_auth, scopes=["token:execution", "token:workload"])], 

118 responses=create_openapi_http_exception_doc( 

119 [ 

120 (status.HTTP_404_NOT_FOUND, "Task Instance not found"), 

121 (status.HTTP_409_CONFLICT, "The TI is already in the requested state"), 

122 (HTTP_422_UNPROCESSABLE_CONTENT, "Invalid payload for the state transition"), 

123 ] 

124 ), 

125 response_model_exclude_unset=True, 

126) 

127def ti_run( 

128 task_instance_id: UUID, 

129 ti_run_payload: Annotated[TIEnterRunningPayload, Body()], 

130 response: Response, 

131 session: SessionDep, 

132 dag_bag: DagBagDep, 

133 services=DepContainer, 

134 token: TIToken = CurrentTIToken, 

135) -> TIRunContext: 

136 """ 

137 Run a TaskInstance. 

138 

139 This endpoint is used to start a TaskInstance that is in the QUEUED state. 

140 """ 

141 bind_contextvars(ti_id=str(task_instance_id)) 

142 log.debug( 

143 "Starting task instance run", 

144 hostname=ti_run_payload.hostname, 

145 unixname=ti_run_payload.unixname, 

146 pid=ti_run_payload.pid, 

147 ) 

148 

149 from sqlalchemy.sql import column 

150 from sqlalchemy.types import JSON 

151 

152 old = ( 

153 select( 

154 TI.state, 

155 TI.dag_id, 

156 TI.run_id, 

157 TI.task_id, 

158 TI.map_index, 

159 TI.try_number, 

160 TI.max_tries, 

161 TI.start_date, 

162 TI.next_method, 

163 TI.hostname, 

164 TI.unixname, 

165 TI.pid, 

166 # This selects the raw JSON value, bypassing the deserialization -- we want that to happen on the 

167 # client 

168 column("next_kwargs", JSON), 

169 DR.logical_date, 

170 DagModel.owners, 

171 ) 

172 .select_from(TI) 

173 .join(DR, and_(TI.dag_id == DR.dag_id, TI.run_id == DR.run_id)) 

174 .join(DagModel, TI.dag_id == DagModel.dag_id) 

175 .where(TI.id == task_instance_id) 

176 .with_for_update(of=TI) 

177 ) 

178 try: 

179 ti = session.execute(old).one() 

180 log.debug("Retrieved task instance details", state=ti.state, dag_id=ti.dag_id, task_id=ti.task_id) 

181 except NoResultFound: 

182 log.error("Task Instance not found") 

183 raise HTTPException( 

184 status_code=status.HTTP_404_NOT_FOUND, 

185 detail={ 

186 "reason": "not_found", 

187 "message": "Task Instance not found", 

188 }, 

189 ) 

190 

191 # We exclude_unset to avoid updating fields that are not set in the payload 

192 data = ti_run_payload.model_dump(exclude_unset=True) 

193 

194 # don't update start date when resuming from deferral 

195 if ti.next_kwargs: 195 ↛ 196line 195 didn't jump to line 196 because the condition on line 195 was never true

196 data.pop("start_date") 

197 log.debug("Removed start_date from update as task is resuming from deferral") 

198 

199 query = update(TI).where(TI.id == task_instance_id).values(data) 

200 

201 previous_state = ti.state 

202 

203 # If we are already running, but this is a duplicate request from the same client return the same OK 

204 # -- it's possible there was a network glitch and they never got the response 

205 if previous_state == TaskInstanceState.RUNNING and (ti.hostname, ti.unixname, ti.pid) == ( 

206 ti_run_payload.hostname, 

207 ti_run_payload.unixname, 

208 ti_run_payload.pid, 

209 ): 

210 log.info("Duplicate start request received", hostname=ti_run_payload.hostname) 

211 elif previous_state not in (TaskInstanceState.QUEUED, TaskInstanceState.RESTARTING): 

212 log.warning( 

213 "Cannot start Task Instance in invalid state", 

214 previous_state=previous_state, 

215 ) 

216 

217 raise HTTPException( 

218 status_code=status.HTTP_409_CONFLICT, 

219 detail={ 

220 "reason": "invalid_state", 

221 "message": "TI was not in a state where it could be marked as running", 

222 "previous_state": previous_state, 

223 }, 

224 ) 

225 else: 

226 log.info("Task started", previous_state=previous_state, hostname=ti_run_payload.hostname) 

227 session.add( 

228 Log( 

229 event=TaskInstanceState.RUNNING.value, 

230 task_id=ti.task_id, 

231 dag_id=ti.dag_id, 

232 run_id=ti.run_id, 

233 map_index=ti.map_index, 

234 try_number=ti.try_number, 

235 logical_date=ti.logical_date, 

236 owner=ti.owners, 

237 extra=json.dumps({"host_name": ti_run_payload.hostname}) if ti_run_payload.hostname else None, 

238 ) 

239 ) 

240 # Ensure there is no end date set and clear retry policy overrides from the previous attempt. 

241 query = query.values( 

242 end_date=None, 

243 hostname=ti_run_payload.hostname, 

244 unixname=ti_run_payload.unixname, 

245 pid=ti_run_payload.pid, 

246 state=TaskInstanceState.RUNNING, 

247 last_heartbeat_at=timezone.utcnow(), 

248 retry_delay_override=None, 

249 retry_reason=None, 

250 ) 

251 

252 try: 

253 result = session.execute(query) 

254 log.info("Task instance state updated", rows_affected=getattr(result, "rowcount", 0)) 

255 

256 dr = ( 

257 session.scalars( 

258 select(DR) 

259 .filter_by(dag_id=ti.dag_id, run_id=ti.run_id) 

260 .options(joinedload(DR.consumed_asset_events)) 

261 ) 

262 .unique() 

263 .one_or_none() 

264 ) 

265 

266 if not dr: 266 ↛ 267line 266 didn't jump to line 267 because the condition on line 266 was never true

267 log.error("DagRun not found", dag_id=ti.dag_id, run_id=ti.run_id) 

268 raise HTTPException( 

269 status_code=status.HTTP_404_NOT_FOUND, 

270 detail={ 

271 "reason": "not_found", 

272 "message": f"DagRun with dag_id={ti.dag_id} and run_id={ti.run_id} not found", 

273 }, 

274 ) 

275 

276 # Send the keys to the SDK so that the client requests to clear those XComs from the server. 

277 # The reason we cannot do this here in the server is because we need to issue a purge on custom XCom backends 

278 # too. With the current assumption, the workers ONLY have access to the custom XCom backends directly and they 

279 # can issue the purge. 

280 

281 # However, do not clear it for deferral 

282 xcom_keys = [] 

283 if not ti.next_method: 283 ↛ 294line 283 didn't jump to line 294 because the condition on line 283 was always true

284 map_index = None if ti.map_index < 0 else ti.map_index 

285 xcom_query = select(XComModel.key).where( 

286 XComModel.dag_id == ti.dag_id, 

287 XComModel.task_id == ti.task_id, 

288 XComModel.run_id == ti.run_id, 

289 ) 

290 if map_index is not None: 290 ↛ 291line 290 didn't jump to line 291 because the condition on line 290 was never true

291 xcom_query = xcom_query.where(XComModel.map_index == map_index) 

292 

293 xcom_keys = list(session.scalars(xcom_query)) 

294 task_reschedule_count = ( 

295 session.scalar( 

296 select(func.count(TaskReschedule.id)).where(TaskReschedule.ti_id == task_instance_id) 

297 ) 

298 or 0 

299 ) 

300 

301 dr.team_name = get_team_name_for_ti(task_instance_id, session) 

302 

303 context = TIRunContext( 

304 dag_run=dr, 

305 task_reschedule_count=task_reschedule_count, 

306 max_tries=ti.max_tries, 

307 # TODO: Add variables and connections that are needed (and has perms) for the task 

308 variables=[], 

309 connections=[], 

310 xcom_keys_to_clear=xcom_keys, 

311 should_retry=_is_eligible_to_retry(previous_state, ti.try_number, ti.max_tries), 

312 ) 

313 

314 # Only set if they are non-null 

315 if ti.next_method: 315 ↛ 316line 315 didn't jump to line 316 because the condition on line 315 was never true

316 context.next_method = ti.next_method 

317 context.next_kwargs = ti.next_kwargs 

318 context.start_date = ti.start_date 

319 except DataError: 

320 # Let the app-level DataErrorHandler return a 422 (not the opaque 500 below). 

321 raise 

322 except SQLAlchemyError: 

323 # Defer to app-level SQLAlchemyError handler (returns HTTP 500). 

324 raise 

325 

326 # JWTReissueMiddleware also writes Refreshed-API-Token but skips workload tokens, so we set it here for the workload→execution swap. 

327 if token.claims.scope == "workload": 327 ↛ 332line 327 didn't jump to line 332 because the condition on line 327 was always true

328 generator: JWTGenerator = services.get(JWTGenerator) 

329 execution_token = generator.generate(extras={"sub": str(task_instance_id), "scope": "execution"}) 

330 response.headers["Refreshed-API-Token"] = execution_token 

331 

332 return context 

333 

334 

335@ti_id_router.patch( 

336 "/{task_instance_id}/state", 

337 status_code=status.HTTP_204_NO_CONTENT, 

338 responses={ 

339 status.HTTP_200_OK: {"description": "The TI was already in the requested state"}, 

340 status.HTTP_404_NOT_FOUND: {"description": "Task Instance not found"}, 

341 status.HTTP_409_CONFLICT: {"description": "The TI is not in a valid state for this transition"}, 

342 HTTP_422_UNPROCESSABLE_CONTENT: {"description": "Invalid payload for the state transition"}, 

343 }, 

344) 

345def ti_update_state( 

346 task_instance_id: UUID, 

347 ti_patch_payload: Annotated[TIStateUpdate, Body()], 

348 session: SessionDep, 

349 dag_bag: DagBagDep, 

350): 

351 """ 

352 Update the state of a TaskInstance. 

353 

354 Not all state transitions are valid, and transitioning to some states requires extra information to be 

355 passed along. (Check out the datamodels for details, the rendered docs might not reflect this accurately) 

356 """ 

357 bind_contextvars(ti_id=str(task_instance_id)) 

358 log.debug("Updating task instance state", new_state=ti_patch_payload.state) 

359 

360 old = ( 

361 select( 

362 TI.state, 

363 TI.try_number, 

364 TI.max_tries, 

365 TI.dag_id, 

366 TI.task_id, 

367 TI.run_id, 

368 TI.map_index, 

369 TI.hostname, 

370 DR.logical_date, 

371 DagModel.owners, 

372 ) 

373 .select_from(TI) 

374 .join(DR, and_(TI.dag_id == DR.dag_id, TI.run_id == DR.run_id)) 

375 .join(DagModel, TI.dag_id == DagModel.dag_id) 

376 .where(TI.id == task_instance_id) 

377 .with_for_update(of=TI) 

378 ) 

379 try: 

380 ( 

381 previous_state, 

382 try_number, 

383 max_tries, 

384 dag_id, 

385 task_id, 

386 run_id, 

387 map_index, 

388 hostname, 

389 logical_date, 

390 owners, 

391 ) = session.execute(old).one() 

392 log.debug( 

393 "Retrieved current task instance state", 

394 previous_state=previous_state, 

395 try_number=try_number, 

396 max_tries=max_tries, 

397 ) 

398 except NoResultFound: 

399 log.error("Task Instance not found") 

400 raise HTTPException( 

401 status_code=status.HTTP_404_NOT_FOUND, 

402 detail={ 

403 "reason": "not_found", 

404 "message": "Task Instance not found", 

405 }, 

406 ) 

407 

408 # TIStateUpdate can include terminal and intermediate states. This idempotency check handles 

409 # duplicate updates when the requested state is already persisted (for example SUCCESS -> 

410 # SUCCESS or DEFERRED -> DEFERRED), including duplicates that would not pass the RUNNING 

411 # transition check below. 

412 if ti_patch_payload.state.value == previous_state: 

413 log.info( 

414 "Duplicate state update request received; state already set", 

415 requested_state=ti_patch_payload.state.value, 

416 previous_state=previous_state, 

417 ) 

418 return Response(status_code=status.HTTP_200_OK) 

419 

420 if previous_state != TaskInstanceState.RUNNING: 420 ↛ 421line 420 didn't jump to line 421 because the condition on line 420 was never true

421 log.warning( 

422 "Cannot update Task Instance in invalid state", 

423 previous_state=previous_state, 

424 ) 

425 raise HTTPException( 

426 status_code=status.HTTP_409_CONFLICT, 

427 detail={ 

428 "reason": "invalid_state", 

429 "message": "TI was not in the running state so it cannot be updated", 

430 "previous_state": previous_state, 

431 }, 

432 ) 

433 

434 # Validate outlet event partition keys early, before entering the catch-all 

435 # except block that would otherwise swallow the HTTPException and mark the TI failed. 

436 if isinstance(ti_patch_payload, TISuccessStatePayload): 

437 try: 

438 _validate_outlet_event_partition_keys(ti_patch_payload.outlet_events) 

439 except InvalidPartitionKeyError as e: 

440 raise HTTPException( 

441 status_code=HTTP_422_UNPROCESSABLE_CONTENT, 

442 detail={"reason": "invalid_partition_key", "message": str(e)}, 

443 ) from e 

444 

445 # We exclude_unset to avoid updating fields that are not set in the payload 

446 data = ti_patch_payload.model_dump( 

447 exclude={"task_outlets", "outlet_events", "retry_delay_seconds", "retry_reason"}, 

448 exclude_unset=True, 

449 ) 

450 if "rendered_map_index" in data: 

451 data["_rendered_map_index"] = data.pop("rendered_map_index") 

452 query = update(TI).where(TI.id == task_instance_id).values(data) 

453 

454 asset_callbacks: Sequence[Callable[[], None]] = () 

455 try: 

456 query, updated_state, asset_callbacks = _create_ti_state_update_query_and_update_state( 

457 ti_patch_payload=ti_patch_payload, 

458 task_instance_id=task_instance_id, 

459 session=session, 

460 query=query, 

461 dag_id=dag_id, 

462 dag_bag=dag_bag, 

463 ) 

464 except DataError: 

465 # Let DataErrorHandler return a 422 instead of silently marking the TI FAILED below. 

466 raise 

467 except Exception: 

468 # Set a task to failed in case any unexpected exception happened during task state update 

469 log.exception( 

470 "Error updating Task Instance state. Setting the task to failed.", 

471 payload=ti_patch_payload, 

472 ) 

473 session.rollback() 

474 ti = session.get(TI, task_instance_id, with_for_update={"of": TI}) 

475 if session.bind is not None: 475 ↛ 477line 475 didn't jump to line 477 because the condition on line 475 was always true

476 query = TI.duration_expression_update(timezone.utcnow(), query, session.bind) 

477 query = query.values(state=(updated_state := TaskInstanceState.FAILED)) 

478 if ti is not None: 478 ↛ 483line 478 didn't jump to line 483 because the condition on line 478 was always true

479 _handle_fail_fast_for_dag(ti=ti, dag_id=dag_id, session=session, dag_bag=dag_bag) 

480 

481 # TODO: Replace this with FastAPI's Custom Exception handling: 

482 # https://fastapi.tiangolo.com/tutorial/handling-errors/#install-custom-exception-handlers 

483 try: 

484 result = session.execute(query) 

485 log.info( 

486 "Task instance state updated", 

487 new_state=updated_state, 

488 rows_affected=getattr(result, "rowcount", 0), 

489 ) 

490 session.add( 

491 Log( 

492 event=updated_state.value, 

493 task_id=task_id, 

494 dag_id=dag_id, 

495 run_id=run_id, 

496 map_index=map_index, 

497 try_number=try_number, 

498 logical_date=logical_date, 

499 owner=owners, 

500 extra=json.dumps({"host_name": hostname}) if hostname else None, 

501 ) 

502 ) 

503 except DataError: 

504 # Let DataErrorHandler return a 422 (not the opaque 500 below). 

505 raise 

506 except SQLAlchemyError: 

507 # Defer to app-level SQLAlchemyError handler (returns HTTP 500). 

508 raise 

509 

510 if updated_state == TaskInstanceState.SUCCESS: 

511 if conf.getboolean("state_store", "clear_on_success"): 511 ↛ 512line 511 didn't jump to line 512 because the condition on line 511 was never true

512 scope = TaskScope( 

513 dag_id=dag_id, 

514 run_id=run_id, 

515 task_id=task_id, 

516 map_index=map_index if map_index is not None else -1, 

517 ) 

518 try: 

519 get_state_backend().clear(scope, session=session) 

520 log.info( 

521 "Cleared task state on success", 

522 dag_id=dag_id, 

523 run_id=run_id, 

524 task_id=task_id, 

525 map_index=map_index, 

526 ) 

527 except Exception: 

528 log.warning( 

529 "Failed to clear task state on success", 

530 dag_id=dag_id, 

531 run_id=run_id, 

532 task_id=task_id, 

533 ) 

534 

535 # Release the task_instance row lock before running listener callbacks. 

536 session.commit() 

537 

538 for callback in asset_callbacks: 

539 callback() 

540 

541 

542def _emit_task_span(ti, state): 

543 # just to be safe 

544 if not ti.dag_run: 544 ↛ 545line 544 didn't jump to line 545 because the condition on line 544 was never true

545 return 

546 if not isinstance(ti.dag_run.context_carrier, dict): 546 ↛ 547line 546 didn't jump to line 547 because the condition on line 546 was never true

547 return 

548 if not isinstance(ti.context_carrier, dict): 548 ↛ 549line 548 didn't jump to line 549 because the condition on line 548 was never true

549 return 

550 dr_ctx = TraceContextTextMapPropagator().extract(ti.dag_run.context_carrier) 

551 

552 # Skip if the run was head-sampled out, so every span in the run agrees with the 

553 # carrier's decision. A parent-based sampler would already drop this child span, 

554 # but the explicit check also covers non-parent-based samplers (which ignore the 

555 # parent and would re-sample it in) and short-circuits before building the span. 

556 # An invalid/empty carrier (legacy/NULL) recorded no decision, so it falls through 

557 # and still emits — preserving prior behavior. 

558 dr_span_context = trace.get_current_span(context=dr_ctx).get_span_context() 

559 if dr_span_context.is_valid and not dr_span_context.trace_flags.sampled: 559 ↛ 562line 559 didn't jump to line 562 because the condition on line 559 was always true

560 return 

561 

562 ti_ctx = TraceContextTextMapPropagator().extract(ti.context_carrier) 

563 ti_span = trace.get_current_span(context=ti_ctx) 

564 span_context = ti_span.get_span_context() 

565 start_time_candidates = (x for x in (ti.queued_dttm, ti.start_date, timezone.utcnow()) if x) 

566 name = f"task_run.{ti.task_id}" 

567 if ti.map_index >= 0: 

568 name += f"[{ti.map_index}]" 

569 with override_ids(span_context.trace_id, span_context.span_id): 

570 span = tracer.start_span( 

571 name=name, 

572 start_time=int(min(start_time_candidates).timestamp() * 1e9), 

573 context=dr_ctx, 

574 ) 

575 

576 span.set_attributes( 

577 { 

578 "airflow.dag_id": ti.dag_id, 

579 "airflow.task_id": ti.task_id, 

580 "airflow.dag_run.run_id": ti.run_id, 

581 "airflow.task_instance.try_number": ti.try_number, 

582 "airflow.task_instance.map_index": ti.map_index if ti.map_index is not None else -1, 

583 "airflow.task_instance.state": state, 

584 "airflow.task_instance.id": str(ti.id), 

585 } 

586 ) 

587 status_code = StatusCode.OK if state == TaskInstanceState.SUCCESS else StatusCode.ERROR 

588 span.set_status(status_code) 

589 span.end() 

590 

591 

592def _handle_fail_fast_for_dag(ti: TI, dag_id: str, session: SessionDep, dag_bag: DagBagDep) -> None: 

593 dr = ti.dag_run 

594 

595 # Check fail_fast from DagModel (simple column lookup) - early exit if False 

596 # This avoids loading 5-50 MB SerializedDAG in 99% of cases 

597 fail_fast = session.scalar(select(DagModel.fail_fast).where(DagModel.dag_id == dag_id)) 

598 if not fail_fast: 598 ↛ 602line 598 didn't jump to line 602 because the condition on line 598 was always true

599 return 

600 

601 # Only load SerializedDAG when fail_fast=True (rare case ~1%) 

602 ser_dag = dag_bag.get_dag_for_run(dag_run=dr, session=session) 

603 if ser_dag: 

604 task_dict = getattr(ser_dag, "task_dict") 

605 task_teardown_map = {k: v.is_teardown for k, v in task_dict.items()} 

606 _stop_remaining_tasks(task_instance=ti, task_teardown_map=task_teardown_map, session=session) 

607 

608 

609def _validate_outlet_event_partition_keys(outlet_events: list[dict[str, Any]]) -> None: 

610 """ 

611 Validate partition_key values embedded in outlet events. 

612 

613 Raises ``InvalidPartitionKeyError`` (which the caller translates to HTTP 422) 

614 if any per-emission partition key is empty/whitespace-only or exceeds the 

615 ``ID_LEN`` column width used in the metadata database. 

616 """ 

617 for event in outlet_events: 

618 if (pk := event.get("partition_key")) is None: 

619 continue 

620 if not pk.strip(): 620 ↛ 621line 620 didn't jump to line 621 because the condition on line 620 was never true

621 raise InvalidPartitionKeyError( 

622 f"partition_key in outlet event must not be empty or whitespace-only; got {pk!r}." 

623 ) 

624 if len(pk) > ID_LEN: 624 ↛ 625line 624 didn't jump to line 625 because the condition on line 624 was never true

625 raise InvalidPartitionKeyError( 

626 f"partition_key in outlet event must be at most {ID_LEN} characters; got {len(pk)}." 

627 ) 

628 

629 

630def _create_ti_state_update_query_and_update_state( 

631 *, 

632 ti_patch_payload: TIStateUpdate, 

633 task_instance_id: UUID, 

634 query: Update, 

635 session: SessionDep, 

636 dag_bag: DagBagDep, 

637 dag_id: str, 

638) -> tuple[Update, TaskInstanceState, Sequence[Callable[[], None]]]: 

639 asset_callbacks: Sequence[Callable[[], None]] = () 

640 if isinstance(ti_patch_payload, (TITerminalStatePayload, TIRetryStatePayload, TISuccessStatePayload)): 

641 ti = session.get(TI, task_instance_id, with_for_update={"of": TI}) 

642 updated_state = TaskInstanceState(ti_patch_payload.state.value) 

643 if session.bind is not None: 643 ↛ 645line 643 didn't jump to line 645 because the condition on line 643 was always true

644 query = TI.duration_expression_update(ti_patch_payload.end_date, query, session.bind) 

645 query = query.values(state=updated_state, next_method=None, next_kwargs=None) 

646 

647 if updated_state == TaskInstanceState.FAILED: 

648 # This is the only case needs extra handling for TITerminalStatePayload 

649 if ti is not None: 649 ↛ 676line 649 didn't jump to line 676 because the condition on line 649 was always true

650 _handle_fail_fast_for_dag(ti=ti, dag_id=dag_id, session=session, dag_bag=dag_bag) 

651 elif isinstance(ti_patch_payload, TIRetryStatePayload): 

652 retry_delay_override = ti_patch_payload.retry_delay_seconds 

653 retry_reason = ti_patch_payload.retry_reason[:500] if ti_patch_payload.retry_reason else None 

654 if ti is not None: 654 ↛ 667line 654 didn't jump to line 667 because the condition on line 654 was always true

655 # Snapshot the finished try onto the TI *before* archiving so record_ti() 

656 # copies the values into task_instance_history (it reads attrs off the 

657 # ti object and cannot see the live-row UPDATE built below). 

658 ti.retry_delay_override = retry_delay_override 

659 ti.retry_reason = retry_reason 

660 ti.end_date = ti_patch_payload.end_date 

661 ti.set_duration() 

662 if "rendered_map_index" in ti_patch_payload.model_fields_set: 662 ↛ 664line 662 didn't jump to line 664 because the condition on line 662 was always true

663 ti._rendered_map_index = ti_patch_payload.rendered_map_index 

664 ti.prepare_db_for_next_try(session=session) 

665 # Store retry policy overrides so next_retry_datetime() can read them. 

666 # These are cleared when the task enters RUNNING (ti_run). 

667 query = query.values(retry_delay_override=retry_delay_override, retry_reason=retry_reason) 

668 elif isinstance(ti_patch_payload, TISuccessStatePayload): 

669 if ti is not None: 669 ↛ 676line 669 didn't jump to line 676 because the condition on line 669 was always true

670 asset_callbacks = TI.register_asset_changes_in_db( 

671 ti, 

672 ti_patch_payload.task_outlets, 

673 ti_patch_payload.outlet_events, 

674 session=session, 

675 ) 

676 try: 

677 _emit_task_span(ti, state=updated_state) 

678 except Exception: 

679 log.warning("Failed to emit task span", exc_info=True) 

680 elif isinstance(ti_patch_payload, TIDeferredStatePayload): 

681 # Calculate timeout if it was passed 

682 timeout = None 

683 if ti_patch_payload.trigger_timeout is not None: 683 ↛ 686line 683 didn't jump to line 686 because the condition on line 683 was always true

684 timeout = timezone.utcnow() + ti_patch_payload.trigger_timeout 

685 

686 trigger_kwargs = ti_patch_payload.trigger_kwargs 

687 if not isinstance(trigger_kwargs, str): 687 ↛ 692line 687 didn't jump to line 692 because the condition on line 687 was always true

688 # If it's passed as a string, assume the client encrypted it, otherwise assume it doesn't need to 

689 # be. Just JSON serialize it 

690 trigger_kwargs = json.dumps(trigger_kwargs) 

691 

692 trigger_row = Trigger( 

693 classpath=ti_patch_payload.classpath, 

694 kwargs={}, 

695 queue=ti_patch_payload.queue, 

696 team_name=get_team_name_for_ti(task_instance_id, session), 

697 ) 

698 trigger_row.encrypted_kwargs = trigger_kwargs 

699 session.add(trigger_row) 

700 session.flush() 

701 

702 # TODO: HANDLE execution timeout later as it requires a call to the DB 

703 # either get it from the serialised DAG or get it from the API 

704 

705 query = update(TI).where(TI.id == task_instance_id) 

706 

707 # Store next_kwargs directly (already serialized by worker) 

708 query = query.values( 

709 state=TaskInstanceState.DEFERRED, 

710 trigger_id=trigger_row.id, 

711 next_method=ti_patch_payload.next_method, 

712 next_kwargs=ti_patch_payload.next_kwargs, 

713 trigger_timeout=timeout, 

714 ) 

715 updated_state = TaskInstanceState.DEFERRED 

716 elif isinstance(ti_patch_payload, TIAwaitingInputStatePayload): 716 ↛ 723line 716 didn't jump to line 723 because the condition on line 716 was never true

717 # Park the task waiting for human input (Human-in-the-loop). No trigger / triggerer is 

718 # created: the task is resumed by the Core API response handler or the scheduler timeout 

719 # sweep. The optional response deadline is stored on the existing trigger_timeout column. 

720 # 

721 # Fixed lock order (TaskInstance -> HITLDetail), matching the Core API response path, so a 

722 # human response racing this park transition cannot deadlock. 

723 ti = session.get(TI, task_instance_id, with_for_update={"of": TI}) 

724 # Lock only the hitl_detail row (of=...): HITLDetail eager-joins task_instance (lazy="joined"), 

725 # and Postgres rejects FOR UPDATE against the nullable side of that outer join. 

726 hitl_detail = session.scalar( 

727 select(HITLDetail).where(HITLDetail.ti_id == task_instance_id).with_for_update(of=HITLDetail) 

728 ) 

729 if ti is not None and hitl_detail is not None and hitl_detail.response_received: 

730 # The human responded in the window between the operator writing the HITL request and 

731 # the worker parking the task. Resume straight to execute_complete instead of parking, 

732 # which would otherwise strand an already-responded task (no trigger/sweep would fire). 

733 # Carry next_method/next_kwargs onto the TI first so the resume dispatches correctly; 

734 # handle_event_submit then injects the response event into next_kwargs. 

735 ti.next_method = ti_patch_payload.next_method 

736 ti.next_kwargs = ti_patch_payload.next_kwargs 

737 handle_event_submit( 

738 TriggerEvent(hitl_detail.as_resume_event_payload()), 

739 task_instance=ti, 

740 session=session, 

741 ) 

742 query = update(TI).where(TI.id == task_instance_id).values(state=TaskInstanceState.SCHEDULED) 

743 updated_state = TaskInstanceState.SCHEDULED 

744 else: 

745 timeout = None 

746 if ti_patch_payload.timeout is not None: 

747 timeout = timezone.utcnow() + ti_patch_payload.timeout 

748 

749 query = update(TI).where(TI.id == task_instance_id) 

750 query = query.values( 

751 state=TaskInstanceState.AWAITING_INPUT, 

752 trigger_id=None, 

753 next_method=ti_patch_payload.next_method, 

754 next_kwargs=ti_patch_payload.next_kwargs, 

755 trigger_timeout=timeout, 

756 ) 

757 updated_state = TaskInstanceState.AWAITING_INPUT 

758 elif isinstance(ti_patch_payload, TIRescheduleStatePayload): 758 ↛ 801line 758 didn't jump to line 801 because the condition on line 758 was always true

759 # Quick check for poke_interval isn't immediately over MySQL's TIMESTAMP limit. 

760 # This check is only rudimentary to catch trivial user errors, e.g. mistakenly 

761 # set the value to milliseconds instead of seconds. There's another check when 

762 # we actually try to reschedule to ensure database coherence. 

763 if get_dialect_name(session) == "mysql": 763 ↛ 765line 763 didn't jump to line 765 because the condition on line 763 was never true

764 # As documented in https://dev.mysql.com/doc/refman/5.7/en/datetime.html. 

765 _MYSQL_TIMESTAMP_MAX = timezone.datetime(2038, 1, 19, 3, 14, 7) 

766 if ti_patch_payload.reschedule_date > _MYSQL_TIMESTAMP_MAX: 

767 # Set a task to failed in case any unexpected exception happened during task state update 

768 log.error( 

769 "Cannot reschedule task past MySQL limit. Setting the task to failed.", 

770 payload=ti_patch_payload, 

771 mysql_timestamp_max=_MYSQL_TIMESTAMP_MAX, 

772 ) 

773 data = ti_patch_payload.model_dump(exclude={"reschedule_date"}, exclude_unset=True) 

774 query = update(TI).where(TI.id == task_instance_id).values(data) 

775 if session.bind is not None: 

776 query = TI.duration_expression_update(timezone.utcnow(), query, session.bind) 

777 query = query.values(state=TaskInstanceState.FAILED) 

778 ti = session.get(TI, task_instance_id, with_for_update={"of": TI}) 

779 if ti is not None: 

780 _handle_fail_fast_for_dag(ti=ti, dag_id=dag_id, session=session, dag_bag=dag_bag) 

781 return query, TaskInstanceState.FAILED, () 

782 

783 actual_start_date = timezone.utcnow() 

784 session.add( 

785 TaskReschedule( 

786 task_instance_id, 

787 actual_start_date, 

788 ti_patch_payload.end_date, 

789 ti_patch_payload.reschedule_date, 

790 ) 

791 ) 

792 

793 query = update(TI).where(TI.id == task_instance_id) 

794 # calculate the duration for TI table too 

795 if session.bind is not None: 795 ↛ 798line 795 didn't jump to line 798 because the condition on line 795 was always true

796 query = TI.duration_expression_update(ti_patch_payload.end_date, query, session.bind) 

797 # clear the next_method and next_kwargs so that none of the retries pick them up 

798 updated_state = TaskInstanceState.UP_FOR_RESCHEDULE 

799 query = query.values(state=updated_state, next_method=None, next_kwargs=None) 

800 else: 

801 raise ValueError(f"Unexpected Payload Type {type(ti_patch_payload)}") 

802 

803 return query, updated_state, asset_callbacks 

804 

805 

806@ti_id_router.patch( 

807 "/{task_instance_id}/skip-downstream", 

808 status_code=status.HTTP_204_NO_CONTENT, 

809 responses=create_openapi_http_exception_doc( 

810 [ 

811 (status.HTTP_404_NOT_FOUND, "Task Instance not found"), 

812 (HTTP_422_UNPROCESSABLE_CONTENT, "Invalid payload for the state transition"), 

813 ] 

814 ), 

815) 

816def ti_skip_downstream( 

817 task_instance_id: UUID, 

818 ti_patch_payload: TISkippedDownstreamTasksStatePayload, 

819 session: SessionDep, 

820): 

821 bind_contextvars(ti_id=str(task_instance_id)) 

822 log.info("Skipping downstream tasks", task_count=len(ti_patch_payload.tasks)) 

823 

824 now = timezone.utcnow() 

825 tasks = ti_patch_payload.tasks 

826 

827 query_result = session.execute(select(TI.dag_id, TI.run_id).where(TI.id == task_instance_id)) 

828 row_result = query_result.fetchone() 

829 if row_result is None: 829 ↛ 830line 829 didn't jump to line 830 because the condition on line 829 was never true

830 raise HTTPException( 

831 status_code=status.HTTP_404_NOT_FOUND, 

832 detail={"reason": "not_found", "message": "Task Instance not found"}, 

833 ) 

834 dag_id, run_id = row_result 

835 log.debug("Retrieved DAG and run info", dag_id=dag_id, run_id=run_id) 

836 

837 task_ids = [task if isinstance(task, tuple) else (task, -1) for task in tasks] 

838 log.debug("Prepared task IDs for skipping", task_ids=task_ids) 

839 

840 # Don't overwrite tasks that are already executing or finished. 

841 # See: https://github.com/apache/airflow/issues/59378 

842 # Note: SQL NULL NOT IN (...) is falsy, so we need an explicit IS NULL check. 

843 skippable_state_clause = or_( 

844 TI.state.is_(None), 

845 TI.state.not_in( 

846 [ 

847 TaskInstanceState.RUNNING, 

848 TaskInstanceState.SUCCESS, 

849 TaskInstanceState.FAILED, 

850 ] 

851 ), 

852 ) 

853 query = ( 

854 update(TI) 

855 .where( 

856 TI.dag_id == dag_id, 

857 TI.run_id == run_id, 

858 tuple_(TI.task_id, TI.map_index).in_(task_ids), 

859 skippable_state_clause, 

860 ) 

861 .values(state=TaskInstanceState.SKIPPED, start_date=now, end_date=now) 

862 .execution_options(synchronize_session=False) 

863 ) 

864 

865 result = session.execute(query) 

866 log.info("Downstream tasks skipped", tasks_skipped=getattr(result, "rowcount", 0)) 

867 

868 

869def _raise_ti_not_in_live_table(task_instance_id: UUID, session: SessionDep) -> NoReturn: 

870 """Raise 410 Gone if the missing TI id was archived to history, else 404 Not Found.""" 

871 if session.scalar( 

872 select(func.count(TIH.task_instance_id)).where(TIH.task_instance_id == task_instance_id) 

873 ): 

874 log.error("TaskInstance not in live table but archived in history", ti_id=str(task_instance_id)) 

875 raise HTTPException( 

876 status_code=status.HTTP_410_GONE, 

877 detail={ 

878 "reason": "not_found", 

879 "message": "Task Instance not found, it may have been moved to the Task Instance History table", 

880 }, 

881 ) 

882 log.error("Task Instance not found", ti_id=str(task_instance_id)) 

883 raise HTTPException( 

884 status_code=status.HTTP_404_NOT_FOUND, 

885 detail={ 

886 "reason": "not_found", 

887 "message": "Task Instance not found", 

888 }, 

889 ) 

890 

891 

892@ti_id_router.put( 

893 "/{task_instance_id}/heartbeat", 

894 status_code=status.HTTP_204_NO_CONTENT, 

895 responses=create_openapi_http_exception_doc( 

896 [ 

897 (status.HTTP_404_NOT_FOUND, "Task Instance not found"), 

898 ( 

899 status.HTTP_409_CONFLICT, 

900 "The TI attempting to heartbeat should be terminated for the given reason", 

901 ), 

902 ( 

903 status.HTTP_410_GONE, 

904 "Task Instance not found in the TI table but exists in the Task Instance History table", 

905 ), 

906 (HTTP_422_UNPROCESSABLE_CONTENT, "Invalid payload for the state transition"), 

907 ] 

908 ), 

909) 

910def ti_heartbeat( 

911 task_instance_id: UUID, 

912 ti_payload: TIHeartbeatInfo, 

913 session: SessionDep, 

914): 

915 """Update the heartbeat of a TaskInstance to mark it as alive & still running.""" 

916 bind_contextvars(ti_id=str(task_instance_id)) 

917 log.debug("Processing heartbeat", hostname=ti_payload.hostname, pid=ti_payload.pid) 

918 

919 # Hot path: in the common case the TI is still running on the same host and pid, 

920 # so we can update last_heartbeat_at directly without first taking a row lock. 

921 fast_path_result = cast( 

922 "CursorResult[Any]", 

923 session.execute( 

924 update(TI) 

925 .where( 

926 TI.id == task_instance_id, 

927 TI.state == TaskInstanceState.RUNNING, 

928 TI.hostname == ti_payload.hostname, 

929 TI.pid == ti_payload.pid, 

930 ) 

931 .values(last_heartbeat_at=timezone.utcnow()) 

932 .execution_options(synchronize_session=False) 

933 ), 

934 ) 

935 if fast_path_result.rowcount is not None and fast_path_result.rowcount > 0: 935 ↛ 939line 935 didn't jump to line 939 because the condition on line 935 was always true

936 log.debug("Heartbeat updated via fast path") 

937 return 

938 

939 log.debug("Heartbeat fast path missed; falling back to diagnostic checks") 

940 

941 old = select(TI.state, TI.hostname, TI.pid).where(TI.id == task_instance_id).with_for_update() 

942 

943 try: 

944 (previous_state, hostname, pid) = session.execute(old).one() 

945 log.debug( 

946 "Retrieved current task state", state=previous_state, current_hostname=hostname, current_pid=pid 

947 ) 

948 except NoResultFound: 

949 # Check if the TI exists in the Task Instance History table. 

950 # If it does, it was likely cleared while running, so return 410 Gone 

951 # instead of 404 Not Found to give the client a more specific signal. 

952 _raise_ti_not_in_live_table(task_instance_id, session) 

953 

954 if hostname != ti_payload.hostname or pid != ti_payload.pid: 

955 log.warning( 

956 "Task running elsewhere", 

957 current_hostname=hostname, 

958 current_pid=pid, 

959 requested_hostname=ti_payload.hostname, 

960 requested_pid=ti_payload.pid, 

961 ) 

962 raise HTTPException( 

963 status_code=status.HTTP_409_CONFLICT, 

964 detail={ 

965 "reason": "running_elsewhere", 

966 "message": "TI is already running elsewhere", 

967 "current_hostname": hostname, 

968 "current_pid": pid, 

969 }, 

970 ) 

971 

972 if previous_state != TaskInstanceState.RUNNING: 

973 log.warning("Task not in running state", current_state=previous_state) 

974 raise HTTPException( 

975 status_code=status.HTTP_409_CONFLICT, 

976 detail={ 

977 "reason": "not_running", 

978 "message": "TI is no longer in the running state and task should terminate", 

979 "current_state": previous_state, 

980 }, 

981 ) 

982 

983 # Update the last heartbeat time! 

984 session.execute(update(TI).where(TI.id == task_instance_id).values(last_heartbeat_at=timezone.utcnow())) 

985 log.debug("Heartbeat updated", state=previous_state) 

986 

987 

988@ti_id_router.put( 

989 "/{task_instance_id}/rtif", 

990 status_code=status.HTTP_201_CREATED, 

991 operation_id="put_rtif", 

992 summary="Set Rendered Task Instance Fields", 

993 description="Store the rendered task instance fields (RTIF) for a task instance. " 

994 "These are the template fields after Jinja rendering has been applied. " 

995 "Called by the worker after task execution begins.", 

996 responses=create_openapi_http_exception_doc( 

997 [ 

998 (status.HTTP_404_NOT_FOUND, "Task Instance not found"), 

999 ( 

1000 status.HTTP_410_GONE, 

1001 "Task Instance not found in the TI table but exists in the Task Instance History table", 

1002 ), 

1003 ( 

1004 HTTP_422_UNPROCESSABLE_CONTENT, 

1005 "Invalid payload for the setting rendered task instance fields", 

1006 ), 

1007 ] 

1008 ), 

1009) 

1010def ti_put_rtif( 

1011 task_instance_id: UUID, 

1012 put_rtif_payload: Annotated[dict[str, JsonValue], Body()], 

1013 session: SessionDep, 

1014): 

1015 """Add an RTIF entry for a task instance, sent by the worker.""" 

1016 bind_contextvars(ti_id=str(task_instance_id)) 

1017 log.info("Updating RenderedTaskInstanceFields", field_count=len(put_rtif_payload)) 

1018 

1019 task_instance = session.scalar(select(TI).where(TI.id == task_instance_id)) 

1020 if not task_instance: 1020 ↛ 1022line 1020 didn't jump to line 1022 because the condition on line 1020 was never true

1021 # On retry/clear, the server regenerates the TI id. Return 410 for the stale id. 

1022 _raise_ti_not_in_live_table(task_instance_id, session) 

1023 task_instance.update_rtif(put_rtif_payload, session=session) 

1024 log.debug("RenderedTaskInstanceFields updated successfully") 

1025 

1026 return {"message": "Rendered task instance fields successfully set"} 

1027 

1028 

1029@ti_id_router.patch( 

1030 "/{task_instance_id}/rendered-map-index", 

1031 status_code=status.HTTP_204_NO_CONTENT, 

1032 responses=create_openapi_http_exception_doc( 

1033 [ 

1034 (status.HTTP_404_NOT_FOUND, "Task Instance not found"), 

1035 ( 

1036 HTTP_422_UNPROCESSABLE_CONTENT, 

1037 "Invalid rendered_map_index value", 

1038 ), 

1039 ] 

1040 ), 

1041) 

1042def ti_patch_rendered_map_index( 

1043 task_instance_id: UUID, 

1044 rendered_map_index: Annotated[str, Body()], 

1045 session: SessionDep, 

1046): 

1047 """Update rendered_map_index for a task instance, sent by the worker during task execution.""" 

1048 bind_contextvars(ti_id=str(task_instance_id)) 

1049 

1050 if not rendered_map_index: 

1051 log.error("rendered_map_index cannot be empty") 

1052 raise HTTPException( 

1053 status_code=HTTP_422_UNPROCESSABLE_CONTENT, 

1054 detail="rendered_map_index cannot be empty", 

1055 ) 

1056 

1057 log.debug("Updating rendered_map_index", length=len(rendered_map_index)) 

1058 

1059 query = update(TI).where(TI.id == task_instance_id).values(_rendered_map_index=rendered_map_index) 

1060 result = session.execute(query) 

1061 

1062 result = cast("CursorResult[Any]", result) 

1063 if result.rowcount == 0: 

1064 log.error("Task Instance not found") 

1065 raise HTTPException( 

1066 status_code=status.HTTP_404_NOT_FOUND, 

1067 detail="Task Instance not found", 

1068 ) 

1069 

1070 

1071@ti_id_router.get( 

1072 "/{task_instance_id}/previous-successful-dagrun", 

1073 status_code=status.HTTP_200_OK, 

1074 responses=create_openapi_http_exception_doc( 

1075 [ 

1076 (status.HTTP_404_NOT_FOUND, "Task Instance or Dag Run not found"), 

1077 ] 

1078 ), 

1079) 

1080def get_previous_successful_dagrun( 

1081 task_instance_id: UUID, session: SessionDep 

1082) -> PrevSuccessfulDagRunResponse: 

1083 """ 

1084 Get the previous successful DagRun for a TaskInstance. 

1085 

1086 The data from this endpoint is used to get values for Task Context. 

1087 """ 

1088 bind_contextvars(ti_id=str(task_instance_id)) 

1089 log.debug("Retrieving previous successful DAG run") 

1090 

1091 task_instance = session.scalar(select(TI).where(TI.id == task_instance_id)) 

1092 if not task_instance or not task_instance.logical_date: 

1093 log.debug("No task instance or logical date found") 

1094 return PrevSuccessfulDagRunResponse() 

1095 

1096 dag_run = session.scalar( 

1097 select(DR) 

1098 .where( 

1099 DR.dag_id == task_instance.dag_id, 

1100 DR.logical_date < task_instance.logical_date, 

1101 DR.state == DagRunState.SUCCESS, 

1102 ) 

1103 .order_by(DR.logical_date.desc()) 

1104 .limit(1) 

1105 ) 

1106 if not dag_run: 

1107 log.debug("No previous successful DAG run found") 

1108 return PrevSuccessfulDagRunResponse() 

1109 

1110 log.debug( 

1111 "Found previous successful DAG run", 

1112 dag_id=dag_run.dag_id, 

1113 run_id=dag_run.run_id, 

1114 logical_date=dag_run.logical_date, 

1115 ) 

1116 return PrevSuccessfulDagRunResponse.model_validate(dag_run) 

1117 

1118 

1119@router.get("/count", status_code=status.HTTP_200_OK) 

1120def get_task_instance_count( 

1121 dag_id: str, 

1122 session: SessionDep, 

1123 dag_bag: DagBagDep, 

1124 map_index: Annotated[int | None, Query()] = None, 

1125 task_ids: Annotated[list[str] | None, Query()] = None, 

1126 task_group_id: Annotated[str | None, Query()] = None, 

1127 logical_dates: Annotated[list[UtcDateTime] | None, Query()] = None, 

1128 run_ids: Annotated[list[str] | None, Query()] = None, 

1129 states: Annotated[list[str] | None, Query()] = None, 

1130) -> int: 

1131 """Get the count of task instances matching the given criteria.""" 

1132 query = select(func.count()).select_from(TI).where(TI.dag_id == dag_id) 

1133 

1134 if task_ids: 1134 ↛ 1137line 1134 didn't jump to line 1137 because the condition on line 1134 was always true

1135 query = query.where(TI.task_id.in_(task_ids)) 

1136 

1137 if map_index is not None: 1137 ↛ 1138line 1137 didn't jump to line 1138 because the condition on line 1137 was never true

1138 query = query.where(TI.map_index == map_index) 

1139 

1140 if logical_dates: 1140 ↛ 1143line 1140 didn't jump to line 1143 because the condition on line 1140 was always true

1141 query = query.where(TI.logical_date.in_(logical_dates)) 

1142 

1143 if run_ids: 1143 ↛ 1144line 1143 didn't jump to line 1144 because the condition on line 1143 was never true

1144 query = query.where(TI.run_id.in_(run_ids)) 

1145 

1146 if task_group_id: 1146 ↛ 1147line 1146 didn't jump to line 1147 because the condition on line 1146 was never true

1147 group_tasks = _get_group_tasks( 

1148 dag_id, task_group_id, session, dag_bag, logical_dates, run_ids, map_index 

1149 ) 

1150 

1151 # Get unique (task_id, map_index) pairs 

1152 task_map_pairs = [(ti.task_id, ti.map_index) for ti in group_tasks] 

1153 

1154 if not task_map_pairs: 

1155 # If no task group tasks found, default to checking the task group ID itself 

1156 # This matches the behavior in _get_external_task_group_task_ids 

1157 task_map_pairs = [(task_group_id, -1)] 

1158 

1159 # Update query to use task_id, map_index pairs 

1160 query = query.where(tuple_(TI.task_id, TI.map_index).in_(task_map_pairs)) 

1161 

1162 if states: 1162 ↛ 1172line 1162 didn't jump to line 1172 because the condition on line 1162 was always true

1163 if "null" in states: 1163 ↛ 1164line 1163 didn't jump to line 1164 because the condition on line 1163 was never true

1164 not_none_states = [s for s in states if s != "null"] 

1165 if not_none_states: 

1166 query = query.where(or_(TI.state.is_(None), TI.state.in_(not_none_states))) 

1167 else: 

1168 query = query.where(TI.state.is_(None)) 

1169 else: 

1170 query = query.where(TI.state.in_(states)) 

1171 

1172 count = session.scalar(query) 

1173 return count or 0 

1174 

1175 

1176@router.get("/previous/{dag_id}/{task_id}", status_code=status.HTTP_200_OK) 

1177def get_previous_task_instance( 

1178 dag_id: str, 

1179 task_id: str, 

1180 session: SessionDep, 

1181 logical_date: Annotated[UtcDateTime | None, Query()] = None, 

1182 map_index: Annotated[int, Query()] = -1, 

1183 state: Annotated[TaskInstanceState | None, Query()] = None, 

1184) -> PreviousTIResponse | None: 

1185 """ 

1186 Get the previous task instance matching the given criteria. 

1187 

1188 :param dag_id: DAG ID (from path) 

1189 :param task_id: Task ID (from path) 

1190 :param logical_date: If provided, finds TI with logical_date < this value (before filter) 

1191 :param map_index: Map index to filter by (defaults to -1 for non-mapped tasks) 

1192 :param state: If provided, filters by TaskInstance state 

1193 """ 

1194 query = ( 

1195 select(TI) 

1196 .join(DR, (TI.dag_id == DR.dag_id) & (TI.run_id == DR.run_id)) 

1197 .options(contains_eager(TI.dag_run).load_only(DR.logical_date)) 

1198 .where(TI.dag_id == dag_id, TI.task_id == task_id, TI.map_index == map_index) 

1199 .order_by(DR.logical_date.desc()) 

1200 ) 

1201 

1202 if logical_date: 

1203 # Find TI with logical_date BEFORE the provided date (previous) 

1204 query = query.where(DR.logical_date < logical_date) 

1205 

1206 if state: 

1207 query = query.where(TI.state == state) 

1208 

1209 ti = session.scalars(query.limit(1)).first() 

1210 

1211 if not ti: 

1212 return None 

1213 

1214 return PreviousTIResponse( 

1215 task_id=ti.task_id, 

1216 dag_id=ti.dag_id, 

1217 run_id=ti.run_id, 

1218 logical_date=ti.dag_run.logical_date, 

1219 start_date=ti.start_date, 

1220 end_date=ti.end_date, 

1221 state=ti.state, 

1222 try_number=ti.try_number, 

1223 map_index=ti.map_index, 

1224 duration=ti.duration, 

1225 ) 

1226 

1227 

1228@router.get("/states", status_code=status.HTTP_200_OK) 

1229def get_task_instance_states( 

1230 dag_id: str, 

1231 session: SessionDep, 

1232 dag_bag: DagBagDep, 

1233 map_index: Annotated[int | None, Query()] = None, 

1234 task_ids: Annotated[list[str] | None, Query()] = None, 

1235 task_group_id: Annotated[str | None, Query()] = None, 

1236 logical_dates: Annotated[list[UtcDateTime] | None, Query()] = None, 

1237 run_ids: Annotated[list[str] | None, Query()] = None, 

1238) -> TaskStatesResponse: 

1239 """Get the states for Task Instances with the given criteria.""" 

1240 run_id_task_state_map: dict[str, dict[str, Any]] = defaultdict(dict) 

1241 

1242 query = select(TI).where(TI.dag_id == dag_id) 

1243 

1244 if task_ids: 

1245 query = query.where(TI.task_id.in_(task_ids)) 

1246 

1247 if logical_dates: 

1248 query = query.where(TI.logical_date.in_(logical_dates)) 

1249 

1250 if run_ids: 

1251 query = query.where(TI.run_id.in_(run_ids)) 

1252 

1253 if map_index is not None: 

1254 query = query.where(TI.map_index == map_index) 

1255 

1256 results = session.scalars(query).all() 

1257 

1258 if task_group_id: 

1259 group_tasks = _get_group_tasks( 

1260 dag_id, task_group_id, session, dag_bag, logical_dates, run_ids, map_index 

1261 ) 

1262 

1263 results = results + group_tasks if task_ids else group_tasks 

1264 

1265 [ 

1266 run_id_task_state_map[task.run_id].update( 

1267 {task.task_id: task.state} 

1268 if task.map_index < 0 

1269 else {f"{task.task_id}_{task.map_index}": task.state} 

1270 ) 

1271 for task in results 

1272 ] 

1273 

1274 return TaskStatesResponse(task_states=run_id_task_state_map) 

1275 

1276 

1277@router.get("/breadcrumbs", status_code=status.HTTP_200_OK) 

1278def get_task_instance_breadcrumbs(dag_id: str, run_id: str, session: SessionDep) -> TaskBreadcrumbsResponse: 

1279 result = session.execute( 

1280 select(TI.task_id, TI.map_index, TI.state, TI.operator, TI.duration) 

1281 .where(TI.dag_id == dag_id, TI.run_id == run_id, TI.state.in_(TerminalTIState)) 

1282 .order_by(TI.task_id, TI.map_index) 

1283 ).mappings() 

1284 

1285 def _iter_breadcrumbs() -> Iterator[dict[str, Any]]: 

1286 for row in result: 

1287 yield {str(k): v for k, v in row.items()} 

1288 

1289 return TaskBreadcrumbsResponse(breadcrumbs=_iter_breadcrumbs()) 

1290 

1291 

1292def _is_eligible_to_retry(state: str, try_number: int, max_tries: int) -> bool: 

1293 """Is task instance is eligible for retry.""" 

1294 if state == TaskInstanceState.RESTARTING: 1294 ↛ 1297line 1294 didn't jump to line 1297 because the condition on line 1294 was never true

1295 # If a task is cleared when running, it goes into RESTARTING state and is always 

1296 # eligible for retry 

1297 return True 

1298 

1299 # max_tries is initialised with the retries defined at task level, we do not need to explicitly ask for 

1300 # retries from the task SDK now, we can handle using max_tries 

1301 return max_tries != 0 and try_number <= max_tries 

1302 

1303 

1304def _get_group_tasks( 

1305 dag_id: str, 

1306 task_group_id: str, 

1307 session: SessionDep, 

1308 dag_bag: DagBagDep, 

1309 logical_dates=None, 

1310 run_ids=None, 

1311 map_index: int | None = None, 

1312): 

1313 # Get all tasks in the task group 

1314 dag = get_latest_version_of_dag(dag_bag, dag_id, session, include_reason=True) 

1315 task_group = dag.task_group_dict.get(task_group_id) 

1316 if not task_group: 

1317 raise HTTPException( 

1318 status.HTTP_404_NOT_FOUND, 

1319 detail={ 

1320 "reason": "not_found", 

1321 "message": f"Task group {task_group_id} not found in DAG {dag_id}", 

1322 }, 

1323 ) 

1324 

1325 # First get all task instances to get the task_id, map_index pairs 

1326 group_tasks = session.scalars( 

1327 select(TI).where( 

1328 TI.dag_id == dag_id, 

1329 TI.task_id.in_(task.task_id for task in task_group.iter_tasks()), 

1330 *([TI.logical_date.in_(logical_dates)] if logical_dates else []), 

1331 *([TI.run_id.in_(run_ids)] if run_ids else []), 

1332 *([TI.map_index == map_index] if map_index is not None else []), 

1333 ) 

1334 ).all() 

1335 

1336 return group_tasks 

1337 

1338 

1339@ti_id_router.get( 

1340 "/{task_instance_id}/validate-inlets-and-outlets", 

1341 status_code=status.HTTP_200_OK, 

1342 responses=create_openapi_http_exception_doc( 

1343 [ 

1344 (status.HTTP_404_NOT_FOUND, "Task Instance not found"), 

1345 ] 

1346 ), 

1347) 

1348def validate_inlets_and_outlets( 

1349 task_instance_id: UUID, 

1350 session: SessionDep, 

1351 dag_bag: DagBagDep, 

1352) -> InactiveAssetsResponse: 

1353 """Validate whether there're inactive assets in inlets and outlets of a given task instance.""" 

1354 bind_contextvars(ti_id=str(task_instance_id)) 

1355 

1356 ti = session.scalar(select(TI).where(TI.id == task_instance_id)) 

1357 if not ti: 1357 ↛ 1358line 1357 didn't jump to line 1358 because the condition on line 1357 was never true

1358 log.error("Task Instance not found") 

1359 raise HTTPException( 

1360 status_code=status.HTTP_404_NOT_FOUND, 

1361 detail={ 

1362 "reason": "not_found", 

1363 "message": "Task Instance not found", 

1364 }, 

1365 ) 

1366 

1367 if not ti.task: 1367 ↛ 1374line 1367 didn't jump to line 1374 because the condition on line 1367 was always true

1368 dr = ti.dag_run 

1369 dag = dag_bag.get_dag_for_run(dag_run=dr, session=session) 

1370 if dag: 1370 ↛ 1374line 1370 didn't jump to line 1374 because the condition on line 1370 was always true

1371 with contextlib.suppress(TaskNotFound): 

1372 ti.task = dag.get_task(ti.task_id) 

1373 

1374 inlets = ( 

1375 [asset.asprofile() for asset in ti.task.inlets if isinstance(asset, SerializedAsset)] 

1376 if ti.task 

1377 else [] 

1378 ) 

1379 outlets = ( 

1380 [asset.asprofile() for asset in ti.task.outlets if isinstance(asset, SerializedAsset)] 

1381 if ti.task 

1382 else [] 

1383 ) 

1384 if not (inlets or outlets): 1384 ↛ 1385line 1384 didn't jump to line 1385 because the condition on line 1384 was never true

1385 return InactiveAssetsResponse(inactive_assets=[]) 

1386 

1387 all_asset_unique_keys: set[SerializedAssetUniqueKey] = { 

1388 SerializedAssetUniqueKey.from_asset(inlet_or_outlet) # type: ignore 

1389 for inlet_or_outlet in itertools.chain(inlets, outlets) 

1390 } 

1391 active_asset_unique_keys = { 

1392 SerializedAssetUniqueKey(name, uri) 

1393 for name, uri in session.execute( 

1394 select(AssetActive.name, AssetActive.uri).where( 

1395 tuple_(AssetActive.name, AssetActive.uri).in_( 

1396 attrs.astuple(key) for key in all_asset_unique_keys 

1397 ) 

1398 ) 

1399 ) 

1400 } 

1401 different = all_asset_unique_keys - active_asset_unique_keys 

1402 

1403 return InactiveAssetsResponse( 

1404 inactive_assets=[asset_unique_key.asprofile() for asset_unique_key in different], 

1405 ) 

1406 

1407 

1408# This line should be at the end of the file to ensure all routes are registered 

1409router.include_router(ti_id_router)