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"])
})
}