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.",
                },
        },
 }

Reply via email to