This is an automated email from the ASF dual-hosted git repository.
henry3260 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 56817294cdd Go SDK: add airflow.TriggerDagRun, a task that triggers a
Dag run (#74003)
56817294cdd is described below
commit 56817294cdd380dc219de2b654b660d129b9a45b
Author: PoAn Yang <[email protected]>
AuthorDate: Thu Oct 1 22:23:35 2026 +0900
Go SDK: add airflow.TriggerDagRun, a task that triggers a Dag run (#74003)
Signed-off-by: PoAn Yang <[email protected]>
---
go-sdk/airflow/dag.go | 79 ++++++++--
go-sdk/airflow/dag_test.go | 6 +
go-sdk/airflow/inputs.go | 6 +
go-sdk/airflow/spec.gen.go | 3 +-
go-sdk/airflow/trigger_dag_run.go | 177 +++++++++++++++++++++
go-sdk/airflow/trigger_dag_run_test.go | 272 +++++++++++++++++++++++++++++++++
go-sdk/internal/genspec/authoring.go | 2 +-
7 files changed, 528 insertions(+), 17 deletions(-)
diff --git a/go-sdk/airflow/dag.go b/go-sdk/airflow/dag.go
index dc0ae28c302..e0b48277909 100644
--- a/go-sdk/airflow/dag.go
+++ b/go-sdk/airflow/dag.go
@@ -84,13 +84,18 @@ type TaskRef struct {
// of them is an upstream task of this one.
inputs []*TaskRef
task bundle.Task
+ // triggerDagRun is the checked copy of the TriggerDagRunSpec of a task
from TriggerDagRun.
+ // 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
}
// Task adds a task that runs fn to the Dag and returns the new task.
//
// fn takes a [Context] first and returns either error or (result, error),
like a function
// passed to [TaskHandler]. The parameters after the Context take the results
of the tasks
-// passed to [Inputs], in order.
+// passed to [Inputs], in order. fn can also be the value that [TriggerDagRun]
returns. The task
+// then runs no Go code and takes no Inputs.
//
// The task_id is the name of fn, spelled exactly as it is in Go.
dag.Task(extractRows) adds the
// task extractRows, and dag.Task(svc.Extract), which passes a method value,
adds the task
@@ -111,12 +116,26 @@ type TaskRef struct {
// Task panics if:
// - fn is not a valid task function
// - fn has no name that Task can read and no TaskSpec sets a TaskID
+// - fn comes from TriggerDagRun and no TaskSpec sets a TaskID
+// - fn comes from TriggerDagRun and its TriggerDagRunSpec is not valid
+// - fn comes from TriggerDagRun and opts holds an Inputs
// - an option is nil or is not one that package airflow defines
// - opts holds more than one TaskSpec or more than one Inputs
// - the tasks passed to Inputs do not match the parameters of fn after the
Context
// - 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 {
+ 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
+ // added a task to this Dag, and a slow one would hold up
Register.
+ copied, err := copyTriggerDagRunSpec(trigger.spec)
+ triggerSpec, triggerErr = &copied, err
+ }
+
d.mu.Lock()
defer d.mu.Unlock()
@@ -127,9 +146,12 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef
{
d.dagID,
))
}
- wrapped, err := newTaskFunction(fn, bundle.NewPositionalTaskFunction)
- if err != nil {
- panic(fmt.Sprintf("airflow.DagRef.Task: Dag %q: %v", d.dagID,
err))
+ var wrapped bundle.Task
+ if !isTrigger {
+ var err error
+ if wrapped, err = newTaskFunction(fn,
bundle.NewPositionalTaskFunction); err != nil {
+ panic(fmt.Sprintf("airflow.DagRef.Task: Dag %q: %v",
d.dagID, err))
+ }
}
var cfg taskConfig
for i, opt := range opts {
@@ -166,6 +188,13 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef
{
spec = cfg.specs[0]
}
taskID := spec.TaskID
+ if taskID == "" && isTrigger {
+ panic(fmt.Sprintf(
+ "airflow.DagRef.Task: 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,
+ ))
+ }
if taskID == "" {
var ok bool
if taskID, ok = taskIDFromFuncName(funcName(fn)); !ok {
@@ -176,6 +205,11 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef
{
))
}
}
+ if triggerErr != nil {
+ panic(fmt.Sprintf(
+ "airflow.DagRef.Task: task %q of Dag %q: %v", taskID,
d.dagID, triggerErr,
+ ))
+ }
if _, exists := d.tasksByID[taskID]; exists {
panic(fmt.Sprintf(
"airflow.DagRef.Task: Dag %q already has a task %q; "+
@@ -183,21 +217,33 @@ func (d *DagRef) Task(fn any, opts ...TaskOption)
*TaskRef {
d.dagID, taskID,
))
}
- fnType := reflect.TypeOf(fn)
- upstreams := d.checkInputs(taskID, fnType, cfg.inputs)
- // newTaskFunction has checked that fn returns either error or (result,
error).
+ var upstreams []*TaskRef
var resultType reflect.Type
- if fnType.NumOut() == 2 {
- resultType = fnType.Out(0)
+ if isTrigger {
+ if len(cfg.inputs) > 0 {
+ panic(fmt.Sprintf(
+ "airflow.DagRef.Task: 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,
+ ))
+ }
+ } else {
+ fnType := reflect.TypeOf(fn)
+ upstreams = d.checkInputs(taskID, fnType, cfg.inputs)
+ // newTaskFunction has checked that fn returns either error or
(result, error).
+ if fnType.NumOut() == 2 {
+ resultType = fnType.Out(0)
+ }
}
task := &TaskRef{
- dag: d,
- taskID: taskID,
- spec: copySpec(spec),
- resultType: resultType,
- inputs: upstreams,
- task: wrapped,
+ dag: d,
+ taskID: taskID,
+ spec: copySpec(spec),
+ resultType: resultType,
+ inputs: upstreams,
+ task: wrapped,
+ triggerDagRun: triggerSpec,
}
if d.tasksByID == nil {
d.tasksByID = make(map[string]*TaskRef)
@@ -223,6 +269,9 @@ func findTaskName(fn any, specs []TaskSpec) string {
return spec.TaskID
}
}
+ if _, ok := fn.(TriggerDagRunTask); ok {
+ return "airflow.TriggerDagRun"
+ }
if taskID, ok := taskIDFromFuncName(funcName(fn)); ok {
return taskID
}
diff --git a/go-sdk/airflow/dag_test.go b/go-sdk/airflow/dag_test.go
index 0ed707ab663..5991283180b 100644
--- a/go-sdk/airflow/dag_test.go
+++ b/go-sdk/airflow/dag_test.go
@@ -251,6 +251,12 @@ func TestTaskRejectsASecondTaskSpec(t *testing.T) {
task: `github\.com/apache/airflow/go-sdk/airflow\.` +
`TestTaskRejectsASecondTaskSpec\.func\d+`,
},
+ {
+ name: "TriggerDagRun and no TaskID",
+ fn: TriggerDagRun(TriggerDagRunSpec{DagID:
"downstream_etl"}),
+ opts: []TaskOption{TaskSpec{}, TaskSpec{}},
+ task: `airflow\.TriggerDagRun`,
+ },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
diff --git a/go-sdk/airflow/inputs.go b/go-sdk/airflow/inputs.go
index 77f6ecb8dbf..3e3ad7f0a0d 100644
--- a/go-sdk/airflow/inputs.go
+++ b/go-sdk/airflow/inputs.go
@@ -107,6 +107,12 @@ func (d *DagRef) checkInputs(taskID string, fnType
reflect.Type, given [][]*Task
param := i + 1
paramType := fnType.In(param)
switch {
+ case upstream.triggerDagRun != nil:
+ panic(fmt.Sprintf(
+ "airflow.DagRef.Task: 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,
+ ))
case upstream.resultType == nil:
panic(fmt.Sprintf(
"airflow.DagRef.Task: task %q of Dag %q takes
parameter %d from task %q, "+
diff --git a/go-sdk/airflow/spec.gen.go b/go-sdk/airflow/spec.gen.go
index f1889915b58..a4674daf848 100644
--- a/go-sdk/airflow/spec.gen.go
+++ b/go-sdk/airflow/spec.gen.go
@@ -151,7 +151,8 @@ type TaskSpec struct {
StartDate time.Time
// TaskID is the task_id of the task. When TaskID is empty, the task_id
is the
- // name of the Go function that the task runs.
+ // name of the Go function that the task runs. A task from
TriggerDagRun runs no
+ // Go function, so it needs a TaskID.
TaskID string
// TriggerRule corresponds to the JSON schema field "trigger_rule".
diff --git a/go-sdk/airflow/trigger_dag_run.go
b/go-sdk/airflow/trigger_dag_run.go
new file mode 100644
index 00000000000..dcec75a9cca
--- /dev/null
+++ b/go-sdk/airflow/trigger_dag_run.go
@@ -0,0 +1,177 @@
+// 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 (
+ "bytes"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "slices"
+ "time"
+
+ "github.com/apache/airflow/go-sdk/pkg/execution/genmodels"
+)
+
+// TriggerDagRunSpec holds the options of a task that triggers a Dag run.
[TriggerDagRun] takes a
+// TriggerDagRunSpec as its argument. Each field sets one parameter of
TriggerDagRunOperator. DagID
+// sets trigger_dag_id and RunID sets trigger_run_id. Every other field sets
the parameter of the
+// same name in snake_case, such as reset_dag_run for ResetDagRun. A field
left at its zero value
+// leaves its parameter at the Python default. PokeInterval and Deferrable are
pointers so that a
+// pointer to 0 or false can set that value instead of leaving the default.
+//
+// DagID, RunID, LogicalDate and the values in Conf are templated. They can
hold Jinja such as
+// "{{ ds }}", which Airflow renders when the task runs.
+type TriggerDagRunSpec struct {
+ // DagID is the dag_id of the Dag to trigger. It is required.
+ DagID string
+ // RunID is the run_id of the new Dag run. When RunID is empty, Airflow
generates one.
+ RunID string
+ // Conf is the conf of the new Dag run. Airflow receives it as JSON, so
each value must
+ // marshal to JSON.
+ Conf map[string]any
+ // LogicalDate is the logical date of the new Dag run, as an ISO 8601
string such as
+ // "2026-09-30T00:00:00+00:00" or a template such as "{{ ds }}". When
LogicalDate is empty and
+ // RunAfter is the zero Time, the logical date is the time the task
runs. When LogicalDate is
+ // empty and RunAfter is set, the new Dag run has no logical date.
+ LogicalDate string
+ // RunAfter is the earliest time at which the new Dag run can start.
When RunAfter is the zero
+ // Time, the new Dag run can start as soon as the task triggers it.
+ RunAfter time.Time
+ // ResetDagRun clears the Dag run if it already exists, instead of
failing the task.
+ ResetDagRun bool
+ // WaitForCompletion makes the task wait until the new Dag run is in a
state that
+ // AllowedStates or FailedStates lists.
+ WaitForCompletion bool
+ // PokeInterval is how often a task that waits checks the state of the
new Dag run. It must be
+ // a whole number of seconds. When PokeInterval is nil, the task checks
every 60 seconds.
+ PokeInterval *time.Duration
+ // AllowedStates are the states of the new Dag run in which a task that
waits succeeds. Each
+ // state is one of queued, running, success and failed. When
AllowedStates is empty, the task
+ // succeeds in the success state.
+ AllowedStates []string
+ // FailedStates are the states of the new Dag run in which a task that
waits fails. Each state
+ // is one of queued, running, success and failed. A nil FailedStates
fails the task in the
+ // failed state. A FailedStates that is empty but not nil means that no
state fails the task.
+ // The task then succeeds once the new Dag run is in a state that
AllowedStates lists. In any
+ // other state, including failed, the task keeps waiting.
+ FailedStates []string
+ // SkipWhenAlreadyExists marks the task skipped if the Dag run already
exists.
+ SkipWhenAlreadyExists bool
+ // FailWhenDagIsPaused fails the task when the Dag to trigger is paused.
+ FailWhenDagIsPaused bool
+ // Note is the note of the new Dag run.
+ Note string
+ // Deferrable makes a task that waits defer instead of holding a worker
slot. When Deferrable
+ // is nil, the task follows the default_deferrable option in the
operators section of the
+ // Airflow configuration.
+ Deferrable *bool
+}
+
+// TriggerDagRunTask is what [TriggerDagRun] returns. [DagRef.Task] takes it
in place of a Go
+// function and adds a task that triggers a Dag run.
+type TriggerDagRunTask struct {
+ spec TriggerDagRunSpec
+}
+
+// TriggerDagRun returns a value to pass to [DagRef.Task] in place of a Go
function. DagRef.Task
+// then adds a task that triggers a run of the Dag that spec.DagID names:
+//
+// dag.Task(
+// airflow.TriggerDagRun(airflow.TriggerDagRunSpec{DagID:
"downstream_etl"}),
+// airflow.TaskSpec{TaskID: "trigger_downstream"},
+// )
+//
+// The task runs no Go code. Once [BundleRef.Serve] serves the Dags from
[Dag], Airflow will run
+// the task as TriggerDagRunOperator on a Python worker.
+//
+// Because the task has no Go function to take a task_id from, DagRef.Task
needs a [TaskSpec]
+// that sets TaskID. For the same reason, the task takes no [Inputs], and it
returns no result
+// that Inputs can pass to another task. DagRef.Task also checks spec and
panics if spec is not
+// valid.
+func TriggerDagRun(spec TriggerDagRunSpec) TriggerDagRunTask {
+ return TriggerDagRunTask{spec: spec}
+}
+
+// validDagRunStates are the values that TriggerDagRunOperator accepts in
allowed_states and
+// failed_states.
+var validDagRunStates = []string{
+ string(genmodels.DagRunStateQueued),
+ string(genmodels.DagRunStateRunning),
+ string(genmodels.DagRunStateSuccess),
+ string(genmodels.DagRunStateFailed),
+}
+
+// copyTriggerDagRunSpec checks spec and returns a deep copy of it, so that
nothing the caller
+// still holds, such as Conf, a state slice or a pointer field, can change the
task that
+// DagRef.Task added.
+func copyTriggerDagRunSpec(spec TriggerDagRunSpec) (TriggerDagRunSpec, error) {
+ if spec.DagID == "" {
+ return TriggerDagRunSpec{},
errors.New("airflow.TriggerDagRunSpec has no DagID")
+ }
+ if poke := spec.PokeInterval; poke != nil && (*poke < 0 ||
*poke%time.Second != 0) {
+ return TriggerDagRunSpec{}, fmt.Errorf(
+ "airflow.TriggerDagRunSpec.PokeInterval is %v; "+
+ "it must be a whole number of seconds and not
negative",
+ *poke,
+ )
+ }
+ for _, field := range []struct {
+ name string
+ states []string
+ }{{"AllowedStates", spec.AllowedStates}, {"FailedStates",
spec.FailedStates}} {
+ for _, state := range field.states {
+ if !slices.Contains(validDagRunStates, state) {
+ return TriggerDagRunSpec{}, fmt.Errorf(
+ "airflow.TriggerDagRunSpec.%s has %q,
which is not a Dag run state; "+
+ "use one of %q",
+ field.name, state, validDagRunStates,
+ )
+ }
+ }
+ }
+ conf, err := copyConf(spec.Conf)
+ if err != nil {
+ return TriggerDagRunSpec{},
fmt.Errorf("airflow.TriggerDagRunSpec.Conf: %w", err)
+ }
+
+ copied := copySpec(spec)
+ // copySpec copies the Conf map but leaves anything nested in its
values shared, such as an
+ // inner map, so Conf is the copy that copyConf made.
+ copied.Conf = conf
+ return copied, nil
+}
+
+// copyConf copies conf by way of JSON, so it also rejects a conf that JSON
cannot hold.
+//
+// UseNumber stores each number as a json.Number, which keeps an integer that
a float64 cannot
+// hold exactly, such as 2^53 + 1. encoding/json writes a json.Number back as
a number, but an
+// encoder that does not know the type, such as msgpack, writes it as a string.
+func copyConf(conf map[string]any) (map[string]any, error) {
+ data, err := json.Marshal(conf)
+ if err != nil {
+ return nil, err
+ }
+ dec := json.NewDecoder(bytes.NewReader(data))
+ dec.UseNumber()
+ var copied map[string]any
+ if err := dec.Decode(&copied); err != nil {
+ return nil, err
+ }
+ return copied, nil
+}
diff --git a/go-sdk/airflow/trigger_dag_run_test.go
b/go-sdk/airflow/trigger_dag_run_test.go
new file mode 100644
index 00000000000..7cc825073f2
--- /dev/null
+++ b/go-sdk/airflow/trigger_dag_run_test.go
@@ -0,0 +1,272 @@
+// 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 (
+ "encoding/json"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+func ptr[T any](v T) *T { return &v }
+
+func TestTriggerDagRunIsATask(t *testing.T) {
+ dag := Dag("etl")
+ gate := dag.Task(extract)
+ spec := TriggerDagRunSpec{
+ DagID: "downstream_etl",
+ RunID: "{{ run_id }}_downstream",
+ Conf: map[string]any{"source": "etl"},
+ LogicalDate: "{{ ds }}",
+ RunAfter: time.Date(2026, 9, 30, 0, 0, 0, 0,
time.UTC),
+ ResetDagRun: true,
+ WaitForCompletion: true,
+ PokeInterval: ptr(30 * time.Second),
+ AllowedStates: []string{"success", "failed"},
+ FailedStates: []string{"queued", "running"},
+ SkipWhenAlreadyExists: true,
+ FailWhenDagIsPaused: true,
+ Note: "triggered by etl",
+ Deferrable: ptr(true),
+ }
+
+ task := dag.Task(TriggerDagRun(spec), TaskSpec{TaskID:
"trigger_downstream"})
+
+ assert.Equal(t, "trigger_downstream", task.taskID)
+ assert.Equal(t, TaskSpec{TaskID: "trigger_downstream"}, task.spec)
+ require.NotNil(t, task.triggerDagRun)
+ assert.Equal(t, spec, *task.triggerDagRun)
+ assert.Equal(t, []*TaskRef{gate, task}, dag.tasks)
+ assert.Nil(t, gate.triggerDagRun, "a task that runs a Go function is
not a TriggerDagRun task")
+}
+
+func TestTriggerDagRunKeepsNilApartFromZero(t *testing.T) {
+ dag := Dag("etl")
+
+ unset := TriggerDagRunSpec{DagID: "downstream_etl"}
+ task := dag.Task(TriggerDagRun(unset), TaskSpec{TaskID: "unset"})
+ assert.Equal(t, unset, *task.triggerDagRun)
+
+ // nil keeps the Python default. An empty slice or a pointer to 0 or
false is a value the task
+ // uses instead. For example, an empty FailedStates means that no state
fails the task.
+ zero := TriggerDagRunSpec{
+ DagID: "downstream_etl",
+ Conf: map[string]any{},
+ PokeInterval: ptr(time.Duration(0)),
+ AllowedStates: []string{},
+ FailedStates: []string{},
+ Deferrable: ptr(false),
+ }
+ task = dag.Task(TriggerDagRun(zero), TaskSpec{TaskID: "zero"})
+ assert.Equal(t, zero, *task.triggerDagRun)
+}
+
+func TestTriggerDagRunNeedsATaskID(t *testing.T) {
+ want := `airflow.DagRef.Task: Dag "etl": a task from
airflow.TriggerDagRun with DagID ` +
+ `"downstream_etl" has no Go function to take a task_id from; ` +
+ `set one with airflow.TaskSpec{TaskID: ...}`
+ trigger := TriggerDagRun(TriggerDagRunSpec{DagID: "downstream_etl"})
+
+ assert.PanicsWithValue(t, want, func() { Dag("etl").Task(trigger) })
+ assert.PanicsWithValue(t, want, func() { Dag("etl").Task(trigger,
TaskSpec{}) })
+}
+
+// buildNestedConf returns a conf in which maps nest depth levels deep.
encoding/json decodes at
+// most 10000 levels of nesting.
+func buildNestedConf(depth int) map[string]any {
+ conf := map[string]any{}
+ for range depth - 1 {
+ conf = map[string]any{"nested": conf}
+ }
+ return conf
+}
+
+func TestTriggerDagRunRejectsAnInvalidSpec(t *testing.T) {
+ tests := []struct {
+ name string
+ trigger TriggerDagRunTask
+ want string
+ }{
+ {
+ name: "no DagID",
+ trigger:
TriggerDagRun(TriggerDagRunSpec{WaitForCompletion: true}),
+ want: "airflow.TriggerDagRunSpec has no DagID",
+ },
+ {
+ name: "negative PokeInterval",
+ trigger: TriggerDagRun(
+ TriggerDagRunSpec{DagID: "downstream_etl",
PokeInterval: ptr(-time.Second)},
+ ),
+ want: "airflow.TriggerDagRunSpec.PokeInterval is -1s; "
+
+ "it must be a whole number of seconds and not
negative",
+ },
+ {
+ name: "PokeInterval with a fraction of a second",
+ trigger: TriggerDagRun(TriggerDagRunSpec{
+ DagID: "downstream_etl", PokeInterval: ptr(1500
* time.Millisecond),
+ }),
+ want: "airflow.TriggerDagRunSpec.PokeInterval is 1.5s;
" +
+ "it must be a whole number of seconds and not
negative",
+ },
+ {
+ name: "unknown state in AllowedStates",
+ trigger: TriggerDagRun(TriggerDagRunSpec{
+ DagID: "downstream_etl", AllowedStates:
[]string{"success", "SUCCESS"},
+ }),
+ want: `airflow.TriggerDagRunSpec.AllowedStates has
"SUCCESS", which is not a Dag ` +
+ `run state; use one of ["queued" "running"
"success" "failed"]`,
+ },
+ {
+ name: "unknown state in FailedStates",
+ trigger: TriggerDagRun(TriggerDagRunSpec{
+ DagID: "downstream_etl", FailedStates:
[]string{"skipped"},
+ }),
+ want: `airflow.TriggerDagRunSpec.FailedStates has
"skipped", which is not a Dag ` +
+ `run state; use one of ["queued" "running"
"success" "failed"]`,
+ },
+ {
+ name: "Conf that JSON cannot hold",
+ trigger: TriggerDagRun(TriggerDagRunSpec{
+ DagID: "downstream_etl", Conf:
map[string]any{"done": make(chan struct{})},
+ }),
+ want: "airflow.TriggerDagRunSpec.Conf: json:
unsupported type: chan struct {}",
+ },
+ {
+ name: "Conf nested deeper than encoding/json decodes",
+ trigger: TriggerDagRun(
+ TriggerDagRunSpec{DagID: "downstream_etl",
Conf: buildNestedConf(10001)},
+ ),
+ want: "airflow.TriggerDagRunSpec.Conf: invalid
character '{' exceeded max depth",
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ dag := Dag("etl")
+ assert.PanicsWithValue(t,
+ `airflow.DagRef.Task: task "trigger_downstream"
of Dag "etl": `+tt.want,
+ func() { dag.Task(tt.trigger, TaskSpec{TaskID:
"trigger_downstream"}) },
+ )
+ assert.NotPanics(t,
+ func() { dag.Task(extract, TaskSpec{TaskID:
"trigger_downstream"}) },
+ "a rejected task does not take its task_id",
+ )
+ })
+ }
+}
+
+func TestTriggerDagRunCopiesTheSpec(t *testing.T) {
+ nested := map[string]any{"table": "rows"}
+ allowed := []string{"success"}
+ failed := []string{"failed"}
+ poke := 30 * time.Second
+ deferrable := true
+ spec := TriggerDagRunSpec{
+ DagID: "downstream_etl",
+ // A float64 cannot hold 2^53 + 1 exactly.
+ Conf: map[string]any{"target": nested, "batch":
int64(9007199254740993)},
+ PokeInterval: &poke,
+ AllowedStates: allowed,
+ FailedStates: failed,
+ Deferrable: &deferrable,
+ }
+ task := Dag("etl").Task(TriggerDagRun(spec), TaskSpec{TaskID:
"trigger_downstream"})
+
+ nested["table"] = "changed"
+ spec.Conf["added"] = true
+ allowed[0] = "running"
+ failed[0] = "queued"
+ poke = time.Minute
+ deferrable = false
+
+ stored := task.triggerDagRun
+ assert.Equal(t,
+ map[string]any{
+ "target": map[string]any{"table": "rows"},
+ "batch": json.Number("9007199254740993"),
+ },
+ stored.Conf,
+ )
+ assert.Equal(t, []string{"success"}, stored.AllowedStates)
+ assert.Equal(t, []string{"failed"}, stored.FailedStates)
+ assert.Equal(t, 30*time.Second, *stored.PokeInterval)
+ assert.True(t, *stored.Deferrable)
+}
+
+// taskAdder adds a task to dag when encoding/json marshals it.
+type taskAdder struct{ dag *DagRef }
+
+func (a taskAdder) MarshalJSON() ([]byte, error) {
+ a.dag.Task(extract)
+ return []byte(`"added"`), nil
+}
+
+func TestTriggerDagRunMarshalsConfOutsideTheLock(t *testing.T) {
+ dag := Dag("etl")
+ spec := TriggerDagRunSpec{
+ DagID: "downstream_etl",
+ Conf: map[string]any{"adder": taskAdder{dag: dag}},
+ }
+ added := make(chan *TaskRef, 1)
+ go func() { added <- dag.Task(TriggerDagRun(spec), TaskSpec{TaskID:
"trigger_downstream"}) }()
+
+ select {
+ case task := <-added:
+ assert.Equal(t, map[string]any{"adder": "added"},
task.triggerDagRun.Conf)
+ assert.Len(t, dag.tasks, 2)
+ case <-time.After(5 * time.Second):
+ t.Fatal("DagRef.Task held the lock of the Dag while it
marshaled Conf")
+ }
+}
+
+func TestTriggerDagRunTakesNoInputs(t *testing.T) {
+ dag := Dag("etl")
+ extracted := dag.Task(readRows)
+ trigger := TriggerDagRun(TriggerDagRunSpec{DagID: "downstream_etl"})
+ want := `airflow.DagRef.Task: task "trigger_downstream" of Dag "etl"
comes from ` +
+ `airflow.TriggerDagRun and takes no airflow.Inputs, because it
has no Go function ` +
+ `to pass the results to`
+
+ assert.PanicsWithValue(t, want, func() {
+ dag.Task(trigger, TaskSpec{TaskID: "trigger_downstream"},
Inputs(extracted))
+ })
+ assert.PanicsWithValue(t, want, func() {
+ dag.Task(trigger, TaskSpec{TaskID: "trigger_downstream"},
Inputs())
+ })
+ assert.NotPanics(t,
+ func() { dag.Task(trigger, TaskSpec{TaskID:
"trigger_downstream"}) },
+ "a rejected task does not take its task_id",
+ )
+}
+
+func TestTriggerDagRunHasNoResultForInputs(t *testing.T) {
+ dag := Dag("etl")
+ trigger := dag.Task(
+ TriggerDagRun(TriggerDagRunSpec{DagID: "downstream_etl"}),
+ TaskSpec{TaskID: "trigger_downstream"},
+ )
+
+ assert.PanicsWithValue(t,
+ `airflow.DagRef.Task: task "countRows" of Dag "etl" takes
parameter 1 from task `+
+ `"trigger_downstream", but that task comes from
airflow.TriggerDagRun and returns `+
+ `no result`,
+ func() { dag.Task(countRows, Inputs(trigger)) },
+ )
+}
diff --git a/go-sdk/internal/genspec/authoring.go
b/go-sdk/internal/genspec/authoring.go
index db0b707e582..4eac1c5ccc1 100644
--- a/go-sdk/internal/genspec/authoring.go
+++ b/go-sdk/internal/genspec/authoring.go
@@ -180,7 +180,7 @@ var taskShape = authoringShape{
"retry_exponential_backoff": {goType: "float64"},
"task_id": {
goType: "string",
- doc: "TaskID is the task_id of the task. When TaskID
is empty, the task_id is the name of the Go function that the task runs.",
+ doc: "TaskID is the task_id of the task. When TaskID
is empty, the task_id is the name of the Go function that the task runs. A task
from TriggerDagRun runs no Go function, so it needs a TaskID.",
},
},
}