Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions airflow-core/src/airflow/api_fastapi/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,8 +142,7 @@ def create_app(apps: str = "all") -> FastAPI:
dag_bag = create_dag_bag()

if "all" in apps_list or "execution" in apps_list:
task_exec_api_app = create_task_execution_api_app()
task_exec_api_app.state.dag_bag = dag_bag
task_exec_api_app = create_task_execution_api_app(dag_bag=dag_bag)
app.mount("/execution", task_exec_api_app)

if "all" in apps_list or "core" in apps_list:
Expand Down
141 changes: 16 additions & 125 deletions airflow-core/src/airflow/api_fastapi/execution_api/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,14 +19,9 @@

import asyncio
import json
import threading
import time
import weakref
from contextlib import AsyncExitStack
from functools import cached_property
from typing import TYPE_CHECKING, Any, cast

import attrs
import svcs
from cadwyn import (
Cadwyn,
Expand All @@ -38,23 +33,24 @@
from opentelemetry import context as otel_context, propagate as otel_propagate
from starlette.middleware.base import BaseHTTPMiddleware

from airflow import settings
from airflow.api_fastapi.auth.tokens import (
JWTGenerator,
JWTValidator,
get_sig_validation_args,
get_signing_args,
)
from airflow.api_fastapi.execution_api.in_process import InProcessExecutionAPI

if TYPE_CHECKING:
import httpx
from airflow.models.dagbag import DBDagBag

import structlog
from structlog.contextvars import bind_contextvars

logger = structlog.get_logger(logger_name=__name__)

__all__ = [
"InProcessExecutionAPI",
"create_task_execution_api_app",
"lifespan",
"CorrelationIdMiddleware",
Expand Down Expand Up @@ -286,8 +282,18 @@ def _inject_trace_context_dep(routes, mode: str) -> None:
route.dependencies.append(dep)


def create_task_execution_api_app(lifespan: svcs.fastapi.lifespan = lifespan) -> FastAPI:
"""Create FastAPI app for task execution API."""
def create_task_execution_api_app(
lifespan: svcs.fastapi.lifespan = lifespan,
*,
dag_bag: DBDagBag | None = None,
) -> FastAPI:
"""
Create the Execution API app.

:param lifespan: Lifespan manager that registers the app's services.
:param dag_bag: Dag bag to use; a new one is created when omitted.
"""
from airflow.api_fastapi.common.dagbag import create_dag_bag
from airflow.api_fastapi.common.exceptions import init_error_handlers
from airflow.api_fastapi.execution_api.routes import execution_api_router
from airflow.api_fastapi.execution_api.versions import bundle
Expand Down Expand Up @@ -327,6 +333,7 @@ def handle_exceptions(request: Request, exc: Exception):
content["correlation-id"] = correlation_id
return JSONResponse(status_code=500, content=content)

app.state.dag_bag = dag_bag if dag_bag is not None else create_dag_bag()
return app


Expand Down Expand Up @@ -354,119 +361,3 @@ def get_extra_schemas() -> dict[str, dict]:
"x-enum-varnames": [DagAttributeTypes.OP.name, DagAttributeTypes.TASK_GROUP.name],
},
}


# Note: _shutdown_loop is used as a finalizer for the WSGI transport returned by
# ``InProcessExecutionAPI.transport``. As such, its arguments must not directly or indirectly reference that
# transport, as this would prevent the transport from being garbage collected.
def _shutdown_loop(
loop: asyncio.AbstractEventLoop,
thread: threading.Thread,
cm: AsyncExitStack,
) -> None:
"""Close the FastAPI lifespan and stop the background event loop + thread."""
try:
asyncio.run_coroutine_threadsafe(cm.aclose(), loop).result(timeout=5)
except Exception:
logger.exception("Error while closing in-process execution API lifespan")
loop.call_soon_threadsafe(loop.stop)
thread.join(timeout=5)


@attrs.define()
class InProcessExecutionAPI:
"""
A helper class to make it possible to run the ExecutionAPI "in-process".

The sync version of this makes use of a2wsgi which runs the async loop in a separate thread. This is
needed so that we can use the sync httpx client
"""

_app: FastAPI | None = None

@cached_property
def app(self):
if not self._app:
from airflow.api_fastapi.common.dagbag import create_dag_bag
from airflow.api_fastapi.execution_api.datamodels.token import TIClaims, TIToken
from airflow.api_fastapi.execution_api.routes.connections import has_connection_access
from airflow.api_fastapi.execution_api.routes.variables import has_variable_access
from airflow.api_fastapi.execution_api.routes.xcoms import has_xcom_access
from airflow.api_fastapi.execution_api.security import _IN_PROCESS_NON_TI_CALLER, _jwt_bearer

# Give this app its own lifespan + services registry so that stubbing services
# (e.g. JWTValidator) doesn't affect the module-level ``lifespan.registry``.
registry = svcs.Registry()
private_lifespan = attrs.evolve(lifespan, registry=registry)
self._app = create_task_execution_api_app(lifespan=private_lifespan)

# In-process callers don't need a real JWTValidator: auth is bypassed below via
# ``dependency_overrides``.
registry.register_value(JWTValidator, None)

# Set up dag_bag in app state for dependency injection
self._app.state.dag_bag = create_dag_bag()

async def always_allow(request: Request):
from uuid import UUID

ti_id = UUID(
request.path_params.get("task_instance_id")
or request.headers.get("X-Airflow-In-Process-Attempt-Id")
or "00000000-0000-0000-0000-000000000000"
)
# Watchers and other trusted in-process callers may have no task identity.
request.scope[_IN_PROCESS_NON_TI_CALLER] = (
"task_instance_id" not in request.path_params
and "X-Airflow-In-Process-Attempt-Id" not in request.headers
)
claims = TIClaims(scope="execution")
return TIToken(id=ti_id, claims=claims)

self._app.dependency_overrides[_jwt_bearer] = always_allow
self._app.dependency_overrides[has_connection_access] = always_allow
self._app.dependency_overrides[has_variable_access] = always_allow
self._app.dependency_overrides[has_xcom_access] = always_allow

return self._app

@cached_property
def transport(self) -> httpx.WSGITransport:
import httpx
from a2wsgi import ASGIMiddleware

# We choose to own the event loop + executor thread here so that we can have explicit control over
# their lifecycle.
loop = asyncio.new_event_loop()
thread = threading.Thread(target=loop.run_forever, name="InProcessExecutionAPI-loop", daemon=True)
thread.start()

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

# https://github.com/abersheeran/a2wsgi/discussions/64
async def start_lifespan(cm: AsyncExitStack, app: FastAPI):
cm.push_async_callback(settings.dispose_async_engine)
await cm.enter_async_context(app.router.lifespan_context(app))

cm = AsyncExitStack()

# Wait for lifespan startup to complete so callers see a ready app and so the finalizer can
# safely aclose() a context whose __aenter__ has actually run.
asyncio.run_coroutine_threadsafe(start_lifespan(cm, self.app), loop).result()

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

# Stop the loop + thread and unwind the lifespan when the *transport* is garbage collected, not
# this InProcessExecutionAPI instance. Callers commonly build a Client from ``.transport`` and drop
# the factory object (e.g. ``Client(transport=InProcessExecutionAPI().transport)``); finalizing on
# ``self`` would stop the loop while the transport is still in use, so every later request would
# hang on the now-dead loop.
weakref.finalize(transport, _shutdown_loop, loop, thread, cm)

return transport

@cached_property
def atransport(self) -> httpx.ASGITransport:
import httpx

return httpx.ASGITransport(app=self.app)
148 changes: 148 additions & 0 deletions airflow-core/src/airflow/api_fastapi/execution_api/in_process.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,148 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
"""In-process Execution API hosting, authentication overrides, and transport lifecycle."""

from __future__ import annotations

import asyncio
import threading
import weakref
from contextlib import AsyncExitStack
from functools import cached_property
from typing import TYPE_CHECKING

import attrs
import structlog
from starlette.requests import Request

if TYPE_CHECKING:
import httpx
from fastapi import FastAPI

logger = structlog.get_logger(logger_name=__name__)


# Finalizer arguments must not reference the transport, or it cannot be garbage collected.
def _shutdown_loop(
loop: asyncio.AbstractEventLoop,
thread: threading.Thread,
cm: AsyncExitStack,
) -> None:
"""Close the FastAPI lifespan and stop the background event loop + thread."""
try:
asyncio.run_coroutine_threadsafe(cm.aclose(), loop).result(timeout=5)
except Exception:
logger.exception("Error while closing in-process execution API lifespan")
loop.call_soon_threadsafe(loop.stop)
thread.join(timeout=5)


def _configure_in_process_auth(app: FastAPI) -> None:
from airflow.api_fastapi.execution_api.datamodels.token import TIClaims, TIToken
from airflow.api_fastapi.execution_api.routes.connections import has_connection_access
from airflow.api_fastapi.execution_api.routes.variables import has_variable_access
from airflow.api_fastapi.execution_api.routes.xcoms import has_xcom_access
from airflow.api_fastapi.execution_api.security import _IN_PROCESS_NON_TI_CALLER, _jwt_bearer

async def always_allow(request: Request):
from uuid import UUID

ti_id = UUID(
request.path_params.get("task_instance_id")
or request.headers.get("X-Airflow-In-Process-Attempt-Id")
or "00000000-0000-0000-0000-000000000000"
)
# Watchers and other trusted in-process callers may have no task identity.
request.scope[_IN_PROCESS_NON_TI_CALLER] = (
"task_instance_id" not in request.path_params
and "X-Airflow-In-Process-Attempt-Id" not in request.headers
)
claims = TIClaims(scope="execution")
return TIToken(id=ti_id, claims=claims)

app.dependency_overrides[_jwt_bearer] = always_allow
app.dependency_overrides[has_connection_access] = always_allow
app.dependency_overrides[has_variable_access] = always_allow
app.dependency_overrides[has_xcom_access] = always_allow


@attrs.define()
class InProcessExecutionAPI:
"""
A helper class to make it possible to run the ExecutionAPI "in-process".

The sync version of this makes use of a2wsgi which runs the async loop in a separate thread. This is
needed so that we can use the sync httpx client
"""

_app: FastAPI | None = None

@cached_property
def app(self):
import svcs

from airflow.api_fastapi.auth.tokens import JWTValidator
from airflow.api_fastapi.execution_api.app import create_task_execution_api_app, lifespan

if not self._app:
# Keep the stub private: a shared None validator would prevent lifespan() from
# registering the real API server's validator.
registry = svcs.Registry()
self._app = create_task_execution_api_app(lifespan=attrs.evolve(lifespan, registry=registry))
# Auth dependency overrides bypass validation, so the stub skips validator setup in lifespan().
registry.register_value(JWTValidator, None)
_configure_in_process_auth(self._app)

return self._app

@cached_property
def transport(self) -> httpx.WSGITransport:
import httpx
from a2wsgi import ASGIMiddleware

from airflow import settings

# Own the event loop and thread so the transport controls their lifecycle.
loop = asyncio.new_event_loop()
thread = threading.Thread(target=loop.run_forever, name="InProcessExecutionAPI-loop", daemon=True)
thread.start()

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

# https://github.com/abersheeran/a2wsgi/discussions/64
async def start_lifespan(cm: AsyncExitStack, app: FastAPI):
cm.push_async_callback(settings.dispose_async_engine)
await cm.enter_async_context(app.router.lifespan_context(app))

cm = AsyncExitStack()

# Wait for startup so callers see a ready app and the finalizer can close an entered context.
asyncio.run_coroutine_threadsafe(start_lifespan(cm, self.app), loop).result()

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

# Callers can retain the transport after dropping this instance; finalizing on self would
# stop the loop while requests still need it.
weakref.finalize(transport, _shutdown_loop, loop, thread, cm)

return transport

@cached_property
def atransport(self) -> httpx.ASGITransport:
import httpx

return httpx.ASGITransport(app=self.app)
14 changes: 2 additions & 12 deletions airflow-core/src/airflow/dag_processing/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
from airflow._shared.observability.metrics import stats
from airflow._shared.observability.metrics.stats import normalize_name_for_stats
from airflow._shared.timezones import timezone
from airflow.api_fastapi.execution_api.in_process import InProcessExecutionAPI
from airflow.callbacks.callback_requests import DagCallbackRequest
from airflow.configuration import conf
from airflow.dag_processing.bundles.base import (
Expand Down Expand Up @@ -97,22 +98,11 @@
from sqlalchemy.orm import Session
from sqlalchemy.sql import Select

from airflow.api_fastapi.execution_api.app import InProcessExecutionAPI
from airflow.callbacks.callback_requests import CallbackRequest
from airflow.dag_processing.bundles.base import BaseDagBundle
from airflow.sdk.api.client import Client


def _make_execution_api() -> InProcessExecutionAPI:
# This is a seriously weighty import, pulling in svcs, cadwyn, fastapi, aiohttp, etc.
#
# Defer it so that an import of this module for types (e.g. DagFileStat, DagFileInto) doesn't need to pay
# that cost.
from airflow.api_fastapi.execution_api.app import InProcessExecutionAPI

return InProcessExecutionAPI()


class BundleState(NamedTuple):
"""Persisted refresh state for a DAG bundle."""

Expand Down Expand Up @@ -304,7 +294,7 @@ class DagFileProcessorManager(LoggingMixin):
)
"""Resolved once per process so file discovery and the deactivation scan use the same value."""

_api_server: InProcessExecutionAPI = attrs.field(init=False, factory=_make_execution_api)
_api_server: InProcessExecutionAPI = attrs.field(init=False, factory=InProcessExecutionAPI)
Comment thread
ashb marked this conversation as resolved.
"""API server to interact with Metadata DB"""

def register_exit_signals(self):
Expand Down
Loading