Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/core_api/routes/public/hitl.py: 58%

95 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 14:22 +0000

1# Licensed to the Apache Software Foundation (ASF) under one 

2# or more contributor license agreements. See the NOTICE file 

3# distributed with this work for additional information 

4# regarding copyright ownership. The ASF licenses this file 

5# to you under the Apache License, Version 2.0 (the 

6# "License"); you may not use this file except in compliance 

7# with the License. You may obtain a copy of the License at 

8# 

9# http://www.apache.org/licenses/LICENSE-2.0 

10# 

11# Unless required by applicable law or agreed to in writing, 

12# software distributed under the License is distributed on an 

13# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY 

14# KIND, either express or implied. See the License for the 

15# specific language governing permissions and limitations 

16# under the License. 

17from __future__ import annotations 

18 

19from typing import Annotated 

20 

21import structlog 

22from fastapi import Depends, HTTPException, status 

23from sqlalchemy import select 

24from sqlalchemy.orm import joinedload 

25 

26from airflow._shared.timezones import timezone 

27from airflow.api_fastapi.auth.managers.models.resource_details import DagAccessEntity 

28from airflow.api_fastapi.common.db.common import SessionDep, paginated_select 

29from airflow.api_fastapi.common.parameters import ( 

30 QueryHITLDetailBodySearch, 

31 QueryHITLDetailDagIdPatternSearch, 

32 QueryHITLDetailDagIdPrefixPatternSearch, 

33 QueryHITLDetailMapIndexFilter, 

34 QueryHITLDetailRespondedUserIdFilter, 

35 QueryHITLDetailRespondedUserNameFilter, 

36 QueryHITLDetailResponseReceivedFilter, 

37 QueryHITLDetailSubjectSearch, 

38 QueryHITLDetailTaskIdFilter, 

39 QueryHITLDetailTaskIdPatternSearch, 

40 QueryHITLDetailTaskIdPrefixPatternSearch, 

41 QueryLimit, 

42 QueryOffset, 

43 QueryTIStateFilter, 

44 RangeFilter, 

45 SortParam, 

46 datetime_range_filter_factory, 

47) 

48from airflow.api_fastapi.common.router import AirflowRouter 

49from airflow.api_fastapi.core_api.datamodels.hitl import ( 

50 HITLDetail, 

51 HITLDetailCollection, 

52 HITLDetailHistory, 

53 HITLDetailResponse, 

54 UpdateHITLDetailPayload, 

55) 

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

57from airflow.api_fastapi.core_api.security import ( 

58 GetUserDep, 

59 ReadableTIFilterDep, 

60 get_auth_manager, 

61 requires_access_dag, 

62) 

63from airflow.api_fastapi.logging.decorators import action_logging 

64from airflow.models.base import Base 

65from airflow.models.dag_version import DagVersion 

66from airflow.models.dagrun import DagRun 

67from airflow.models.hitl import HITLDetail as HITLDetailModel, HITLUser 

68from airflow.models.taskinstance import TaskInstance as TI 

69from airflow.models.taskinstancehistory import TaskInstanceHistory as TIH 

70from airflow.models.trigger import handle_event_submit 

71from airflow.triggers.base import TriggerEvent 

72from airflow.utils.state import TaskInstanceState 

73 

74task_instances_hitl_router = AirflowRouter( 

75 tags=["Task Instance"], 

76 prefix="/dags/{dag_id}/dagRuns/{dag_run_id}", 

77) 

78task_instance_hitl_path = "/taskInstances/{task_id}/{map_index}/hitlDetails" 

79 

80log = structlog.get_logger(__name__) 

81 

82 

83def _get_task_instance_with_hitl_detail( 

84 dag_id: str, 

85 dag_run_id: str, 

86 task_id: str, 

87 session: SessionDep, 

88 map_index: int, 

89 try_number: int | None = None, 

90) -> TI | TIH: 

91 def _query(orm_object: Base) -> TI | TIH | None: 

92 options = [joinedload(orm_object.hitl_detail)] 

93 if orm_object is TI: 

94 options.append(joinedload(TI.rendered_task_instance_fields)) 

95 query = ( 

96 select(orm_object) 

97 .where( 

98 orm_object.dag_id == dag_id, 

99 orm_object.run_id == dag_run_id, 

100 orm_object.task_id == task_id, 

101 orm_object.map_index == map_index, 

102 ) 

103 .options(*options) 

104 ) 

105 

106 if try_number is not None: 

107 query = query.where(orm_object.try_number == try_number) 

108 

109 ti_or_tih = session.scalar(query) 

110 return ti_or_tih 

111 

112 if try_number is None: 

113 ti_or_tih = _query(TI) 

114 else: 

115 ti_or_tih = _query(TIH) or _query(TI) 

116 

117 if ti_or_tih is None: 117 ↛ 126line 117 didn't jump to line 126 because the condition on line 117 was always true

118 raise HTTPException( 

119 status_code=status.HTTP_404_NOT_FOUND, 

120 detail=( 

121 f"The Task Instance with dag_id: `{dag_id}`, run_id: `{dag_run_id}`, " 

122 f"task_id: `{task_id}` and map_index: `{map_index}` was not found" 

123 ), 

124 ) 

125 

126 if not ti_or_tih.hitl_detail: 

127 raise HTTPException( 

128 status_code=status.HTTP_404_NOT_FOUND, 

129 detail=f"Human-in-the-loop detail does not exist for Task Instance with id {ti_or_tih.id}", 

130 ) 

131 

132 return ti_or_tih 

133 

134 

135@task_instances_hitl_router.patch( 

136 task_instance_hitl_path, 

137 responses=create_openapi_http_exception_doc( 

138 [ 

139 status.HTTP_400_BAD_REQUEST, 

140 status.HTTP_403_FORBIDDEN, 

141 status.HTTP_404_NOT_FOUND, 

142 status.HTTP_409_CONFLICT, 

143 ] 

144 ), 

145 dependencies=[ 

146 Depends(requires_access_dag(method="PUT", access_entity=DagAccessEntity.HITL_DETAIL)), 

147 Depends(action_logging()), 

148 ], 

149) 

150def update_hitl_detail( 

151 dag_id: str, 

152 dag_run_id: str, 

153 task_id: str, 

154 update_hitl_detail_payload: UpdateHITLDetailPayload, 

155 user: GetUserDep, 

156 session: SessionDep, 

157 map_index: int = -1, 

158) -> HITLDetailResponse: 

159 """Update a Human-in-the-loop detail.""" 

160 task_instance = _get_task_instance_with_hitl_detail( 

161 dag_id=dag_id, 

162 dag_run_id=dag_run_id, 

163 task_id=task_id, 

164 session=session, 

165 map_index=map_index, 

166 ) 

167 

168 # Acquire row locks in a fixed order -- TaskInstance first, then the HITL row -- matching the 

169 # Execution API park transition, so a human response racing the worker's park cannot deadlock. 

170 # Locking the TI also serializes respond-vs-clear (the clear path locks the TI, not the HITL row). 

171 locked_ti = ( 

172 session.get(TI, task_instance.id, with_for_update={"of": TI}) 

173 if isinstance(task_instance, TI) 

174 else None 

175 ) 

176 # Lock the hitl_detail row (FOR UPDATE OF hitl_detail). of= scopes the lock to hitl_detail, which 

177 # eager-joins task_instance (lazy="joined"); a bare with_for_update() would emit FOR UPDATE against 

178 # the nullable side of that outer join, which Postgres rejects. The joinedloaded relationship object 

179 # reused below is the same identity-mapped row, now locked for this transaction. 

180 session.execute( 

181 select(HITLDetailModel) 

182 .where(HITLDetailModel.ti_id == task_instance.id) 

183 .with_for_update(of=HITLDetailModel) 

184 ) 

185 hitl_detail_model = task_instance.hitl_detail 

186 if hitl_detail_model.response_received: 

187 raise HTTPException( 

188 status_code=status.HTTP_409_CONFLICT, 

189 detail=( 

190 f"Human-in-the-loop detail has already been updated for Task Instance with id {task_instance.id} " 

191 "and is not allowed to write again." 

192 ), 

193 ) 

194 

195 user_id = user.get_id() 

196 user_name = user.get_name() 

197 if isinstance(user_id, int): 

198 # FabAuthManager (ab_user) store user id as integer, but common interface is string type 

199 user_id = str(user_id) 

200 hitl_user = HITLUser(id=user_id, name=user_name) 

201 if hitl_detail_model.assigned_users: 

202 # Convert assigned_users list to set of user IDs for authorization check 

203 assigned_user_ids = {assigned_user["id"] for assigned_user in hitl_detail_model.assigned_users} 

204 if not get_auth_manager().is_authorized_hitl_task(assigned_users=assigned_user_ids, user=user): 

205 log.error("User=%s (id=%s) is not a respondent for the task", user_name, user_id) 

206 raise HTTPException( 

207 status.HTTP_403_FORBIDDEN, 

208 f"User={user_name} (id={user_id}) is not a respondent for the task.", 

209 ) 

210 

211 # Write-side validation: reject an invalid response here (400) instead of accepting it and 

212 # failing the task later on resume. Mirrors HITLOperator.validate_chosen_options + cardinality. 

213 allowed_options = set(hitl_detail_model.options or []) 

214 invalid_options = set(update_hitl_detail_payload.chosen_options) - allowed_options 

215 if invalid_options: 

216 raise HTTPException( 

217 status.HTTP_400_BAD_REQUEST, 

218 f"Invalid options {sorted(invalid_options)}; allowed options are {sorted(allowed_options)}.", 

219 ) 

220 if not hitl_detail_model.multiple and len(update_hitl_detail_payload.chosen_options) > 1: 

221 raise HTTPException( 

222 status.HTTP_400_BAD_REQUEST, 

223 "Multiple options chosen but this Human-in-the-loop task accepts only a single option.", 

224 ) 

225 

226 hitl_detail_model.responded_by = hitl_user 

227 hitl_detail_model.responded_at = timezone.utcnow() 

228 hitl_detail_model.chosen_options = update_hitl_detail_payload.chosen_options 

229 hitl_detail_model.params_input = update_hitl_detail_payload.params_input 

230 session.add(hitl_detail_model) 

231 

232 # Event-driven resume: if the task is parked waiting for this input, transition it directly, 

233 # without a trigger. handle_event_submit packs the response into next_kwargs["event"], sets 

234 # state=SCHEDULED + scheduled_dttm; the scheduler then re-queues execute_complete. Gated on the 

235 # parked states so a finished/cleared TI is never resurrected. `locked_ti` was locked at the top 

236 # (TI-before-HITLDetail order), so a concurrent clear cannot interleave between this state check 

237 # and the resume and have its reset silently overwritten. 

238 if locked_ti is not None and locked_ti.state in ( 

239 TaskInstanceState.AWAITING_INPUT, 

240 TaskInstanceState.DEFERRED, 

241 ): 

242 handle_event_submit( 

243 TriggerEvent(hitl_detail_model.as_resume_event_payload()), 

244 task_instance=locked_ti, 

245 session=session, 

246 ) 

247 return HITLDetailResponse.model_validate(hitl_detail_model) 

248 

249 

250@task_instances_hitl_router.get( 

251 task_instance_hitl_path, 

252 status_code=status.HTTP_200_OK, 

253 responses=create_openapi_http_exception_doc([status.HTTP_404_NOT_FOUND]), 

254 dependencies=[Depends(requires_access_dag(method="GET", access_entity=DagAccessEntity.HITL_DETAIL))], 

255) 

256def get_hitl_detail( 

257 dag_id: str, 

258 dag_run_id: str, 

259 task_id: str, 

260 session: SessionDep, 

261 map_index: int = -1, 

262) -> HITLDetail: 

263 """Get a Human-in-the-loop detail of a specific task instance.""" 

264 task_instance = _get_task_instance_with_hitl_detail( 

265 dag_id=dag_id, 

266 dag_run_id=dag_run_id, 

267 task_id=task_id, 

268 session=session, 

269 map_index=map_index, 

270 try_number=None, 

271 ) 

272 return HITLDetail.model_validate(task_instance.hitl_detail) 

273 

274 

275@task_instances_hitl_router.get( 

276 task_instance_hitl_path + "/tries/{try_number}", 

277 status_code=status.HTTP_200_OK, 

278 responses=create_openapi_http_exception_doc([status.HTTP_404_NOT_FOUND]), 

279 dependencies=[Depends(requires_access_dag(method="GET", access_entity=DagAccessEntity.HITL_DETAIL))], 

280) 

281def get_hitl_detail_try_detail( 

282 dag_id: str, 

283 dag_run_id: str, 

284 task_id: str, 

285 session: SessionDep, 

286 map_index: int = -1, 

287 try_number: int | None = None, 

288) -> HITLDetailHistory: 

289 """Get a Human-in-the-loop detail of a specific task instance.""" 

290 task_instance_history = _get_task_instance_with_hitl_detail( 

291 dag_id=dag_id, 

292 dag_run_id=dag_run_id, 

293 task_id=task_id, 

294 session=session, 

295 map_index=map_index, 

296 try_number=try_number, 

297 ) 

298 return task_instance_history.hitl_detail 

299 

300 

301@task_instances_hitl_router.get( 

302 "/hitlDetails", 

303 status_code=status.HTTP_200_OK, 

304 dependencies=[Depends(requires_access_dag(method="GET", access_entity=DagAccessEntity.HITL_DETAIL))], 

305) 

306def get_hitl_details( 

307 dag_id: str, 

308 dag_run_id: str, 

309 limit: QueryLimit, 

310 offset: QueryOffset, 

311 order_by: Annotated[ 

312 SortParam, 

313 Depends( 

314 SortParam( 

315 allowed_attrs=[ 

316 "ti_id", 

317 "subject", 

318 "responded_at", 

319 "created_at", 

320 "responded_by_user_id", 

321 "responded_by_user_name", 

322 ], 

323 model=HITLDetailModel, 

324 to_replace={ 

325 "dag_id": TI.dag_id, 

326 "run_id": TI.run_id, 

327 "task_display_name": TI.task_display_name, 

328 "run_after": DagRun.run_after, 

329 "rendered_map_index": TI.rendered_map_index, 

330 "task_instance_operator": TI.operator, 

331 "task_instance_state": TI.state, 

332 }, 

333 ).dynamic_depends(), 

334 ), 

335 ], 

336 session: SessionDep, 

337 # permission filter 

338 readable_ti_filter: ReadableTIFilterDep, 

339 # ti related filter 

340 dag_id_pattern: QueryHITLDetailDagIdPatternSearch, 

341 dag_id_prefix_pattern: QueryHITLDetailDagIdPrefixPatternSearch, 

342 task_id: QueryHITLDetailTaskIdFilter, 

343 task_id_pattern: QueryHITLDetailTaskIdPatternSearch, 

344 task_id_prefix_pattern: QueryHITLDetailTaskIdPrefixPatternSearch, 

345 map_index: QueryHITLDetailMapIndexFilter, 

346 ti_state: QueryTIStateFilter, 

347 # hitl detail related filter 

348 response_received: QueryHITLDetailResponseReceivedFilter, 

349 responded_by_user_id: QueryHITLDetailRespondedUserIdFilter, 

350 responded_by_user_name: QueryHITLDetailRespondedUserNameFilter, 

351 subject_patten: QueryHITLDetailSubjectSearch, 

352 body_patten: QueryHITLDetailBodySearch, 

353 created_at: Annotated[RangeFilter, Depends(datetime_range_filter_factory("created_at", HITLDetailModel))], 

354) -> HITLDetailCollection: 

355 """Get Human-in-the-loop details.""" 

356 query = ( 

357 select(HITLDetailModel) 

358 .join(TI, HITLDetailModel.ti_id == TI.id) 

359 .join(TI.dag_run) 

360 .options( 

361 joinedload(HITLDetailModel.task_instance).options( 

362 joinedload(TI.dag_run).joinedload(DagRun.dag_model), 

363 joinedload(TI.task_instance_note), 

364 joinedload(TI.dag_version).joinedload(DagVersion.bundle), 

365 joinedload(TI.rendered_task_instance_fields), 

366 ), 

367 ) 

368 ) 

369 if dag_id != "~": 

370 query = query.where(TI.dag_id == dag_id) 

371 if dag_run_id != "~": 371 ↛ 373line 371 didn't jump to line 373 because the condition on line 371 was always true

372 query = query.where(TI.run_id == dag_run_id) 

373 hitl_detail_select, total_entries = paginated_select( 

374 statement=query, 

375 filters=[ 

376 # permission filter 

377 readable_ti_filter, 

378 # ti related filter 

379 dag_id_pattern, 

380 dag_id_prefix_pattern, 

381 task_id, 

382 task_id_pattern, 

383 task_id_prefix_pattern, 

384 map_index, 

385 ti_state, 

386 # hitl detail related filter 

387 response_received, 

388 responded_by_user_id, 

389 responded_by_user_name, 

390 subject_patten, 

391 body_patten, 

392 created_at, 

393 ], 

394 offset=offset, 

395 limit=limit, 

396 order_by=order_by, 

397 session=session, 

398 ) 

399 

400 hitl_details = session.scalars(hitl_detail_select) 

401 

402 return HITLDetailCollection( 

403 hitl_details=hitl_details, 

404 total_entries=total_entries, 

405 )