This is an automated email from the ASF dual-hosted git repository.
jason810496 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new e08216d0081 Go SDK: delete the XComs of earlier tries before a task
runs (#74178)
e08216d0081 is described below
commit e08216d008101e471421da992c5e20999aabdcc5
Author: PoAn Yang <[email protected]>
AuthorDate: Mon Oct 5 20:33:14 2026 +0800
Go SDK: delete the XComs of earlier tries before a task runs (#74178)
* Go SDK: delete the XComs of earlier tries before a task runs
Signed-off-by: PoAn Yang <[email protected]>
* Go SDK: share the map_index omission rule between PushXCom and deleteXCom
Signed-off-by: PoAn Yang <[email protected]>
---------
Signed-off-by: PoAn Yang <[email protected]>
---
go-sdk/airflow/if.go | 3 -
go-sdk/airflow/if_test.go | 33 ++--
go-sdk/internal/bundle/task.go | 22 +--
go-sdk/internal/bundle/task_test.go | 35 +---
go-sdk/pkg/execution/client.go | 32 ++-
go-sdk/pkg/execution/client_test.go | 59 ++++++
go-sdk/pkg/execution/integration_test.go | 326 ++++++++++++++++++++++++++++++-
go-sdk/pkg/execution/task_runner.go | 64 ++++--
8 files changed, 471 insertions(+), 103 deletions(-)
diff --git a/go-sdk/airflow/if.go b/go-sdk/airflow/if.go
index 2d5b79d7912..a5c7bba0c45 100644
--- a/go-sdk/airflow/if.go
+++ b/go-sdk/airflow/if.go
@@ -185,9 +185,6 @@ func (g *IfRef) setTask(side string, task *TaskRef) {
// does not take.
func (g *IfRef) wrapCondition(fn any) (bundle.Task, error) {
fnType := reflect.TypeOf(fn)
- // DagRef.Task also takes a function whose last result has a concrete
type that implements
- // error. A nil value of that type becomes a non-nil error when the
runtime reads it, so the
- // condition task would always fail.
if fnType.NumOut() != 2 ||
fnType.Out(0) != reflect.TypeFor[bool]() ||
fnType.Out(1) != reflect.TypeFor[error]() {
diff --git a/go-sdk/airflow/if_test.go b/go-sdk/airflow/if_test.go
index 2f5730547c3..c61dae205e7 100644
--- a/go-sdk/airflow/if_test.go
+++ b/go-sdk/airflow/if_test.go
@@ -21,7 +21,6 @@ import (
"context"
"errors"
"fmt"
- "maps"
"reflect"
"strings"
"testing"
@@ -493,16 +492,13 @@ func (c *conditionClient) skip(_ context.Context, taskIDs
[]string) error {
// runCondition runs the task of gate through its Execute method, as the
runtime runs a task. It
// passes one XCom binding per input, named after the arg tag of rowSet, so
that a struct
// parameter shows whether it takes the whole result. results maps the task_id
of each upstream
-// task to its result. earlier maps each key to an XCom that an earlier try of
the task left.
-func runCondition(gate *IfRef, results, earlier map[string]any)
(*conditionClient, error) {
+// task to its result.
+func runCondition(gate *IfRef, results map[string]any) (*conditionClient,
error) {
args := make([]binding.Arg, len(gate.task.inputs))
for i, upstream := range gate.task.inputs {
args[i] = binding.XComArg{Kind: "xcom", Name: "rows", TaskID:
upstream.taskID}
}
- client := &conditionClient{results: results, xcoms: maps.Clone(earlier)}
- if client.xcoms == nil {
- client.xcoms = map[string]any{}
- }
+ client := &conditionClient{results: results, xcoms: map[string]any{}}
ti := sdk.TaskInstance{DagID: "etl", RunID: "run1", TaskID:
gate.task.taskID}
ctx := context.WithValue(
context.Background(),
@@ -586,19 +582,20 @@ func TestConditionSkipsTheSideThatItDoesNotTake(t
*testing.T) {
client, err := runCondition(gate, map[string]any{
"readRows": map[string]any{"rows": tt.rows},
- }, nil)
+ })
require.NoError(t, err)
assert.Equal(t, tt.want, client.xcoms["return_value"])
if len(tt.skipped) == 0 {
assert.Empty(t, client.skipped)
+ assert.NotContains(t, client.xcoms,
"skipmixin_key")
} else {
assert.Equal(t, [][]string{tt.skipped},
client.skipped)
+ assert.Equal(t,
+ map[string][]string{"skipped":
tt.skipped},
+ client.xcoms["skipmixin_key"],
+ )
}
- assert.Equal(t,
- map[string][]string{"skipped": tt.skipped},
- client.xcoms["skipmixin_key"],
- )
})
}
}
@@ -611,19 +608,11 @@ func TestConditionThatFailsSkipsNothing(t *testing.T) {
).Then(dag.Task(load)).Else(dag.Task(reportEmpty))
Bundle().Register(dag)
- // An earlier try of the condition skipped load, and this try fails
before it decides which
- // side to skip.
- client, err := runCondition(gate, nil, map[string]any{
- "skipmixin_key": map[string][]string{"skipped": {"load"}},
- })
+ client, err := runCondition(gate, nil)
require.EqualError(t, err, "cannot reach the table")
assert.Empty(t, client.skipped)
- assert.Equal(t,
- map[string][]string{"skipped": {}},
- client.xcoms["skipmixin_key"],
- "the list of the earlier try must not be left behind",
- )
+ assert.NotContains(t, client.xcoms, "skipmixin_key")
}
// Else and Register both take the lock of the Dag, so each Else call either
names its task before
diff --git a/go-sdk/internal/bundle/task.go b/go-sdk/internal/bundle/task.go
index 9db3626a1ff..73fba021653 100644
--- a/go-sdk/internal/bundle/task.go
+++ b/go-sdk/internal/bundle/task.go
@@ -78,9 +78,9 @@ func NewPositionalTaskFunction(fn any) (Task, error) {
// NewPositionalBranchFunction is like NewPositionalTaskFunction, but the Task
also skips tasks
// that are downstream of it. fn must return a result and an error. Before fn
runs, Execute checks
-// that the runtime can skip tasks, and writes an empty list to the
skipmixin_key XCom of the task.
-// When fn returns a nil error, Execute passes the result to findSkipped. If
findSkipped returns
-// task_ids, Execute records them in that XCom and skips those tasks.
+// that the runtime can skip tasks. When fn returns a nil error, Execute
passes the result to
+// findSkipped. If findSkipped returns task_ids, Execute records them in the
skipmixin_key XCom of
+// the task and skips those tasks.
func NewPositionalBranchFunction(fn any, findSkipped func(result any)
[]string) (Task, error) {
return newTaskFunction(fn, binding.AnalyzePositional, findSkipped)
}
@@ -118,7 +118,7 @@ func (f *taskFunction) Execute(
}
var branch *branchRun
if f.findSkipped != nil {
- if branch, err = startBranch(ctx, sdkClient); err != nil {
+ if branch, err = startBranch(ctx); err != nil {
return err
}
}
@@ -200,11 +200,8 @@ type branchRun struct {
}
// startBranch runs before fn, so that a runtime that cannot skip tasks fails
the task before fn
-// has any effect. It writes an empty list to the skipmixin_key XCom because
the Go runtime does
-// not delete the XComs of earlier tries. Without the write,
NotPreviouslySkippedDep could read
-// the list of an earlier try after this try fails or skips nothing. The list
must be empty rather
-// than null, because NotPreviouslySkippedDep cannot read null.
-func startBranch(ctx context.Context, client sdk.Client) (*branchRun, error) {
+// has any effect.
+func startBranch(ctx context.Context) (*branchRun, error) {
skip, ok := ctx.Value(skipDownstreamTasksKey{}).(func(context.Context,
[]string) error)
if !ok {
return nil, errors.New("the task runtime cannot skip downstream
tasks")
@@ -213,12 +210,7 @@ func startBranch(ctx context.Context, client sdk.Client)
(*branchRun, error) {
if !ok {
return nil, errors.New("task runtime context is missing")
}
- ti := runtimeContext.TaskInstance()
- err := client.PushXCom(ctx, ti, skipMixinXComKey,
map[string][]string{"skipped": {}})
- if err != nil {
- return nil, fmt.Errorf("clearing the %s XCom: %w",
skipMixinXComKey, err)
- }
- return &branchRun{skip: skip, ti: ti}, nil
+ return &branchRun{skip: skip, ti: runtimeContext.TaskInstance()}, nil
}
// skipDownstream records taskIDs in the skipmixin_key XCom, and then skips
those tasks.
diff --git a/go-sdk/internal/bundle/task_test.go
b/go-sdk/internal/bundle/task_test.go
index 40e023f784a..850d5ccaf72 100644
--- a/go-sdk/internal/bundle/task_test.go
+++ b/go-sdk/internal/bundle/task_test.go
@@ -293,8 +293,6 @@ func runBranch(task Task, client *branchClient, ti, canSkip
bool) error {
return task.Execute(ctx, slog.New(logging.NewTeeLogger()), nil)
}
-const clearCall = "PushXCom decide skipmixin_key map[skipped:[]]"
-
func (s *TaskSuite) TestBranchFunctionSkipsTheTasksThatFindSkippedReturns() {
var got any
task, err := NewPositionalBranchFunction(
@@ -311,7 +309,6 @@ func (s *TaskSuite)
TestBranchFunctionSkipsTheTasksThatFindSkippedReturns() {
s.Equal(true, got)
s.Equal([]string{
- clearCall,
"PushXCom decide return_value true",
"PushXCom decide skipmixin_key map[skipped:[load report]]",
"SkipDownstreamTasks [load report]",
@@ -330,20 +327,13 @@ func (s *TaskSuite) TestBranchFunctionWithNothingToSkip()
{
client := &branchClient{}
s.Require().NoError(runBranch(task, client, true, true))
- s.Equal([]string{clearCall, "PushXCom decide
return_value true"}, client.calls)
- s.Equal(
- map[string][]string{"skipped": {}},
- client.values["skipmixin_key"],
- "the list must not be nil",
- )
+ s.Equal([]string{"PushXCom decide return_value true"},
client.calls)
})
}
}
-// An earlier try of the task may have left a list of skipped tasks in the
XCom. A try that fails
-// before it skips anything still replaces that list with an empty list,
whether fn returns an
-// error or panics.
-func (s *TaskSuite) TestBranchFunctionClearsTheListOfAnEarlierTry() {
+// A try that fails neither records nor skips any task, whether fn returns an
error or panics.
+func (s *TaskSuite) TestBranchFunctionThatFailsSkipsNothing() {
cases := map[string]func(contexttest.Context) (bool, error){
"error": func(contexttest.Context) (bool, error) { return
false, errors.New("no table") },
"panic": func(contexttest.Context) (bool, error) { panic("no
table") },
@@ -357,21 +347,15 @@ func (s *TaskSuite)
TestBranchFunctionClearsTheListOfAnEarlierTry() {
})
s.Require().NoError(err)
- client := &branchClient{values: map[string]any{
- "skipmixin_key": map[string][]string{"skipped":
{"load"}},
- }}
+ client := &branchClient{}
func() {
defer func() { _ = recover() }()
s.Error(runBranch(task, client, true, true))
}()
s.False(called)
- s.Equal(
- map[string][]string{"skipped": {}},
- client.values["skipmixin_key"],
- "a try that skipped nothing must not leave the
list of an earlier try",
- )
for _, call := range client.calls {
+ s.NotContains(call, "skipmixin_key")
s.NotContains(call, "SkipDownstreamTasks")
}
})
@@ -397,13 +381,6 @@ func (s *TaskSuite)
TestBranchFunctionFailsWhenItCannotSkip() {
canSkip: true,
wantErr: "task runtime context is missing",
},
- "clearing the skipmixin_key XCom fails": {
- client: &branchClient{failOn: clearCall},
- withTI: true,
- canSkip: true,
- wantErr: "clearing the skipmixin_key XCom: xcom
refused",
- wantCalls: []string{clearCall},
- },
"recording the skipped tasks fails": {
client: &branchClient{
failOn: "PushXCom decide skipmixin_key
map[skipped:[load]]",
@@ -413,7 +390,6 @@ func (s *TaskSuite)
TestBranchFunctionFailsWhenItCannotSkip() {
wantErr: "recording the skipped tasks in the
skipmixin_key XCom: xcom refused",
wantRun: true,
wantCalls: []string{
- clearCall,
"PushXCom decide return_value false",
"PushXCom decide skipmixin_key
map[skipped:[load]]",
},
@@ -425,7 +401,6 @@ func (s *TaskSuite)
TestBranchFunctionFailsWhenItCannotSkip() {
wantErr: `skipping the downstream tasks ["load"]:
supervisor refused`,
wantRun: true,
wantCalls: []string{
- clearCall,
"PushXCom decide return_value false",
"PushXCom decide skipmixin_key
map[skipped:[load]]",
"SkipDownstreamTasks [load]",
diff --git a/go-sdk/pkg/execution/client.go b/go-sdk/pkg/execution/client.go
index b98d8883ed9..1f09da1e1ba 100644
--- a/go-sdk/pkg/execution/client.go
+++ b/go-sdk/pkg/execution/client.go
@@ -242,18 +242,38 @@ func (c *CoordinatorClient) PushXCom(
TaskID: ti.TaskID,
RunID: ti.RunID,
}
- // map_index mirrors Python's SetXCom.map_index (int | None): -1 is the
- // unmapped sentinel, omitted from the payload rather than sent. Assign
the
- // pointer, not the dereferenced int, so an explicit index 0 survives
omitempty
- // (see GetXCom).
- if ti.MapIndex != nil && *ti.MapIndex != -1 {
- msg.MapIndex = ti.MapIndex
+ msg.MapIndex = omittedMapIndex(ti.MapIndex)
+
+ _, err := c.comm.Communicate(ctx, msg)
+ return err
+}
+
+// deleteXCom asks the supervisor to delete the XCom of ti with the given key.
Like PushXCom, it
+// leaves map_index out for an unmapped task instance, and the Execution API
then deletes the XCom
+// with map_index -1.
+func (c *CoordinatorClient) deleteXCom(ctx context.Context, ti
sdk.TaskInstance, key string) error {
+ msg := genmodels.DeleteXCom{
+ Key: key,
+ DagID: ti.DagID,
+ TaskID: ti.TaskID,
+ RunID: ti.RunID,
}
+ msg.MapIndex = omittedMapIndex(ti.MapIndex)
_, err := c.comm.Communicate(ctx, msg)
return err
}
+// omittedMapIndex returns mapIndex, or nil for the unmapped sentinel -1, so
that msgpack omits
+// map_index from the payload instead of sending it. An explicit index 0
survives omitempty because
+// the pointer, not the dereferenced int, is returned (see GetXCom).
+func omittedMapIndex(mapIndex *int) *int {
+ if mapIndex == nil || *mapIndex == -1 {
+ return nil
+ }
+ return mapIndex
+}
+
// skipDownstreamTasks asks the supervisor to mark the tasks with the given
task_ids as skipped
// in the Dag run of the running task. Airflow does not change a task instance
that is running,
// has succeeded or has failed.
diff --git a/go-sdk/pkg/execution/client_test.go
b/go-sdk/pkg/execution/client_test.go
index b7c8d39e7e8..da15844c621 100644
--- a/go-sdk/pkg/execution/client_test.go
+++ b/go-sdk/pkg/execution/client_test.go
@@ -516,3 +516,62 @@ func TestCoordinatorClientSkipDownstreamTasks(t
*testing.T) {
})
}
}
+
+// TestCoordinatorClientDeleteXCom verifies the DeleteXCom frame sent to the
supervisor, and that
+// deleteXCom returns a supervisor ErrorResponse as an error. Like PushXCom,
it leaves map_index out
+// for an unmapped task instance and sends index 0.
+func TestCoordinatorClientDeleteXCom(t *testing.T) {
+ tests := []struct {
+ name string
+ mapIndex *int
+ errBody map[string]any
+ wantMapIndex any
+ }{
+ {name: "nil map_index is omitted"},
+ {name: "-1 map_index is omitted", mapIndex: ptr(-1)},
+ {name: "map_index 0 is sent", mapIndex: ptr(0), wantMapIndex:
int8(0)},
+ {name: "map_index 3 is sent", mapIndex: ptr(3), wantMapIndex:
int8(3)},
+ {
+ name: "supervisor error",
+ errBody: map[string]any{
+ "type": "ErrorResponse",
+ "error": "API_SERVER_ERROR",
+ "detail": map[string]any{"status_code": 500},
+ },
+ },
+ }
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf,
encodeResponseFrame(t, 0, nil, tc.errBody)))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ client :=
NewCoordinatorClient(NewCoordinatorComm(&responseBuf, &requestBuf, logger))
+
+ ti := sdk.TaskInstance{DagID: "d", RunID: "r", TaskID:
"t", MapIndex: tc.mapIndex}
+ err := client.deleteXCom(context.Background(), ti,
"skipmixin_key")
+ if tc.errBody != nil {
+ var apiErr *APIError
+ require.ErrorAs(t, err, &apiErr)
+ assert.Equal(t, "API_SERVER_ERROR", apiErr.Err)
+ } else {
+ require.NoError(t, err)
+ }
+
+ sent, err := readFrame(&requestBuf)
+ require.NoError(t, err)
+ want := map[string]any{
+ "type": "DeleteXCom",
+ "dag_id": "d",
+ "run_id": "r",
+ "task_id": "t",
+ "key": "skipmixin_key",
+ }
+ if tc.wantMapIndex != nil {
+ want["map_index"] = tc.wantMapIndex
+ }
+ assert.Equal(t, want, rawToMap(t, sent.Body))
+ })
+ }
+}
diff --git a/go-sdk/pkg/execution/integration_test.go
b/go-sdk/pkg/execution/integration_test.go
index 66d6611f084..2ee83df8f8e 100644
--- a/go-sdk/pkg/execution/integration_test.go
+++ b/go-sdk/pkg/execution/integration_test.go
@@ -22,6 +22,7 @@ import (
"context"
"encoding/json"
"errors"
+ "fmt"
"io"
"log/slog"
"net"
@@ -179,12 +180,15 @@ func TestTaskRunnerTaskNotFound(t *testing.T) {
})
details := newStartupDetails("nonexistent")
+ details.TIContext.XcomKeysToClear = []string{"return_value"}
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
- comm := NewCoordinatorComm(bytes.NewReader(nil), io.Discard, logger)
+ var sent bytes.Buffer
+ comm := NewCoordinatorComm(bytes.NewReader(nil), &sent, logger)
result := RunTask(context.Background(), bundle, details, comm, logger)
assertTaskState(t, result, genmodels.TaskStateStateRemoved)
+ assert.Zero(t, sent.Len(), "a task that is not in the bundle must not
delete any XCom")
}
func TestTaskRunnerPanic(t *testing.T) {
@@ -852,11 +856,9 @@ func TestServeClientRoundTripEndToEnd(t *testing.T) {
}
// TestServeSkipsDownstreamTasksEndToEnd drives a task that skips downstream
tasks through the
-// real Serve. Before the terminal SucceedTask frame, the supervisor gets an
empty skipmixin_key
-// XCom and the return value XCom. If there is a task to skip, the
skipmixin_key XCom with its
-// task_id and the SkipDownstreamTasks request follow, in that order. The
empty list replaces any
-// list that an earlier try of the task left in the XCom. It has to arrive as
a list, because
-// NotPreviouslySkippedDep raises a TypeError on null.
+// real Serve. Before the terminal SucceedTask frame, the supervisor gets the
return value XCom.
+// If there is a task to skip, the skipmixin_key XCom with its task_id and the
SkipDownstreamTasks
+// request follow, in that order.
func TestServeSkipsDownstreamTasksEndToEnd(t *testing.T) {
xcom := func(key string, value any) map[string]any {
return map[string]any{
@@ -879,7 +881,6 @@ func TestServeSkipsDownstreamTasksEndToEnd(t *testing.T) {
name: "false skips load",
result: false,
wantRequests: []map[string]any{
- xcom("skipmixin_key", map[string]any{"skipped":
[]any{}}),
xcom("return_value", false),
xcom("skipmixin_key", map[string]any{"skipped":
[]any{"load"}}),
{"type": "SkipDownstreamTasks", "tasks":
[]any{"load"}},
@@ -890,7 +891,6 @@ func TestServeSkipsDownstreamTasksEndToEnd(t *testing.T) {
name: "true skips nothing",
result: true,
wantRequests: []map[string]any{
- xcom("skipmixin_key", map[string]any{"skipped":
[]any{}}),
xcom("return_value", true),
},
},
@@ -975,6 +975,316 @@ func TestServeSkipsDownstreamTasksEndToEnd(t *testing.T) {
}
}
+// serveTask runs task as the task instance taskID of dag1 through the real
Serve. The
+// StartupDetails frame carries mapIndex, unless it is nil, and tiContext as
ti_context. The fake
+// supervisor passes each runtime request to answer, which returns the
ErrorResponse to reply with,
+// or nil to reply with an empty response. serveTask returns the requests in
the order they
+// arrived, and the body of the terminal frame.
+func serveTask(
+ t *testing.T,
+ taskID string,
+ task bundle.Task,
+ mapIndex *int,
+ tiContext map[string]any,
+ answer func(request map[string]any) map[string]any,
+) (requests []map[string]any, terminal map[string]any) {
+ t.Helper()
+ commAddr, logsAddr, commCh, logsCh, cleanup := startSupervisor(t)
+ defer cleanup()
+
+ done := make(chan error, 1)
+ go func() { done <- Serve(testBundle{"dag1": testDag{taskID: task}},
commAddr, logsAddr) }()
+
+ commConn := <-commCh
+ defer commConn.Close()
+ logsConn := <-logsCh
+ defer logsConn.Close()
+ go func() { _, _ = io.Copy(io.Discard, logsConn) }()
+ require.NoError(t, commConn.SetDeadline(time.Now().Add(10*time.Second)))
+
+ ti := map[string]any{
+ "id": "550e8400-e29b-41d4-a716-446655440000",
+ "dag_id": "dag1",
+ "task_id": taskID,
+ "run_id": "run1",
+ "try_number": 2,
+ }
+ if mapIndex != nil {
+ ti["map_index"] = *mapIndex
+ }
+ startup, err := encodeRequest(0, map[string]any{
+ "type": "StartupDetails",
+ "ti": ti,
+ "ti_context": tiContext,
+ "bundle_info": map[string]any{"name": "fake", "version": "1.0"},
+ })
+ require.NoError(t, err)
+ require.NoError(t, writeFrame(commConn, startup))
+
+ for {
+ frame, err := readFrame(commConn)
+ require.NoError(t, err)
+ require.True(t, isNilRaw(frame.Err))
+ body := rawToMap(t, frame.Body)
+ switch body["type"] {
+ case "SucceedTask", "TaskState", "RetryTask":
+ terminal = body
+ }
+ if terminal != nil {
+ break
+ }
+ requests = append(requests, body)
+ var reply []byte
+ if errBody := answer(body); errBody != nil {
+ reply = encodeResponseFrame(t, frame.ID, nil, errBody)
+ } else {
+ reply, err = encodeRequest(frame.ID, map[string]any{})
+ require.NoError(t, err)
+ }
+ require.NoError(t, writeFrame(commConn, reply))
+ }
+
+ select {
+ case err := <-done:
+ require.NoError(t, err)
+ case <-time.After(2 * time.Second):
+ t.Fatal("Serve did not return after task completion")
+ }
+ return requests, terminal
+}
+
+// answerAll replies to every runtime request with an empty response.
+func answerAll(map[string]any) map[string]any { return nil }
+
+// TestServeClearsTheXComsOfEarlierTriesEndToEnd pins that the runtime deletes
each XCom that
+// ti_context.xcom_keys_to_clear lists before the task sends anything. The
DeleteXCom frame leaves
+// map_index out for an unmapped task instance and carries the index of a
mapped one, 0 included.
+func TestServeClearsTheXComsOfEarlierTriesEndToEnd(t *testing.T) {
+ deleteXCom := func(key string, mapIndex any) map[string]any {
+ frame := map[string]any{
+ "type": "DeleteXCom",
+ "dag_id": "dag1",
+ "run_id": "run1",
+ "task_id": "extract",
+ "key": key,
+ }
+ if mapIndex != nil {
+ frame["map_index"] = mapIndex
+ }
+ return frame
+ }
+ returnValue := func(mapIndex any) map[string]any {
+ frame := map[string]any{
+ "type": "SetXCom",
+ "dag_id": "dag1",
+ "run_id": "run1",
+ "task_id": "extract",
+ "key": "return_value",
+ "value": "rows",
+ }
+ if mapIndex != nil {
+ frame["map_index"] = mapIndex
+ }
+ return frame
+ }
+ tests := []struct {
+ name string
+ mapIndex *int
+ keys []any
+ wantRequests []map[string]any
+ }{
+ {
+ name: "unmapped",
+ mapIndex: ptr(-1),
+ keys: []any{"return_value", "skipmixin_key"},
+ wantRequests: []map[string]any{
+ deleteXCom("return_value", nil),
+ deleteXCom("skipmixin_key", nil),
+ returnValue(nil),
+ },
+ },
+ {
+ name: "map index 0",
+ mapIndex: ptr(0),
+ keys: []any{"return_value"},
+ wantRequests: []map[string]any{
+ deleteXCom("return_value", int8(0)),
+ returnValue(int8(0)),
+ },
+ },
+ {
+ name: "map index 2",
+ mapIndex: ptr(2),
+ keys: []any{"return_value"},
+ wantRequests: []map[string]any{
+ deleteXCom("return_value", int8(2)),
+ returnValue(int8(2)),
+ },
+ },
+ {
+ name: "no keys",
+ mapIndex: ptr(-1),
+ wantRequests: []map[string]any{returnValue(nil)},
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ extract, err := bundle.NewTaskFunction(
+ func(contexttest.Context) (string, error) {
return "rows", nil },
+ )
+ require.NoError(t, err)
+ tiContext := map[string]any{}
+ if tt.keys != nil {
+ tiContext["xcom_keys_to_clear"] = tt.keys
+ }
+
+ requests, terminal := serveTask(
+ t,
+ "extract",
+ extract,
+ tt.mapIndex,
+ tiContext,
+ answerAll,
+ )
+
+ assert.Equal(t, tt.wantRequests, requests)
+ assert.Equal(t, "SucceedTask", terminal["type"])
+ })
+ }
+}
+
+// TestServeFailsWhenItCannotClearAnXComEndToEnd pins that the task does not
run when the runtime
+// cannot delete an XCom that ti_context.xcom_keys_to_clear lists. The task
then ends like a failed
+// task: as RetryTask when ti_context.should_retry is set, and as a FAILED
TaskState otherwise.
+func TestServeFailsWhenItCannotClearAnXComEndToEnd(t *testing.T) {
+ for _, shouldRetry := range []bool{true, false} {
+ t.Run(fmt.Sprintf("should_retry=%t", shouldRetry), func(t
*testing.T) {
+ ran := false
+ extract, err :=
bundle.NewTaskFunction(func(contexttest.Context) error {
+ ran = true
+ return nil
+ })
+ require.NoError(t, err)
+
+ requests, terminal := serveTask(t, "extract", extract,
nil, map[string]any{
+ "xcom_keys_to_clear": []any{"return_value",
"skipmixin_key"},
+ "should_retry": shouldRetry,
+ }, func(map[string]any) map[string]any {
+ return map[string]any{
+ "type": "ErrorResponse",
+ "error": "API_SERVER_ERROR",
+ "detail": map[string]any{"status_code":
500},
+ }
+ })
+
+ assert.False(t, ran, "the task must not run")
+ assert.Equal(t, []map[string]any{{
+ "type": "DeleteXCom",
+ "dag_id": "dag1",
+ "run_id": "run1",
+ "task_id": "extract",
+ "key": "return_value",
+ }}, requests, "the runtime stops at the first XCom it
cannot delete")
+ if shouldRetry {
+ assert.Equal(t, "RetryTask", terminal["type"])
+ assert.Contains(t, terminal["retry_reason"],
"API_SERVER_ERROR")
+ } else {
+ assert.Equal(t, "TaskState", terminal["type"])
+ assert.Equal(t, "failed", terminal["state"])
+ }
+ })
+ }
+}
+
+// TestServeClearsXComsBeforeItBindsArgumentsEndToEnd pins that the runtime
deletes the XComs
+// before it converts arg_bindings. When arg_bindings is invalid, the task
fails without running,
+// and the XComs are still gone, as they are for a Python task whose templates
fail to render.
+func TestServeClearsXComsBeforeItBindsArgumentsEndToEnd(t *testing.T) {
+ ran := false
+ extract, err := bundle.NewTaskFunction(func(contexttest.Context) error {
+ ran = true
+ return nil
+ })
+ require.NoError(t, err)
+
+ requests, terminal := serveTask(t, "extract", extract, nil,
map[string]any{
+ "xcom_keys_to_clear": []any{"skipmixin_key"},
+ "arg_bindings": []any{map[string]any{"kind": "bogus",
"name": "x"}},
+ }, answerAll)
+
+ assert.False(t, ran, "the task must not run")
+ assert.Equal(t, []map[string]any{{
+ "type": "DeleteXCom",
+ "dag_id": "dag1",
+ "run_id": "run1",
+ "task_id": "extract",
+ "key": "skipmixin_key",
+ }}, requests)
+ assert.Equal(t, "TaskState", terminal["type"])
+ assert.Equal(t, "failed", terminal["state"])
+}
+
+// TestServeConditionDoesNotLeaveTheListOfAnEarlierTryEndToEnd covers a task
that skips downstream
+// tasks and whose earlier try skipped load. When this try fails or skips
nothing, it writes no
+// list, so only the deletion keeps NotPreviouslySkippedDep from reading the
old list and skipping
+// load again. The fake supervisor keeps the XComs of the task instance to
show what is left.
+func TestServeConditionDoesNotLeaveTheListOfAnEarlierTryEndToEnd(t *testing.T)
{
+ tests := []struct {
+ name string
+ fn func(contexttest.Context) (bool, error)
+ wantTerminal string
+ }{
+ {
+ name: "fails",
+ fn: func(contexttest.Context) (bool, error) {
+ return false, errors.New("cannot reach the
table")
+ },
+ wantTerminal: "TaskState",
+ },
+ {
+ name: "skips nothing",
+ fn: func(contexttest.Context) (bool, error) {
return true, nil },
+ wantTerminal: "SucceedTask",
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ decide, err :=
bundle.NewPositionalBranchFunction(tt.fn, func(result any) []string {
+ if result.(bool) {
+ return nil
+ }
+ return []string{"load"}
+ })
+ require.NoError(t, err)
+
+ xcoms := map[string]any{
+ "return_value": false,
+ "skipmixin_key": map[string]any{"skipped":
[]any{"load"}},
+ }
+ requests, terminal := serveTask(t, "decide", decide,
nil, map[string]any{
+ "xcom_keys_to_clear": []any{"return_value",
"skipmixin_key"},
+ }, func(request map[string]any) map[string]any {
+ switch request["type"] {
+ case "DeleteXCom":
+ delete(xcoms, request["key"].(string))
+ case "SetXCom":
+ xcoms[request["key"].(string)] =
request["value"]
+ }
+ return nil
+ })
+
+ require.GreaterOrEqual(t, len(requests), 2)
+ assert.Equal(t, "DeleteXCom", requests[0]["type"])
+ assert.Equal(t, "DeleteXCom", requests[1]["type"])
+ for _, request := range requests {
+ assert.NotEqual(t, "SkipDownstreamTasks",
request["type"])
+ }
+ assert.NotContains(t, xcoms, "skipmixin_key")
+ assert.Equal(t, tt.wantTerminal, terminal["type"])
+ })
+ }
+}
+
// TestServeFailureAfterConnectClosesComm asserts the failure-signaling
// contract: when Serve fails after the sockets are connected, it returns the
// error (so the caller exits non-zero) without writing a terminal frame. The
diff --git a/go-sdk/pkg/execution/task_runner.go
b/go-sdk/pkg/execution/task_runner.go
index 13e2d2d0f05..509e9d0e7ba 100644
--- a/go-sdk/pkg/execution/task_runner.go
+++ b/go-sdk/pkg/execution/task_runner.go
@@ -34,8 +34,10 @@ import (
// RunTask executes a task based on StartupDetails received from the
supervisor.
//
// It looks up the task in the bundle, creates a CoordinatorClient for SDK
-// calls, executes the task, and returns the terminal body to ship as the final
-// response frame: one of genmodels.SucceedTask, TaskState, or RetryTask.
+// calls, deletes the XComs that ti_context.xcom_keys_to_clear lists, executes
+// the task, and returns the terminal body to ship as the final response frame:
+// one of genmodels.SucceedTask, TaskState, or RetryTask. When RunTask cannot
+// delete an XCom, or arg_bindings is invalid, the task fails without running.
//
// The supervisor owns the Execution-API state transitions, so the runtime only
// invokes the user's task function and returns its terminal response.
@@ -69,16 +71,17 @@ func RunTask(
// context is a placeholder, because binding reads only the task
instance
// and Dag run from this value and builds the airflow.Context around the
// live task context.
+ ti := sdk.TaskInstance{
+ DagID: details.TI.DagID,
+ RunID: details.TI.RunID,
+ TaskID: details.TI.TaskID,
+ MapIndex: mapIndexPtr(details.TI.MapIndex),
+ TryNumber: details.TI.TryNumber,
+ }
dagRun := details.TIContext.DagRun
runtimeContext := sdk.NewTIRunContext(
context.Background(),
- sdk.TaskInstance{
- DagID: details.TI.DagID,
- RunID: details.TI.RunID,
- TaskID: details.TI.TaskID,
- MapIndex: mapIndexPtr(details.TI.MapIndex),
- TryNumber: details.TI.TryNumber,
- },
+ ti,
sdk.DagRun{
DagID: details.TI.DagID,
RunID: details.TI.RunID,
@@ -92,6 +95,23 @@ func RunTask(
ctx = context.WithValue(ctx, sdkcontext.RuntimeContextKey,
runtimeContext)
ctx = bundle.WithSkipDownstreamTasks(ctx, client.skipDownstreamTasks)
+ // Airflow keeps the XComs that a task instance already has when it
starts a new try. It lists
+ // their keys in xcom_keys_to_clear, and the runtime deletes them
before the task runs, as the
+ // Python task runner does before it renders templates. The list is
empty when the task resumes
+ // from a deferral.
+ for _, key := range details.TIContext.XcomKeysToClear {
+ logger.Debug("Clearing XCom with key", "key", key)
+ if err := client.deleteXCom(ctx, ti, key); err != nil {
+ logger.Error("Unable to clear XCom",
+ "dag_id", details.TI.DagID,
+ "task_id", details.TI.TaskID,
+ "key", key,
+ "error", err,
+ )
+ return failTask(details.TIContext.ShouldRetry, err)
+ }
+ }
+
args, err := convertArgBindings(details.TIContext.ArgBindings)
if err != nil {
logger.Error("Invalid arg_bindings spec from supervisor",
@@ -99,21 +119,27 @@ func RunTask(
"task_id", details.TI.TaskID,
"error", err,
)
- if details.TIContext.ShouldRetry {
- return genmodels.RetryTask{
- EndDate: time.Now().UTC(),
- RetryReason: err.Error(),
- }
- }
- return genmodels.TaskState{
- State: genmodels.TaskStateStateFailed,
- EndDate: time.Now().UTC(),
- }
+ return failTask(details.TIContext.ShouldRetry, err)
}
return executeTask(ctx, task, args, details.TIContext.ShouldRetry,
logger)
}
+// failTask returns the terminal body for a task that fails before it runs:
RetryTask when
+// ti_context.should_retry is set, otherwise a FAILED TaskState.
+func failTask(shouldRetry bool, err error) any {
+ if shouldRetry {
+ return genmodels.RetryTask{
+ EndDate: time.Now().UTC(),
+ RetryReason: err.Error(),
+ }
+ }
+ return genmodels.TaskState{
+ State: genmodels.TaskStateStateFailed,
+ EndDate: time.Now().UTC(),
+ }
+}
+
func convertArgBindings(specsPtr *genmodels.ArgBindings) ([]binding.Arg,
error) {
if specsPtr == nil || len(*specsPtr) == 0 {
return nil, nil