Skip to content

Commit cab90c1

Browse files
committed
Add a POST /xcoms/{dag_id}/{run_id}/{task_id}/keys execution API endpoint
XComIterable, the result of an iterated task, stores one XCom per iteration under distinct keys (return_value_0, return_value_1, ...) of the same task instance. The existing slice endpoint for a mapped task's XComs ranges over map_index for a single key, the inverse shape, so iterating or slicing an XComIterable cost one GET per value. The new endpoint takes a list of keys and returns their values in that order, None for a key without an XCom, filtered by map_index (-1 by default), in one database query. Around it: * XComKeysRequest body model, and an AddXComKeysEndpoint execution API version change so older clients keep their contract. * has_xcom_access reads the optional key from the request's path parameters, so the endpoint needs no separate router or dependency. * GetXComByKeys supervisor message, XComOperations.get_by_keys client method, handle_get_xcom_by_keys handler registered with the task supervisor and the DAG processor, an AddGetXComByKeys supervisor schema version entry and the regenerated schema snapshot. * BaseXCom.get_by_keys sends the message and deserializes the values, and XComIterable iterates and slices through it: one request instead of one per value. A single index stays one XCom read. Stacked on the task iteration PR, which owns XComIterable; only this commit belongs to the endpoint. Squashed from the earlier history of this branch and re-based on that PR, with XComIterable's batched reads moved to its current home in bases/xcom.py.
1 parent 6b51e64 commit cab90c1

19 files changed

Lines changed: 431 additions & 41 deletions

File tree

‎airflow-core/src/airflow/api_fastapi/execution_api/datamodels/xcom.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,3 +40,9 @@ class XComSequenceSliceResponse(RootModel):
4040
"""XCom schema with minimal structure for slice-based access."""
4141

4242
root: list[JsonValue]
43+
44+
45+
class XComKeysRequest(BaseModel):
46+
"""Request body for fetching multiple XCom values by key list in a single query."""
47+
48+
keys: list[str]

‎airflow-core/src/airflow/api_fastapi/execution_api/routes/xcoms.py‎

Lines changed: 38 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
from airflow.api_fastapi.core_api.base import BaseModel
3131
from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc
3232
from airflow.api_fastapi.execution_api.datamodels.xcom import (
33+
XComKeysRequest,
3334
XComResponse,
3435
XComSequenceIndexResponse,
3536
XComSequenceSliceResponse,
@@ -45,7 +46,6 @@ def has_xcom_access(
4546
dag_id: str,
4647
run_id: str,
4748
task_id: str,
48-
xcom_key: Annotated[str, Path(alias="key", min_length=1)],
4949
request: Request,
5050
session: SessionDep,
5151
token=CurrentTIToken,
@@ -71,6 +71,7 @@ def has_xcom_access(
7171
from airflow.configuration import conf
7272

7373
write = request.method not in {"GET", "HEAD", "OPTIONS"}
74+
xcom_key = request.path_params.get("key")
7475

7576
log.debug(
7677
"Checking %s XCom access for task instance '%s' to XCom '%s' on dag '%s'",
@@ -116,6 +117,7 @@ def has_xcom_access(
116117
dependencies=[Depends(has_xcom_access)],
117118
)
118119

120+
119121
log = logging.getLogger(__name__)
120122

121123

@@ -405,6 +407,41 @@ def get_xcom(
405407
return XComResponse(key=key, value=(result[0] if isinstance(result, tuple) else result).value)
406408

407409

410+
@router.post(
411+
"/{dag_id}/{run_id}/{task_id}/keys",
412+
summary="Get multiple XCom values by keys",
413+
description=(
414+
"Fetch multiple XCom values by key list in a single database query. "
415+
"Optimised for XComIterable iteration, reducing N round-trips to one."
416+
),
417+
)
418+
def get_xcom_by_keys(
419+
dag_id: str,
420+
run_id: str,
421+
task_id: str,
422+
request_body: XComKeysRequest,
423+
session: SessionDep,
424+
map_index: Annotated[int, Query()] = -1,
425+
) -> XComSequenceSliceResponse:
426+
"""Fetch multiple XCom values by different keys in a single database query."""
427+
key_list = request_body.keys
428+
if not key_list:
429+
return XComSequenceSliceResponse([])
430+
431+
xcom_read = XComModel.get_many(
432+
run_id=run_id,
433+
task_ids=task_id,
434+
dag_ids=dag_id,
435+
map_indexes=map_index,
436+
)
437+
entity = xcom_entity(xcom_read)
438+
query = (
439+
xcom_read.with_only_columns(entity.key, entity.value).where(entity.key.in_(key_list)).order_by(None)
440+
)
441+
rows = {row.key: row.value for row in session.execute(query)}
442+
return XComSequenceSliceResponse([rows.get(key) for key in key_list])
443+
444+
408445
# TODO: once we have JWT tokens, then remove dag_id/run_id/task_id from the URL and just use the info in
409446
# the token
410447
@router.post(

‎airflow-core/src/airflow/api_fastapi/execution_api/versions/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -59,6 +59,7 @@
5959
AddMultiTeamToTIRunContext,
6060
AddStoppedTaskReport,
6161
AddTerminalStateRetryReasonField,
62+
AddXComKeysEndpoint,
6263
IdentifyArchivedTaskStateUpdates,
6364
)
6465

@@ -73,6 +74,7 @@
7374
AddMultiTeamToTIRunContext,
7475
AddStoppedTaskReport,
7576
IdentifyArchivedTaskStateUpdates,
77+
AddXComKeysEndpoint,
7678
),
7779
Version(
7880
"2026-06-30",

‎airflow-core/src/airflow/api_fastapi/execution_api/versions/v2026_10_30.py‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -108,3 +108,13 @@ class AddMultiTeamToTIRunContext(VersionChange):
108108
def remove_multi_team_field(response: ResponseInfo) -> None: # type: ignore[misc]
109109
"""Strip ``multi_team`` from the run context for older clients."""
110110
response.body.pop("multi_team", None)
111+
112+
113+
class AddXComKeysEndpoint(VersionChange):
114+
"""Add the ``xcoms/{dag_id}/{run_id}/{task_id}/keys`` endpoint that fetches several XCom values by key in one query."""
115+
116+
description = __doc__
117+
118+
instructions_to_migrate_to_previous_version = (
119+
endpoint("/xcoms/{dag_id}/{run_id}/{task_id}/keys", ["POST"]).didnt_exist,
120+
)

‎airflow-core/src/airflow/dag_processing/processor.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -55,6 +55,7 @@
5555
GetVariable,
5656
GetVariableKeys,
5757
GetXCom,
58+
GetXComByKeys,
5859
GetXComCount,
5960
GetXComSequenceItem,
6061
GetXComSequenceSlice,
@@ -150,6 +151,7 @@ class DagFileParsingResult(BaseModel):
150151
| GetXCom
151152
| GetXComCount
152153
| GetXComSequenceItem
154+
| GetXComByKeys
153155
| GetXComSequenceSlice
154156
| MaskSecret,
155157
Field(discriminator="type"),
@@ -636,6 +638,7 @@ def _handle_parsing_result(
636638
GetVariable,
637639
GetVariableKeys,
638640
GetXCom,
641+
GetXComByKeys,
639642
GetXComCount,
640643
GetXComSequenceItem,
641644
GetXComSequenceSlice,

‎airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_xcoms.py‎

Lines changed: 107 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -79,11 +79,10 @@ def _(
7979
dag_id: str = Path(),
8080
run_id: str = Path(),
8181
task_id: str = Path(),
82-
xcom_key: str = Path(alias="key"),
8382
token=CurrentTIToken,
8483
):
8584
with create_session() as session:
86-
has_xcom_access(dag_id, run_id, task_id, xcom_key, request, session, token)
85+
has_xcom_access(dag_id, run_id, task_id, request, session, token)
8786
raise HTTPException(
8887
status_code=status.HTTP_403_FORBIDDEN,
8988
detail={
@@ -875,3 +874,109 @@ def test_teamless_requester_scoping(self, client, session, dag_maker):
875874

876875
assert forbidden.status_code == 403, forbidden.json()
877876
assert allowed.status_code == 200, allowed.json()
877+
878+
879+
class TestGetXComByKeys:
880+
"""Tests for the POST /{dag_id}/{run_id}/{task_id}/keys batched-fetch endpoint."""
881+
882+
def _insert_xcom(self, session, ti, key, value):
883+
session.add(XComModelV2(key=key, value=value, task_instance_id=ti.id))
884+
885+
def test_returns_values_in_request_key_order(self, client, create_task_instance, session):
886+
"""Values are returned in the same order as the requested keys."""
887+
ti = create_task_instance()
888+
self._insert_xcom(session, ti, "return_value_0", "alpha")
889+
self._insert_xcom(session, ti, "return_value_1", "beta")
890+
self._insert_xcom(session, ti, "return_value_2", "gamma")
891+
session.commit()
892+
893+
response = client.post(
894+
f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/keys",
895+
json={"keys": ["return_value_2", "return_value_0", "return_value_1"]},
896+
)
897+
898+
assert response.status_code == 200
899+
assert response.json() == ["gamma", "alpha", "beta"]
900+
901+
def test_missing_key_returned_as_none(self, client, create_task_instance, session):
902+
"""Keys not found in the database are returned as null (None)."""
903+
ti = create_task_instance()
904+
self._insert_xcom(session, ti, "return_value_0", "exists")
905+
session.commit()
906+
907+
response = client.post(
908+
f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/keys",
909+
json={"keys": ["return_value_0", "return_value_1"]},
910+
)
911+
912+
assert response.status_code == 200
913+
assert response.json() == ["exists", None]
914+
915+
def test_empty_key_list_returns_empty_list(self, client, create_task_instance, session):
916+
ti = create_task_instance()
917+
session.commit()
918+
919+
response = client.post(
920+
f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/keys",
921+
json={"keys": []},
922+
)
923+
924+
assert response.status_code == 200
925+
assert response.json() == []
926+
927+
def test_map_index_filter(self, client, dag_maker, session):
928+
"""Only XCom rows matching the requested map_index are returned."""
929+
930+
class MyOperator(EmptyOperator):
931+
def __init__(self, *, x, **kwargs):
932+
super().__init__(**kwargs)
933+
self.x = x
934+
935+
with dag_maker(dag_id="dag"):
936+
MyOperator.partial(task_id="task").expand(x=["a", "b"])
937+
dag_run = dag_maker.create_dagrun(run_id="run")
938+
tis = {ti.map_index: ti for ti in dag_run.task_instances}
939+
940+
# Insert XComs for map_index=0 and map_index=1 (both have real TI rows).
941+
# map_index=-1 has no TI row for a mapped task, so it is not inserted.
942+
session.add(XComModelV2(key="return_value_0", value="for_index_0", task_instance_id=tis[0].id))
943+
session.add(XComModelV2(key="return_value_0", value="for_index_1", task_instance_id=tis[1].id))
944+
session.commit()
945+
946+
dag_id = tis[0].dag_id
947+
task_id = tis[0].task_id
948+
949+
# Default map_index=-1: no rows stored there, so null is returned.
950+
response_default = client.post(
951+
f"/execution/xcoms/{dag_id}/{dag_run.run_id}/{task_id}/keys",
952+
json={"keys": ["return_value_0"]},
953+
)
954+
response_index_0 = client.post(
955+
f"/execution/xcoms/{dag_id}/{dag_run.run_id}/{task_id}/keys",
956+
params={"map_index": 0},
957+
json={"keys": ["return_value_0"]},
958+
)
959+
response_index_1 = client.post(
960+
f"/execution/xcoms/{dag_id}/{dag_run.run_id}/{task_id}/keys",
961+
params={"map_index": 1},
962+
json={"keys": ["return_value_0"]},
963+
)
964+
965+
assert response_default.json() == [None]
966+
assert response_index_0.json() == ["for_index_0"]
967+
assert response_index_1.json() == ["for_index_1"]
968+
969+
def test_does_not_return_other_tasks_xcoms(self, client, create_task_instance, session):
970+
"""Keys from a different task_id are not returned."""
971+
ti = create_task_instance()
972+
other_ti = create_task_instance(dag_id="other_dag", task_id="other_task")
973+
self._insert_xcom(session, other_ti, "return_value_0", "not_mine")
974+
session.commit()
975+
976+
response = client.post(
977+
f"/execution/xcoms/{ti.dag_id}/{ti.run_id}/{ti.task_id}/keys",
978+
json={"keys": ["return_value_0"]},
979+
)
980+
981+
assert response.status_code == 200
982+
assert response.json() == [None]

‎airflow-core/tests/unit/jobs/test_triggerer_job.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3054,6 +3054,7 @@ def get_type_names(union_type):
30543054
"GetPreviousDagRun",
30553055
"GetTaskBreadcrumbs",
30563056
"GetTaskRescheduleStartDate",
3057+
"GetXComByKeys",
30573058
"GetXComCount",
30583059
"GetXComSequenceItem",
30593060
"GetXComSequenceSlice",

‎task-sdk/src/airflow/sdk/api/client.py‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -746,6 +746,22 @@ def get_sequence_slice(
746746
resp = self.client.get(f"xcoms/{dag_id}/{run_id}/{task_id}/{key}/slice", params=params)
747747
return XComSequenceSliceResponse.model_validate_json(resp.read())
748748

749+
def get_by_keys(
750+
self,
751+
dag_id: str,
752+
run_id: str,
753+
task_id: str,
754+
keys: list[str],
755+
map_index: int = -1,
756+
) -> XComSequenceSliceResponse:
757+
"""Fetch multiple XCom values by key list in a single round-trip."""
758+
resp = self.client.post(
759+
f"xcoms/{dag_id}/{run_id}/{task_id}/keys",
760+
params={"map_index": map_index} if map_index >= 0 else {},
761+
json={"keys": keys},
762+
)
763+
return XComSequenceSliceResponse.model_validate_json(resp.read())
764+
749765

750766
class TaskStateStoreOperations:
751767
__slots__ = ("client",)

‎task-sdk/src/airflow/sdk/api/datamodels/_generated.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -529,6 +529,14 @@ class VariableResponse(BaseModel):
529529
value: Annotated[str | None, Field(title="Value")]
530530

531531

532+
class XComKeysRequest(BaseModel):
533+
"""
534+
Request body for fetching multiple XCom values by key list in a single query.
535+
"""
536+
537+
keys: Annotated[list[str], Field(title="Keys")]
538+
539+
532540
class XComResponse(BaseModel):
533541
"""
534542
XCom schema for responses with fields that are needed for Runtime.

0 commit comments

Comments
 (0)