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 bf558f6f29a Go SDK: add dag.If with Then and Else so native Dags can 
branch (#74075)
bf558f6f29a is described below

commit bf558f6f29a70922f86a0ade6293b67c097dde10
Author: PoAn Yang <[email protected]>
AuthorDate: Sat Oct 3 19:19:32 2026 +0800

    Go SDK: add dag.If with Then and Else so native Dags can branch (#74075)
    
    * Go SDK: add dag.If with Then and Else so native Dags can branch
    
    Signed-off-by: PoAn Yang <[email protected]>
    
    * Go SDK: address review of dag.If
    
    Signed-off-by: PoAn Yang <[email protected]>
    
    ---------
    
    Signed-off-by: PoAn Yang <[email protected]>
---
 go-sdk/airflow/bundle.go                 |   7 +-
 go-sdk/airflow/dag.go                    |  81 ++--
 go-sdk/airflow/if.go                     | 221 +++++++++++
 go-sdk/airflow/if_test.go                | 643 +++++++++++++++++++++++++++++++
 go-sdk/airflow/inputs.go                 |  51 +--
 go-sdk/airflow/task_option.go            |  12 +-
 go-sdk/internal/bundle/task.go           | 112 +++++-
 go-sdk/internal/bundle/task_test.go      | 217 +++++++++++
 go-sdk/pkg/execution/client.go           |   8 +
 go-sdk/pkg/execution/client_test.go      |  47 +++
 go-sdk/pkg/execution/integration_test.go | 124 ++++++
 go-sdk/pkg/execution/task_runner.go      |   1 +
 12 files changed, 1459 insertions(+), 65 deletions(-)

diff --git a/go-sdk/airflow/bundle.go b/go-sdk/airflow/bundle.go
index 8f82b071211..91b2773fd75 100644
--- a/go-sdk/airflow/bundle.go
+++ b/go-sdk/airflow/bundle.go
@@ -76,14 +76,15 @@ type Registerable interface{ registerable() }
 //
 //     bundle.Register(reports.Handlers()...)
 //
-// Add every task to a Dag before registering the Dag. [DagRef.Task] panics 
once the Dag is
-// registered.
+// Add every task to a Dag before registering the Dag. [DagRef.Task], 
[DagRef.If], [IfRef.Then]
+// and [IfRef.Else] panic once the Dag is registered.
 //
 // Register panics if a task handler with the same dag_id and task_id is 
already registered,
 // if a Dag with the same dag_id is already registered, if a task handler and 
a Dag have the
 // same dag_id, or if [BundleRef.Serve] has already been called: registration 
closes when
 // serving starts. A task handler runs a task of a Python Dag, so its dag_id 
cannot also belong
-// to a Dag authored in Go.
+// to a Dag authored in Go. Register also panics if a Dag has a condition from 
[DagRef.If]
+// without a task from [IfRef.Then].
 func (b *BundleRef) Register(items ...Registerable) {
        if b.closed.Load() {
                panic(
diff --git a/go-sdk/airflow/dag.go b/go-sdk/airflow/dag.go
index 5f301676d15..bbfbf022dc6 100644
--- a/go-sdk/airflow/dag.go
+++ b/go-sdk/airflow/dag.go
@@ -50,7 +50,8 @@ type DagRef struct {
 //
 //     bundle.Register(dag)
 //
-// Add every task before Register. [DagRef.Task] panics once the Dag is 
registered.
+// Add every task before Register. [DagRef.Task], [DagRef.If], [IfRef.Then] 
and [IfRef.Else]
+// panic once the Dag is registered.
 //
 // [BundleRef.Serve] does not yet serve the Dags that Dag returns. It leaves 
them out of the
 // --airflow-metadata manifest and cannot run their tasks.
@@ -72,7 +73,8 @@ func Dag(dagID string, spec ...DagSpec) *DagRef {
 func (*DagRef) registerable() {}
 
 // TaskRef is a task that [DagRef.Task] added to a Dag. Pass it to [Inputs] to 
give its result
-// to a task that DagRef.Task adds later.
+// to a task that DagRef.Task or [DagRef.If] adds later. Pass it to 
[IfRef.Then] or [IfRef.Else]
+// to run it on one side of a condition.
 type TaskRef struct {
        dag    *DagRef
        taskID string
@@ -88,6 +90,9 @@ type TaskRef struct {
        // It is nil for a task that runs a Go function. A task from 
TriggerDagRun runs no Go
        // function, so its resultType, inputs and task are nil.
        triggerDagRun *TriggerDagRunSpec
+       // ifRef is the IfRef that DagRef.If returned for the task. It is nil 
for a task from
+       // DagRef.Task.
+       ifRef *IfRef
 }
 
 // Task adds a task that runs fn to the Dag and returns the new task.
@@ -127,12 +132,18 @@ type TaskRef struct {
 //   - the Dag already has a task with the same task_id
 //   - the Dag is already registered
 func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef {
+       return d.addTask("airflow.DagRef.Task", fn, opts, nil)
+}
+
+// addTask adds a task for Task and If. method names the caller in panic 
messages. ifRef is the
+// IfRef that If returns, and nil when Task calls addTask.
+func (d *DagRef) addTask(method string, fn any, opts []TaskOption, ifRef 
*IfRef) *TaskRef {
        trigger, isTrigger := fn.(TriggerDagRunTask)
        var triggerSpec *TriggerDagRunSpec
        var triggerErr error
        if isTrigger {
                // Copying the spec marshals Conf and so runs the MarshalJSON 
methods of the caller's
-               // values. Task copies before it locks d.mu. Otherwise such a 
method would deadlock if it
+               // values. addTask copies before it locks d.mu. Otherwise such 
a method would deadlock if it
                // added a task to this Dag, and a slow one would hold up 
Register.
                copied, err := copyTriggerDagRunSpec(trigger.spec)
                triggerSpec, triggerErr = &copied, err
@@ -143,36 +154,39 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) 
*TaskRef {
 
        if d.registered {
                panic(fmt.Sprintf(
-                       "airflow.DagRef.Task: Dag %q has already been 
registered; "+
-                               "add every task before Register",
-                       d.dagID,
+                       "%s: Dag %q has already been registered; add every task 
before Register",
+                       method, d.dagID,
                ))
        }
        var wrapped bundle.Task
        if !isTrigger {
+               wrap := bundle.NewPositionalTaskFunction
+               if ifRef != nil {
+                       wrap = ifRef.wrapCondition
+               }
                var err error
-               if wrapped, err = newTaskFunction(fn, 
bundle.NewPositionalTaskFunction); err != nil {
-                       panic(fmt.Sprintf("airflow.DagRef.Task: Dag %q: %v", 
d.dagID, err))
+               if wrapped, err = newTaskFunction(fn, wrap); err != nil {
+                       panic(fmt.Sprintf("%s: Dag %q: %v", method, d.dagID, 
err))
                }
        }
        var cfg taskConfig
        for i, opt := range opts {
                switch opt := opt.(type) {
                case nil:
-                       panic(fmt.Sprintf("airflow.DagRef.Task: Dag %q: 
opts[%d] is nil", d.dagID, i))
+                       panic(fmt.Sprintf("%s: Dag %q: opts[%d] is nil", 
method, d.dagID, i))
                case *TaskSpec:
                        if opt == nil {
                                panic(fmt.Sprintf(
-                                       "airflow.DagRef.Task: Dag %q: opts[%d] 
is a nil *airflow.TaskSpec", d.dagID, i,
+                                       "%s: Dag %q: opts[%d] is a nil 
*airflow.TaskSpec", method, d.dagID, i,
                                ))
                        }
                case TaskSpec, inputs:
                default:
                        // Only a struct that embeds a TaskSpec or a TaskOption 
gets here.
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: Dag %q: opts[%d] has type 
%T, "+
+                               "%s: Dag %q: opts[%d] has type %T, "+
                                        "which is not an option that package 
airflow defines",
-                               d.dagID, i, opt,
+                               method, d.dagID, i, opt,
                        ))
                }
                if err := opt.applyTask(&cfg); err != nil {
@@ -183,7 +197,7 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef {
                                continue
                        }
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q: %v", 
findTaskName(fn, opts), d.dagID, err,
+                               "%s: task %q of Dag %q: %v", method, 
findTaskName(fn, opts), d.dagID, err,
                        ))
                }
        }
@@ -191,34 +205,32 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) 
*TaskRef {
        taskID := cfg.spec.TaskID
        if taskID == "" && isTrigger {
                panic(fmt.Sprintf(
-                       "airflow.DagRef.Task: Dag %q: a task from 
airflow.TriggerDagRun with DagID %q has no "+
+                       "%s: Dag %q: a task from airflow.TriggerDagRun with 
DagID %q has no "+
                                "Go function to take a task_id from; set one 
with airflow.TaskSpec{TaskID: ...}",
-                       d.dagID, trigger.spec.DagID,
+                       method, d.dagID, trigger.spec.DagID,
                ))
        }
        if taskID == "" {
                var ok bool
                if taskID, ok = taskIDFromFuncName(funcName(fn)); !ok {
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: Dag %q: %s has no name to 
use as the task_id; "+
+                               "%s: Dag %q: %s has no name to use as the 
task_id; "+
                                        "set one with airflow.TaskSpec{TaskID: 
...}",
-                               d.dagID, funcName(fn),
+                               method, d.dagID, funcName(fn),
                        ))
                }
        }
        if triggerErr != nil {
-               panic(fmt.Sprintf(
-                       "airflow.DagRef.Task: task %q of Dag %q: %v", taskID, 
d.dagID, triggerErr,
-               ))
+               panic(fmt.Sprintf("%s: task %q of Dag %q: %v", method, taskID, 
d.dagID, triggerErr))
        }
        if err := checkTaskSpec(cfg.spec); err != nil {
-               panic(fmt.Sprintf("airflow.DagRef.Task: task %q of Dag %q: %v", 
taskID, d.dagID, err))
+               panic(fmt.Sprintf("%s: task %q of Dag %q: %v", method, taskID, 
d.dagID, err))
        }
        if _, exists := d.tasksByID[taskID]; exists {
                panic(fmt.Sprintf(
-                       "airflow.DagRef.Task: Dag %q already has a task %q; "+
+                       "%s: Dag %q already has a task %q; "+
                                "set another task_id with 
airflow.TaskSpec{TaskID: ...}",
-                       d.dagID, taskID,
+                       method, d.dagID, taskID,
                ))
        }
        var upstreams []*TaskRef
@@ -226,14 +238,14 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) 
*TaskRef {
        if isTrigger {
                if cfg.hasInputs {
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q comes 
from airflow.TriggerDagRun and "+
+                               "%s: task %q of Dag %q comes from 
airflow.TriggerDagRun and "+
                                        "takes no airflow.Inputs, because it 
has no Go function to pass the results to",
-                               taskID, d.dagID,
+                               method, taskID, d.dagID,
                        ))
                }
        } else {
                fnType := reflect.TypeOf(fn)
-               upstreams = d.checkInputs(taskID, fnType, cfg.inputs)
+               upstreams = d.checkInputs(method, taskID, fnType, cfg.inputs)
                // newTaskFunction has checked that fn returns either error or 
(result, error).
                if fnType.NumOut() == 2 {
                        resultType = fnType.Out(0)
@@ -248,6 +260,10 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef 
{
                inputs:        upstreams,
                task:          wrapped,
                triggerDagRun: triggerSpec,
+               ifRef:         ifRef,
+       }
+       if ifRef != nil {
+               ifRef.task = task
        }
        if d.tasksByID == nil {
                d.tasksByID = make(map[string]*TaskRef)
@@ -257,16 +273,27 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) 
*TaskRef {
        return task
 }
 
+// markRegistered marks d as registered, which stops any further change to d. 
It panics instead
+// when a condition from If has no task from Then.
 func (d *DagRef) markRegistered() {
        d.mu.Lock()
        defer d.mu.Unlock()
 
+       for _, task := range d.tasks {
+               if task.ifRef != nil && task.ifRef.thenTask == nil {
+                       panic(fmt.Sprintf(
+                               "airflow.BundleRef.Register: condition %q of 
Dag %q has no task from Then; "+
+                                       "name the task that runs when the 
condition is true with IfRef.Then",
+                               task.taskID, d.dagID,
+                       ))
+               }
+       }
        d.registered = true
 }
 
 func funcName(fn any) string { return 
runtime.FuncForPC(reflect.ValueOf(fn).Pointer()).Name() }
 
-// findTaskName names a task in an error that Task raises before it settles 
the task_id.
+// findTaskName names a task in an error that addTask raises before it settles 
the task_id.
 func findTaskName(fn any, opts []TaskOption) string {
        for _, opt := range opts {
                var spec TaskSpec
diff --git a/go-sdk/airflow/if.go b/go-sdk/airflow/if.go
new file mode 100644
index 00000000000..2e49f7409d6
--- /dev/null
+++ b/go-sdk/airflow/if.go
@@ -0,0 +1,221 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements.  See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership.  The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License.  You may obtain a copy of the License at
+//
+//   http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied.  See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+package airflow
+
+import (
+       "fmt"
+       "reflect"
+       "strings"
+
+       "github.com/apache/airflow/go-sdk/internal/bundle"
+)
+
+// IfRef is a condition that [DagRef.If] added to a Dag. [IfRef.Then] names 
the task that runs
+// when the condition is true, and [IfRef.Else] names the task that runs when 
it is false.
+type IfRef struct {
+       // task is the task that runs the condition function. thenTask and 
elseTask run after task,
+       // but neither takes the result of task.
+       task     *TaskRef
+       thenTask *TaskRef
+       elseTask *TaskRef
+}
+
+// If adds a task that runs the condition function fn, and returns the 
condition. Use
+// [IfRef.Then] to name the task that runs when fn returns true, and 
[IfRef.Else] to name the task
+// that runs when fn returns false:
+//
+//     extracted := dag.Task(extract)
+//     loaded := dag.Task(load)
+//     reportedEmpty := dag.Task(reportEmpty)
+//
+//     gate := dag.If(hasRows, airflow.Inputs(extracted))
+//     gate.Then(loaded)
+//     gate.Else(reportedEmpty)
+//
+// fn returns (bool, error). In every other way it follows the rules of a 
function passed to
+// [DagRef.Task]: it takes a [Context] first, and [Inputs] fills the 
parameters after the Context.
+// For the Dag above, the functions could be:
+//
+//     func extract(actx airflow.Context) ([]string, error)
+//     func hasRows(actx airflow.Context, rows []string) (bool, error)
+//     func load(actx airflow.Context) error
+//     func reportEmpty(actx airflow.Context) error
+//
+// The task that If adds gets its task_id in the same way as a task from 
DagRef.Task, and a
+// [TaskSpec] sets its attributes.
+//
+// When the task runs, the result of fn decides what it skips:
+//   - true skips the task from Else, or nothing if there is no task from Else
+//   - false skips the task from Then
+//   - an error fails the task, which then skips nothing
+//
+// The condition skips only that one task. In Python, a branch skips every 
task directly after it
+// that it does not follow, and a short circuit skips every task after it. 
After the skip, the
+// trigger rule of each task after the skipped task decides whether that task 
runs. A task that
+// runs after both the task from Then and the task from Else needs a trigger 
rule that runs it when
+// one of the two is skipped, such as [TriggerRuleNoneFailedMinOneSuccess].
+//
+// [BundleRef.Register] panics if a condition has no task from Then.
+//
+// If panics for the same reasons as DagRef.Task does for a Go function. It 
also panics if fn
+// comes from [TriggerDagRun] or does not return (bool, error).
+func (d *DagRef) If(fn any, opts ...TaskOption) *IfRef {
+       if _, ok := fn.(TriggerDagRunTask); ok {
+               panic(fmt.Sprintf(
+                       "airflow.DagRef.If: Dag %q: fn comes from 
airflow.TriggerDagRun, "+
+                               "but a condition function is a Go function that 
returns (bool, error)",
+                       d.dagID,
+               ))
+       }
+       ifRef := &IfRef{}
+       d.addTask("airflow.DagRef.If", fn, opts, ifRef)
+       return ifRef
+}
+
+// Then names the task that runs when the condition of g is true. When the 
condition is false,
+// the task is skipped. Then returns g, so that a call to Else can follow. 
With the tasks from the
+// example of [DagRef.If], the condition fits in one statement:
+//
+//     dag.If(hasRows, 
airflow.Inputs(extracted)).Then(loaded).Else(reportedEmpty)
+//
+// The task runs after the condition, but its function does not take the 
result of the
+// condition as a parameter.
+//
+// Then panics if:
+//   - g is not the IfRef that DagRef.If returned, for example a copy of that 
IfRef
+//   - task is nil, or DagRef.Task did not add task to the Dag of the condition
+//   - the condition already has a task from Then
+//   - task is the task from Else
+//   - the Dag is already registered
+func (g *IfRef) Then(task *TaskRef) *IfRef {
+       g.setTask("Then", task)
+       return g
+}
+
+// Else names the task that runs when the condition of g is false. When the 
condition is true,
+// the task is skipped. Like Then, Else returns g. A condition does not need a 
task from Else.
+//
+// Else panics for the reasons that Then lists, with Then and Else swapped.
+func (g *IfRef) Else(task *TaskRef) *IfRef {
+       g.setTask("Else", task)
+       return g
+}
+
+// setTask makes task the task from Then or from Else of g. side is "Then" or 
"Else".
+func (g *IfRef) setTask(side string, task *TaskRef) {
+       method := "airflow.IfRef." + side
+       // When the condition task runs, it reads the tasks from Then and Else 
from the IfRef that If
+       // returned. A task given to a copy of that IfRef would never reach the 
condition task.
+       if g == nil || g.task == nil || g.task.ifRef != g {
+               panic(method + ": DagRef.If did not return the *airflow.IfRef")
+       }
+       condition, d := g.task.taskID, g.task.dag
+       d.mu.Lock()
+       defer d.mu.Unlock()
+
+       if d.registered {
+               panic(fmt.Sprintf(
+                       "%s: Dag %q has already been registered; "+
+                               "name the tasks of every condition before 
Register",
+                       method, d.dagID,
+               ))
+       }
+       switch {
+       case task == nil:
+               panic(fmt.Sprintf(
+                       "%s: condition %q of Dag %q got a nil 
*airflow.TaskRef", method, condition, d.dagID,
+               ))
+       case task.dag != nil && task.dag != d:
+               panic(fmt.Sprintf(
+                       "%s: condition %q of Dag %q cannot run task %q of 
another Dag, %q; "+
+                               "pass a task of the same Dag",
+                       method, condition, d.dagID, task.taskID, task.dag.dagID,
+               ))
+       // A zero TaskRef and a copy of a TaskRef get here.
+       case d.tasksByID[task.taskID] != task:
+               panic(fmt.Sprintf(
+                       "%s: condition %q of Dag %q got a *airflow.TaskRef that 
DagRef.Task did not return",
+                       method, condition, d.dagID,
+               ))
+       }
+
+       slot, other, otherSide := &g.thenTask, g.elseTask, "Else"
+       if side == "Else" {
+               slot, other, otherSide = &g.elseTask, g.thenTask, "Then"
+       }
+       switch {
+       case *slot != nil:
+               panic(fmt.Sprintf(
+                       "%s: condition %q of Dag %q already has task %q from 
%s; call %s once",
+                       method, condition, d.dagID, (*slot).taskID, side, side,
+               ))
+       case task == other:
+               panic(fmt.Sprintf(
+                       "%s: condition %q of Dag %q already has task %q from 
%s, "+
+                               "and a task cannot be on both sides of a 
condition",
+                       method, condition, d.dagID, task.taskID, otherSide,
+               ))
+       }
+       *slot = task
+}
+
+// wrapCondition wraps fn as the task of g. The task skips the side of g that 
the result of fn
+// 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]() {
+               return nil, fmt.Errorf(
+                       "%s returns %s, but a condition function must return 
(bool, error)",
+                       funcName(fn), describeResults(fnType),
+               )
+       }
+       return bundle.NewPositionalBranchFunction(fn, g.findSkipped)
+}
+
+// findSkipped returns the task_id of the task on the side of g that result 
does not take. It
+// returns nil when that side has no task.
+func (g *IfRef) findSkipped(result any) []string {
+       notTaken := g.thenTask
+       if result.(bool) {
+               notTaken = g.elseTask
+       }
+       if notTaken == nil {
+               return nil
+       }
+       return []string{notTaken.taskID}
+}
+
+func describeResults(fnType reflect.Type) string {
+       results := make([]string, fnType.NumOut())
+       for i := range results {
+               results[i] = fnType.Out(i).String()
+       }
+       switch len(results) {
+       case 0:
+               return "nothing"
+       case 1:
+               return results[0]
+       default:
+               return "(" + strings.Join(results, ", ") + ")"
+       }
+}
diff --git a/go-sdk/airflow/if_test.go b/go-sdk/airflow/if_test.go
new file mode 100644
index 00000000000..cf51c330b1e
--- /dev/null
+++ b/go-sdk/airflow/if_test.go
@@ -0,0 +1,643 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements.  See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership.  The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License.  You may obtain a copy of the License at
+//
+//   http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied.  See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+package airflow
+
+import (
+       "context"
+       "errors"
+       "fmt"
+       "maps"
+       "reflect"
+       "strings"
+       "testing"
+       "time"
+
+       "github.com/stretchr/testify/assert"
+       "github.com/stretchr/testify/require"
+
+       "github.com/apache/airflow/go-sdk/internal/bundle"
+       "github.com/apache/airflow/go-sdk/pkg/binding"
+       "github.com/apache/airflow/go-sdk/pkg/sdkcontext"
+       "github.com/apache/airflow/go-sdk/sdk"
+)
+
+type ready bool
+
+type conditionError struct{}
+
+func (*conditionError) Error() string { return "condition failed" }
+
+func isReady(Context) (bool, error)                       { return true, nil }
+func hasRows(_ Context, rows rowSet) (bool, error)        { return 
len(rows.Rows) > 0, nil }
+func isReadyOnlyError(Context) error                      { return nil }
+func isReadyAsInt(Context) (int, error)                   { return 0, nil }
+func isReadyAsPointer(Context) (*bool, error)             { return nil, nil }
+func isReadyAsNamedBool(Context) (ready, error)           { return false, nil }
+func isReadyWithoutResults(Context)                       {}
+func isReadyWithOwnError(Context) (bool, *conditionError) { return true, nil }
+func load(Context) error                                  { return nil }
+func reportEmpty(Context) error                           { return nil }
+
+func TestIfAddsTheConditionAsATask(t *testing.T) {
+       dag := Dag("etl")
+       read := dag.Task(readRows)
+
+       gate := dag.If(hasRows, Inputs(read))
+
+       require.NotNil(t, gate.task)
+       assert.Equal(t, "hasRows", gate.task.taskID)
+       assert.Equal(t, []*TaskRef{read, gate.task}, dag.tasks)
+       assert.Same(t, gate.task, dag.tasksByID["hasRows"])
+       assertInputs(t, gate.task, read)
+       assert.Equal(t, reflect.TypeFor[bool](), gate.task.resultType)
+       assert.Same(t, gate, gate.task.ifRef)
+       assert.Nil(t, read.ifRef, "a task from DagRef.Task is not a condition")
+
+       named := dag.If(isReady, TaskSpec{TaskID: "is_ready"})
+       assert.Equal(t, "is_ready", named.task.taskID)
+       assert.Equal(t, TaskSpec{TaskID: "is_ready"}, named.task.spec)
+}
+
+func TestIfNeedsAFunctionThatReturnsABool(t *testing.T) {
+       tests := []struct {
+               name    string
+               fn      any
+               fnName  string
+               returns string
+       }{
+               {"only an error", isReadyOnlyError, "isReadyOnlyError", 
"error"},
+               {"another result type", isReadyAsInt, "isReadyAsInt", "(int, 
error)"},
+               {"pointer to bool", isReadyAsPointer, "isReadyAsPointer", 
"(*bool, error)"},
+               {"named bool type", isReadyAsNamedBool, "isReadyAsNamedBool", 
"(airflow.ready, error)"},
+               {"no results", isReadyWithoutResults, "isReadyWithoutResults", 
"nothing"},
+               {
+                       "another error type",
+                       isReadyWithOwnError,
+                       "isReadyWithOwnError",
+                       "(bool, *airflow.conditionError)",
+               },
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       dag := Dag("etl")
+                       assert.PanicsWithValue(t,
+                               `airflow.DagRef.If: Dag "etl": 
github.com/apache/airflow/go-sdk/airflow.`+
+                                       tt.fnName+` returns `+tt.returns+
+                                       `, but a condition function must return 
(bool, error)`,
+                               func() { dag.If(tt.fn) },
+                       )
+                       assert.Empty(t, dag.tasks)
+               })
+       }
+}
+
+func TestIfRejectsTriggerDagRun(t *testing.T) {
+       dag := Dag("etl")
+       assert.PanicsWithValue(t,
+               `airflow.DagRef.If: Dag "etl": fn comes from 
airflow.TriggerDagRun, `+
+                       `but a condition function is a Go function that returns 
(bool, error)`,
+               func() {
+                       dag.If(
+                               TriggerDagRun(TriggerDagRunSpec{DagID: 
"downstream_etl"}),
+                               TaskSpec{TaskID: "trigger"},
+                       )
+               },
+       )
+       assert.Empty(t, dag.tasks)
+}
+
+// If shares its checks with DagRef.Task, and each panic message has to name 
DagRef.If, the
+// function that the caller called.
+func TestIfPanicsUnderItsOwnName(t *testing.T) {
+       read := func(dag *DagRef) *TaskRef { return dag.Task(readRows) }
+       literal := func(Context) (bool, error) { return true, nil }
+
+       tests := []struct {
+               name string
+               add  func(dag *DagRef)
+               // want is a part of the panic message after 
"airflow.DagRef.If: ".
+               want string
+       }{
+               {
+                       name: "registered Dag",
+                       add: func(dag *DagRef) {
+                               Bundle().Register(dag)
+                               dag.If(isReady)
+                       },
+                       want: "has already been registered",
+               },
+               {
+                       name: "not a function",
+                       add:  func(dag *DagRef) { dag.If("isReady") },
+                       want: "fn is string, not a function",
+               },
+               {
+                       name: "nil option",
+                       add:  func(dag *DagRef) { dag.If(isReady, nil) },
+                       want: "opts[0] is nil",
+               },
+               {
+                       name: "nil *TaskSpec",
+                       add:  func(dag *DagRef) { dag.If(isReady, 
(*TaskSpec)(nil)) },
+                       want: "opts[0] is a nil *airflow.TaskSpec",
+               },
+               {
+                       name: "option from a struct that embeds a TaskSpec",
+                       add:  func(dag *DagRef) { dag.If(isReady, 
wrappedSpec{}) },
+                       want: "opts[0] has type airflow.wrappedSpec",
+               },
+               {
+                       name: "second TaskSpec",
+                       add:  func(dag *DagRef) { dag.If(isReady, TaskSpec{}, 
TaskSpec{}) },
+                       want: "got more than one airflow.TaskSpec",
+               },
+               {
+                       name: "no name for the task_id",
+                       add:  func(dag *DagRef) { dag.If(literal) },
+                       want: "has no name to use as the task_id",
+               },
+               {
+                       name: "unknown trigger rule",
+                       add:  func(dag *DagRef) { dag.If(isReady, 
TaskSpec{TriggerRule: "bogus"}) },
+                       want: `airflow.TaskSpec.TriggerRule is "bogus", which 
is not a trigger rule`,
+               },
+               {
+                       name: "duplicate task_id",
+                       add: func(dag *DagRef) {
+                               dag.Task(extract, TaskSpec{TaskID: "isReady"})
+                               dag.If(isReady)
+                       },
+                       want: `already has a task "isReady"`,
+               },
+               {
+                       name: "second Inputs",
+                       add:  func(dag *DagRef) { dag.If(hasRows, 
Inputs(read(dag)), Inputs()) },
+                       want: "got more than one airflow.Inputs",
+               },
+               {
+                       name: "nil input",
+                       add:  func(dag *DagRef) { dag.If(hasRows, Inputs(nil)) 
},
+                       want: "airflow.Inputs got a nil *airflow.TaskRef",
+               },
+               {
+                       name: "input from another Dag",
+                       add:  func(dag *DagRef) { dag.If(hasRows, 
Inputs(read(Dag("reports")))) },
+                       want: "cannot take an input from task",
+               },
+               {
+                       name: "input that DagRef.Task did not return",
+                       add:  func(dag *DagRef) { dag.If(hasRows, 
Inputs(&TaskRef{})) },
+                       want: "that DagRef.Task did not return",
+               },
+               {
+                       name: "missing input",
+                       add:  func(dag *DagRef) { dag.If(hasRows) },
+                       want: "but airflow.Inputs passes no task",
+               },
+               {
+                       name: "input from TriggerDagRun",
+                       add: func(dag *DagRef) {
+                               trigger := dag.Task(
+                                       TriggerDagRun(TriggerDagRunSpec{DagID: 
"downstream_etl"}),
+                                       TaskSpec{TaskID: "trigger"},
+                               )
+                               dag.If(hasRows, Inputs(trigger))
+                       },
+                       want: "comes from airflow.TriggerDagRun and returns no 
result",
+               },
+               {
+                       name: "input with no result",
+                       add:  func(dag *DagRef) { dag.If(hasRows, 
Inputs(dag.Task(extract))) },
+                       want: "but that task returns only an error",
+               },
+               {
+                       name: "input of the wrong type",
+                       add:  func(dag *DagRef) { dag.If(hasRows, 
Inputs(dag.Task(readNames))) },
+                       want: "which cannot be assigned to airflow.rowSet",
+               },
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       msg := panicMessage(t, func() { tt.add(Dag("etl")) })
+                       assert.True(t, strings.HasPrefix(msg, 
"airflow.DagRef.If: "), msg)
+                       assert.Contains(t, msg, tt.want)
+               })
+       }
+}
+
+func TestThenAndElseNameTheTasksOfTheCondition(t *testing.T) {
+       dag := Dag("etl")
+       loaded := dag.Task(load)
+       reported := dag.Task(reportEmpty)
+
+       gate := dag.If(isReady)
+       assert.Same(t, gate, gate.Then(loaded))
+       assert.Same(t, gate, gate.Else(reported))
+
+       assert.Same(t, loaded, gate.thenTask)
+       assert.Same(t, reported, gate.elseTask)
+       assert.Empty(t, loaded.inputs, "the condition passes no result to the 
task from Then")
+       assert.Empty(t, reported.inputs, "the condition passes no result to the 
task from Else")
+
+       reversed := dag.If(isReady, TaskSpec{TaskID: 
"reversed"}).Else(reported).Then(loaded)
+       assert.Same(t, loaded, reversed.thenTask, "Else can come before Then")
+       assert.Same(t, reported, reversed.elseTask)
+}
+
+func TestThenAndElseRejectATaskOutsideTheDag(t *testing.T) {
+       dag := Dag("etl")
+       loaded := dag.Task(load)
+       loadedCopy := *loaded
+       fromReports := Dag("reports").Task(load)
+       fromAnotherEtl := Dag("etl").Task(load)
+
+       tests := []struct {
+               name string
+               task *TaskRef
+               want string
+       }{
+               {name: "nil", task: nil, want: `got a nil *airflow.TaskRef`},
+               {
+                       name: "zero TaskRef",
+                       task: &TaskRef{},
+                       want: `got a *airflow.TaskRef that DagRef.Task did not 
return`,
+               },
+               {
+                       name: "copy of a task of the Dag",
+                       task: &loadedCopy,
+                       want: `got a *airflow.TaskRef that DagRef.Task did not 
return`,
+               },
+               {
+                       name: "task of another Dag",
+                       task: fromReports,
+                       want: `cannot run task "load" of another Dag, 
"reports"; pass a task of the same Dag`,
+               },
+               {
+                       name: "task of another Dag with the same dag_id",
+                       task: fromAnotherEtl,
+                       want: `cannot run task "load" of another Dag, "etl"; 
pass a task of the same Dag`,
+               },
+       }
+       for i, tt := range tests {
+               for _, side := range []struct {
+                       name string
+                       call func(gate *IfRef, task *TaskRef)
+               }{
+                       {"Then", func(gate *IfRef, task *TaskRef) { 
gate.Then(task) }},
+                       {"Else", func(gate *IfRef, task *TaskRef) { 
gate.Else(task) }},
+               } {
+                       t.Run(tt.name+"/"+side.name, func(t *testing.T) {
+                               condition := fmt.Sprintf("ready_%d_%s", i, 
side.name)
+                               gate := dag.If(isReady, TaskSpec{TaskID: 
condition})
+                               assert.PanicsWithValue(t,
+                                       "airflow.IfRef."+side.name+`: condition 
"`+condition+`" of Dag "etl" `+tt.want,
+                                       func() { side.call(gate, tt.task) },
+                               )
+                               assert.Nil(t, gate.thenTask)
+                               assert.Nil(t, gate.elseTask)
+                       })
+               }
+       }
+}
+
+func TestThenAndElseNameOneTaskEach(t *testing.T) {
+       dag := Dag("etl")
+       loaded := dag.Task(load)
+       reported := dag.Task(reportEmpty)
+
+       gate := dag.If(isReady).Then(loaded)
+       assert.PanicsWithValue(t,
+               `airflow.IfRef.Then: condition "isReady" of Dag "etl" already 
has task "load" from Then; `+
+                       `call Then once`,
+               func() { gate.Then(reported) },
+       )
+       assert.PanicsWithValue(t,
+               `airflow.IfRef.Else: condition "isReady" of Dag "etl" already 
has task "load" from Then, `+
+                       `and a task cannot be on both sides of a condition`,
+               func() { gate.Else(loaded) },
+       )
+       gate.Else(reported)
+       assert.PanicsWithValue(t,
+               `airflow.IfRef.Else: condition "isReady" of Dag "etl" already 
has task "reportEmpty" `+
+                       `from Else; call Else once`,
+               func() { gate.Else(loaded) },
+       )
+
+       elseFirst := dag.If(isReady, TaskSpec{TaskID: 
"else_first"}).Else(reported)
+       assert.PanicsWithValue(t,
+               `airflow.IfRef.Then: condition "else_first" of Dag "etl" 
already has task "reportEmpty" `+
+                       `from Else, and a task cannot be on both sides of a 
condition`,
+               func() { elseFirst.Then(reported) },
+       )
+
+       assert.Same(t, loaded, gate.thenTask)
+       assert.Same(t, reported, gate.elseTask)
+       assert.Nil(t, elseFirst.thenTask)
+}
+
+func TestThenAndElseAfterRegisterPanic(t *testing.T) {
+       dag := Dag("etl")
+       loaded := dag.Task(load)
+       reported := dag.Task(reportEmpty)
+       gate := dag.If(isReady).Then(loaded)
+       Bundle().Register(dag)
+
+       for _, side := range []struct {
+               name string
+               call func()
+       }{
+               {"Then", func() { gate.Then(reported) }},
+               {"Else", func() { gate.Else(reported) }},
+       } {
+               assert.PanicsWithValue(t,
+                       "airflow.IfRef."+side.name+`: Dag "etl" has already 
been registered; `+
+                               `name the tasks of every condition before 
Register`,
+                       side.call,
+               )
+       }
+       assert.Same(t, loaded, gate.thenTask)
+       assert.Nil(t, gate.elseTask)
+}
+
+func TestIfRefThatDagRefIfDidNotReturnPanics(t *testing.T) {
+       dag := Dag("etl")
+       loaded := dag.Task(load)
+       gate := dag.If(isReady)
+       gateCopy := *gate
+       var nilRef *IfRef
+
+       for name, ref := range map[string]*IfRef{"nil": nilRef, "zero": {}, 
"copy": &gateCopy} {
+               t.Run(name, func(t *testing.T) {
+                       assert.PanicsWithValue(t,
+                               "airflow.IfRef.Then: DagRef.If did not return 
the *airflow.IfRef",
+                               func() { ref.Then(loaded) },
+                       )
+                       assert.PanicsWithValue(t,
+                               "airflow.IfRef.Else: DagRef.If did not return 
the *airflow.IfRef",
+                               func() { ref.Else(loaded) },
+                       )
+               })
+       }
+       assert.Nil(t, gate.thenTask)
+       assert.Nil(t, gate.elseTask)
+}
+
+func TestRegisterRejectsAConditionWithoutThen(t *testing.T) {
+       tests := []struct {
+               name string
+               add  func(dag *DagRef)
+       }{
+               {name: "no Then or Else", add: func(dag *DagRef) { 
dag.If(isReady) }},
+               {
+                       name: "only Else",
+                       add:  func(dag *DagRef) { 
dag.If(isReady).Else(dag.Task(reportEmpty)) },
+               },
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       dag := Dag("etl")
+                       dag.If(isReady, TaskSpec{TaskID: 
"complete"}).Then(dag.Task(load))
+                       tt.add(dag)
+                       b := Bundle()
+
+                       assert.PanicsWithValue(t,
+                               `airflow.BundleRef.Register: condition 
"isReady" of Dag "etl" has no task from `+
+                                       `Then; name the task that runs when the 
condition is true with IfRef.Then`,
+                               func() { b.Register(dag) },
+                       )
+                       assert.False(t, dag.registered, "a Dag that Register 
rejects can still take tasks")
+                       assert.NotContains(t, b.dags.dags, "etl")
+               })
+       }
+}
+
+// conditionClient answers GetXCom from results, which maps a task_id to the 
result of that task.
+// It records the XComs that a condition pushes. Its skip method stands in for 
the function that
+// the runtime passes through bundle.WithSkipDownstreamTasks, and records the 
task_ids that the
+// condition skips.
+type conditionClient struct {
+       sdk.Client
+       results map[string]any
+       xcoms   map[string]any
+       skipped [][]string
+}
+
+func (c *conditionClient) GetXCom(
+       _ context.Context,
+       _, _, taskID string,
+       _ *int,
+       _ string,
+       _ any,
+) (any, error) {
+       return c.results[taskID], nil
+}
+
+func (c *conditionClient) PushXCom(_ context.Context, _ sdk.TaskInstance, key 
string, v any) error {
+       c.xcoms[key] = v
+       return nil
+}
+
+func (c *conditionClient) skip(_ context.Context, taskIDs []string) error {
+       c.skipped = append(c.skipped, taskIDs)
+       return nil
+}
+
+// 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) {
+       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{}
+       }
+       ti := sdk.TaskInstance{DagID: "etl", RunID: "run1", TaskID: 
gate.task.taskID}
+       ctx := context.WithValue(
+               context.Background(),
+               sdkcontext.SdkClientContextKey,
+               sdk.Client(client),
+       )
+       ctx = context.WithValue(
+               ctx,
+               sdkcontext.RuntimeContextKey,
+               sdk.NewTIRunContext(context.Background(), ti, sdk.DagRun{DagID: 
"etl", RunID: "run1"}),
+       )
+       ctx = bundle.WithSkipDownstreamTasks(ctx, client.skip)
+       return client, gate.task.task.Execute(ctx, discardLogger(), args)
+}
+
+func TestConditionSkipsTheSideThatItDoesNotTake(t *testing.T) {
+       tests := []struct {
+               name     string
+               rows     []any
+               withElse bool
+               want     bool
+               skipped  []string
+       }{
+               {
+                       name:     "true skips the task from Else",
+                       rows:     []any{"a"},
+                       withElse: true,
+                       want:     true,
+                       skipped:  []string{"report_empty"},
+               },
+               {
+                       name:     "false skips the task from Then",
+                       rows:     []any{},
+                       withElse: true,
+                       want:     false,
+                       skipped:  []string{"load_rows"},
+               },
+               {
+                       name:    "true without Else skips nothing",
+                       rows:    []any{"a"},
+                       want:    true,
+                       skipped: []string{},
+               },
+               {
+                       name:    "false without Else skips the task from Then",
+                       rows:    []any{},
+                       want:    false,
+                       skipped: []string{"load_rows"},
+               },
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       dag := Dag("etl")
+                       read := dag.Task(readRows)
+                       // The TaskIDs differ from the function names, so the 
test shows that the condition
+                       // skips a task by its task_id.
+                       loaded := dag.Task(load, TaskSpec{TaskID: "load_rows"})
+                       reported := dag.Task(reportEmpty, TaskSpec{TaskID: 
"report_empty"})
+                       gate := dag.If(hasRows, Inputs(read)).Then(loaded)
+                       if tt.withElse {
+                               gate.Else(reported)
+                       }
+                       Bundle().Register(dag)
+
+                       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)
+                       } else {
+                               assert.Equal(t, [][]string{tt.skipped}, 
client.skipped)
+                       }
+                       assert.Equal(t,
+                               map[string][]string{"skipped": tt.skipped},
+                               client.xcoms["skipmixin_key"],
+                       )
+               })
+       }
+}
+
+func TestConditionThatFailsSkipsNothing(t *testing.T) {
+       dag := Dag("etl")
+       gate := dag.If(
+               func(Context) (bool, error) { return false, errors.New("cannot 
reach the table") },
+               TaskSpec{TaskID: "has_rows"},
+       ).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"}},
+       })
+
+       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",
+       )
+}
+
+// Else and Register both take the lock of the Dag, so each Else call either 
names its task before
+// the Dag is registered or panics. If Else read the registered flag without 
the lock, only a run
+// with -race would fail.
+func TestRegisterWhileElseIsCalled(t *testing.T) {
+       const conditions = 1000
+
+       dag := Dag("etl")
+       gates := make([]*IfRef, conditions)
+       reported := make([]*TaskRef, conditions)
+       for i := range conditions {
+               gates[i] = dag.If(isReady, TaskSpec{TaskID: 
fmt.Sprintf("ready_%d", i)}).
+                       Then(dag.Task(load, TaskSpec{TaskID: 
fmt.Sprintf("load_%d", i)}))
+               reported[i] = dag.Task(reportEmpty, TaskSpec{TaskID: 
fmt.Sprintf("report_%d", i)})
+       }
+       started := make(chan struct{})
+       registered := make(chan struct{})
+       type outcome struct {
+               named     int
+               recovered any
+       }
+       done := make(chan outcome)
+       go func() {
+               var o outcome
+               defer func() {
+                       o.recovered = recover()
+                       done <- o
+               }()
+               close(started)
+               for ; o.named < conditions-1; o.named++ {
+                       select {
+                       case <-registered:
+                               // Register has returned, so this Else call has 
to panic.
+                               gates[o.named].Else(reported[o.named])
+                               return
+                       default:
+                               gates[o.named].Else(reported[o.named])
+                       }
+               }
+               select {
+               case <-registered:
+               case <-time.After(10 * time.Second):
+                       return
+               }
+               gates[o.named].Else(reported[o.named])
+       }()
+       // Wait until the goroutine is running before calling Register, as 
TestRegisterWhileTasksAreAdded
+       // explains.
+       <-started
+       Bundle().Register(dag)
+       close(registered)
+       o := <-done
+
+       assert.Equal(t,
+               `airflow.IfRef.Else: Dag "etl" has already been registered; `+
+                       `name the tasks of every condition before Register`,
+               o.recovered,
+       )
+       for i, gate := range gates {
+               if i < o.named {
+                       assert.Same(t, reported[i], gate.elseTask, "condition 
%d", i)
+               } else {
+                       assert.Nil(t, gate.elseTask, "condition %d", i)
+               }
+       }
+}
diff --git a/go-sdk/airflow/inputs.go b/go-sdk/airflow/inputs.go
index 62430dba3c8..baba9d5f027 100644
--- a/go-sdk/airflow/inputs.go
+++ b/go-sdk/airflow/inputs.go
@@ -38,9 +38,9 @@ func (in inputs) applyTask(c *taskConfig) error {
        return nil
 }
 
-// Inputs passes the results of tasks to the task that [DagRef.Task] adds, and 
makes each of
-// those tasks an upstream task of the new one. It is the Go form of a Python 
TaskFlow call such
-// as transform(extract()):
+// Inputs passes the results of tasks to the task that [DagRef.Task] or 
[DagRef.If] adds, and
+// makes each of those tasks an upstream task of the new one. It is the Go 
form of a Python
+// TaskFlow call such as transform(extract()):
 //
 //     extracted := dag.Task(extract)
 //     transformed := dag.Task(transform, airflow.Inputs(extracted))
@@ -58,35 +58,39 @@ func (in inputs) applyTask(c *taskConfig) error {
 // fills. So each field of a struct parameter comes from the matching JSON 
key. A parameter of
 // type any gets a map[string]any when the result is a struct.
 //
-// When [DagRef.Task] adds the task, it panics unless each parameter after the 
Context gets
-// exactly one task of the same Dag, and the result type of that task is 
assignable to the
-// parameter type, as in a Go function call. Pass at most one Inputs to a task.
+// When [DagRef.Task] or [DagRef.If] adds the task, it panics unless each 
parameter after the
+// Context gets exactly one task of the same Dag, and the result type of that 
task is assignable
+// to the parameter type, as in a Go function call. Pass at most one Inputs to 
a task.
 func Inputs(refs ...*TaskRef) TaskOption { return inputs(refs) }
 
 // checkInputs returns the tasks that task taskID got through Inputs, in a new 
slice. It panics if
 // any of those tasks was not added to d by DagRef.Task, or if the tasks do 
not match the
-// parameters that a function of type fnType takes after the Context.
-func (d *DagRef) checkInputs(taskID string, fnType reflect.Type, tasks 
[]*TaskRef) []*TaskRef {
+// parameters that a function of type fnType takes after the Context. method 
names the caller in
+// panic messages.
+func (d *DagRef) checkInputs(
+       method, taskID string,
+       fnType reflect.Type,
+       tasks []*TaskRef,
+) []*TaskRef {
        for i, upstream := range tasks {
                switch {
                case upstream == nil:
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q: "+
-                                       "airflow.Inputs got a nil 
*airflow.TaskRef at index %d",
-                               taskID, d.dagID, i,
+                               "%s: task %q of Dag %q: airflow.Inputs got a 
nil *airflow.TaskRef at index %d",
+                               method, taskID, d.dagID, i,
                        ))
                case upstream.dag != nil && upstream.dag != d:
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q cannot 
take an input from task %q of "+
+                               "%s: task %q of Dag %q cannot take an input 
from task %q of "+
                                        "another Dag, %q; pass tasks of the 
same Dag to airflow.Inputs",
-                               taskID, d.dagID, upstream.taskID, 
upstream.dag.dagID,
+                               method, taskID, d.dagID, upstream.taskID, 
upstream.dag.dagID,
                        ))
                // A zero TaskRef and a copy of a TaskRef get here.
                case d.tasksByID[upstream.taskID] != upstream:
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q: 
airflow.Inputs got a *airflow.TaskRef "+
+                               "%s: task %q of Dag %q: airflow.Inputs got a 
*airflow.TaskRef "+
                                        "at index %d that DagRef.Task did not 
return",
-                               taskID, d.dagID, i,
+                               method, taskID, d.dagID, i,
                        ))
                }
        }
@@ -95,9 +99,9 @@ func (d *DagRef) checkInputs(taskID string, fnType 
reflect.Type, tasks []*TaskRe
        // after it.
        if params := fnType.NumIn() - 1; len(tasks) != params {
                panic(fmt.Sprintf(
-                       "airflow.DagRef.Task: task %q of Dag %q has %d 
parameter(s) after airflow.Context, "+
+                       "%s: task %q of Dag %q has %d parameter(s) after 
airflow.Context, "+
                                "but airflow.Inputs passes %s",
-                       taskID, d.dagID, params, describeTasks(tasks),
+                       method, taskID, d.dagID, params, describeTasks(tasks),
                ))
        }
        for i, upstream := range tasks {
@@ -106,15 +110,15 @@ func (d *DagRef) checkInputs(taskID string, fnType 
reflect.Type, tasks []*TaskRe
                switch {
                case upstream.triggerDagRun != nil:
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q takes 
parameter %d from task %q, "+
+                               "%s: task %q of Dag %q takes parameter %d from 
task %q, "+
                                        "but that task comes from 
airflow.TriggerDagRun and returns no result",
-                               taskID, d.dagID, param, upstream.taskID,
+                               method, taskID, d.dagID, param, upstream.taskID,
                        ))
                case upstream.resultType == nil:
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q takes 
parameter %d from task %q, "+
+                               "%s: task %q of Dag %q takes parameter %d from 
task %q, "+
                                        "but that task returns only an error",
-                               taskID, d.dagID, param, upstream.taskID,
+                               method, taskID, d.dagID, param, upstream.taskID,
                        ))
                case !upstream.resultType.AssignableTo(paramType):
                        var sameName string
@@ -122,9 +126,10 @@ func (d *DagRef) checkInputs(taskID string, fnType 
reflect.Type, tasks []*TaskRe
                                sameName = ", a different type with the same 
name"
                        }
                        panic(fmt.Sprintf(
-                               "airflow.DagRef.Task: task %q of Dag %q takes 
parameter %d from task %q, "+
+                               "%s: task %q of Dag %q takes parameter %d from 
task %q, "+
                                        "but that task returns %s, which cannot 
be assigned to %s%s",
-                               taskID, d.dagID, param, upstream.taskID, 
upstream.resultType, paramType, sameName,
+                               method, taskID, d.dagID, param, 
upstream.taskID, upstream.resultType, paramType,
+                               sameName,
                        ))
                }
        }
diff --git a/go-sdk/airflow/task_option.go b/go-sdk/airflow/task_option.go
index 91efff4ff0d..248feec9641 100644
--- a/go-sdk/airflow/task_option.go
+++ b/go-sdk/airflow/task_option.go
@@ -19,19 +19,19 @@ package airflow
 
 import "errors"
 
-// TaskOption is an option to [DagRef.Task]. There are two kinds: a [TaskSpec] 
sets the
-// attributes of the task that DagRef.Task adds, and [Inputs] passes the 
results of other tasks
-// to that task.
+// TaskOption is an option to [DagRef.Task] and [DagRef.If]. There are two 
kinds: a [TaskSpec]
+// sets the attributes of the task that DagRef.Task or DagRef.If adds, and 
[Inputs] passes the
+// results of other tasks to that task.
 //
 // Its only method is unexported, so a type outside this package cannot 
declare it.
-// A struct that embeds a TaskSpec or a TaskOption still satisfies the 
interface, and Task
-// panics when it is given one.
+// A struct that embeds a TaskSpec or a TaskOption still satisfies the 
interface, but DagRef.Task
+// and DagRef.If panic when they get such a struct.
 type TaskOption interface{ applyTask(*taskConfig) error }
 
 type taskConfig struct {
        spec   TaskSpec
        inputs []*TaskRef
-       // hasSpec is true once DagRef.Task has applied a TaskSpec, and 
hasInputs is true once it has
+       // hasSpec is true once addTask has applied a TaskSpec, and hasInputs 
is true once it has
        // applied an Inputs. The spec and inputs fields cannot show that an 
option was applied,
        // because TaskSpec{} leaves spec at its zero value and Inputs() leaves 
inputs nil.
        hasSpec   bool
diff --git a/go-sdk/internal/bundle/task.go b/go-sdk/internal/bundle/task.go
index 8a1743eabf0..9db3626a1ff 100644
--- a/go-sdk/internal/bundle/task.go
+++ b/go-sdk/internal/bundle/task.go
@@ -31,8 +31,8 @@ import (
 )
 
 // Task is one registered task that the coordinator runtime can execute. Bundle
-// authors do not implement this directly. airflow.TaskHandler and
-// airflow.DagRef.Task wrap a plain Go function into a Task.
+// authors do not implement this directly. airflow.TaskHandler, 
airflow.DagRef.Task
+// and airflow.DagRef.If wrap a plain Go function into a Task.
 type Task interface {
        Execute(ctx context.Context, logger *slog.Logger, args []binding.Arg) 
error
 }
@@ -61,29 +61,45 @@ type taskFunction struct {
        fn       reflect.Value
        fullName string
        plan     *binding.Plan
+       // findSkipped is nil unless the task comes from 
NewPositionalBranchFunction.
+       findSkipped func(result any) []string
 }
 
 var _ Task = (*taskFunction)(nil)
 
 // NewTaskFunction validates and wraps a Go function as a Task.
-func NewTaskFunction(fn any) (Task, error) { return newTaskFunction(fn, 
binding.Analyze) }
+func NewTaskFunction(fn any) (Task, error) { return newTaskFunction(fn, 
binding.Analyze, nil) }
 
 // NewPositionalTaskFunction is like NewTaskFunction, but the Task binds each 
argument to one
 // parameter, in order, as binding.AnalyzePositional describes.
 func NewPositionalTaskFunction(fn any) (Task, error) {
-       return newTaskFunction(fn, binding.AnalyzePositional)
+       return newTaskFunction(fn, binding.AnalyzePositional, nil)
+}
+
+// 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.
+func NewPositionalBranchFunction(fn any, findSkipped func(result any) 
[]string) (Task, error) {
+       return newTaskFunction(fn, binding.AnalyzePositional, findSkipped)
 }
 
 func newTaskFunction(
        fn any,
        analyze func(fnType reflect.Type, fnName string) (*binding.Plan, error),
+       findSkipped func(result any) []string,
 ) (Task, error) {
        // The kind comes first: Value.Pointer panics on an int, and Value.Type 
on an untyped nil.
        v := reflect.ValueOf(fn)
        if v.Kind() != reflect.Func {
                return nil, fmt.Errorf("expected a func as input but was %s", 
v.Kind())
        }
-       f := &taskFunction{fn: v, fullName: 
runtime.FuncForPC(v.Pointer()).Name()}
+       f := &taskFunction{
+               fn:          v,
+               fullName:    runtime.FuncForPC(v.Pointer()).Name(),
+               findSkipped: findSkipped,
+       }
        if err := f.validateFn(v.Type(), analyze); err != nil {
                return nil, err
        }
@@ -100,11 +116,17 @@ func (f *taskFunction) Execute(
        if err != nil {
                return err
        }
+       var branch *branchRun
+       if f.findSkipped != nil {
+               if branch, err = startBranch(ctx, sdkClient); err != nil {
+                       return err
+               }
+       }
        reflectArgs, err := f.plan.Resolve(ctx, logger, sdkClient, args)
        if err != nil {
                return err
        }
-       return f.call(ctx, sdkClient, reflectArgs, logger)
+       return f.call(ctx, sdkClient, reflectArgs, logger, branch)
 }
 
 func clientFrom(ctx context.Context) (sdk.Client, error) {
@@ -120,6 +142,7 @@ func (f *taskFunction) call(
        sdkClient sdk.Client,
        reflectArgs []reflect.Value,
        logger *slog.Logger,
+       branch *branchRun,
 ) error {
        slog.Debug("Attempting to call fn", "fn", f.fn, "args", reflectArgs)
        retValues := f.fn.Call(reflectArgs)
@@ -139,9 +162,86 @@ func (f *taskFunction) call(
                res := retValues[0].Interface()
                f.sendXcom(ctx, res, sdkClient, logger)
        }
+       if err == nil && branch != nil {
+               return branch.skipDownstream(
+                       ctx,
+                       sdkClient,
+                       f.findSkipped(retValues[0].Interface()),
+                       logger,
+               )
+       }
        return err
 }
 
+// skipMixinXComKey is the key of the XCom that lists the tasks a task 
skipped. When one of those
+// tasks is cleared, NotPreviouslySkippedDep in Airflow core reads the XCom 
and skips the cleared
+// task again instead of running it. SkipMixin in the standard provider writes 
the same key.
+const skipMixinXComKey = "skipmixin_key"
+
+type skipDownstreamTasksKey struct{}
+
+// WithSkipDownstreamTasks returns a copy of ctx that carries skip, the 
function that a task from
+// NewPositionalBranchFunction calls to skip tasks downstream of it. The 
runtime passes skip in the
+// context and not in sdk.Client, so that task functions cannot call it. 
Airflow core reads the
+// skipmixin_key XCom only from a task that has _can_skip_downstream set in 
the serialized Dag. A
+// task skipped by a task without that flag would run when someone clears it.
+func WithSkipDownstreamTasks(
+       ctx context.Context,
+       skip func(ctx context.Context, taskIDs []string) error,
+) context.Context {
+       return context.WithValue(ctx, skipDownstreamTasksKey{}, skip)
+}
+
+// branchRun holds the skip function and the task instance that a run of a 
task from
+// NewPositionalBranchFunction uses to skip tasks after fn returns.
+type branchRun struct {
+       skip func(ctx context.Context, taskIDs []string) error
+       ti   sdk.TaskInstance
+}
+
+// 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) {
+       skip, ok := ctx.Value(skipDownstreamTasksKey{}).(func(context.Context, 
[]string) error)
+       if !ok {
+               return nil, errors.New("the task runtime cannot skip downstream 
tasks")
+       }
+       runtimeContext, ok := 
ctx.Value(sdkcontext.RuntimeContextKey).(sdk.TIRunContext)
+       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
+}
+
+// skipDownstream records taskIDs in the skipmixin_key XCom, and then skips 
those tasks.
+func (b *branchRun) skipDownstream(
+       ctx context.Context,
+       client sdk.Client,
+       taskIDs []string,
+       logger *slog.Logger,
+) error {
+       if len(taskIDs) == 0 {
+               return nil
+       }
+       err := client.PushXCom(ctx, b.ti, skipMixinXComKey, 
map[string][]string{"skipped": taskIDs})
+       if err != nil {
+               return fmt.Errorf("recording the skipped tasks in the %s XCom: 
%w", skipMixinXComKey, err)
+       }
+       logger.InfoContext(ctx, "Skipping downstream tasks", "task_ids", 
taskIDs)
+       if err := b.skip(ctx, taskIDs); err != nil {
+               return fmt.Errorf("skipping the downstream tasks %q: %w", 
taskIDs, err)
+       }
+       return nil
+}
+
 func (f *taskFunction) sendXcom(
        ctx context.Context,
        value any,
diff --git a/go-sdk/internal/bundle/task_test.go 
b/go-sdk/internal/bundle/task_test.go
index d1d302ab0ec..40e023f784a 100644
--- a/go-sdk/internal/bundle/task_test.go
+++ b/go-sdk/internal/bundle/task_test.go
@@ -19,6 +19,8 @@ package bundle
 
 import (
        "context"
+       "errors"
+       "fmt"
        "log/slog"
        "testing"
 
@@ -233,3 +235,218 @@ func (s *TaskSuite) 
TestExecuteRequiresCoordinatorClient() {
                s.Contains(err.Error(), "coordinator SDK client is missing")
        }
 }
+
+// branchClient records each XCom that a task pushes and keeps the last value 
of each key.
+// PushXCom fails for the call whose record equals failOn. Its skip method 
stands in for the
+// function that the runtime passes through WithSkipDownstreamTasks. skip 
records each call in the
+// same list as the XComs, so that the tests see the order of the calls. It 
returns skipErr.
+type branchClient struct {
+       sdk.Client
+       calls   []string
+       values  map[string]any
+       failOn  string
+       skipErr error
+}
+
+func (c *branchClient) PushXCom(
+       _ context.Context,
+       ti sdk.TaskInstance,
+       key string,
+       value any,
+) error {
+       call := fmt.Sprintf("PushXCom %s %s %v", ti.TaskID, key, value)
+       c.calls = append(c.calls, call)
+       if c.values == nil {
+               c.values = map[string]any{}
+       }
+       c.values[key] = value
+       if call == c.failOn {
+               return errors.New("xcom refused")
+       }
+       return nil
+}
+
+func (c *branchClient) skip(_ context.Context, taskIDs []string) error {
+       c.calls = append(c.calls, fmt.Sprintf("SkipDownstreamTasks %v", 
taskIDs))
+       return c.skipErr
+}
+
+// runBranch runs task as the task instance decide of dag1 with client. When 
ti is false, the
+// context has no task instance. When canSkip is false, the context has no 
function to skip
+// tasks with.
+func runBranch(task Task, client *branchClient, ti, canSkip bool) error {
+       ctx := context.WithValue(
+               context.Background(),
+               sdkcontext.SdkClientContextKey,
+               sdk.Client(client),
+       )
+       if ti {
+               ctx = context.WithValue(ctx, sdkcontext.RuntimeContextKey, 
sdk.NewTIRunContext(
+                       context.Background(),
+                       sdk.TaskInstance{DagID: "dag1", RunID: "run1", TaskID: 
"decide"},
+                       sdk.DagRun{DagID: "dag1", RunID: "run1"},
+               ))
+       }
+       if canSkip {
+               ctx = WithSkipDownstreamTasks(ctx, client.skip)
+       }
+       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(
+               func(contexttest.Context) (bool, error) { return true, nil },
+               func(result any) []string {
+                       got = result
+                       return []string{"load", "report"}
+               },
+       )
+       s.Require().NoError(err)
+
+       client := &branchClient{}
+       s.Require().NoError(runBranch(task, client, true, true))
+
+       s.Equal(true, got)
+       s.Equal([]string{
+               clearCall,
+               "PushXCom decide return_value true",
+               "PushXCom decide skipmixin_key map[skipped:[load report]]",
+               "SkipDownstreamTasks [load report]",
+       }, client.calls)
+}
+
+func (s *TaskSuite) TestBranchFunctionWithNothingToSkip() {
+       for name, skipped := range map[string][]string{"nil": nil, "empty": {}} 
{
+               s.Run(name, func() {
+                       task, err := NewPositionalBranchFunction(
+                               func(contexttest.Context) (bool, error) { 
return true, nil },
+                               func(any) []string { return skipped },
+                       )
+                       s.Require().NoError(err)
+
+                       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",
+                       )
+               })
+       }
+}
+
+// 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() {
+       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") },
+       }
+       for name, fn := range cases {
+               s.Run(name, func() {
+                       called := false
+                       task, err := NewPositionalBranchFunction(fn, func(any) 
[]string {
+                               called = true
+                               return []string{"load"}
+                       })
+                       s.Require().NoError(err)
+
+                       client := &branchClient{values: map[string]any{
+                               "skipmixin_key": map[string][]string{"skipped": 
{"load"}},
+                       }}
+                       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, "SkipDownstreamTasks")
+                       }
+               })
+       }
+}
+
+func (s *TaskSuite) TestBranchFunctionFailsWhenItCannotSkip() {
+       cases := map[string]struct {
+               client    *branchClient
+               withTI    bool
+               canSkip   bool
+               wantErr   string
+               wantRun   bool
+               wantCalls []string
+       }{
+               "runtime cannot skip": {
+                       client:  &branchClient{},
+                       withTI:  true,
+                       wantErr: "the task runtime cannot skip downstream 
tasks",
+               },
+               "no task instance": {
+                       client:  &branchClient{},
+                       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]]",
+                       },
+                       withTI:  true,
+                       canSkip: true,
+                       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]]",
+                       },
+               },
+               "skip fails": {
+                       client:  &branchClient{skipErr: errors.New("supervisor 
refused")},
+                       withTI:  true,
+                       canSkip: true,
+                       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]",
+                       },
+               },
+       }
+       for name, tt := range cases {
+               s.Run(name, func() {
+                       ran := false
+                       task, err := NewPositionalBranchFunction(
+                               func(contexttest.Context) (bool, error) {
+                                       ran = true
+                                       return false, nil
+                               },
+                               func(any) []string { return []string{"load"} },
+                       )
+                       s.Require().NoError(err)
+
+                       s.EqualError(runBranch(task, tt.client, tt.withTI, 
tt.canSkip), tt.wantErr)
+                       s.Equal(tt.wantRun, ran, "whether fn ran")
+                       s.Equal(tt.wantCalls, tt.client.calls)
+               })
+       }
+}
diff --git a/go-sdk/pkg/execution/client.go b/go-sdk/pkg/execution/client.go
index 118ffd893da..b98d8883ed9 100644
--- a/go-sdk/pkg/execution/client.go
+++ b/go-sdk/pkg/execution/client.go
@@ -253,3 +253,11 @@ func (c *CoordinatorClient) PushXCom(
        _, err := c.comm.Communicate(ctx, msg)
        return err
 }
+
+// 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.
+func (c *CoordinatorClient) skipDownstreamTasks(ctx context.Context, taskIDs 
[]string) error {
+       _, err := c.comm.Communicate(ctx, genmodels.SkipDownstreamTasks{Tasks: 
taskIDs})
+       return err
+}
diff --git a/go-sdk/pkg/execution/client_test.go 
b/go-sdk/pkg/execution/client_test.go
index 0c0a09d7cdb..b7c8d39e7e8 100644
--- a/go-sdk/pkg/execution/client_test.go
+++ b/go-sdk/pkg/execution/client_test.go
@@ -469,3 +469,50 @@ func (f assertNoWriteWriter) Write(p []byte) (int, error) {
        f.t.Fatalf("unexpected Write on comm socket: env override should have 
short-circuited")
        return 0, nil
 }
+
+// TestCoordinatorClientSkipDownstreamTasks verifies the SkipDownstreamTasks 
frame sent to the
+// supervisor, and that skipDownstreamTasks returns a supervisor ErrorResponse 
as an error.
+func TestCoordinatorClientSkipDownstreamTasks(t *testing.T) {
+       tests := []struct {
+               name    string
+               errBody map[string]any
+               wantErr bool
+       }{
+               {name: "skipped"},
+               {
+                       name: "supervisor error",
+                       errBody: map[string]any{
+                               "type":   "ErrorResponse",
+                               "error":  "API_SERVER_ERROR",
+                               "detail": map[string]any{"status_code": 404},
+                       },
+                       wantErr: true,
+               },
+       }
+       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))
+
+                       err := client.skipDownstreamTasks(context.Background(), 
[]string{"load", "report"})
+                       if tc.wantErr {
+                               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)
+                       assert.Equal(t, map[string]any{
+                               "type":  "SkipDownstreamTasks",
+                               "tasks": []any{"load", "report"},
+                       }, rawToMap(t, sent.Body))
+               })
+       }
+}
diff --git a/go-sdk/pkg/execution/integration_test.go 
b/go-sdk/pkg/execution/integration_test.go
index c53bed89695..66d6611f084 100644
--- a/go-sdk/pkg/execution/integration_test.go
+++ b/go-sdk/pkg/execution/integration_test.go
@@ -851,6 +851,130 @@ func TestServeClientRoundTripEndToEnd(t *testing.T) {
        assert.Equal(t, "hello", gotVar)
 }
 
+// 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.
+func TestServeSkipsDownstreamTasksEndToEnd(t *testing.T) {
+       xcom := func(key string, value any) map[string]any {
+               return map[string]any{
+                       "type":    "SetXCom",
+                       "dag_id":  "dag1",
+                       "run_id":  "run1",
+                       "task_id": "decide",
+                       "key":     key,
+                       "value":   value,
+               }
+       }
+       tests := []struct {
+               name         string
+               result       bool
+               wantRequests []map[string]any
+               // wantSkipLogs holds the task_ids of each "Skipping downstream 
tasks" log entry.
+               wantSkipLogs []any
+       }{
+               {
+                       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"}},
+                       },
+                       wantSkipLogs: []any{[]any{"load"}},
+               },
+               {
+                       name:   "true skips nothing",
+                       result: true,
+                       wantRequests: []map[string]any{
+                               xcom("skipmixin_key", map[string]any{"skipped": 
[]any{}}),
+                               xcom("return_value", true),
+                       },
+               },
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       commAddr, logsAddr, commCh, logsCh, cleanup := 
startSupervisor(t)
+                       defer cleanup()
+
+                       decide, err := bundle.NewPositionalBranchFunction(
+                               func(contexttest.Context) (bool, error) { 
return tt.result, nil },
+                               func(result any) []string {
+                                       if result.(bool) {
+                                               return nil
+                                       }
+                                       return []string{"load"}
+                               },
+                       )
+                       require.NoError(t, err)
+                       tasks := testBundle{"dag1": testDag{"decide": decide}}
+
+                       done := make(chan error, 1)
+                       go func() { done <- Serve(tasks, commAddr, logsAddr) }()
+
+                       commConn := <-commCh
+                       defer commConn.Close()
+                       logsConn := <-logsCh
+                       defer logsConn.Close()
+                       require.NoError(t, 
commConn.SetDeadline(time.Now().Add(10*time.Second)))
+                       require.NoError(t, 
logsConn.SetDeadline(time.Now().Add(10*time.Second)))
+
+                       startup, err := encodeRequest(0, map[string]any{
+                               "type": "StartupDetails",
+                               "ti": map[string]any{
+                                       "id":         
"550e8400-e29b-41d4-a716-446655440000",
+                                       "dag_id":     "dag1",
+                                       "task_id":    "decide",
+                                       "run_id":     "run1",
+                                       "try_number": 1,
+                               },
+                               "bundle_info": map[string]any{"name": "fake", 
"version": "1.0"},
+                       })
+                       require.NoError(t, err)
+                       require.NoError(t, writeFrame(commConn, startup))
+
+                       // Until the terminal frame arrives, answer each 
runtime request with an empty response
+                       // so that the task goes on.
+                       var requests []map[string]any
+                       for {
+                               frame, err := readFrame(commConn)
+                               require.NoError(t, err)
+                               require.True(t, isNilRaw(frame.Err))
+                               if peekBodyType(frame.Body) == "SucceedTask" {
+                                       break
+                               }
+                               requests = append(requests, rawToMap(t, 
frame.Body))
+                               reply, err := encodeRequest(frame.ID, 
map[string]any{})
+                               require.NoError(t, err)
+                               require.NoError(t, writeFrame(commConn, reply))
+                       }
+                       assert.Equal(t, tt.wantRequests, requests)
+
+                       select {
+                       case err := <-done:
+                               require.NoError(t, err)
+                       case <-time.After(2 * time.Second):
+                               t.Fatal("Serve did not return after task 
completion")
+                       }
+
+                       output, err := io.ReadAll(logsConn)
+                       require.NoError(t, err)
+                       var skipLogs []any
+                       for line := range strings.Lines(string(output)) {
+                               var entry map[string]any
+                               require.NoError(t, json.Unmarshal([]byte(line), 
&entry))
+                               if entry["event"] == "Skipping downstream 
tasks" {
+                                       skipLogs = append(skipLogs, 
entry["task_ids"])
+                               }
+                       }
+                       assert.Equal(t, tt.wantSkipLogs, skipLogs)
+               })
+       }
+}
+
 // 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 82defbbfe91..13e2d2d0f05 100644
--- a/go-sdk/pkg/execution/task_runner.go
+++ b/go-sdk/pkg/execution/task_runner.go
@@ -90,6 +90,7 @@ func RunTask(
 
        ctx = context.WithValue(ctx, sdkcontext.SdkClientContextKey, 
sdk.Client(client))
        ctx = context.WithValue(ctx, sdkcontext.RuntimeContextKey, 
runtimeContext)
+       ctx = bundle.WithSkipDownstreamTasks(ctx, client.skipDownstreamTasks)
 
        args, err := convertArgBindings(details.TIContext.ArgBindings)
        if err != nil {

Reply via email to