Skip to content
Open
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
6 changes: 6 additions & 0 deletions airflow-core/docs/tutorial/hitl.rst
Original file line number Diff line number Diff line change
Expand Up @@ -221,6 +221,12 @@ calls involved (``~`` works as a wildcard for ``dag_id`` and ``dag_run_id``):
PATCH /api/v2/dags/{dag_id}/dagRuns/{dag_run_id}/taskInstances/{task_id}/{map_index}/hitlDetails
{"chosen_options": ["Approve"], "params_input": {}}

.. note::

Keys in ``params_input`` may not be Airflow's reserved serialization keys (``__classname__`` or
``__id__``). A response containing one is rejected with ``400`` at submission time, because it
could not be serialized when the task resumes.

.. note::

``response_timeout`` and timeout defaults are enforced by the scheduler, which does not run
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,13 +16,15 @@
# under the License.
from __future__ import annotations

from collections.abc import Mapping
from typing import Annotated

import structlog
from fastapi import Depends, HTTPException, status
from sqlalchemy import select
from sqlalchemy.orm import joinedload

from airflow._shared.serialization import CLASSNAME, SCHEMA_ID
from airflow._shared.timezones import timezone
from airflow.api_fastapi.auth.managers.models.resource_details import DagAccessEntity
from airflow.api_fastapi.common.db.common import SessionDep, paginated_select
Expand Down Expand Up @@ -80,6 +82,27 @@
log = structlog.get_logger(__name__)


# Keys that ``serde.serialize`` refuses at any depth. A ``params_input`` carrying one could not be
# serialized into the resume event, so the response is rejected before it is stored. Extend this
# tuple if serde gains another reserved key; the write-side check and its tests read from here.
_SERDE_RESERVED_KEYS = (CLASSNAME, SCHEMA_ID)


def _find_serde_reserved_key(value: object) -> str | None:
"""Return the first ``_SERDE_RESERVED_KEYS`` entry found at any depth in ``value``, else ``None``."""
stack: list[object] = [value]
while stack:
current = stack.pop()
if isinstance(current, Mapping):
for key in _SERDE_RESERVED_KEYS:
if key in current:
return key
stack.extend(current.values())
elif isinstance(current, (list, tuple)):
stack.extend(current)
return None


def _get_task_instance_with_hitl_detail(
dag_id: str,
dag_run_id: str,
Expand Down Expand Up @@ -221,6 +244,13 @@ def update_hitl_detail(
status.HTTP_400_BAD_REQUEST,
"Multiple options chosen but this Human-in-the-loop task accepts only a single option.",
)
reserved_key = _find_serde_reserved_key(update_hitl_detail_payload.params_input)
if reserved_key is not None:
raise HTTPException(
status.HTTP_400_BAD_REQUEST,
f"params_input may not contain the reserved key {reserved_key!r}, which cannot be "
"serialized when the task resumes.",
)

hitl_detail_model.responded_by = hitl_user
hitl_detail_model.responded_at = timezone.utcnow()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from sqlalchemy.orm import Session

from airflow._shared.timezones.timezone import utc, utcnow
from airflow.api_fastapi.core_api.routes.public.hitl import _SERDE_RESERVED_KEYS
from airflow.models.hitl import HITLDetail
from airflow.models.log import Log
from airflow.models.taskinstance import TaskInstance as TIModel
Expand Down Expand Up @@ -408,6 +409,75 @@ def test_should_respond_400_for_invalid_option(
assert response.status_code == 400
assert "Invalid options" in response.json()["detail"]

@pytest.mark.usefixtures("sample_hitl_detail")
@pytest.mark.parametrize("reserved_key", _SERDE_RESERVED_KEYS)
@pytest.mark.parametrize(
"make_params_input",
[
pytest.param(lambda key: {key: "x"}, id="top-level"),
pytest.param(lambda key: {"nested": {key: "x"}}, id="nested-dict"),
pytest.param(lambda key: {"items": [{key: "x"}]}, id="inside-list"),
],
)
def test_should_respond_400_for_serde_reserved_key_in_params_input(
self,
test_client: TestClient,
sample_ti_url_identifier: str,
reserved_key: str,
make_params_input: Callable[[str], dict[str, Any]],
) -> None:
"""A params_input carrying a serde-reserved key at any depth is rejected (400) at submission time."""
response = test_client.patch(
f"{sample_ti_url_identifier}/hitlDetails",
json={"chosen_options": ["Approve"], "params_input": make_params_input(reserved_key)},
)
assert response.status_code == 400
assert reserved_key in response.json()["detail"]

@time_machine.travel(datetime(2025, 7, 3, 0, 0, 0), tick=False)
@pytest.mark.usefixtures("sample_hitl_detail")
def test_rejected_reserved_key_leaves_task_resumable(
self,
test_client: TestClient,
sample_ti_url_identifier: str,
sample_ti: TaskInstance,
session: Session,
) -> None:
"""A rejected reserved-key response records nothing, so the parked task is still resumable by a corrected resubmission."""
ti = session.get(TIModel, sample_ti.id)
assert ti is not None
ti.state = TaskInstanceState.AWAITING_INPUT
ti.next_method = "execute_complete"
ti.next_kwargs = {}
ti.trigger_id = None
session.commit()

rejected = test_client.patch(
f"{sample_ti_url_identifier}/hitlDetails",
json={"chosen_options": ["Approve"], "params_input": {_SERDE_RESERVED_KEYS[0]: "x"}},
)
assert rejected.status_code == 400

session.expire_all()
parked = session.get(TIModel, sample_ti.id)
assert parked is not None
assert parked.state == TaskInstanceState.AWAITING_INPUT
detail = session.scalar(select(HITLDetail).where(HITLDetail.ti_id == sample_ti.id))
assert detail is not None
assert detail.response_received is False

accepted = test_client.patch(
f"{sample_ti_url_identifier}/hitlDetails",
json={"chosen_options": ["Approve"], "params_input": {"input_1": 2}},
)
assert accepted.status_code == 200

session.expire_all()
resumed = session.get(TIModel, sample_ti.id)
assert resumed is not None
assert resumed.state == TaskInstanceState.SCHEDULED
assert "event" in (resumed.next_kwargs or {})

@time_machine.travel(datetime(2025, 7, 3, 0, 0, 0), tick=False)
@pytest.mark.usefixtures("sample_hitl_detail_respondent")
def test_should_respond_200_to_assigned_users(
Expand Down
Loading