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
« 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 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
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
41from airflow.api_fastapi.auth.tokens import (
42 JWTGenerator,
43 JWTValidator,
44 get_sig_validation_args,
45 get_signing_args,
46)
48if TYPE_CHECKING: 48 ↛ 49line 48 didn't jump to line 49 because the condition on line 48 was never true
49 import httpx
51import structlog
52from structlog.contextvars import bind_contextvars
54logger = structlog.get_logger(logger_name=__name__)
56__all__ = [
57 "create_task_execution_api_app",
58 "lifespan",
59 "CorrelationIdMiddleware",
60]
63def _jwt_validator():
64 from airflow.configuration import conf
66 required_claims = frozenset(["aud", "exp", "iat"])
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
79def _jwt_generator():
80 from airflow.configuration import conf
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
93@svcs.fastapi.lifespan
94async def lifespan(app: FastAPI, registry: svcs.Registry):
95 app.state.lifespan_called = True
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
101 registry.register_factory(JWTGenerator, _jwt_generator)
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)
108 yield
111class CorrelationIdMiddleware(BaseHTTPMiddleware):
112 """
113 Middleware to handle correlation-id for request tracing.
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
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 """
124 async def dispatch(self, request: Request, call_next):
125 correlation_id = request.headers.get("correlation-id")
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)
130 response: Response = await call_next(request)
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
135 return response
138class JWTReissueMiddleware(BaseHTTPMiddleware):
139 async def dispatch(self, request: Request, call_next):
140 response: Response = await call_next(request)
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, {})
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
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 )
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
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)
182 resp.body = resp.render(open_apischema)
184 return resp
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.
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.
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
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
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 }
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
223 openapi_schema = replace_any_of_with_one_of(openapi_schema)
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]
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]
244 return openapi_schema
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)
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
284 if any(getattr(d, "dependency", None) is require_auth for d in route.dependencies):
285 route.dependencies.append(dep)
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
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
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 )
310 # Add correlation-id middleware for request tracing
311 app.add_middleware(CorrelationIdMiddleware)
312 app.add_middleware(JWTReissueMiddleware)
314 mode = conf.get("execution_api", "otel_trace_propagation", fallback="only-authenticated")
315 _inject_trace_context_dep(execution_api_router.routes, mode)
317 app.generate_and_include_versioned_routers(execution_api_router)
318 init_error_handlers(app)
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)
329 return app
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
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 }
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)
375@attrs.define()
376class InProcessExecutionAPI:
377 """
378 A helper class to make it possible to run the ExecutionAPI "in-process".
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 """
384 _app: FastAPI | None = None
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
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)
402 # In-process callers don't need a real JWTValidator: auth is bypassed below via
403 # ``dependency_overrides``.
404 registry.register_value(JWTValidator, None)
406 # Set up dag_bag in app state for dependency injection
407 self._app.state.dag_bag = create_dag_bag()
409 async def always_allow(request: Request):
410 from uuid import UUID
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)
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
423 return self._app
425 @cached_property
426 def transport(self) -> httpx.WSGITransport:
427 import httpx
428 from a2wsgi import ASGIMiddleware
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()
436 middleware = ASGIMiddleware(self.app, loop=loop)
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))
442 cm = AsyncExitStack()
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()
448 transport = httpx.WSGITransport(app=middleware) # type: ignore[arg-type]
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)
457 return transport
459 @cached_property
460 def atransport(self) -> httpx.ASGITransport:
461 import httpx
463 return httpx.ASGITransport(app=self.app)