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 a13edb47aef Go SDK: add dag.Switch with Case for multi-way branching 
(#74321)
a13edb47aef is described below

commit a13edb47aef230311253220dd4e522872d1a4232
Author: PoAn Yang <[email protected]>
AuthorDate: Tue Oct 6 19:17:06 2026 +0800

    Go SDK: add dag.Switch with Case for multi-way branching (#74321)
    
    * Go SDK: add dag.Switch with Case for multi-way branching
    
    Signed-off-by: PoAn Yang <[email protected]>
    
    * Go SDK: address review of dag.Switch
    
    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                    |  83 +++--
 go-sdk/airflow/if.go                     |  24 +-
 go-sdk/airflow/if_test.go                |  48 +--
 go-sdk/airflow/inputs.go                 |   6 +-
 go-sdk/airflow/spec.gen.go               |   4 +-
 go-sdk/airflow/switch.go                 | 253 ++++++++++++++
 go-sdk/airflow/switch_test.go            | 583 +++++++++++++++++++++++++++++++
 go-sdk/airflow/task_group.go             |  11 +
 go-sdk/airflow/task_group_test.go        |  66 +++-
 go-sdk/airflow/task_option.go            |   6 +-
 go-sdk/internal/bundle/task.go           |  57 +--
 go-sdk/internal/bundle/task_test.go      |  56 ++-
 go-sdk/internal/genspec/authoring.go     |   4 +-
 go-sdk/pkg/execution/integration_test.go |  26 +-
 15 files changed, 1105 insertions(+), 129 deletions(-)

diff --git a/go-sdk/airflow/bundle.go b/go-sdk/airflow/bundle.go
index 6ac967a0de1..9abb02717a0 100644
--- a/go-sdk/airflow/bundle.go
+++ b/go-sdk/airflow/bundle.go
@@ -77,8 +77,8 @@ type Registerable interface{ registerable() }
 //     bundle.Register(reports.Handlers()...)
 //
 // Add every task to a Dag before registering the Dag. [DagRef.Task], 
[DagRef.If],
-// [DagRef.TaskGroup], [IfRef.Then], [IfRef.Else] and the methods of 
[TaskGroupRef] panic once the
-// Dag is registered.
+// [DagRef.Switch], [DagRef.TaskGroup], [IfRef.Then], [IfRef.Else], 
[SwitchRef.Case] and the
+// methods of [TaskGroupRef] panic once the Dag is registered.
 //
 // Register is where a Dag's task dependencies are checked for a cycle, over 
the whole graph at
 // once: [TaskRef.Before], [TaskRef.After] and [Inputs] each record an edge 
without walking the
@@ -93,7 +93,8 @@ type Registerable interface{ registerable() }
 // same dag_id, if the task dependencies of a Dag contain a cycle, 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. 
Register also panics if a
-// Dag has a condition from [DagRef.If] without a task from [IfRef.Then].
+// Dag has a condition from [DagRef.If] without a task from [IfRef.Then], or a 
switch from
+// [DagRef.Switch] without a case from [SwitchRef.Case].
 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 710abe863ea..97a72a472ff 100644
--- a/go-sdk/airflow/dag.go
+++ b/go-sdk/airflow/dag.go
@@ -67,8 +67,9 @@ type DagRef struct {
 //
 //     bundle.Register(dag)
 //
-// Add every task before Register. [DagRef.Task], [DagRef.If], 
[DagRef.TaskGroup], [IfRef.Then],
-// [IfRef.Else] and the methods of [TaskGroupRef] panic once the Dag is 
registered.
+// Add every task before Register. [DagRef.Task], [DagRef.If], 
[DagRef.Switch], [DagRef.TaskGroup],
+// [IfRef.Then], [IfRef.Else], [SwitchRef.Case] and the methods of 
[TaskGroupRef] 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.
@@ -91,12 +92,12 @@ func (*DagRef) registerable() {}
 
 // TaskRef is a task that [DagRef.Task] or [TaskGroupRef.Task] added to a Dag. 
Pass it to [Inputs]
 // to give its result to a task added later. Pass it to [IfRef.Then] or 
[IfRef.Else] to run it on
-// one side of a condition. A TaskRef is a [Node], so [TaskRef.Before] and 
[TaskRef.After] order it
-// against another task or a task group.
+// one side of a condition. Pass it to [SwitchRef.Case] to make it a case of a 
switch. A TaskRef is
+// a [Node], so [TaskRef.Before] and [TaskRef.After] order it against another 
task or a task group.
 type TaskRef struct {
        dag *DagRef
        // group is the task group that the task was added through. It is nil 
for a task that
-       // DagRef.Task or DagRef.If added.
+       // DagRef.Task, DagRef.If or DagRef.Switch added.
        group  *TaskGroupRef
        taskID string
        spec   TaskSpec
@@ -116,9 +117,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 or TaskGroupRef.If returned for 
the task. It is nil for a
-       // task from DagRef.Task or TaskGroupRef.Task.
-       ifRef *IfRef
+       // decider is the IfRef or the SwitchRef that an If or a Switch method 
returned for the task. It
+       // is nil for a task from DagRef.Task or TaskGroupRef.Task.
+       decider decider
 }
 
 // Task adds a task that runs fn to the Dag and returns the new task.
@@ -165,11 +166,22 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) 
*TaskRef {
        return d.addTask("airflow.DagRef.Task", nil, fn, opts, nil)
 }
 
-// addTask adds a task for Task and If, of the Dag or of a task group. method 
names the caller in
-// panic messages. group is the task group that the task is added through, and 
nil for a task of
-// the Dag itself. ifRef is the IfRef that If returns, and nil when Task calls 
addTask.
+// decider is the IfRef that If returns or the SwitchRef that Switch returns. 
addTask uses it to
+// add the task that decides which tasks after it to skip.
+type decider interface {
+       // wrap checks the result types of fn and wraps fn as a bundle.Task 
that skips the tasks that fn
+       // does not choose.
+       wrap(fn any) (bundle.Task, error)
+       // bind records task as the task that runs the function of the decider.
+       bind(task *TaskRef)
+}
+
+// addTask adds a task for Task, If and Switch, of the Dag or of a task group. 
method names the
+// caller in panic messages. group is the task group that the task is added 
through, and nil for a
+// task of the Dag itself. decider is the IfRef or the SwitchRef of the task, 
and nil when Task
+// calls addTask.
 func (d *DagRef) addTask(
-       method string, group *TaskGroupRef, fn any, opts []TaskOption, ifRef 
*IfRef,
+       method string, group *TaskGroupRef, fn any, opts []TaskOption, decider 
decider,
 ) *TaskRef {
        trigger, isTrigger := fn.(TriggerDagRunTask)
        var triggerSpec *TriggerDagRunSpec
@@ -195,8 +207,8 @@ func (d *DagRef) addTask(
        var wrapped bundle.Task
        if !isTrigger {
                wrap := bundle.NewPositionalTaskFunction
-               if ifRef != nil {
-                       wrap = ifRef.wrapCondition
+               if decider != nil {
+                       wrap = decider.wrap
                }
                var err error
                if wrapped, err = newTaskFunction(fn, wrap); err != nil {
@@ -307,10 +319,10 @@ func (d *DagRef) addTask(
                inputs:        upstreams,
                task:          wrapped,
                triggerDagRun: triggerSpec,
-               ifRef:         ifRef,
+               decider:       decider,
        }
-       if ifRef != nil {
-               ifRef.task = task
+       if decider != nil {
+               decider.bind(task)
        }
        if d.tasksByID == nil {
                d.tasksByID = make(map[string]*TaskRef)
@@ -329,13 +341,13 @@ func (d *DagRef) addTask(
 }
 
 // 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, or when the edges of d 
close a cycle. The Dag is
-// whole by then, so markRegistered can expand the group edges in the order 
they were first
-// declared, and a walk of the whole graph answers for every edge. It walks 
the graph before the
-// expansion too, so that a cycle between the edges the author declared is 
reported as declared.
-// A Dag that fails a check stays unregistered and holds only the edges its 
author declared. The
-// author can still give a condition its task from Then, but cannot undo a 
cycle, since a Dag only
-// ever gains edges.
+// when a condition from If has no task from Then, when a switch from Switch 
has no case, or when
+// the edges of d close a cycle. The Dag is whole by then, so markRegistered 
can expand the group
+// edges in the order they were first declared, and a walk of the whole graph 
answers for every
+// edge. It walks the graph before the expansion too, so that a cycle between 
the edges the author
+// declared is reported as declared. A Dag that fails a check stays 
unregistered and holds only the
+// edges its author declared. The author can still complete the Dag with 
IfRef.Then or
+// SwitchRef.Case, but cannot undo a cycle, since a Dag only ever gains edges.
 func (d *DagRef) markRegistered() {
        d.mu.Lock()
        defer d.mu.Unlock()
@@ -345,12 +357,23 @@ func (d *DagRef) markRegistered() {
                return
        }
        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,
-                       ))
+               switch decider := task.decider.(type) {
+               case *IfRef:
+                       if decider.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,
+                               ))
+                       }
+               case *SwitchRef:
+                       if len(decider.cases) == 0 {
+                               panic(fmt.Sprintf(
+                                       "airflow.BundleRef.Register: switch %q 
of Dag %q has no case; "+
+                                               "name each task that the switch 
can choose with SwitchRef.Case",
+                                       task.taskID, d.dagID,
+                               ))
+                       }
                }
        }
        if cycle := d.cycleLocked(); cycle != nil {
diff --git a/go-sdk/airflow/if.go b/go-sdk/airflow/if.go
index 42b71e90ed0..a153237a0c2 100644
--- a/go-sdk/airflow/if.go
+++ b/go-sdk/airflow/if.go
@@ -63,7 +63,7 @@ type IfRef struct {
 // 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
+//   - an error fails the task, which then pushes no XCom and skips nothing
 //
 // A task from Then or Else runs after the condition, so naming it records an 
edge from the
 // condition to it, as [TaskRef.Before] would. Declaring that edge as well 
changes nothing.
@@ -132,7 +132,7 @@ 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 {
+       if g == nil || g.task == nil || g.task.decider != g {
                panic(method + ": DagRef.If or TaskGroupRef.If did not return 
the *airflow.IfRef")
        }
        condition, d := g.task.taskID, g.task.dag
@@ -190,9 +190,9 @@ func (g *IfRef) setTask(side string, task *TaskRef) {
        d.addEdgeLocked(g.task, 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) {
+// wrap 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) wrap(fn any) (bundle.Task, error) {
        fnType := reflect.TypeOf(fn)
        if fnType.NumOut() != 2 ||
                fnType.Out(0) != reflect.TypeFor[bool]() ||
@@ -202,20 +202,22 @@ func (g *IfRef) wrapCondition(fn any) (bundle.Task, 
error) {
                        funcName(fn), describeResults(fnType),
                )
        }
-       return bundle.NewPositionalBranchFunction(fn, g.findSkipped)
+       return bundle.NewPositionalBranchFunction(fn, g.decide)
 }
 
-// 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 {
+func (g *IfRef) bind(task *TaskRef) { g.task = task }
+
+// decide returns result as the value to push. It also returns the task_id of 
the task on the side
+// of g that result does not take, when that side has a task.
+func (g *IfRef) decide(result any) (any, []string, error) {
        notTaken := g.thenTask
        if result.(bool) {
                notTaken = g.elseTask
        }
        if notTaken == nil {
-               return nil
+               return result, nil, nil
        }
-       return []string{notTaken.taskID}
+       return result, []string{notTaken.taskID}, nil
 }
 
 func describeResults(fnType reflect.Type) string {
diff --git a/go-sdk/airflow/if_test.go b/go-sdk/airflow/if_test.go
index 059e2fa21dc..819ee60f58d 100644
--- a/go-sdk/airflow/if_test.go
+++ b/go-sdk/airflow/if_test.go
@@ -64,8 +64,8 @@ func TestIfAddsTheConditionAsATask(t *testing.T) {
        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")
+       assert.Same(t, gate, gate.task.decider)
+       assert.Nil(t, read.decider, "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)
@@ -460,18 +460,18 @@ func TestRegisterRejectsAConditionWithoutThen(t 
*testing.T) {
        }
 }
 
-// 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 {
+// deciderClient answers GetXCom from results, which maps a task_id to the 
result of that task. It
+// records the XComs that a condition or a switch pushes. Its skip method 
stands in for the
+// function that the runtime passes through bundle.WithSkipDownstreamTasks, 
and records the
+// task_ids that the condition or the switch skips.
+type deciderClient struct {
        sdk.Client
        results map[string]any
        xcoms   map[string]any
        skipped [][]string
 }
 
-func (c *conditionClient) GetXCom(
+func (c *deciderClient) GetXCom(
        _ context.Context,
        _, _, taskID string,
        _ *int,
@@ -481,27 +481,27 @@ func (c *conditionClient) GetXCom(
        return c.results[taskID], nil
 }
 
-func (c *conditionClient) PushXCom(_ context.Context, _ sdk.TaskInstance, key 
string, v any) error {
+func (c *deciderClient) 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 {
+func (c *deciderClient) 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.
-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 {
+// runDecider runs task, the task of a condition or a switch, 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.
+func runDecider(task *TaskRef, results map[string]any) (*deciderClient, error) 
{
+       args := make([]binding.Arg, len(task.inputs))
+       for i, upstream := range task.inputs {
                args[i] = binding.XComArg{Kind: "xcom", Name: "rows", TaskID: 
upstream.taskID}
        }
-       client := &conditionClient{results: results, xcoms: map[string]any{}}
-       ti := sdk.TaskInstance{DagID: "etl", RunID: "run1", TaskID: 
gate.task.taskID}
+       client := &deciderClient{results: results, xcoms: map[string]any{}}
+       ti := sdk.TaskInstance{DagID: "etl", RunID: "run1", TaskID: task.taskID}
        ctx := context.WithValue(
                context.Background(),
                sdkcontext.SdkClientContextKey,
@@ -513,7 +513,7 @@ func runCondition(gate *IfRef, results map[string]any) 
(*conditionClient, error)
                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)
+       return client, task.task.Execute(ctx, discardLogger(), args)
 }
 
 // TestRegisterRejectsACycleThroughACondition pins that the edge Then records 
is one the cycle
@@ -582,7 +582,7 @@ func TestConditionSkipsTheSideThatItDoesNotTake(t 
*testing.T) {
                        }
                        Bundle().Register(dag)
 
-                       client, err := runCondition(gate, map[string]any{
+                       client, err := runDecider(gate.task, map[string]any{
                                "readRows": map[string]any{"rows": tt.rows},
                        })
                        require.NoError(t, err)
@@ -602,7 +602,7 @@ func TestConditionSkipsTheSideThatItDoesNotTake(t 
*testing.T) {
        }
 }
 
-func TestConditionThatFailsSkipsNothing(t *testing.T) {
+func TestConditionThatFailsPushesAndSkipsNothing(t *testing.T) {
        dag := Dag("etl")
        gate := dag.If(
                func(Context) (bool, error) { return false, errors.New("cannot 
reach the table") },
@@ -610,11 +610,11 @@ func TestConditionThatFailsSkipsNothing(t *testing.T) {
        ).Then(dag.Task(load)).Else(dag.Task(reportEmpty))
        Bundle().Register(dag)
 
-       client, err := runCondition(gate, nil)
+       client, err := runDecider(gate.task, nil)
 
        require.EqualError(t, err, "cannot reach the table")
        assert.Empty(t, client.skipped)
-       assert.NotContains(t, client.xcoms, "skipmixin_key")
+       assert.Empty(t, client.xcoms)
 }
 
 // Else and Register both take the lock of the Dag, so each Else call either 
names its task before
diff --git a/go-sdk/airflow/inputs.go b/go-sdk/airflow/inputs.go
index 75775b3580c..41390729821 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], 
[DagRef.If], or one of the
-// methods of the same names on [TaskGroupRef] 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 a task that [DagRef.Task], 
[DagRef.If], [DagRef.Switch] or
+// a method of the same name on [TaskGroupRef] adds. Each of those tasks 
becomes 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))
diff --git a/go-sdk/airflow/spec.gen.go b/go-sdk/airflow/spec.gen.go
index d5eda56f3ad..5b7cccd829d 100644
--- a/go-sdk/airflow/spec.gen.go
+++ b/go-sdk/airflow/spec.gen.go
@@ -99,8 +99,8 @@ type TaskGroupSpec struct {
        UIFgColor string
 }
 
-// TaskSpec holds the attributes of a task. DagRef.Task, DagRef.If and the 
methods
-// of the same names on TaskGroupRef take at most one per task.
+// TaskSpec holds the attributes of a task. DagRef.Task, DagRef.If, 
DagRef.Switch
+// and the methods of the same names on TaskGroupRef take at most one per task.
 type TaskSpec struct {
        // TaskDisplayName corresponds to the JSON schema field 
"_task_display_name".
        TaskDisplayName string
diff --git a/go-sdk/airflow/switch.go b/go-sdk/airflow/switch.go
new file mode 100644
index 00000000000..d70c1508577
--- /dev/null
+++ b/go-sdk/airflow/switch.go
@@ -0,0 +1,253 @@
+// 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"
+       "slices"
+       "strconv"
+       "strings"
+
+       "github.com/apache/airflow/go-sdk/internal/bundle"
+)
+
+// SwitchRef is a switch that [DagRef.Switch] or [TaskGroupRef.Switch] added 
to a Dag. Each call to
+// [SwitchRef.Case] names a task that the switch can choose.
+type SwitchRef struct {
+       // task is the task that runs the decider function. The cases run after 
task, but none of them
+       // takes the result of task.
+       task *TaskRef
+       // cases holds the tasks that Case named, in the order Case named them.
+       cases []*TaskRef
+}
+
+// Switch adds a task that runs the decider function fn, and returns the 
switch. Use
+// [SwitchRef.Case] to name each task that fn can choose. fn returns 
(*airflow.TaskRef, error), and
+// the TaskRef that it returns is the case that runs. So fn needs the TaskRefs 
of the cases, for
+// example through the receiver of a method value:
+//
+//     type paths struct{ long, short *airflow.TaskRef }
+//
+//     func (p *paths) pickPath(actx airflow.Context, rows []string) 
(*airflow.TaskRef, error) {
+//             if len(rows) > 1000 {
+//                     return p.long, nil
+//             }
+//             return p.short, nil
+//     }
+//
+// Where the Dag is built:
+//
+//     extracted := dag.Task(extract)
+//     p := &paths{long: dag.Task(handleLong), short: dag.Task(handleShort)}
+//
+//     dag.Switch(p.pickPath, 
airflow.Inputs(extracted)).Case(p.long).Case(p.short)
+//
+// The task that Switch adds then gets the task_id pickPath, the name of the 
method. fn can also be
+// a function literal that closes over the TaskRefs. A function literal has no 
name to take the
+// task_id from, so pass Switch a [TaskSpec] that sets TaskID. A function that 
reads the TaskRefs
+// from package-level variables works only when the code that builds the Dag 
runs once: a second
+// Dag built by the same code would overwrite the variables.
+//
+// In every other way fn follows the rules of a function passed to 
[DagRef.Task]: it takes a
+// [Context] first, and [Inputs] fills the parameters after the Context. The 
task that Switch 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:
+//   - a case skips every other case
+//   - a TaskRef that is not a case, or a nil one, fails the task with an 
error that names the
+//     TaskRef and the cases
+//   - an error fails the task
+//
+// When the task fails, it pushes no XCom and skips nothing. Otherwise its 
return_value XCom holds
+// the task_id of the case that fn chose.
+//
+// fn chooses exactly one case. A Python branch can choose several tasks by 
returning a list of
+// task_ids, but a switch cannot. To run several tasks together, order them 
after one task and make
+// that task a case, or give each of them a condition of its own with 
[DagRef.If].
+//
+// A switch has no default case, because a Python branch operator has none. To 
run a task when no
+// other case fits, make it a case and have fn return it.
+//
+// A case runs after the switch, so naming it records an edge from the switch 
to it, as
+// [TaskRef.Before] would. When fn chooses a case, the switch skips only the 
cases that fn did not
+// choose. Whether a task after the cases runs is then up to its trigger rule. 
With the default
+// trigger rule, a task is skipped when one of its upstream tasks is skipped. 
So a task after both
+// cases of the example above would never run. Give that task a trigger rule 
like
+// [TriggerRuleNoneFailedMinOneSuccess], which runs a task when at least one 
of its upstream tasks
+// succeeds and none fails:
+//
+//     report := dag.Task(writeReport, airflow.TaskSpec{
+//             TriggerRule: airflow.TriggerRuleNoneFailedMinOneSuccess,
+//     })
+//     report.After(p.long, p.short)
+//
+// Do not make report a case as well. The switch skips every case that fn does 
not choose, so when
+// fn chooses p.short, the switch would skip report too. A Python branch works 
differently here: it
+// does not skip a task that runs after the task it chooses. So the Python 
pattern
+// branch >> [optional, join] with optional >> join needs a change in Go. Make 
join a task like
+// report, and give the path without optional a case of its own, such as a 
task that does nothing.
+//
+// [BundleRef.Register] panics if a switch has no case.
+//
+// Switch 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 (*airflow.TaskRef, error).
+func (d *DagRef) Switch(fn any, opts ...TaskOption) *SwitchRef {
+       return d.addSwitch("airflow.DagRef.Switch", nil, fn, opts)
+}
+
+// addSwitch adds a switch for DagRef.Switch and TaskGroupRef.Switch. group is 
the task group that
+// the switch is added through, and nil for DagRef.Switch.
+func (d *DagRef) addSwitch(
+       method string, group *TaskGroupRef, fn any, opts []TaskOption,
+) *SwitchRef {
+       if _, ok := fn.(TriggerDagRunTask); ok {
+               panic(fmt.Sprintf(
+                       "%s: Dag %q: fn comes from airflow.TriggerDagRun, "+
+                               "but a decider function is a Go function that 
returns (*airflow.TaskRef, error)",
+                       method, d.dagID,
+               ))
+       }
+       switchRef := &SwitchRef{}
+       d.addTask(method, group, fn, opts, switchRef)
+       return switchRef
+}
+
+// Case names task as a task that the switch s can choose. When the decider 
function of s chooses
+// another case, the task is skipped. Case returns s, so that a switch and all 
of its cases fit in
+// one statement, as in the example of [DagRef.Switch]:
+//
+//     dag.Switch(p.pickPath, 
airflow.Inputs(extracted)).Case(p.long).Case(p.short)
+//
+// The task runs after the decider, but its function does not take the result 
of the decider as a
+// parameter.
+//
+// Case panics if:
+//   - s is not the SwitchRef that DagRef.Switch or TaskGroupRef.Switch 
returned, for example a
+//     copy of it
+//   - task is nil, or neither DagRef.Task nor TaskGroupRef.Task added task to 
the Dag of the
+//     switch
+//   - task is already a case of s
+//   - the Dag is already registered
+func (s *SwitchRef) Case(task *TaskRef) *SwitchRef {
+       const method = "airflow.SwitchRef.Case"
+       // When the decider runs, it reads the cases from the SwitchRef that 
Switch returned. A case
+       // given to a copy of that SwitchRef would never reach the decider.
+       if s == nil || s.task == nil || s.task.decider != s {
+               panic(method + ": DagRef.Switch or TaskGroupRef.Switch did not 
return the " +
+                       "*airflow.SwitchRef")
+       }
+       decider, d := s.task.taskID, s.task.dag
+       d.mu.Lock()
+       defer d.mu.Unlock()
+
+       if d.registered {
+               panic(fmt.Sprintf(
+                       "%s: Dag %q has already been registered; "+
+                               "name the cases of every switch before 
Register",
+                       method, d.dagID,
+               ))
+       }
+       switch {
+       case task == nil:
+               panic(fmt.Sprintf(
+                       "%s: switch %q of Dag %q got a nil *airflow.TaskRef", 
method, decider, d.dagID,
+               ))
+       case task.dag != nil && task.dag != d:
+               panic(fmt.Sprintf(
+                       "%s: switch %q of Dag %q cannot choose task %q of 
another Dag, %q; "+
+                               "pass a task of the same Dag",
+                       method, decider, 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: switch %q of Dag %q got a *airflow.TaskRef that 
DagRef.Task or "+
+                               "TaskGroupRef.Task did not return",
+                       method, decider, d.dagID,
+               ))
+       case slices.Contains(s.cases, task):
+               panic(fmt.Sprintf(
+                       "%s: switch %q of Dag %q already has task %q as a case; 
name each case once",
+                       method, decider, d.dagID, task.taskID,
+               ))
+       }
+       s.cases = append(s.cases, task)
+       // The case runs after the decider, so the case is a downstream task of 
the decider. Recording
+       // the edge puts the decider in the serialized Dag as the upstream of 
the case, and lets
+       // registration see a cycle that runs through a switch.
+       d.addEdgeLocked(s.task, task, "")
+       return s
+}
+
+// wrap wraps fn as the task of s. The task skips the cases of s that fn does 
not choose.
+func (s *SwitchRef) wrap(fn any) (bundle.Task, error) {
+       fnType := reflect.TypeOf(fn)
+       if fnType.NumOut() != 2 ||
+               fnType.Out(0) != reflect.TypeFor[*TaskRef]() ||
+               fnType.Out(1) != reflect.TypeFor[error]() {
+               return nil, fmt.Errorf(
+                       "%s returns %s, but a decider function must return 
(*airflow.TaskRef, error)",
+                       funcName(fn), describeResults(fnType),
+               )
+       }
+       return bundle.NewPositionalBranchFunction(fn, s.decide)
+}
+
+func (s *SwitchRef) bind(task *TaskRef) { s.task = task }
+
+// decide returns the task_id of result, the case that the decider chose, as 
the value to push, and
+// the task_ids of the other cases of s to skip. It returns an error when 
result is not a case of s.
+func (s *SwitchRef) decide(result any) (any, []string, error) {
+       chosen := result.(*TaskRef)
+       if !slices.Contains(s.cases, chosen) {
+               cases := make([]string, len(s.cases))
+               for i, task := range s.cases {
+                       cases[i] = strconv.Quote(task.taskID)
+               }
+               return nil, nil, fmt.Errorf(
+                       "switch %q of Dag %q returned %s, which is not one of 
its cases: %s",
+                       s.task.taskID, s.task.dag.dagID, 
s.describeChoice(chosen), strings.Join(cases, ", "),
+               )
+       }
+       skipped := make([]string, 0, len(s.cases)-1)
+       for _, task := range s.cases {
+               if task != chosen {
+                       skipped = append(skipped, task.taskID)
+               }
+       }
+       return chosen.taskID, skipped, nil
+}
+
+// describeChoice names chosen, which the decider of s returned and which is 
not a case of s.
+func (s *SwitchRef) describeChoice(chosen *TaskRef) string {
+       d := s.task.dag
+       switch {
+       case chosen == nil:
+               return "a nil *airflow.TaskRef"
+       // A zero TaskRef gets here.
+       case chosen.dag == nil:
+               return "a *airflow.TaskRef that DagRef.Task or 
TaskGroupRef.Task did not return"
+       case chosen.dag != d:
+               return fmt.Sprintf("task %q of another Dag, %q", chosen.taskID, 
chosen.dag.dagID)
+       case d.tasksByID[chosen.taskID] != chosen:
+               return fmt.Sprintf("a copy of the *airflow.TaskRef of task %q", 
chosen.taskID)
+       default:
+               return fmt.Sprintf("task %q", chosen.taskID)
+       }
+}
diff --git a/go-sdk/airflow/switch_test.go b/go-sdk/airflow/switch_test.go
new file mode 100644
index 00000000000..9055d19d406
--- /dev/null
+++ b/go-sdk/airflow/switch_test.go
@@ -0,0 +1,583 @@
+// 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 (
+       "errors"
+       "fmt"
+       "reflect"
+       "strings"
+       "testing"
+       "time"
+
+       "github.com/stretchr/testify/assert"
+       "github.com/stretchr/testify/require"
+)
+
+type deciderError struct{}
+
+func (*deciderError) Error() string { return "decider failed" }
+
+func pickPath(Context) (*TaskRef, error)                     { return nil, nil 
}
+func pickPathFromRows(Context, rowSet) (*TaskRef, error)     { return nil, nil 
}
+func pickPathOnlyError(Context) error                        { return nil }
+func pickPathAsTaskRef(Context) (TaskRef, error)             { return 
TaskRef{}, nil }
+func pickPathAsNode(Context) (Node, error)                   { return nil, nil 
}
+func pickPathAsTaskID(Context) (string, error)               { return "", nil }
+func pickPathWithoutError(Context) *TaskRef                  { return nil }
+func pickPathWithoutResults(Context)                         {}
+func pickPathWithOwnError(Context) (*TaskRef, *deciderError) { return nil, nil 
}
+func handleLong(Context) error                               { return nil }
+func handleShort(Context) error                              { return nil }
+func handleOther(Context) error                              { return nil }
+
+func TestSwitchAddsTheDeciderAsATask(t *testing.T) {
+       dag := Dag("etl")
+       read := dag.Task(readRows)
+
+       pick := dag.Switch(pickPathFromRows, Inputs(read))
+
+       require.NotNil(t, pick.task)
+       assert.Equal(t, "pickPathFromRows", pick.task.taskID)
+       assert.Equal(t, []*TaskRef{read, pick.task}, dag.tasks)
+       assert.Same(t, pick.task, dag.tasksByID["pickPathFromRows"])
+       assertInputs(t, pick.task, read)
+       assert.Equal(t, reflect.TypeFor[*TaskRef](), pick.task.resultType)
+       assert.Same(t, pick, pick.task.decider)
+       assert.Nil(t, read.decider, "a task from DagRef.Task is not a switch")
+
+       named := dag.Switch(pickPath, TaskSpec{TaskID: "pick_path"})
+       assert.Equal(t, "pick_path", named.task.taskID)
+       assert.Equal(t, TaskSpec{TaskID: "pick_path"}, named.task.spec)
+}
+
+func TestSwitchNeedsAFunctionThatReturnsATaskRef(t *testing.T) {
+       tests := []struct {
+               name    string
+               fn      any
+               fnName  string
+               returns string
+       }{
+               {"only an error", pickPathOnlyError, "pickPathOnlyError", 
"error"},
+               {"a TaskRef value", pickPathAsTaskRef, "pickPathAsTaskRef", 
"(airflow.TaskRef, error)"},
+               {"a Node", pickPathAsNode, "pickPathAsNode", "(airflow.Node, 
error)"},
+               {"a task_id", pickPathAsTaskID, "pickPathAsTaskID", "(string, 
error)"},
+               {"no error", pickPathWithoutError, "pickPathWithoutError", 
"*airflow.TaskRef"},
+               {"no results", pickPathWithoutResults, 
"pickPathWithoutResults", "nothing"},
+               {
+                       "another error type",
+                       pickPathWithOwnError,
+                       "pickPathWithOwnError",
+                       "(*airflow.TaskRef, *airflow.deciderError)",
+               },
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       dag := Dag("etl")
+                       assert.PanicsWithValue(t,
+                               `airflow.DagRef.Switch: Dag "etl": 
github.com/apache/airflow/go-sdk/airflow.`+
+                                       tt.fnName+` returns `+tt.returns+
+                                       `, but a decider function must return 
(*airflow.TaskRef, error)`,
+                               func() { dag.Switch(tt.fn) },
+                       )
+                       assert.Empty(t, dag.tasks)
+               })
+       }
+}
+
+func TestSwitchRejectsTriggerDagRun(t *testing.T) {
+       dag := Dag("etl")
+       assert.PanicsWithValue(t,
+               `airflow.DagRef.Switch: Dag "etl": fn comes from 
airflow.TriggerDagRun, `+
+                       `but a decider function is a Go function that returns 
(*airflow.TaskRef, error)`,
+               func() {
+                       dag.Switch(
+                               TriggerDagRun(TriggerDagRunSpec{DagID: 
"downstream_etl"}),
+                               TaskSpec{TaskID: "trigger"},
+                       )
+               },
+       )
+       assert.Empty(t, dag.tasks)
+}
+
+// Switch shares its checks with DagRef.Task, and each panic message has to 
name DagRef.Switch, the
+// function that the caller called.
+func TestSwitchPanicsUnderItsOwnName(t *testing.T) {
+       literal := func(Context) (*TaskRef, error) { return nil, nil }
+
+       tests := []struct {
+               name string
+               add  func(dag *DagRef)
+               // want is a part of the panic message after 
"airflow.DagRef.Switch: ".
+               want string
+       }{
+               {
+                       name: "registered Dag",
+                       add: func(dag *DagRef) {
+                               Bundle().Register(dag)
+                               dag.Switch(pickPath)
+                       },
+                       want: "has already been registered",
+               },
+               {
+                       name: "not a function",
+                       add:  func(dag *DagRef) { dag.Switch("pickPath") },
+                       want: "fn is string, not a function",
+               },
+               {
+                       name: "no Context first",
+                       add: func(dag *DagRef) {
+                               dag.Switch(func(rowSet) (*TaskRef, error) { 
return nil, nil })
+                       },
+                       want: "but the first parameter must be airflow.Context",
+               },
+               {
+                       name: "nil option",
+                       add:  func(dag *DagRef) { dag.Switch(pickPath, nil) },
+                       want: "opts[0] is nil",
+               },
+               {
+                       name: "second TaskSpec",
+                       add:  func(dag *DagRef) { dag.Switch(pickPath, 
TaskSpec{}, TaskSpec{}) },
+                       want: "got more than one airflow.TaskSpec",
+               },
+               {
+                       name: "no name for the task_id",
+                       add:  func(dag *DagRef) { dag.Switch(literal) },
+                       want: "has no name to use as the task_id",
+               },
+               {
+                       name: "unknown trigger rule",
+                       add:  func(dag *DagRef) { dag.Switch(pickPath, 
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: "pickPath"})
+                               dag.Switch(pickPath)
+                       },
+                       want: `already has a task "pickPath"`,
+               },
+               {
+                       name: "missing input",
+                       add:  func(dag *DagRef) { dag.Switch(pickPathFromRows) 
},
+                       want: "but airflow.Inputs passes no task",
+               },
+               {
+                       name: "input of the wrong type",
+                       add:  func(dag *DagRef) { dag.Switch(pickPathFromRows, 
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.Switch: "), msg)
+                       assert.Contains(t, msg, tt.want)
+               })
+       }
+}
+
+func TestCaseNamesTheTasksOfTheSwitch(t *testing.T) {
+       dag := Dag("etl")
+       long := dag.Task(handleLong)
+       short := dag.Task(handleShort)
+
+       pick := dag.Switch(pickPath)
+       assert.Same(t, pick, pick.Case(long))
+       assert.Same(t, pick, pick.Case(short))
+
+       assert.Equal(t, []*TaskRef{long, short}, pick.cases)
+       assert.Empty(t, long.inputs, "the decider passes no result to a case")
+       assert.Empty(t, short.inputs, "the decider passes no result to a case")
+}
+
+// TestCaseRecordsTheEdgeFromTheDecider pins that naming a case orders it 
after the decider. The
+// edge puts the decider in the serialized Dag as the upstream of the case.
+func TestCaseRecordsTheEdgeFromTheDecider(t *testing.T) {
+       dag := Dag("etl")
+       long := dag.Task(handleLong)
+       short := dag.Task(handleShort)
+
+       pick := dag.Switch(pickPath).Case(long).Case(short)
+
+       assertTasks(t, pick.task.downstreams, long, short)
+       assertTasks(t, long.upstreams, pick.task)
+       assertTasks(t, short.upstreams, pick.task)
+       assertEdgeLabel(t, dag, "pickPath", "handleLong", "")
+       assertEdgeLabel(t, dag, "pickPath", "handleShort", "")
+}
+
+func TestCaseRejectsATaskOutsideTheDag(t *testing.T) {
+       dag := Dag("etl")
+       long := dag.Task(handleLong)
+       longCopy := *long
+       fromReports := Dag("reports").Task(handleLong)
+       fromAnotherEtl := Dag("etl").Task(handleLong)
+
+       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 or 
TaskGroupRef.Task did not return`,
+               },
+               {
+                       name: "copy of a task of the Dag",
+                       task: &longCopy,
+                       want: `got a *airflow.TaskRef that DagRef.Task or 
TaskGroupRef.Task did not return`,
+               },
+               {
+                       name: "task of another Dag",
+                       task: fromReports,
+                       want: `cannot choose task "handleLong" 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 choose task "handleLong" of another Dag, 
"etl"; ` +
+                               `pass a task of the same Dag`,
+               },
+       }
+       for i, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       decider := fmt.Sprintf("pick_%d", i)
+                       pick := dag.Switch(pickPath, TaskSpec{TaskID: decider})
+                       assert.PanicsWithValue(t,
+                               `airflow.SwitchRef.Case: switch "`+decider+`" 
of Dag "etl" `+tt.want,
+                               func() { pick.Case(tt.task) },
+                       )
+                       assert.Empty(t, pick.cases)
+                       assert.Empty(t, pick.task.downstreams)
+               })
+       }
+}
+
+func TestCaseRejectsATaskThatIsAlreadyACase(t *testing.T) {
+       dag := Dag("etl")
+       long := dag.Task(handleLong)
+       short := dag.Task(handleShort)
+       pick := dag.Switch(pickPath).Case(long).Case(short)
+       before := snapshot(dag)
+
+       assert.PanicsWithValue(t,
+               `airflow.SwitchRef.Case: switch "pickPath" of Dag "etl" already 
has task "handleLong" `+
+                       `as a case; name each case once`,
+               func() { pick.Case(long) },
+       )
+       assert.Equal(t, before, snapshot(dag))
+}
+
+func TestCaseAfterRegisterPanics(t *testing.T) {
+       dag := Dag("etl")
+       long := dag.Task(handleLong)
+       short := dag.Task(handleShort)
+       pick := dag.Switch(pickPath).Case(long)
+       Bundle().Register(dag)
+
+       assert.PanicsWithValue(t,
+               `airflow.SwitchRef.Case: Dag "etl" has already been registered; 
`+
+                       `name the cases of every switch before Register`,
+               func() { pick.Case(short) },
+       )
+       assert.Equal(t, []*TaskRef{long}, pick.cases)
+}
+
+func TestSwitchRefThatDagRefSwitchDidNotReturnPanics(t *testing.T) {
+       dag := Dag("etl")
+       long := dag.Task(handleLong)
+       pick := dag.Switch(pickPath)
+       pickCopy := *pick
+       var nilRef *SwitchRef
+
+       for name, ref := range map[string]*SwitchRef{"nil": nilRef, "zero": {}, 
"copy": &pickCopy} {
+               t.Run(name, func(t *testing.T) {
+                       assert.PanicsWithValue(t,
+                               "airflow.SwitchRef.Case: DagRef.Switch or 
TaskGroupRef.Switch did not return "+
+                                       "the *airflow.SwitchRef",
+                               func() { ref.Case(long) },
+                       )
+               })
+       }
+       assert.Empty(t, pick.cases)
+}
+
+func TestRegisterRejectsASwitchWithoutACase(t *testing.T) {
+       dag := Dag("etl")
+       dag.Switch(pickPath, TaskSpec{TaskID: 
"complete"}).Case(dag.Task(handleLong))
+       dag.Switch(pickPath)
+       b := Bundle()
+
+       assert.PanicsWithValue(t,
+               `airflow.BundleRef.Register: switch "pickPath" of Dag "etl" has 
no case; `+
+                       `name each task that the switch can choose with 
SwitchRef.Case`,
+               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")
+}
+
+// TestRegisterTakesASwitchWithOneCase pins that one case is enough. The 
decider then has to
+// return that case, and skips nothing.
+func TestRegisterTakesASwitchWithOneCase(t *testing.T) {
+       dag := Dag("etl")
+       dag.Switch(pickPath).Case(dag.Task(handleLong))
+
+       assert.NotPanics(t, func() { Bundle().Register(dag) })
+       assert.True(t, dag.registered)
+}
+
+// TestRegisterRejectsACycleThroughASwitch pins that the cycle check walks the 
edge that Case
+// records, so Register rejects a cycle that runs through a switch like any 
other cycle.
+func TestRegisterRejectsACycleThroughASwitch(t *testing.T) {
+       dag := Dag("etl")
+       read := dag.Task(readRows)
+       long := dag.Task(handleLong)
+       dag.Switch(pickPathFromRows, Inputs(read)).Case(long)
+       long.Before(read)
+
+       assert.PanicsWithValue(t,
+               `airflow.BundleRef.Register: the task dependencies of Dag "etl" 
contain a cycle: `+
+                       `readRows -> pickPathFromRows -> handleLong -> 
readRows`,
+               func() { Bundle().Register(dag) },
+       )
+       assert.False(t, dag.registered)
+}
+
+// switchDag builds and registers a Dag whose switch pick_path has the given 
number of cases, taken
+// in order from handle_long, handle_short and handle_other. The decider 
returns *choice. The
+// task_ids differ from the function names, so that a test shows that the 
switch pushes and skips a
+// case by its task_id.
+func switchDag(t *testing.T, cases int, choice **TaskRef) (*SwitchRef, 
[]*TaskRef) {
+       t.Helper()
+       dag := Dag("etl")
+       tasks := []*TaskRef{
+               dag.Task(handleLong, TaskSpec{TaskID: "handle_long"}),
+               dag.Task(handleShort, TaskSpec{TaskID: "handle_short"}),
+               dag.Task(handleOther, TaskSpec{TaskID: "handle_other"}),
+       }
+       pick := dag.Switch(
+               func(Context) (*TaskRef, error) { return *choice, nil },
+               TaskSpec{TaskID: "pick_path"},
+       )
+       for _, task := range tasks[:cases] {
+               pick.Case(task)
+       }
+       Bundle().Register(dag)
+       return pick, tasks
+}
+
+func TestSwitchSkipsTheCasesThatItDoesNotChoose(t *testing.T) {
+       tests := []struct {
+               name    string
+               cases   int
+               choose  int
+               skipped []string
+       }{
+               {
+                       name:    "the first case skips the others",
+                       cases:   3,
+                       choose:  0,
+                       skipped: []string{"handle_short", "handle_other"},
+               },
+               {
+                       name:    "the last case skips the others",
+                       cases:   3,
+                       choose:  2,
+                       skipped: []string{"handle_long", "handle_short"},
+               },
+               {name: "the only case skips nothing", cases: 1, choose: 0},
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       var choice *TaskRef
+                       pick, tasks := switchDag(t, tt.cases, &choice)
+                       choice = tasks[tt.choose]
+
+                       client, err := runDecider(pick.task, nil)
+                       require.NoError(t, err)
+
+                       assert.Equal(t, tasks[tt.choose].taskID, 
client.xcoms["return_value"],
+                               "the switch pushes the task_id of the case, not 
the *airflow.TaskRef")
+                       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"],
+                               )
+                       }
+               })
+       }
+}
+
+// TestSwitchSkipsACaseThatRunsAfterTheChosenCase pins the difference from a 
Python branch that the
+// documentation of DagRef.Switch describes.
+func TestSwitchSkipsACaseThatRunsAfterTheChosenCase(t *testing.T) {
+       dag := Dag("etl")
+       optional := dag.Task(handleLong, TaskSpec{TaskID: "optional"})
+       join := dag.Task(handleShort, TaskSpec{TaskID: "join"})
+       optional.Before(join)
+       pick := dag.Switch(
+               func(Context) (*TaskRef, error) { return optional, nil },
+               TaskSpec{TaskID: "pick_path"},
+       ).Case(optional).Case(join)
+       Bundle().Register(dag)
+
+       client, err := runDecider(pick.task, nil)
+
+       require.NoError(t, err)
+       assert.Equal(t, [][]string{{"join"}}, client.skipped)
+}
+
+func TestSwitchFailsWhenTheDeciderReturnsATaskRefThatIsNotACase(t *testing.T) {
+       var choice *TaskRef
+       pick, tasks := switchDag(t, 2, &choice)
+       longCopy := *tasks[0]
+       fromReports := Dag("reports").Task(handleLong, TaskSpec{TaskID: 
"handle_long"})
+
+       tests := []struct {
+               name   string
+               choice *TaskRef
+               want   string
+       }{
+               {name: "a task of the Dag", choice: tasks[2], want: `task 
"handle_other"`},
+               {name: "nil", choice: nil, want: "a nil *airflow.TaskRef"},
+               {
+                       name:   "a task of another Dag",
+                       choice: fromReports,
+                       want:   `task "handle_long" of another Dag, "reports"`,
+               },
+               {
+                       name:   "a zero TaskRef",
+                       choice: &TaskRef{},
+                       want:   "a *airflow.TaskRef that DagRef.Task or 
TaskGroupRef.Task did not return",
+               },
+               {
+                       name:   "a copy of a case",
+                       choice: &longCopy,
+                       want:   `a copy of the *airflow.TaskRef of task 
"handle_long"`,
+               },
+       }
+       for _, tt := range tests {
+               t.Run(tt.name, func(t *testing.T) {
+                       choice = tt.choice
+
+                       client, err := runDecider(pick.task, nil)
+
+                       require.EqualError(t, err,
+                               `switch "pick_path" of Dag "etl" returned 
`+tt.want+
+                                       `, which is not one of its cases: 
"handle_long", "handle_short"`,
+                       )
+                       assert.Empty(t, client.xcoms, "a switch that fails 
pushes no XCom")
+                       assert.Empty(t, client.skipped)
+               })
+       }
+}
+
+// TestSwitchThatFailsPushesAndSkipsNothing covers a decider that returns a 
case together with an
+// error. The task fails without pushing the task_id of that case or skipping 
the other cases.
+func TestSwitchThatFailsPushesAndSkipsNothing(t *testing.T) {
+       dag := Dag("etl")
+       long := dag.Task(handleLong)
+       pick := dag.Switch(
+               func(Context) (*TaskRef, error) { return long, 
errors.New("cannot reach the table") },
+               TaskSpec{TaskID: "pick_path"},
+       ).Case(long).Case(dag.Task(handleShort))
+       Bundle().Register(dag)
+
+       client, err := runDecider(pick.task, nil)
+
+       require.EqualError(t, err, "cannot reach the table")
+       assert.Empty(t, client.xcoms)
+       assert.Empty(t, client.skipped)
+}
+
+// Case and Register both take the lock of the Dag, so each Case call either 
names its case before
+// the Dag is registered or panics. If Case read the registered flag without 
the lock, only a run
+// with -race would fail.
+func TestRegisterWhileCaseIsCalled(t *testing.T) {
+       const switches = 1000
+
+       dag := Dag("etl")
+       picks := make([]*SwitchRef, switches)
+       shorts := make([]*TaskRef, switches)
+       for i := range switches {
+               picks[i] = dag.Switch(pickPath, TaskSpec{TaskID: 
fmt.Sprintf("pick_%d", i)}).
+                       Case(dag.Task(handleLong, TaskSpec{TaskID: 
fmt.Sprintf("long_%d", i)}))
+               shorts[i] = dag.Task(handleShort, TaskSpec{TaskID: 
fmt.Sprintf("short_%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 < switches-1; o.named++ {
+                       select {
+                       case <-registered:
+                               // Register has returned, so this Case call has 
to panic.
+                               picks[o.named].Case(shorts[o.named])
+                               return
+                       default:
+                               picks[o.named].Case(shorts[o.named])
+                       }
+               }
+               select {
+               case <-registered:
+               case <-time.After(10 * time.Second):
+                       return
+               }
+               picks[o.named].Case(shorts[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.SwitchRef.Case: Dag "etl" has already been registered; 
`+
+                       `name the cases of every switch before Register`,
+               o.recovered,
+       )
+       for i, pick := range picks {
+               if i < o.named {
+                       assert.Len(t, pick.cases, 2, "switch %d", i)
+               } else {
+                       assert.Len(t, pick.cases, 1, "switch %d", i)
+               }
+       }
+}
diff --git a/go-sdk/airflow/task_group.go b/go-sdk/airflow/task_group.go
index 177e04467c5..ec3327ff1aa 100644
--- a/go-sdk/airflow/task_group.go
+++ b/go-sdk/airflow/task_group.go
@@ -136,6 +136,17 @@ func (g *TaskGroupRef) If(fn any, opts ...TaskOption) 
*IfRef {
        return g.groupDag(method).addIf(method, g, fn, opts)
 }
 
+// Switch adds a switch to the Dag of g, inside g, and returns it. It is 
[DagRef.Switch] for a
+// switch of the group, and the group_id of g prefixes the task_id of the 
decider task as
+// [TaskGroupRef.Task] describes. [SwitchRef.Case] takes any task of the Dag, 
inside the group or
+// not.
+//
+// Switch panics for the reasons that DagRef.Switch and TaskGroupRef.Task list.
+func (g *TaskGroupRef) Switch(fn any, opts ...TaskOption) *SwitchRef {
+       const method = "airflow.TaskGroupRef.Switch"
+       return g.groupDag(method).addSwitch(method, g, fn, opts)
+}
+
 func (*TaskGroupRef) node() {}
 
 // Before makes the group an upstream of every node, which is Python's
diff --git a/go-sdk/airflow/task_group_test.go 
b/go-sdk/airflow/task_group_test.go
index 20a236545b3..7887d9ad4e0 100644
--- a/go-sdk/airflow/task_group_test.go
+++ b/go-sdk/airflow/task_group_test.go
@@ -233,6 +233,32 @@ func TestTaskGroupIfRejectsATriggerDagRun(t *testing.T) {
        )
 }
 
+func TestTaskGroupAddsASwitchInsideTheGroup(t *testing.T) {
+       dag := Dag("etl")
+       group := dag.TaskGroup("route")
+       long := group.Task(handleLong)
+       short := dag.Task(handleShort)
+
+       pick := group.Switch(pickPath).Case(long).Case(short)
+
+       assert.Equal(t, "route.pickPath", pick.task.taskID)
+       assert.Same(t, group, pick.task.group)
+       assertChildren(t, group, long, pick.task)
+       assertTasks(t, long.upstreams, pick.task)
+       assertTasks(t, short.upstreams, pick.task)
+}
+
+func TestTaskGroupSwitchRejectsATriggerDagRun(t *testing.T) {
+       dag := Dag("etl")
+       group := dag.TaskGroup("route")
+
+       assert.PanicsWithValue(t,
+               `airflow.TaskGroupRef.Switch: Dag "etl": fn comes from 
airflow.TriggerDagRun, `+
+                       `but a decider function is a Go function that returns 
(*airflow.TaskRef, error)`,
+               func() { group.Switch(TriggerDagRun(TriggerDagRunSpec{DagID: 
"other"})) },
+       )
+}
+
 func TestTaskGroupTakesAtMostOneTaskGroupSpec(t *testing.T) {
        dag := Dag("etl")
 
@@ -424,6 +450,18 @@ func TestTaskGroupMethodsRejectAGroupTheDagDidNotReturn(t 
*testing.T) {
                        want: `airflow.TaskGroupRef.If: Dag "etl" got a 
*airflow.TaskGroupRef ` +
                                `that DagRef.TaskGroup or 
TaskGroupRef.TaskGroup did not return`,
                },
+               {
+                       name: "Switch on a copy",
+                       call: func() { copied.Switch(pickPath) },
+                       want: `airflow.TaskGroupRef.Switch: Dag "etl" got a 
*airflow.TaskGroupRef ` +
+                               `that DagRef.TaskGroup or 
TaskGroupRef.TaskGroup did not return`,
+               },
+               {
+                       name: "Switch on a nil group",
+                       call: func() { (*TaskGroupRef)(nil).Switch(pickPath) },
+                       want: "airflow.TaskGroupRef.Switch: DagRef.TaskGroup or 
TaskGroupRef.TaskGroup " +
+                               "did not return the *airflow.TaskGroupRef",
+               },
                {
                        name: "Task on a zero group",
                        call: func() { (&TaskGroupRef{}).Task(cleanRows) },
@@ -473,6 +511,11 @@ func TestTaskGroupMethodsAfterRegisterPanic(t *testing.T) {
                        `add every task before Register`,
                func() { group.If(isReady) },
        )
+       assert.PanicsWithValue(t,
+               `airflow.TaskGroupRef.Switch: Dag "etl" has already been 
registered; `+
+                       `add every task before Register`,
+               func() { group.Switch(pickPath) },
+       )
        assert.Equal(t, registered, snapshot(dag))
 }
 
@@ -1116,7 +1159,7 @@ func 
TestTasksOfAGroupTakeInputsAndConditionsFromOutsideIt(t *testing.T) {
 }
 
 // dagSnapshot is what a Dag records, by ID, so that two snapshots are equal 
when the Dag records
-// the same tasks, groups, inputs, conditions and edges in the same order.
+// the same tasks, groups, inputs, conditions, switches and edges in the same 
order.
 type dagSnapshot struct {
        registered             bool
        tasks, groups          []string
@@ -1124,6 +1167,7 @@ type dagSnapshot struct {
        children               map[string][]string
        inputs                 map[string][]string
        sides                  map[string][2]string
+       cases                  map[string][]string
        upstreams, downstreams map[string][]string
        edgeLabels             map[edgeKey]string
        groupEdges             []edgeKey
@@ -1141,6 +1185,7 @@ func snapshot(dag *DagRef) dagSnapshot {
                children:        make(map[string][]string),
                inputs:          make(map[string][]string),
                sides:           make(map[string][2]string),
+               cases:           make(map[string][]string),
                upstreams:       make(map[string][]string),
                downstreams:     make(map[string][]string),
                edgeLabels:      maps.Clone(dag.edgeLabels),
@@ -1151,15 +1196,18 @@ func snapshot(dag *DagRef) dagSnapshot {
                s.inputs[task.taskID] = taskIDs(task.inputs)
                s.upstreams[task.taskID] = taskIDs(task.upstreams)
                s.downstreams[task.taskID] = taskIDs(task.downstreams)
-               if task.ifRef != nil {
+               switch decider := task.decider.(type) {
+               case *IfRef:
                        var sides [2]string
-                       if task.ifRef.thenTask != nil {
-                               sides[0] = task.ifRef.thenTask.taskID
+                       if decider.thenTask != nil {
+                               sides[0] = decider.thenTask.taskID
                        }
-                       if task.ifRef.elseTask != nil {
-                               sides[1] = task.ifRef.elseTask.taskID
+                       if decider.elseTask != nil {
+                               sides[1] = decider.elseTask.taskID
                        }
                        s.sides[task.taskID] = sides
+               case *SwitchRef:
+                       s.cases[task.taskID] = taskIDs(decider.cases)
                }
        }
        for _, group := range dag.groups {
@@ -1253,6 +1301,12 @@ func TestAPanicLeavesTheDagAsItWas(t *testing.T) {
                        call:    func(p dagParts) { Bundle().Register(p.dag) },
                        want:    "has no task from Then",
                },
+               {
+                       name:    "a switch without a case at Register",
+                       prepare: func(p dagParts) { 
p.transform.Switch(pickPath) },
+                       call:    func(p dagParts) { Bundle().Register(p.dag) },
+                       want:    "has no case",
+               },
        } {
                t.Run(tc.name, func(t *testing.T) {
                        dag := Dag("etl")
diff --git a/go-sdk/airflow/task_option.go b/go-sdk/airflow/task_option.go
index f24e59487d3..1a3cb5868e7 100644
--- a/go-sdk/airflow/task_option.go
+++ b/go-sdk/airflow/task_option.go
@@ -19,9 +19,9 @@ package airflow
 
 import "errors"
 
-// TaskOption is an option to [DagRef.Task], [DagRef.If] and the methods of 
the same names on
-// [TaskGroupRef]. There are two kinds: a [TaskSpec] sets the attributes of 
the task that the
-// method adds, and [Inputs] passes the results of other tasks to that task.
+// TaskOption is an option to [DagRef.Task], [DagRef.If], [DagRef.Switch] and 
the methods of the
+// same names on [TaskGroupRef]. There are two kinds: a [TaskSpec] sets the 
attributes of the task
+// that the method 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, but the methods
diff --git a/go-sdk/internal/bundle/task.go b/go-sdk/internal/bundle/task.go
index 73fba021653..c0cccb95697 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, 
airflow.DagRef.Task
-// and airflow.DagRef.If wrap a plain Go function into a Task.
+// authors do not implement this directly. airflow.TaskHandler, 
airflow.DagRef.Task,
+// airflow.DagRef.If and airflow.DagRef.Switch wrap a plain Go function into a 
Task.
 type Task interface {
        Execute(ctx context.Context, logger *slog.Logger, args []binding.Arg) 
error
 }
@@ -61,8 +61,8 @@ 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
+       // decide is nil unless the task comes from NewPositionalBranchFunction.
+       decide DecideFunc
 }
 
 var _ Task = (*taskFunction)(nil)
@@ -76,19 +76,26 @@ func NewPositionalTaskFunction(fn any) (Task, error) {
        return newTaskFunction(fn, binding.AnalyzePositional, nil)
 }
 
+// DecideFunc takes the result of the function of a task from 
NewPositionalBranchFunction. It
+// returns value, which the task pushes as its return_value XCom, and the 
task_ids of the tasks to
+// skip. When value is a nil pointer, the task does not push the return_value 
XCom. A non-nil error
+// fails the task.
+type DecideFunc func(result any) (value any, skipped []string, err 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. 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)
+// decide and pushes the value that decide returns. If decide also returns 
task_ids, Execute
+// records them in the skipmixin_key XCom of the task and skips those tasks. 
When fn or decide
+// returns an error, the task fails without pushing an XCom.
+func NewPositionalBranchFunction(fn any, decide DecideFunc) (Task, error) {
+       return newTaskFunction(fn, binding.AnalyzePositional, decide)
 }
 
 func newTaskFunction(
        fn any,
        analyze func(fnType reflect.Type, fnName string) (*binding.Plan, error),
-       findSkipped func(result any) []string,
+       decide DecideFunc,
 ) (Task, error) {
        // The kind comes first: Value.Pointer panics on an int, and Value.Type 
on an untyped nil.
        v := reflect.ValueOf(fn)
@@ -96,9 +103,9 @@ func newTaskFunction(
                return nil, fmt.Errorf("expected a func as input but was %s", 
v.Kind())
        }
        f := &taskFunction{
-               fn:          v,
-               fullName:    runtime.FuncForPC(v.Pointer()).Name(),
-               findSkipped: findSkipped,
+               fn:       v,
+               fullName: runtime.FuncForPC(v.Pointer()).Name(),
+               decide:   decide,
        }
        if err := f.validateFn(v.Type(), analyze); err != nil {
                return nil, err
@@ -117,7 +124,7 @@ func (f *taskFunction) Execute(
                return err
        }
        var branch *branchRun
-       if f.findSkipped != nil {
+       if f.decide != nil {
                if branch, err = startBranch(ctx); err != nil {
                        return err
                }
@@ -157,19 +164,27 @@ func (f *taskFunction) call(
                        )
                }
        }
+       if branch != nil {
+               // The task pushes the value that decide returns, and decide 
runs only on a result that fn
+               // returned without an error. So the task pushes no XCom when 
fn or decide fails.
+               if err != nil {
+                       return err
+               }
+               value, skipped, err := f.decide(retValues[0].Interface())
+               if err != nil {
+                       return err
+               }
+               rv := reflect.ValueOf(value)
+               if rv.Kind() != reflect.Ptr || !rv.IsNil() {
+                       f.sendXcom(ctx, value, sdkClient, logger)
+               }
+               return branch.skipDownstream(ctx, sdkClient, skipped, logger)
+       }
        // If there are two results, convert the first only if it's not a nil 
pointer
        if len(retValues) > 1 && (retValues[0].Kind() != reflect.Ptr || 
!retValues[0].IsNil()) {
                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
 }
 
diff --git a/go-sdk/internal/bundle/task_test.go 
b/go-sdk/internal/bundle/task_test.go
index 850d5ccaf72..4d3aeadf307 100644
--- a/go-sdk/internal/bundle/task_test.go
+++ b/go-sdk/internal/bundle/task_test.go
@@ -293,13 +293,14 @@ func runBranch(task Task, client *branchClient, ti, 
canSkip bool) error {
        return task.Execute(ctx, slog.New(logging.NewTeeLogger()), nil)
 }
 
-func (s *TaskSuite) TestBranchFunctionSkipsTheTasksThatFindSkippedReturns() {
+// The task pushes the value that decide returns, not the result of fn.
+func (s *TaskSuite) 
TestBranchFunctionPushesTheValueAndSkipsTheTasksThatDecideReturns() {
        var got any
        task, err := NewPositionalBranchFunction(
                func(contexttest.Context) (bool, error) { return true, nil },
-               func(result any) []string {
+               func(result any) (any, []string, error) {
                        got = result
-                       return []string{"load", "report"}
+                       return "chosen", []string{"load", "report"}, nil
                },
        )
        s.Require().NoError(err)
@@ -309,18 +310,34 @@ func (s *TaskSuite) 
TestBranchFunctionSkipsTheTasksThatFindSkippedReturns() {
 
        s.Equal(true, got)
        s.Equal([]string{
-               "PushXCom decide return_value true",
+               "PushXCom decide return_value chosen",
                "PushXCom decide skipmixin_key map[skipped:[load report]]",
                "SkipDownstreamTasks [load report]",
        }, client.calls)
 }
 
+func (s *TaskSuite) TestBranchFunctionPushesNoNilPointer() {
+       task, err := NewPositionalBranchFunction(
+               func(contexttest.Context) (bool, error) { return true, nil },
+               func(any) (any, []string, error) { return (*int)(nil), 
[]string{"load"}, nil },
+       )
+       s.Require().NoError(err)
+
+       client := &branchClient{}
+       s.Require().NoError(runBranch(task, client, true, true))
+
+       s.Equal([]string{
+               "PushXCom decide skipmixin_key map[skipped:[load]]",
+               "SkipDownstreamTasks [load]",
+       }, 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 },
+                               func(result any) (any, []string, error) { 
return result, skipped, nil },
                        )
                        s.Require().NoError(err)
 
@@ -332,8 +349,8 @@ func (s *TaskSuite) TestBranchFunctionWithNothingToSkip() {
        }
 }
 
-// A try that fails neither records nor skips any task, whether fn returns an 
error or panics.
-func (s *TaskSuite) TestBranchFunctionThatFailsSkipsNothing() {
+// A try whose fn fails pushes no XCom and skips nothing, whether fn returns 
an error or panics.
+func (s *TaskSuite) TestBranchFunctionThatFailsPushesAndSkipsNothing() {
        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") },
@@ -341,9 +358,9 @@ func (s *TaskSuite) 
TestBranchFunctionThatFailsSkipsNothing() {
        for name, fn := range cases {
                s.Run(name, func() {
                        called := false
-                       task, err := NewPositionalBranchFunction(fn, func(any) 
[]string {
+                       task, err := NewPositionalBranchFunction(fn, 
func(result any) (any, []string, error) {
                                called = true
-                               return []string{"load"}
+                               return result, []string{"load"}, nil
                        })
                        s.Require().NoError(err)
 
@@ -354,14 +371,25 @@ func (s *TaskSuite) 
TestBranchFunctionThatFailsSkipsNothing() {
                        }()
 
                        s.False(called)
-                       for _, call := range client.calls {
-                               s.NotContains(call, "skipmixin_key")
-                               s.NotContains(call, "SkipDownstreamTasks")
-                       }
+                       s.Empty(client.calls)
                })
        }
 }
 
+func (s *TaskSuite) TestBranchFunctionFailsWhenDecideReturnsAnError() {
+       task, err := NewPositionalBranchFunction(
+               func(contexttest.Context) (bool, error) { return true, nil },
+               func(any) (any, []string, error) {
+                       return "chosen", []string{"load"}, errors.New("not one 
of the cases")
+               },
+       )
+       s.Require().NoError(err)
+
+       client := &branchClient{}
+       s.EqualError(runBranch(task, client, true, true), "not one of the 
cases")
+       s.Empty(client.calls)
+}
+
 func (s *TaskSuite) TestBranchFunctionFailsWhenItCannotSkip() {
        cases := map[string]struct {
                client    *branchClient
@@ -415,7 +443,7 @@ func (s *TaskSuite) 
TestBranchFunctionFailsWhenItCannotSkip() {
                                        ran = true
                                        return false, nil
                                },
-                               func(any) []string { return []string{"load"} },
+                               func(result any) (any, []string, error) { 
return result, []string{"load"}, nil },
                        )
                        s.Require().NoError(err)
 
diff --git a/go-sdk/internal/genspec/authoring.go 
b/go-sdk/internal/genspec/authoring.go
index 85660dfef81..655c0285c45 100644
--- a/go-sdk/internal/genspec/authoring.go
+++ b/go-sdk/internal/genspec/authoring.go
@@ -116,8 +116,8 @@ var dagShape = authoringShape{
 }
 
 var taskShape = authoringShape{
-       doc: "TaskSpec holds the attributes of a task. DagRef.Task, DagRef.If 
and the methods of " +
-               "the same names on TaskGroupRef take at most one per task.",
+       doc: "TaskSpec holds the attributes of a task. DagRef.Task, DagRef.If, 
DagRef.Switch and " +
+               "the methods of the same names on TaskGroupRef take at most one 
per task.",
        exclude: map[string]string{
                "task_type":                     "the operator class name, 
which the SDK fills in",
                "_task_module":                  "the operator's Python module, 
which the SDK fills in",
diff --git a/go-sdk/pkg/execution/integration_test.go 
b/go-sdk/pkg/execution/integration_test.go
index 2ee83df8f8e..edcc07145fa 100644
--- a/go-sdk/pkg/execution/integration_test.go
+++ b/go-sdk/pkg/execution/integration_test.go
@@ -902,11 +902,11 @@ func TestServeSkipsDownstreamTasksEndToEnd(t *testing.T) {
 
                        decide, err := bundle.NewPositionalBranchFunction(
                                func(contexttest.Context) (bool, error) { 
return tt.result, nil },
-                               func(result any) []string {
+                               func(result any) (any, []string, error) {
                                        if result.(bool) {
-                                               return nil
+                                               return result, nil, nil
                                        }
-                                       return []string{"load"}
+                                       return result, []string{"load"}, nil
                                },
                        )
                        require.NoError(t, err)
@@ -1233,6 +1233,7 @@ func 
TestServeConditionDoesNotLeaveTheListOfAnEarlierTryEndToEnd(t *testing.T) {
                name         string
                fn           func(contexttest.Context) (bool, error)
                wantTerminal string
+               wantXComs    map[string]any
        }{
                {
                        name: "fails",
@@ -1240,21 +1241,26 @@ func 
TestServeConditionDoesNotLeaveTheListOfAnEarlierTryEndToEnd(t *testing.T) {
                                return false, errors.New("cannot reach the 
table")
                        },
                        wantTerminal: "TaskState",
+                       wantXComs:    map[string]any{},
                },
                {
                        name:         "skips nothing",
                        fn:           func(contexttest.Context) (bool, error) { 
return true, nil },
                        wantTerminal: "SucceedTask",
+                       wantXComs:    map[string]any{"return_value": true},
                },
        }
        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"}
-                       })
+                       decide, err := bundle.NewPositionalBranchFunction(
+                               tt.fn,
+                               func(result any) (any, []string, error) {
+                                       if result.(bool) {
+                                               return result, nil, nil
+                                       }
+                                       return result, []string{"load"}, nil
+                               },
+                       )
                        require.NoError(t, err)
 
                        xcoms := map[string]any{
@@ -1279,7 +1285,7 @@ func 
TestServeConditionDoesNotLeaveTheListOfAnEarlierTryEndToEnd(t *testing.T) {
                        for _, request := range requests {
                                assert.NotEqual(t, "SkipDownstreamTasks", 
request["type"])
                        }
-                       assert.NotContains(t, xcoms, "skipmixin_key")
+                       assert.Equal(t, tt.wantXComs, xcoms)
                        assert.Equal(t, tt.wantTerminal, terminal["type"])
                })
        }

Reply via email to