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