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

78 statements  

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

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

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

3# distributed with this work for additional information 

4# regarding copyright ownership. The ASF licenses this file 

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

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

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

8# 

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

10# 

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

12# software distributed under the License is distributed on an 

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

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

15# specific language governing permissions and limitations 

16# under the License. 

17from __future__ import annotations 

18 

19import json 

20from datetime import datetime, timedelta, timezone 

21from typing import TYPE_CHECKING, Annotated, Literal 

22 

23from fastapi import Depends, HTTPException, Query, status 

24from sqlalchemy import select 

25 

26from airflow._shared.state import TaskScope 

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 QueryLimit, QueryOffset 

30from airflow.api_fastapi.common.router import AirflowRouter 

31from airflow.api_fastapi.core_api.datamodels.task_state_store import ( 

32 TaskStateStoreBody, 

33 TaskStateStoreCollectionResponse, 

34 TaskStateStorePatchBody, 

35 TaskStateStoreResponse, 

36) 

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

38from airflow.api_fastapi.core_api.security import requires_access_dag 

39from airflow.configuration import conf 

40from airflow.models.task_state_store import TaskStateStoreModel 

41from airflow.models.taskinstance import TaskInstance as TI 

42from airflow.state.metastore import _get_db_backend 

43 

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

45 from sqlalchemy.orm import Session 

46 

47task_state_store_router = AirflowRouter( 

48 tags=["Task State Store"], 

49 prefix="/dags/{dag_id}/dagRuns/{dag_run_id}/taskInstances/{task_id}/state-store", 

50) 

51 

52 

53def _require_task_instance( 

54 dag_id: str, 

55 dag_run_id: str, 

56 task_id: str, 

57 map_index: int | None, 

58 session: Session, 

59) -> None: 

60 """Raise 404 unless the task instance exists. ``map_index=None`` matches any map index.""" 

61 statement = select(TI.task_id).where( 

62 TI.dag_id == dag_id, 

63 TI.run_id == dag_run_id, 

64 TI.task_id == task_id, 

65 ) 

66 if map_index is not None: 

67 statement = statement.where(TI.map_index == map_index) 

68 if session.scalar(statement.limit(1)) is None: 

69 addressed_by = "all_map_indices=True" if map_index is None else f"map_index={map_index}" 

70 raise HTTPException( 

71 status_code=status.HTTP_404_NOT_FOUND, 

72 detail=( 

73 f"Task instance not found for dag_id={dag_id!r}, run_id={dag_run_id!r}, " 

74 f"task_id={task_id!r}, {addressed_by}" 

75 ), 

76 ) 

77 

78 

79def _resolve_scope( 

80 dag_id: str, 

81 dag_run_id: str, 

82 task_id: str, 

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

84) -> TaskScope: 

85 """Map the path and query parameters onto the task instance they address.""" 

86 return TaskScope(dag_id=dag_id, run_id=dag_run_id, task_id=task_id, map_index=map_index) 

87 

88 

89TaskScopeDep = Annotated[TaskScope, Depends(_resolve_scope)] 

90 

91 

92def _validate_scope(scope: TaskScopeDep, session: SessionDep) -> TaskScope: 

93 """Resolve the scope, 404ing when the task instance it addresses does not exist.""" 

94 _require_task_instance(scope.dag_id, scope.run_id, scope.task_id, scope.map_index, session) 

95 return scope 

96 

97 

98ValidatedTaskScopeDep = Annotated[TaskScope, Depends(_validate_scope)] 

99 

100 

101def _validate_clear_scope( 

102 scope: TaskScopeDep, 

103 session: SessionDep, 

104 all_map_indices: Annotated[bool, Query()] = False, 

105) -> TaskScope: 

106 """ 

107 Resolve the scope for a clear request, 404ing when the task instance does not exist. 

108 

109 ``all_map_indices`` addresses the task across every index, so it is validated against any 

110 instance -- an expanded mapped task has no ``map_index=-1`` instance to check. 

111 """ 

112 _require_task_instance( 

113 scope.dag_id, scope.run_id, scope.task_id, None if all_map_indices else scope.map_index, session 

114 ) 

115 return scope 

116 

117 

118ValidatedClearTaskScopeDep = Annotated[TaskScope, Depends(_validate_clear_scope)] 

119 

120 

121def _resolve_expires_at(expires_at: datetime | None | Literal["default"]) -> datetime | None: 

122 """ 

123 Resolve the expires_at value from the request body. 

124 

125 - ``"default"``: apply configured ``[state_store] default_retention_days``. 

126 ``0`` means never expire. Negative values raise HTTP 400. 

127 - ``None``: never expire 

128 - datetime: use as-is 

129 """ 

130 if expires_at == "default": 130 ↛ 139line 130 didn't jump to line 139 because the condition on line 130 was always true

131 days = conf.getint("state_store", "default_retention_days") 

132 if days < 0: 132 ↛ 133line 132 didn't jump to line 133 because the condition on line 132 was never true

133 raise HTTPException( 

134 status_code=status.HTTP_400_BAD_REQUEST, 

135 detail=f"[state_store] default_retention_days must be >= 0, got {days}. " 

136 "Set to 0 to disable expiry.", 

137 ) 

138 return None if days == 0 else datetime.now(tz=timezone.utc) + timedelta(days=days) 

139 return expires_at 

140 

141 

142@task_state_store_router.get( 

143 "", 

144 dependencies=[Depends(requires_access_dag(method="GET", access_entity=DagAccessEntity.TASK_INSTANCE))], 

145) 

146def list_task_state_store( 

147 scope: TaskScopeDep, 

148 limit: QueryLimit, 

149 offset: QueryOffset, 

150 session: SessionDep, 

151) -> TaskStateStoreCollectionResponse: 

152 """List all task state store entries for a task instance.""" 

153 base = ( 

154 select( 

155 TaskStateStoreModel.key, 

156 TaskStateStoreModel.value, 

157 TaskStateStoreModel.updated_at, 

158 TaskStateStoreModel.expires_at, 

159 ) 

160 .where( 

161 TaskStateStoreModel.dag_id == scope.dag_id, 

162 TaskStateStoreModel.run_id == scope.run_id, 

163 TaskStateStoreModel.task_id == scope.task_id, 

164 TaskStateStoreModel.map_index == scope.map_index, 

165 ) 

166 .order_by(TaskStateStoreModel.key.asc()) 

167 ) 

168 paginated, total_entries = paginated_select( 

169 statement=base, 

170 filters=None, 

171 order_by=None, 

172 offset=offset, 

173 limit=limit, 

174 session=session, 

175 ) 

176 rows = session.execute(paginated).all() 

177 entries = [ 

178 TaskStateStoreResponse( 

179 key=r.key, value=json.loads(r.value), updated_at=r.updated_at, expires_at=r.expires_at 

180 ) 

181 for r in rows 

182 ] 

183 return TaskStateStoreCollectionResponse(task_state_store=entries, total_entries=total_entries) 

184 

185 

186@task_state_store_router.get( 

187 "/{key:path}", 

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

189 dependencies=[Depends(requires_access_dag(method="GET", access_entity=DagAccessEntity.TASK_INSTANCE))], 

190) 

191def get_task_state_store( 

192 scope: TaskScopeDep, 

193 key: str, 

194 session: SessionDep, 

195) -> TaskStateStoreResponse: 

196 """Get a single task state store entry.""" 

197 row = session.execute( 

198 select( 

199 TaskStateStoreModel.key, 

200 TaskStateStoreModel.value, 

201 TaskStateStoreModel.updated_at, 

202 TaskStateStoreModel.expires_at, 

203 ).where( 

204 TaskStateStoreModel.dag_id == scope.dag_id, 

205 TaskStateStoreModel.run_id == scope.run_id, 

206 TaskStateStoreModel.task_id == scope.task_id, 

207 TaskStateStoreModel.map_index == scope.map_index, 

208 TaskStateStoreModel.key == key, 

209 ) 

210 ).one_or_none() 

211 if row is None: 

212 raise HTTPException( 

213 status_code=status.HTTP_404_NOT_FOUND, 

214 detail=f"Task state store key {key!r} not found", 

215 ) 

216 return TaskStateStoreResponse( 

217 key=row.key, value=json.loads(row.value), updated_at=row.updated_at, expires_at=row.expires_at 

218 ) 

219 

220 

221@task_state_store_router.put( 

222 "/{key:path}", 

223 status_code=status.HTTP_204_NO_CONTENT, 

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

225 dependencies=[Depends(requires_access_dag(method="PUT", access_entity=DagAccessEntity.TASK_INSTANCE))], 

226) 

227def set_task_state_store( 

228 scope: ValidatedTaskScopeDep, 

229 key: str, 

230 body: TaskStateStoreBody, 

231 session: SessionDep, 

232) -> None: 

233 """Set a task state store value. Creates or overwrites the key.""" 

234 expires_at = _resolve_expires_at(body.expires_at) 

235 try: 

236 _get_db_backend().set(scope, key, json.dumps(body.value), expires_at=expires_at, session=session) 

237 except ValueError as e: 

238 raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(e)) from e 

239 

240 

241@task_state_store_router.patch( 

242 "/{key:path}", 

243 status_code=status.HTTP_200_OK, 

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

245 dependencies=[Depends(requires_access_dag(method="PUT", access_entity=DagAccessEntity.TASK_INSTANCE))], 

246) 

247def patch_task_state_store( 

248 scope: ValidatedTaskScopeDep, 

249 key: str, 

250 body: TaskStateStorePatchBody, 

251 session: SessionDep, 

252) -> None: 

253 """Update the value of an existing task state store key.""" 

254 existing = session.execute( 

255 select(TaskStateStoreModel.expires_at).where( 

256 TaskStateStoreModel.dag_id == scope.dag_id, 

257 TaskStateStoreModel.run_id == scope.run_id, 

258 TaskStateStoreModel.task_id == scope.task_id, 

259 TaskStateStoreModel.map_index == scope.map_index, 

260 TaskStateStoreModel.key == key, 

261 ) 

262 ).one_or_none() 

263 

264 if existing is None: 

265 raise HTTPException( 

266 status_code=status.HTTP_404_NOT_FOUND, 

267 detail=f"Task state store key {key!r} not found", 

268 ) 

269 

270 _get_db_backend().set(scope, key, json.dumps(body.value), expires_at=existing.expires_at, session=session) 

271 

272 

273@task_state_store_router.delete( 

274 "/{key:path}", 

275 status_code=status.HTTP_204_NO_CONTENT, 

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

277 dependencies=[Depends(requires_access_dag(method="DELETE", access_entity=DagAccessEntity.TASK_INSTANCE))], 

278) 

279def delete_task_state_store( 

280 scope: ValidatedTaskScopeDep, 

281 key: str, 

282 session: SessionDep, 

283) -> None: 

284 """Delete a single task state store key. No-op if the key does not exist.""" 

285 _get_db_backend().delete(scope, key, session=session) 

286 

287 

288@task_state_store_router.delete( 

289 "", 

290 status_code=status.HTTP_204_NO_CONTENT, 

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

292 dependencies=[Depends(requires_access_dag(method="DELETE", access_entity=DagAccessEntity.TASK_INSTANCE))], 

293) 

294def clear_task_state_store( 

295 scope: ValidatedClearTaskScopeDep, 

296 session: SessionDep, 

297 all_map_indices: Annotated[bool, Query()] = False, 

298) -> None: 

299 """ 

300 Delete all task state store keys for a task instance. 

301 

302 When ``all_map_indices=true``, state store is cleared for every map index of the task and 

303 the ``map_index`` parameter is ignored. 

304 """ 

305 _get_db_backend().clear(scope, all_map_indices=all_map_indices, session=session)