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

Reply via email to