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
« 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.
18from __future__ import annotations
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
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
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
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
100router = VersionedAPIRouter()
102ti_id_router = VersionedAPIRouter(
103 route_class=ExecutionAPIRoute,
104 dependencies=[
105 Security(require_auth, scopes=["ti:self"]),
106 ],
107)
110log = structlog.get_logger(__name__)
111tracer = trace.get_tracer(__name__)
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.
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 )
149 from sqlalchemy.sql import column
150 from sqlalchemy.types import JSON
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 )
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)
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")
199 query = update(TI).where(TI.id == task_instance_id).values(data)
201 previous_state = ti.state
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 )
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 )
252 try:
253 result = session.execute(query)
254 log.info("Task instance state updated", rows_affected=getattr(result, "rowcount", 0))
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 )
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 )
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.
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)
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 )
301 dr.team_name = get_team_name_for_ti(task_instance_id, session)
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 )
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
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
332 return context
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.
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)
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 )
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)
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 )
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
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)
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)
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
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 )
535 # Release the task_instance row lock before running listener callbacks.
536 session.commit()
538 for callback in asset_callbacks:
539 callback()
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)
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
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 )
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()
592def _handle_fail_fast_for_dag(ti: TI, dag_id: str, session: SessionDep, dag_bag: DagBagDep) -> None:
593 dr = ti.dag_run
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
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)
609def _validate_outlet_event_partition_keys(outlet_events: list[dict[str, Any]]) -> None:
610 """
611 Validate partition_key values embedded in outlet events.
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 )
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)
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
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)
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()
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
705 query = update(TI).where(TI.id == task_instance_id)
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
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, ()
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 )
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)}")
803 return query, updated_state, asset_callbacks
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))
824 now = timezone.utcnow()
825 tasks = ti_patch_payload.tasks
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)
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)
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 )
865 result = session.execute(query)
866 log.info("Downstream tasks skipped", tasks_skipped=getattr(result, "rowcount", 0))
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 )
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)
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
939 log.debug("Heartbeat fast path missed; falling back to diagnostic checks")
941 old = select(TI.state, TI.hostname, TI.pid).where(TI.id == task_instance_id).with_for_update()
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)
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 )
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 )
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)
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))
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")
1026 return {"message": "Rendered task instance fields successfully set"}
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))
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 )
1057 log.debug("Updating rendered_map_index", length=len(rendered_map_index))
1059 query = update(TI).where(TI.id == task_instance_id).values(_rendered_map_index=rendered_map_index)
1060 result = session.execute(query)
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 )
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.
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")
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()
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()
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)
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)
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))
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)
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))
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))
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 )
1151 # Get unique (task_id, map_index) pairs
1152 task_map_pairs = [(ti.task_id, ti.map_index) for ti in group_tasks]
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)]
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))
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))
1172 count = session.scalar(query)
1173 return count or 0
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.
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 )
1202 if logical_date:
1203 # Find TI with logical_date BEFORE the provided date (previous)
1204 query = query.where(DR.logical_date < logical_date)
1206 if state:
1207 query = query.where(TI.state == state)
1209 ti = session.scalars(query.limit(1)).first()
1211 if not ti:
1212 return None
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 )
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)
1242 query = select(TI).where(TI.dag_id == dag_id)
1244 if task_ids:
1245 query = query.where(TI.task_id.in_(task_ids))
1247 if logical_dates:
1248 query = query.where(TI.logical_date.in_(logical_dates))
1250 if run_ids:
1251 query = query.where(TI.run_id.in_(run_ids))
1253 if map_index is not None:
1254 query = query.where(TI.map_index == map_index)
1256 results = session.scalars(query).all()
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 )
1263 results = results + group_tasks if task_ids else group_tasks
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 ]
1274 return TaskStatesResponse(task_states=run_id_task_state_map)
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()
1285 def _iter_breadcrumbs() -> Iterator[dict[str, Any]]:
1286 for row in result:
1287 yield {str(k): v for k, v in row.items()}
1289 return TaskBreadcrumbsResponse(breadcrumbs=_iter_breadcrumbs())
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
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
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 )
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()
1336 return group_tasks
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))
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 )
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)
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=[])
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
1403 return InactiveAssetsResponse(
1404 inactive_assets=[asset_unique_key.asprofile() for asset_unique_key in different],
1405 )
1408# This line should be at the end of the file to ensure all routes are registered
1409router.include_router(ti_id_router)