Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/execution_api/app.py: 53%

215 statements  

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

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

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

3# distributed with this work for additional information 

4# regarding copyright ownership. The ASF licenses this file 

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

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

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

8# 

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

10# 

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

12# software distributed under the License is distributed on an 

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

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

15# specific language governing permissions and limitations 

16# under the License. 

17 

18from __future__ import annotations 

19 

20import asyncio 

21import json 

22import threading 

23import time 

24import weakref 

25from contextlib import AsyncExitStack 

26from functools import cached_property 

27from typing import TYPE_CHECKING, Any, cast 

28 

29import attrs 

30import svcs 

31from cadwyn import ( 

32 Cadwyn, 

33 current_dependency_solver, 

34) 

35from fastapi import Depends, FastAPI, Request, Response 

36from fastapi.responses import JSONResponse 

37from fastapi.routing import APIRoute 

38from opentelemetry import context as otel_context, propagate as otel_propagate 

39from starlette.middleware.base import BaseHTTPMiddleware 

40 

41from airflow.api_fastapi.auth.tokens import ( 

42 JWTGenerator, 

43 JWTValidator, 

44 get_sig_validation_args, 

45 get_signing_args, 

46) 

47 

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

49 import httpx 

50 

51import structlog 

52from structlog.contextvars import bind_contextvars 

53 

54logger = structlog.get_logger(logger_name=__name__) 

55 

56__all__ = [ 

57 "create_task_execution_api_app", 

58 "lifespan", 

59 "CorrelationIdMiddleware", 

60] 

61 

62 

63def _jwt_validator(): 

64 from airflow.configuration import conf 

65 

66 required_claims = frozenset(["aud", "exp", "iat"]) 

67 

68 if issuer := conf.get("api_auth", "jwt_issuer", fallback=None): 68 ↛ 69line 68 didn't jump to line 69 because the condition on line 68 was never true

69 required_claims = required_claims | {"iss"} 

70 validator = JWTValidator( 

71 required_claims=required_claims, 

72 issuer=issuer, 

73 audience=conf.get_mandatory_list_value("execution_api", "jwt_audience"), 

74 **get_sig_validation_args(make_secret_key_if_needed=False), 

75 ) 

76 return validator 

77 

78 

79def _jwt_generator(): 

80 from airflow.configuration import conf 

81 

82 generator = JWTGenerator( 

83 valid_for=conf.getint("execution_api", "jwt_expiration_time"), 

84 audience=conf.get_mandatory_list_value("execution_api", "jwt_audience")[0], 

85 issuer=conf.get("api_auth", "jwt_issuer", fallback=None), 

86 # Since this one is used across components/server, there is no point trying to generate one, error 

87 # instead 

88 **get_signing_args(make_secret_key_if_needed=False), 

89 ) 

90 return generator 

91 

92 

93@svcs.fastapi.lifespan 

94async def lifespan(app: FastAPI, registry: svcs.Registry): 

95 app.state.lifespan_called = True 

96 

97 # According to svcs's docs this shouldn't be needed, but something about SubApps is odd, and we need to 

98 # record this here 

99 app.state.svcs_registry = registry 

100 

101 registry.register_factory(JWTGenerator, _jwt_generator) 

102 

103 # InProcessExecutionAPI stubs out JWTValidator: don't re-register in that case. 

104 if JWTValidator not in registry: 104 ↛ 108line 104 didn't jump to line 108 because the condition on line 104 was always true

105 # Create an app scoped validator, so that we don't have to fetch it every time 

106 registry.register_value(JWTValidator, _jwt_validator(), ping=JWTValidator.status) 

107 

108 yield 

109 

110 

111class CorrelationIdMiddleware(BaseHTTPMiddleware): 

112 """ 

113 Middleware to handle correlation-id for request tracing. 

114 

115 This middleware: 

116 1. Extracts correlation-id from request headers 

117 2. Binds it to structlog context for all logs within the request 

118 3. Echoes correlation-id back in response headers for tracing 

119 

120 Note: Context variables are automatically isolated per async task in Python, 

121 so manual cleanup is not necessary. Each request gets its own context copy. 

122 """ 

123 

124 async def dispatch(self, request: Request, call_next): 

125 correlation_id = request.headers.get("correlation-id") 

126 

127 if correlation_id: 127 ↛ 130line 127 didn't jump to line 130 because the condition on line 127 was always true

128 bind_contextvars(correlation_id=correlation_id) 

129 

130 response: Response = await call_next(request) 

131 

132 if correlation_id: 132 ↛ 135line 132 didn't jump to line 135 because the condition on line 132 was always true

133 response.headers["correlation-id"] = correlation_id 

134 

135 return response 

136 

137 

138class JWTReissueMiddleware(BaseHTTPMiddleware): 

139 async def dispatch(self, request: Request, call_next): 

140 response: Response = await call_next(request) 

141 

142 refreshed_token: str | None = None 

143 auth_header = request.headers.get("authorization") 

144 if auth_header and auth_header.lower().startswith("bearer "): 144 ↛ 170line 144 didn't jump to line 170 because the condition on line 144 was always true

145 token = auth_header.split(" ", 1)[1] 

146 try: 

147 async with svcs.Container(request.app.state.svcs_registry) as services: 

148 validator: JWTValidator = await services.aget(JWTValidator) 

149 claims = await validator.avalidated_claims(token, {}) 

150 

151 # Workload tokens are long-lived and meant to survive queue 

152 # wait times so avoid refreshing them. If avalidated_claims 

153 # raises for a workload token, the outer except handles it. 

154 if claims.get("scope") == "workload": 

155 return response 

156 

157 now = int(time.time()) 

158 token_lifetime = int(claims.get("exp", 0)) - int(claims.get("iat", 0)) 

159 refresh_when_less_than = max(int(token_lifetime * 0.20), 30) 

160 valid_left = int(claims.get("exp", 0)) - now 

161 if valid_left <= refresh_when_less_than: 161 ↛ 162line 161 didn't jump to line 162 because the condition on line 161 was never true

162 generator: JWTGenerator = await services.aget(JWTGenerator) 

163 refreshed_token = generator.generate(claims) 

164 except Exception as err: 

165 # Do not block the response if refreshing fails; log a warning for visibility 

166 logger.warning( 

167 "JWT reissue middleware failed to refresh token", error=str(err), exc_info=True 

168 ) 

169 

170 if refreshed_token: 170 ↛ 171line 170 didn't jump to line 171 because the condition on line 170 was never true

171 response.headers["Refreshed-API-Token"] = refreshed_token 

172 return response 

173 

174 

175class CadwynWithOpenAPICustomization(Cadwyn): 

176 # Workaround lack of customzation https://github.com/zmievsa/cadwyn/issues/255 

177 async def openapi_jsons(self, req: Request) -> JSONResponse: 

178 resp = await super().openapi_jsons(req) 

179 open_apischema = json.loads(cast("bytes", resp.body)) 

180 open_apischema = self.customize_openapi(open_apischema) 

181 

182 resp.body = resp.render(open_apischema) 

183 

184 return resp 

185 

186 def customize_openapi(self, openapi_schema: dict[str, Any]) -> dict[str, Any]: 

187 """ 

188 Customize the OpenAPI schema to include additional schemas not tied to specific endpoints. 

189 

190 This is particularly useful for client SDKs that require models for types 

191 not directly exposed in any endpoint's request or response schema. 

192 

193 We also replace ``anyOf`` with ``oneOf`` in the API spec as this produces better results for the code 

194 generators. This is because anyOf can technically be more than of the given schemas, but 99.9% of the 

195 time (perhaps 100% in this API) the types are mutually exclusive, so oneOf is more correct 

196 

197 References: 

198 - https://fastapi.tiangolo.com/how-to/extending-openapi/#modify-the-openapi-schema 

199 """ 

200 extra_schemas = get_extra_schemas() 

201 for schema_name, schema in extra_schemas.items(): 

202 if schema_name not in openapi_schema["components"]["schemas"]: 

203 openapi_schema["components"]["schemas"][schema_name] = schema 

204 

205 # The `JsonValue` component is missing any info. causes issues when generating models 

206 openapi_schema["components"]["schemas"]["JsonValue"] = { 

207 "title": "Any valid JSON value", 

208 "oneOf": [ 

209 {"type": t} for t in ("string", "number", "integer", "object", "array", "boolean", "null") 

210 ], 

211 } 

212 

213 def replace_any_of_with_one_of(spec): 

214 if isinstance(spec, dict): 

215 return { 

216 ("oneOf" if key == "anyOf" else key): replace_any_of_with_one_of(value) 

217 for key, value in spec.items() 

218 } 

219 if isinstance(spec, list): 

220 return [replace_any_of_with_one_of(item) for item in spec] 

221 return spec 

222 

223 openapi_schema = replace_any_of_with_one_of(openapi_schema) 

224 

225 for comp in openapi_schema["components"]["schemas"].values(): 

226 for prop in comp.get("properties", {}).values(): 

227 # {"type": "string", "const": "deferred"} 

228 # to 

229 # {"type": "string", "enum": ["deferred"]} 

230 # 

231 # this produces better results in the code generator 

232 if prop.get("type") == "string" and (const := prop.pop("const", None)): 

233 prop["enum"] = [const] 

234 

235 # Remove internal x-airflow-* extension fields from OpenAPI spec 

236 # These are used for runtime validation but shouldn't be exposed in the public API 

237 for path_item in openapi_schema.get("paths", {}).values(): 

238 for operation in path_item.values(): 

239 if isinstance(operation, dict): 

240 keys_to_remove = [key for key in operation.keys() if key.startswith("x-airflow-")] 

241 for key in keys_to_remove: 

242 del operation[key] 

243 

244 return openapi_schema 

245 

246 

247async def _extract_w3c_trace_context( 

248 request: Request, 

249 dependency_solver=Depends(current_dependency_solver), 

250): 

251 # Cadwyn solves dependencies twice (the real request, then again to migrate the 

252 # request body). Only act in the real "fastapi" pass so we attach/detach exactly 

253 # once, in the context the endpoint runs in. 

254 if dependency_solver != "fastapi": 

255 yield 

256 return 

257 ctx = otel_propagate.extract(request.headers) 

258 attached_in = asyncio.current_task() 

259 token = otel_context.attach(ctx) 

260 try: 

261 yield 

262 finally: 

263 if asyncio.current_task() is attached_in: 

264 otel_context.detach(token) 

265 

266 

267def _inject_trace_context_dep(routes, mode: str) -> None: 

268 dep = Depends(_extract_w3c_trace_context) 

269 for route in routes: 

270 if not isinstance(route, APIRoute): 270 ↛ 271line 270 didn't jump to line 271 because the condition on line 270 was never true

271 continue 

272 # Idempotent: create_task_execution_api_app() runs more than once per process 

273 # (cached_app + InProcessExecutionAPI), and execution_api_router is shared 

274 # module state, so strip any prior injection first. 

275 route.dependencies[:] = [ 

276 d for d in route.dependencies if getattr(d, "dependency", None) is not _extract_w3c_trace_context 

277 ] 

278 match mode: 

279 case "unsafe-always": 279 ↛ 280line 279 didn't jump to line 280 because the pattern on line 279 never matched

280 route.dependencies.insert(0, dep) 

281 case "only-authenticated": 281 ↛ 269line 281 didn't jump to line 269 because the pattern on line 281 always matched

282 from airflow.api_fastapi.execution_api.security import require_auth 

283 

284 if any(getattr(d, "dependency", None) is require_auth for d in route.dependencies): 

285 route.dependencies.append(dep) 

286 

287 

288def create_task_execution_api_app(lifespan: svcs.fastapi.lifespan = lifespan) -> FastAPI: 

289 """Create FastAPI app for task execution API.""" 

290 from airflow.api_fastapi.common.exceptions import init_error_handlers 

291 from airflow.api_fastapi.execution_api.routes import execution_api_router 

292 from airflow.api_fastapi.execution_api.versions import bundle 

293 from airflow.configuration import conf 

294 

295 def custom_generate_unique_id(route: APIRoute): 

296 # This is called only if the route doesn't provide an explicit operation ID 

297 return route.name 

298 

299 # See https://docs.cadwyn.dev/concepts/version_changes/ for info about API versions 

300 app = CadwynWithOpenAPICustomization( 

301 title="Airflow Task Execution API", 

302 description="The private Airflow Task Execution API.", 

303 lifespan=lifespan, 

304 generate_unique_id_function=custom_generate_unique_id, 

305 api_version_parameter_name="Airflow-API-Version", 

306 api_version_default_value=bundle.versions[0].value, 

307 versions=bundle, 

308 ) 

309 

310 # Add correlation-id middleware for request tracing 

311 app.add_middleware(CorrelationIdMiddleware) 

312 app.add_middleware(JWTReissueMiddleware) 

313 

314 mode = conf.get("execution_api", "otel_trace_propagation", fallback="only-authenticated") 

315 _inject_trace_context_dep(execution_api_router.routes, mode) 

316 

317 app.generate_and_include_versioned_routers(execution_api_router) 

318 init_error_handlers(app) 

319 

320 # As we are mounted as a sub app, we don't get any logs for unhandled exceptions without this! 

321 @app.exception_handler(Exception) 

322 def handle_exceptions(request: Request, exc: Exception): 

323 logger.exception("Handle died with an error", exc_info=(type(exc), exc, exc.__traceback__)) 

324 content = {"message": "Internal server error"} 

325 if correlation_id := request.headers.get("correlation-id"): 

326 content["correlation-id"] = correlation_id 

327 return JSONResponse(status_code=500, content=content) 

328 

329 return app 

330 

331 

332def get_extra_schemas() -> dict[str, dict]: 

333 """Get all the extra schemas that are not part of the main FastAPI app.""" 

334 from airflow.api_fastapi.execution_api.datamodels.taskinstance import TaskInstance 

335 from airflow.executors.workloads import BundleInfo 

336 from airflow.serialization.enums import DagAttributeTypes 

337 from airflow.task.trigger_rule import TriggerRule 

338 from airflow.task.weight_rule import WeightRule 

339 from airflow.utils.state import TaskInstanceState, TerminalTIState 

340 

341 return { 

342 "TaskInstance": TaskInstance.model_json_schema(), 

343 "BundleInfo": BundleInfo.model_json_schema(), 

344 # Include the combined state enum too. In the datamodels we separate out SUCCESS from the other states 

345 # as that has different payload requirements 

346 "TerminalTIState": {"type": "string", "enum": list(TerminalTIState)}, 

347 "TaskInstanceState": {"type": "string", "enum": list(TaskInstanceState)}, 

348 "WeightRule": {"type": "string", "enum": list(WeightRule)}, 

349 "TriggerRule": {"type": "string", "enum": list(TriggerRule)}, 

350 "DagAttributeTypes": { 

351 "type": "string", 

352 "enum": [DagAttributeTypes.OP.value, DagAttributeTypes.TASK_GROUP.value], 

353 "x-enum-varnames": [DagAttributeTypes.OP.name, DagAttributeTypes.TASK_GROUP.name], 

354 }, 

355 } 

356 

357 

358# Note: _shutdown_loop is used as a finalizer for the WSGI transport returned by 

359# ``InProcessExecutionAPI.transport``. As such, its arguments must not directly or indirectly reference that 

360# transport, as this would prevent the transport from being garbage collected. 

361def _shutdown_loop( 

362 loop: asyncio.AbstractEventLoop, 

363 thread: threading.Thread, 

364 cm: AsyncExitStack, 

365) -> None: 

366 """Close the FastAPI lifespan and stop the background event loop + thread.""" 

367 try: 

368 asyncio.run_coroutine_threadsafe(cm.aclose(), loop).result(timeout=5) 

369 except Exception: 

370 logger.exception("Error while closing in-process execution API lifespan") 

371 loop.call_soon_threadsafe(loop.stop) 

372 thread.join(timeout=5) 

373 

374 

375@attrs.define() 

376class InProcessExecutionAPI: 

377 """ 

378 A helper class to make it possible to run the ExecutionAPI "in-process". 

379 

380 The sync version of this makes use of a2wsgi which runs the async loop in a separate thread. This is 

381 needed so that we can use the sync httpx client 

382 """ 

383 

384 _app: FastAPI | None = None 

385 

386 @cached_property 

387 def app(self): 

388 if not self._app: 

389 from airflow.api_fastapi.common.dagbag import create_dag_bag 

390 from airflow.api_fastapi.execution_api.datamodels.token import TIClaims, TIToken 

391 from airflow.api_fastapi.execution_api.routes.connections import has_connection_access 

392 from airflow.api_fastapi.execution_api.routes.variables import has_variable_access 

393 from airflow.api_fastapi.execution_api.routes.xcoms import has_xcom_access 

394 from airflow.api_fastapi.execution_api.security import _jwt_bearer 

395 

396 # Give this app its own lifespan + services registry so that stubbing services 

397 # (e.g. JWTValidator) doesn't affect the module-level ``lifespan.registry``. 

398 registry = svcs.Registry() 

399 private_lifespan = attrs.evolve(lifespan, registry=registry) 

400 self._app = create_task_execution_api_app(lifespan=private_lifespan) 

401 

402 # In-process callers don't need a real JWTValidator: auth is bypassed below via 

403 # ``dependency_overrides``. 

404 registry.register_value(JWTValidator, None) 

405 

406 # Set up dag_bag in app state for dependency injection 

407 self._app.state.dag_bag = create_dag_bag() 

408 

409 async def always_allow(request: Request): 

410 from uuid import UUID 

411 

412 ti_id = UUID( 

413 request.path_params.get("task_instance_id", "00000000-0000-0000-0000-000000000000") 

414 ) 

415 claims = TIClaims(scope="execution") 

416 return TIToken(id=ti_id, claims=claims) 

417 

418 self._app.dependency_overrides[_jwt_bearer] = always_allow 

419 self._app.dependency_overrides[has_connection_access] = always_allow 

420 self._app.dependency_overrides[has_variable_access] = always_allow 

421 self._app.dependency_overrides[has_xcom_access] = always_allow 

422 

423 return self._app 

424 

425 @cached_property 

426 def transport(self) -> httpx.WSGITransport: 

427 import httpx 

428 from a2wsgi import ASGIMiddleware 

429 

430 # We choose to own the event loop + executor thread here so that we can have explicit control over 

431 # their lifecycle. 

432 loop = asyncio.new_event_loop() 

433 thread = threading.Thread(target=loop.run_forever, name="InProcessExecutionAPI-loop", daemon=True) 

434 thread.start() 

435 

436 middleware = ASGIMiddleware(self.app, loop=loop) 

437 

438 # https://github.com/abersheeran/a2wsgi/discussions/64 

439 async def start_lifespan(cm: AsyncExitStack, app: FastAPI): 

440 await cm.enter_async_context(app.router.lifespan_context(app)) 

441 

442 cm = AsyncExitStack() 

443 

444 # Wait for lifespan startup to complete so callers see a ready app and so the finalizer can 

445 # safely aclose() a context whose __aenter__ has actually run. 

446 asyncio.run_coroutine_threadsafe(start_lifespan(cm, self.app), loop).result() 

447 

448 transport = httpx.WSGITransport(app=middleware) # type: ignore[arg-type] 

449 

450 # Stop the loop + thread and unwind the lifespan when the *transport* is garbage collected, not 

451 # this InProcessExecutionAPI instance. Callers commonly build a Client from ``.transport`` and drop 

452 # the factory object (e.g. ``Client(transport=InProcessExecutionAPI().transport)``); finalizing on 

453 # ``self`` would stop the loop while the transport is still in use, so every later request would 

454 # hang on the now-dead loop. 

455 weakref.finalize(transport, _shutdown_loop, loop, thread, cm) 

456 

457 return transport 

458 

459 @cached_property 

460 def atransport(self) -> httpx.ASGITransport: 

461 import httpx 

462 

463 return httpx.ASGITransport(app=self.app)