Skip to content
Draft
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
34 changes: 32 additions & 2 deletions airflow-core/src/airflow/api_fastapi/execution_api/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,10 +32,11 @@
Cadwyn,
current_dependency_solver,
)
from fastapi import Depends, FastAPI, Request, Response
from fastapi import Depends, FastAPI, HTTPException, Request, Response, status
from fastapi.responses import JSONResponse
from fastapi.routing import APIRoute
from opentelemetry import context as otel_context, propagate as otel_propagate
from sqlalchemy import select
from starlette.middleware.base import BaseHTTPMiddleware

from airflow.api_fastapi.auth.tokens import (
Expand All @@ -44,6 +45,7 @@
get_sig_validation_args,
get_signing_args,
)
from airflow.api_fastapi.common.db.common import SessionDep

if TYPE_CHECKING:
import httpx
Expand Down Expand Up @@ -390,8 +392,9 @@ def app(self):
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.routes.xcoms import get_xcom_write_ti, has_xcom_access
from airflow.api_fastapi.execution_api.security import _jwt_bearer
from airflow.models.taskinstance import TaskInstance

# Give this app its own lifespan + services registry so that stubbing services
# (e.g. JWTValidator) doesn't affect the module-level ``lifespan.registry``.
Expand All @@ -415,10 +418,37 @@ async def always_allow(request: Request):
claims = TIClaims(scope="execution")
return TIToken(id=ti_id, claims=claims)

def resolve_xcom_write_ti(
dag_id: str,
run_id: str,
task_id: str,
map_index: int = -1,
*,
session: SessionDep,
) -> TaskInstance:
ti = session.scalar(
select(TaskInstance).where(
TaskInstance.dag_id == dag_id,
TaskInstance.run_id == run_id,
TaskInstance.task_id == task_id,
TaskInstance.map_index == map_index,
)
)
if ti is None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"reason": "access_denied",
"message": "Task may only set XComs for its own task instance",
},
)
return ti

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
self._app.dependency_overrides[get_xcom_write_ti] = resolve_xcom_write_ti

return self._app

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,12 +27,14 @@

from airflow.api_fastapi.common.db.common import SessionDep
from airflow.api_fastapi.core_api.base import BaseModel
from airflow.api_fastapi.execution_api.datamodels.token import TIToken
from airflow.api_fastapi.execution_api.datamodels.xcom import (
XComResponse,
XComSequenceIndexResponse,
XComSequenceSliceResponse,
)
from airflow.api_fastapi.execution_api.security import CurrentTIToken
from airflow.models.taskinstance import TaskInstance
from airflow.models.taskmap import TaskMap
from airflow.models.xcom import XComModel
from airflow.utils.db import get_query_count
Expand Down Expand Up @@ -116,6 +118,33 @@ def has_xcom_access(
log = logging.getLogger(__name__)


def get_xcom_write_ti(
dag_id: str,
run_id: str,
task_id: str,
map_index: int = -1,
token: TIToken = CurrentTIToken,
*,
session: SessionDep,
) -> TaskInstance:
"""Resolve and authorize the task instance that owns an XCom write."""
ti = session.get(TaskInstance, token.id)
if ti is None or (dag_id, run_id, task_id, map_index) != (
ti.dag_id,
ti.run_id,
ti.task_id,
ti.map_index,
):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={
"reason": "access_denied",
"message": "Task may only set XComs for its own task instance",
},
)
return ti


async def xcom_query(
dag_id: str,
run_id: str,
Expand Down Expand Up @@ -356,8 +385,6 @@ def get_xcom(
return XComResponse(key=key, value=(result[0] if isinstance(result, tuple) else result).value)


# TODO: once we have JWT tokens, then remove dag_id/run_id/task_id from the URL and just use the info in
# the token
@router.post(
"/{dag_id}/{run_id}/{task_id}/{key:path}",
status_code=status.HTTP_201_CREATED,
Expand All @@ -368,6 +395,7 @@ def set_xcom(
task_id: str,
key: Annotated[str, Path(min_length=1)],
session: SessionDep,
ti: Annotated[TaskInstance, Depends(get_xcom_write_ti)],
value: Annotated[
JsonValue,
Body(
Expand Down Expand Up @@ -410,10 +438,10 @@ def set_xcom(

if mapped_length is not None:
task_map = TaskMap(
dag_id=dag_id,
task_id=task_id,
run_id=run_id,
map_index=map_index,
dag_id=ti.dag_id,
task_id=ti.task_id,
run_id=ti.run_id,
map_index=ti.map_index,
length=mapped_length,
keys=None,
)
Expand All @@ -437,10 +465,10 @@ def set_xcom(
XComModel.set(
key=key,
value=value,
run_id=run_id,
task_id=task_id,
dag_id=dag_id,
map_index=map_index,
run_id=ti.run_id,
task_id=ti.task_id,
dag_id=ti.dag_id,
map_index=ti.map_index,
serialize=False,
dag_result=dag_result,
session=session,
Expand Down
46 changes: 46 additions & 0 deletions airflow-core/tests/unit/api_fastapi/execution_api/test_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from fastapi.routing import APIRoute
from fastapi.testclient import TestClient
from opentelemetry import context as otel_context, propagate as otel_propagate
from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError

from airflow.api_fastapi.execution_api.app import (
Expand All @@ -40,6 +41,8 @@
from airflow.api_fastapi.execution_api.datamodels.token import TIClaims, TIToken
from airflow.api_fastapi.execution_api.security import require_auth
from airflow.api_fastapi.execution_api.versions import bundle
from airflow.models.xcom import XComModel
from airflow.sdk.serde import serialize

from tests_common.test_utils.config import conf_vars

Expand Down Expand Up @@ -163,6 +166,49 @@ def test_in_process_execution_api_runs_without_jwt_secret():
assert response.status_code == 200


@pytest.mark.parametrize("map_index", [-1, 2])
def test_in_process_execution_api_sets_xcom_for_route_task_instance(create_task_instance, session, map_index):
ti = create_task_instance(map_index=map_index)
session.commit()
params = {"map_index": map_index} if map_index >= 0 else None

api = InProcessExecutionAPI()
with httpx.Client(transport=api.transport, base_url="http://localhost") as client:
response = client.post(
f"/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/in_process",
params=params,
json=serialize('"value"'),
)

assert response.status_code == status.HTTP_201_CREATED, response.json()
xcom = session.scalar(
select(XComModel).where(
XComModel.dag_id == ti.dag_id,
XComModel.run_id == ti.run_id,
XComModel.task_id == ti.task_id,
XComModel.map_index == map_index,
XComModel.key == "in_process",
)
)
assert xcom is not None
assert xcom.value == '"value"'


def test_in_process_execution_api_rejects_xcom_for_missing_route_task_instance(session):
api = InProcessExecutionAPI()
with httpx.Client(transport=api.transport, base_url="http://localhost") as client:
response = client.post("/xcoms/missing/run/task/key", json=serialize('"value"'))

assert response.status_code == status.HTTP_403_FORBIDDEN, response.json()
assert response.json() == {
"detail": {
"reason": "access_denied",
"message": "Task may only set XComs for its own task instance",
}
}
assert session.scalar(select(XComModel).where(XComModel.key == "key")) is None


def test_in_process_execution_api_transport_lifecycle():
"""The background loop + thread lifecycle is tied to the ``.transport``, not the factory instance.
Expand Down
Loading