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 ae082ab651f Go SDK: add dag.TaskGroup with nesting and group-level
edges (#74214)
ae082ab651f is described below
commit ae082ab651f16ed6d91322d065a711fcf9d50cd5
Author: PoAn Yang <[email protected]>
AuthorDate: Mon Oct 5 20:41:56 2026 +0800
Go SDK: add dag.TaskGroup with nesting and group-level edges (#74214)
Signed-off-by: PoAn Yang <[email protected]>
---
go-sdk/Justfile | 2 +-
go-sdk/README.md | 11 +-
go-sdk/adr/0008-native-dag-interface.md | 21 +-
go-sdk/airflow/bundle.go | 11 +-
go-sdk/airflow/dag.go | 179 +++-
go-sdk/airflow/if.go | 27 +-
go-sdk/airflow/if_test.go | 14 +-
go-sdk/airflow/inputs.go | 14 +-
go-sdk/airflow/inputs_test.go | 4 +-
go-sdk/airflow/node.go | 172 ++-
go-sdk/airflow/node_test.go | 11 +-
go-sdk/airflow/spec.gen.go | 31 +-
go-sdk/airflow/spec.go | 6 +-
go-sdk/airflow/task_group.go | 553 ++++++++++
go-sdk/airflow/task_group_test.go | 1274 +++++++++++++++++++++++
go-sdk/airflow/task_option.go | 10 +-
go-sdk/airflow/trigger_dag_run.go | 4 +-
go-sdk/internal/genspec/authoring.go | 35 +-
go-sdk/internal/genspec/authoring_test.go | 22 +-
go-sdk/internal/genspec/main.go | 4 +-
go-sdk/internal/genspec/normalize.go | 14 +-
go-sdk/internal/genspec/normalize_test.go | 33 +-
scripts/ci/prek/check_go_sdk_generated_drift.py | 5 +-
23 files changed, 2305 insertions(+), 152 deletions(-)
diff --git a/go-sdk/Justfile b/go-sdk/Justfile
index 8445b9e1b92..f62260424d1 100644
--- a/go-sdk/Justfile
+++ b/go-sdk/Justfile
@@ -43,6 +43,6 @@ docs port="6060":
generate-models:
go generate ./pkg/execution/genmodels/...
-# Regenerate the Dag and task spec structs from the core Dag serialization
schema
+# Regenerate the Dag, task and task group spec structs from the core Dag
serialization schema
generate-specs:
go generate ./airflow/...
diff --git a/go-sdk/README.md b/go-sdk/README.md
index 57e6d4378de..3b82e2c405b 100644
--- a/go-sdk/README.md
+++ b/go-sdk/README.md
@@ -433,12 +433,13 @@ a Dag author does not import `genmodels`.
`TestDagRunStateMatchesGenmodels` fail
the models adds, renames or removes a `DagRunState` constant. The test keeps
failing until
`airflow/enums.go` declares the same constants as `genmodels`.
-## Regenerating the Dag and task specs
+## Regenerating the Dag, task and task group specs
-`airflow.DagSpec` and `airflow.TaskSpec` in
[`airflow/spec.gen.go`](./airflow/spec.gen.go) are
-generated from `schema/dag-schema.json`, this module's vendored copy of
airflow-core's Dag
-serialization schema (`airflow-core/src/airflow/serialization/schema.json`),
which Python owns; do
-not edit either by hand. Refresh the copy with `prek run sync-go-sdk-schemas
--hook-stage manual`,
+`airflow.DagSpec`, `airflow.TaskSpec` and `airflow.TaskGroupSpec` in
+[`airflow/spec.gen.go`](./airflow/spec.gen.go) are generated from
`schema/dag-schema.json`, this
+module's vendored copy of airflow-core's Dag serialization schema
+(`airflow-core/src/airflow/serialization/schema.json`), which Python owns; do
not edit any of them
+by hand. Refresh the copy with `prek run sync-go-sdk-schemas --hook-stage
manual`,
then run `just generate-specs` after changing the schema or the generator.
The schema is the serialized shape rather than the authoring one, so
diff --git a/go-sdk/adr/0008-native-dag-interface.md
b/go-sdk/adr/0008-native-dag-interface.md
index 974218268a5..97bb8c54546 100644
--- a/go-sdk/adr/0008-native-dag-interface.md
+++ b/go-sdk/adr/0008-native-dag-interface.md
@@ -39,7 +39,7 @@ Proposed.
9. **A user-facing enum carries its type in the constant name** — e.g.
`airflow.TriggerRuleAllDone`.
10. **Everything an author writes comes from one `airflow` package.**
11. **No Go-native deferral**, and none is needed: the constructs that defer
are DSL tasks Python executes.
-12. **`DagSpec` and `TaskSpec` are generated from Airflow core's serialization
schema** (`airflow-core/src/airflow/serialization/schema.json`) into the
`airflow` package itself and committed, the way `models.gen.go` already is for
the supervisor schema.
+12. **`DagSpec`, `TaskSpec` and `TaskGroupSpec` are generated from Airflow
core's serialization schema**
(`airflow-core/src/airflow/serialization/schema.json`) into the `airflow`
package itself and committed, the way `models.gen.go` already is for the
supervisor schema.
`TaskSpec` implements `airflow.TaskOption`, so a generated struct travels
in the same variadic as `airflow.Inputs`.
## Context
@@ -110,12 +110,12 @@ package airflow
func Dag(dagID string, spec ...DagSpec) *DagRef
func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef
-func (d *DagRef) TaskGroup(groupID string, opts ...TaskGroupOption)
*TaskGroupRef
+func (d *DagRef) TaskGroup(groupID string, spec ...TaskGroupSpec) *TaskGroupRef
func (g *TaskGroupRef) Task(fn any, opts ...TaskOption) *TaskRef
-func (g *TaskGroupRef) TaskGroup(groupID string, opts ...TaskGroupOption)
*TaskGroupRef
+func (g *TaskGroupRef) TaskGroup(groupID string, spec ...TaskGroupSpec)
*TaskGroupRef
-// DagSpec and TaskSpec are generated into this package from
+// DagSpec, TaskSpec and TaskGroupSpec are generated into this package from
// airflow-core/src/airflow/serialization/schema.json and committed.
type DagSpec struct {
Schedule string
@@ -171,7 +171,7 @@ const (
Returning the receiver would read like a chain and mean a second fan-out
from `a`.
- **The specs generate into the `airflow` package, not a `gen` package beside
it.** An unexported method belongs to the package that declares it, so a
generated type living elsewhere could not implement the sealed `TaskOption`,
and a type alias cannot gain methods either.
Generating in place is what keeps both `airflow.TaskSpec` and the seal.
-- **The generated names need a mapping.** The core schema carries no `title`
fields, unlike the supervisor schema `models.gen.go` reads, so its `dag` and
`operator` definitions would generate as `Dag`, a name the constructor already
takes, and `Operator`, which is not the SDK's vocabulary.
+- **The generated names need a mapping.** The core schema carries no `title`
fields, unlike the supervisor schema `models.gen.go` reads, so its `dag`,
`operator` and `task_group` definitions would generate as `Dag`, a name the
constructor already takes, `Operator`, which is not the SDK's vocabulary, and
`TaskGroup`, the name of the method that adds a group.
Either the schema gains titles or the generate step keeps the map.
- **The schema is the serialized shape, not the authoring shape.** It requires
`fileloc` and `tasks` on a Dag, and `task_type`, `_task_module`, `ui_color`,
`ui_fgcolor`, and `template_fields` on an operator, all of which the SDK fills
in, and it carries a serialized `timetable` object where an author writes a
schedule.
Generation needs an exclusion list and a hand-written field or two, the same
kind of rule
[ADR-0009](../../airflow-core/adr/lang-sdk/0009-provider-operators-as-generated-dsl.md)
states for provider operators.
@@ -183,6 +183,17 @@ const (
- **A data edge is labelled by redeclaring it.** Declaring an edge that
already exists is idempotent, so `extracted.Before(airflow.Label(transformed,
"rows"))` labels the edge `Inputs` created.
- **Renaming a Go function renames the task.** The id is derived, and history,
clears, and the UI all key on task_id, so renaming a function whose task
carries no `TaskSpec` id is a Dag change.
`airflow.TaskSpec{TaskID: ...}` pins an id that has to outlive the
function's name, and it is also how a Dag gets snake_case ids, since nothing
transforms a Go name.
+ Renaming a task group renames every task whose task_id the group_id
prefixes, for the same reason.
+- **A group edge is expanded at registration, in the order the group edges
were first declared.** `group.Before(loaded)` stands for an edge from each last
task of the group, and `extracted.Before(group)` for an edge to each first task.
+ Registration expands the group edges one at a time. Each expansion reads
every task, every edge declared between two tasks, and the task edges that
earlier group edges expanded into, which is the rule the TypeScript SDK's
serializer applies.
+ So an author can add the tasks of a group, and the edges between them, after
putting the group on an edge. Python expands a group edge when `>>` runs, so
the two agree when a Dag declares its group edges after its tasks and the edges
between them, every group at an end of a group edge holds a task, and no label
sits on an edge whose receiver is inside a task group.
+ The cycle check runs on the declared edges first and on the expanded graph
after, so a cycle that only group edges close is caught, and the message names
those group edges.
+- **A group with no task is stepped over.** An edge to it continues along each
edge from it, and an edge from it back along each edge to it, whenever those
were declared, as the TypeScript SDK does. So
`extracted.Before(empty).Before(loaded)` runs `load` after `extract`, as
Python's `extract >> empty >> load` does.
+ What Python does with an empty group depends on how and when its edges are
declared: `load << empty << extract` adds no edge, and an empty group with no
task before it falls back to the last tasks of the enclosing group or of the
whole Dag, which can make a task depend on itself.
+- **An edge cannot connect a group to what it holds.** Python's `group >>
node_in_group` orders the last tasks of the group before the first tasks of a
node inside it, which fails as a cycle whenever that node holds a task. The
edge verb rejects the edge when it is declared.
+- **A label stays on the edge it is declared on.** A label on a group edge
labels none of the task edges that the group edge stands for, and a labelled
edge keeps its ends.
+ Python's behavior depends on which end is the receiver and which groups hold
the ends. When no group holds `extract`, `extract >> Label("rows") >> group`
labels each of those task edges as well. When `a` and `b` are in different
groups, `a >> Label("x") >> b` replaces `a` with its group, so every last task
of `a`'s group runs before `b`. The Go SDK does neither.
+- **`TaskGroup` takes `spec ...TaskGroupSpec`, as `Dag` takes `spec
...DagSpec`.** A group has one kind of option, so it needs no sealed option
interface. A mapped task group would come through a method of its own, not
through this variadic.
## Alternatives
diff --git a/go-sdk/airflow/bundle.go b/go-sdk/airflow/bundle.go
index d100231d832..6ac967a0de1 100644
--- a/go-sdk/airflow/bundle.go
+++ b/go-sdk/airflow/bundle.go
@@ -76,12 +76,17 @@ type Registerable interface{ registerable() }
//
// bundle.Register(reports.Handlers()...)
//
-// 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.
+// Add every task to a Dag before registering the Dag. [DagRef.Task],
[DagRef.If],
+// [DagRef.TaskGroup], [IfRef.Then], [IfRef.Else] and the methods of
[TaskGroupRef] panic once the
+// Dag is registered.
//
// Register is where a Dag's task dependencies are checked for a cycle, over
the whole graph at
// once: [TaskRef.Before], [TaskRef.After] and [Inputs] each record an edge
without walking the
-// graph, so building a Dag stays linear in its edges however many a task has.
+// graph, so building a Dag stays linear in its edges however many a task has.
Register also turns
+// each edge to or from a task group, which [TaskGroupRef.Before] describes,
into edges between
+// tasks. It expands those edges one at a time, in the order they were first
declared, each
+// against every task, every edge declared between two tasks, and the task
edges that earlier
+// group edges expanded into.
//
// 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
diff --git a/go-sdk/airflow/dag.go b/go-sdk/airflow/dag.go
index 8a45c867eb4..710abe863ea 100644
--- a/go-sdk/airflow/dag.go
+++ b/go-sdk/airflow/dag.go
@@ -23,6 +23,8 @@ import (
"runtime"
"strings"
"sync"
+ "unicode"
+ "unicode/utf8"
"github.com/apache/airflow/go-sdk/internal/bundle"
)
@@ -36,12 +38,23 @@ type DagRef struct {
mu sync.Mutex
registered bool
- tasks []*TaskRef
- tasksByID map[string]*TaskRef
- // edgeLabels holds every edge of the Dag, whether Inputs, Before or
After declared it, and
- // the label that Label put on it. An edge with no label maps to the
empty string, so a
- // lookup reports whether the edge has been declared.
+ // tasks holds every task of the Dag, inside a task group or not, in
the order they were added.
+ tasks []*TaskRef
+ tasksByID map[string]*TaskRef
+ // groups holds every task group of the Dag, nested or not, in the
order they were added.
+ // Tasks and task groups share one namespace of IDs, so tasksByID and
groupsByID hold no key
+ // in common.
+ groups []*TaskGroupRef
+ groupsByID map[string]*TaskGroupRef
+ // edgeLabels holds every edge between two tasks of the Dag, whether
Inputs, Before or After
+ // declared it, and the label that Label put on it. An edge with no
label maps to the empty
+ // string, so a lookup reports whether the edge has been declared.
edgeLabels map[edgeKey]string
+ // groupEdges holds the edges that have a task group at one end or
both, in the order they
+ // were first declared, and groupEdgeLabels holds their labels as
edgeLabels does. Registration
+ // adds the edges between tasks that they stand for to edgeLabels.
+ groupEdges []groupEdge
+ groupEdgeLabels map[edgeKey]string
}
// Dag returns an empty Dag with the given dag_id. An optional [DagSpec] holds
the rest of the
@@ -54,8 +67,8 @@ type DagRef struct {
//
// bundle.Register(dag)
//
-// Add every task before Register. [DagRef.Task], [DagRef.If], [IfRef.Then]
and [IfRef.Else]
-// panic once the Dag is registered.
+// Add every task before Register. [DagRef.Task], [DagRef.If],
[DagRef.TaskGroup], [IfRef.Then],
+// [IfRef.Else] and the methods of [TaskGroupRef] panic once the Dag is
registered.
//
// [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.
@@ -76,12 +89,15 @@ 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 or [DagRef.If] adds later. Pass it to
[IfRef.Then] or [IfRef.Else]
-// to run it on one side of a condition. A TaskRef is a [Node], so
[TaskRef.Before] and
-// [TaskRef.After] order it against another task.
+// TaskRef is a task that [DagRef.Task] or [TaskGroupRef.Task] added to a Dag.
Pass it to [Inputs]
+// to give its result to a task added later. Pass it to [IfRef.Then] or
[IfRef.Else] to run it on
+// one side of a condition. A TaskRef is a [Node], so [TaskRef.Before] and
[TaskRef.After] order it
+// against another task or a task group.
type TaskRef struct {
- dag *DagRef
+ dag *DagRef
+ // group is the task group that the task was added through. It is nil
for a task that
+ // DagRef.Task or DagRef.If added.
+ group *TaskGroupRef
taskID string
spec TaskSpec
// resultType is the type of the result that the task function returns
with its error. It is
@@ -100,8 +116,8 @@ 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 is the IfRef that DagRef.If or TaskGroupRef.If returned for
the task. It is nil for a
+ // task from DagRef.Task or TaskGroupRef.Task.
ifRef *IfRef
}
@@ -139,15 +155,22 @@ type TaskRef struct {
// - the TaskSpec sets TriggerRule to a value that is not a TriggerRule
constant
// - the TaskSpec sets WeightRule to a value that is not a WeightRule
constant
// - 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 task_id, with the group_ids that prefix it, is longer than 250
characters, or holds a
+// character other than a letter, a digit, an underscore, a dash or a dot,
as Python's
+// validate_key requires
+// - the Dag already has a task with the same task_id, or a task group or
the join node of one
+// takes it
// - the Dag is already registered
func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef {
- return d.addTask("airflow.DagRef.Task", fn, opts, nil)
+ return d.addTask("airflow.DagRef.Task", nil, 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 {
+// addTask adds a task for Task and If, of the Dag or of a task group. method
names the caller in
+// panic messages. group is the task group that the task is added through, and
nil for a task of
+// the Dag itself. ifRef is the IfRef that If returns, and nil when Task calls
addTask.
+func (d *DagRef) addTask(
+ method string, group *TaskGroupRef, fn any, opts []TaskOption, ifRef
*IfRef,
+) *TaskRef {
trigger, isTrigger := fn.(TriggerDagRunTask)
var triggerSpec *TriggerDagRunSpec
var triggerErr error
@@ -168,6 +191,7 @@ func (d *DagRef) addTask(method string, fn any, opts
[]TaskOption, ifRef *IfRef)
method, d.dagID,
))
}
+ d.checkGroupLocked(method, group)
var wrapped bundle.Task
if !isTrigger {
wrap := bundle.NewPositionalTaskFunction
@@ -207,7 +231,7 @@ func (d *DagRef) addTask(method string, fn any, opts
[]TaskOption, ifRef *IfRef)
continue
}
panic(fmt.Sprintf(
- "%s: task %q of Dag %q: %v", method,
findTaskName(fn, opts), d.dagID, err,
+ "%s: task %q of Dag %q: %v", method,
findTaskName(group, fn, opts), d.dagID, err,
))
}
}
@@ -230,6 +254,11 @@ func (d *DagRef) addTask(method string, fn any, opts
[]TaskOption, ifRef *IfRef)
))
}
}
+ unprefixed := taskID
+ taskID = group.childID(taskID)
+ if err := checkTaskID(taskID, taskID != unprefixed); err != nil {
+ panic(fmt.Sprintf("%s: Dag %q: %v", method, d.dagID, err))
+ }
if triggerErr != nil {
panic(fmt.Sprintf("%s: task %q of Dag %q: %v", method, taskID,
d.dagID, triggerErr))
}
@@ -243,6 +272,13 @@ func (d *DagRef) addTask(method string, fn any, opts
[]TaskOption, ifRef *IfRef)
method, d.dagID, taskID,
))
}
+ if taken := d.describeIDLocked(taskID); taken != "" {
+ panic(fmt.Sprintf(
+ "%s: Dag %q cannot add task %q, because %s already
takes the ID; "+
+ "set another task_id with
airflow.TaskSpec{TaskID: ...}",
+ method, d.dagID, taskID, taken,
+ ))
+ }
var upstreams []*TaskRef
var resultType reflect.Type
if isTrigger {
@@ -264,6 +300,7 @@ func (d *DagRef) addTask(method string, fn any, opts
[]TaskOption, ifRef *IfRef)
task := &TaskRef{
dag: d,
+ group: group,
taskID: taskID,
spec: copySpec(cfg.spec),
resultType: resultType,
@@ -280,6 +317,9 @@ func (d *DagRef) addTask(method string, fn any, opts
[]TaskOption, ifRef *IfRef)
}
d.tasksByID[taskID] = task
d.tasks = append(d.tasks, task)
+ if group != nil {
+ group.children = append(group.children, task)
+ }
// Inputs passes a task once per parameter it fills, so the same task
can arrive twice. The
// edge is one either way, and the task is new, so no edge to it
carries a label to settle.
for _, upstream := range upstreams {
@@ -289,13 +329,21 @@ func (d *DagRef) addTask(method string, fn any, opts
[]TaskOption, ifRef *IfRef)
}
// markRegistered marks d as registered, which stops any further change to d.
It panics instead
-// when a condition from If has no task from Then, or when the edges of d
close a cycle. The Dag
-// is whole by then, so one walk of the graph answers for every edge its tasks
declared, and a Dag
-// that fails a check stays unregistered and can still be corrected.
+// when a condition from If has no task from Then, or when the edges of d
close a cycle. The Dag is
+// whole by then, so markRegistered can expand the group edges in the order
they were first
+// declared, and a walk of the whole graph answers for every edge. It walks
the graph before the
+// expansion too, so that a cycle between the edges the author declared is
reported as declared.
+// A Dag that fails a check stays unregistered and holds only the edges its
author declared. The
+// author can still give a condition its task from Then, but cannot undo a
cycle, since a Dag only
+// ever gains edges.
func (d *DagRef) markRegistered() {
d.mu.Lock()
defer d.mu.Unlock()
+ // A Dag that another bundle registered has been checked and expanded
already.
+ if d.registered {
+ return
+ }
for _, task := range d.tasks {
if task.ifRef != nil && task.ifRef.thenTask == nil {
panic(fmt.Sprintf(
@@ -306,18 +354,41 @@ func (d *DagRef) markRegistered() {
}
}
if cycle := d.cycleLocked(); cycle != nil {
- panic(fmt.Sprintf(
- "airflow.BundleRef.Register: the task dependencies of
Dag %q contain a cycle: %s",
- d.dagID, strings.Join(cycle, " -> "),
- ))
+ panic(cycleMessage(d.dagID, cycle, nil))
+ }
+ expanded := d.expandGroupEdgesLocked()
+ if cycle := d.cycleLocked(); cycle != nil {
+ through := groupEdgesOn(cycle, expanded)
+ d.removeEdgesLocked(expanded)
+ panic(cycleMessage(d.dagID, cycle, through))
}
d.registered = true
}
+// cycleMessage reports a cycle in the task dependencies of Dag dagID. through
names the group
+// edges that edges on the cycle stand for, so that a cycle that only group
edges close points at
+// them.
+func cycleMessage(dagID string, cycle, through []string) string {
+ message := fmt.Sprintf(
+ "airflow.BundleRef.Register: the task dependencies of Dag %q
contain a cycle: %s",
+ dagID, strings.Join(cycle, " -> "),
+ )
+ switch len(through) {
+ case 0:
+ return message
+ case 1:
+ return message + ", through the group edge " + through[0]
+ default:
+ return message + ", through the group edges " +
strings.Join(through, ", ")
+ }
+}
+
func funcName(fn any) string { return
runtime.FuncForPC(reflect.ValueOf(fn).Pointer()).Name() }
-// findTaskName names a task in an error that addTask raises before it settles
the task_id.
-func findTaskName(fn any, opts []TaskOption) string {
+// findTaskName names a task in an error that addTask raises before it settles
the task_id. group
+// is the task group that the task is added through, and prefixes the task_id
that findTaskName
+// finds.
+func findTaskName(group *TaskGroupRef, fn any, opts []TaskOption) string {
for _, opt := range opts {
var spec TaskSpec
switch opt := opt.(type) {
@@ -329,18 +400,64 @@ func findTaskName(fn any, opts []TaskOption) string {
}
}
if spec.TaskID != "" {
- return spec.TaskID
+ return group.childID(spec.TaskID)
}
}
if _, ok := fn.(TriggerDagRunTask); ok {
return "airflow.TriggerDagRun"
}
if taskID, ok := taskIDFromFuncName(funcName(fn)); ok {
- return taskID
+ return group.childID(taskID)
}
return funcName(fn)
}
+// taskIDMaxLength is the longest task_id that Python's validate_key accepts,
counted in
+// characters.
+const taskIDMaxLength = 250
+
+// checkTaskID checks a task_id as Python's validate_key does when an operator
is constructed,
+// which is after the group_ids of its task groups prefix it. Python matches
the ID against
+// ^[\w.-]+$, whose $ also lets a trailing newline through, and checkTaskID
does not. prefixed
+// reports whether a group_id prefixes the task_id, so that the error can say
what to shorten.
+func checkTaskID(taskID string, prefixed bool) error {
+ if !utf8.ValidString(taskID) {
+ return fmt.Errorf(
+ "task_id %q is not valid UTF-8; set another one with
airflow.TaskSpec{TaskID: ...}",
+ taskID,
+ )
+ }
+ if n := utf8.RuneCountInString(taskID); n > taskIDMaxLength {
+ if prefixed {
+ return fmt.Errorf(
+ "task_id %q has %d characters, counting the
group_ids that prefix it, and a "+
+ "task_id has at most %d; shorten a
group_id, or set a shorter task_id with "+
+ "airflow.TaskSpec{TaskID: ...}",
+ taskID, n, taskIDMaxLength,
+ )
+ }
+ return fmt.Errorf(
+ "task_id %q has %d characters, and a task_id has at
most %d; "+
+ "set a shorter one with
airflow.TaskSpec{TaskID: ...}",
+ taskID, n, taskIDMaxLength,
+ )
+ }
+ for _, r := range taskID {
+ if !isWordRune(r) && r != '-' && r != '.' {
+ return fmt.Errorf(
+ "task_id %q holds %q, and a task_id holds only
letters, digits, underscores, "+
+ "dashes and dots; set another one with
airflow.TaskSpec{TaskID: ...}",
+ taskID, r,
+ )
+ }
+ }
+ return nil
+}
+
+// isWordRune reports whether r matches \w in a Python regular expression: a
character for which
+// str.isalnum is true, or the underscore.
+func isWordRune(r rune) bool { return unicode.IsLetter(r) ||
unicode.IsNumber(r) || r == '_' }
+
// taskIDFromFuncName takes the runtime name of a function and returns the
name that the
// function is declared with. It reports false when the runtime name does not
carry one.
func taskIDFromFuncName(name string) (string, bool) {
diff --git a/go-sdk/airflow/if.go b/go-sdk/airflow/if.go
index a5c7bba0c45..42b71e90ed0 100644
--- a/go-sdk/airflow/if.go
+++ b/go-sdk/airflow/if.go
@@ -25,8 +25,9 @@ import (
"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.
+// IfRef is a condition that [DagRef.If] or [TaskGroupRef.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.
@@ -78,15 +79,21 @@ type IfRef struct {
// 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 {
+ return d.addIf("airflow.DagRef.If", nil, fn, opts)
+}
+
+// addIf adds a condition for DagRef.If and TaskGroupRef.If. group is the task
group that the
+// condition is added through, and nil for DagRef.If.
+func (d *DagRef) addIf(method string, group *TaskGroupRef, fn any, opts
[]TaskOption) *IfRef {
if _, ok := fn.(TriggerDagRunTask); ok {
panic(fmt.Sprintf(
- "airflow.DagRef.If: Dag %q: fn comes from
airflow.TriggerDagRun, "+
+ "%s: Dag %q: fn comes from airflow.TriggerDagRun, "+
"but a condition function is a Go function that
returns (bool, error)",
- d.dagID,
+ method, d.dagID,
))
}
ifRef := &IfRef{}
- d.addTask("airflow.DagRef.If", fn, opts, ifRef)
+ d.addTask(method, group, fn, opts, ifRef)
return ifRef
}
@@ -100,8 +107,9 @@ func (d *DagRef) If(fn any, opts ...TaskOption) *IfRef {
// 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
+// - g is not the IfRef that DagRef.If or TaskGroupRef.If returned, for
example a copy of it
+// - task is nil, or neither DagRef.Task nor TaskGroupRef.Task added 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
@@ -125,7 +133,7 @@ func (g *IfRef) setTask(side string, task *TaskRef) {
// 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")
+ panic(method + ": DagRef.If or TaskGroupRef.If did not return
the *airflow.IfRef")
}
condition, d := g.task.taskID, g.task.dag
d.mu.Lock()
@@ -152,7 +160,8 @@ func (g *IfRef) setTask(side string, task *TaskRef) {
// 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",
+ "%s: condition %q of Dag %q got a *airflow.TaskRef that
DagRef.Task or "+
+ "TaskGroupRef.Task did not return",
method, condition, d.dagID,
))
}
diff --git a/go-sdk/airflow/if_test.go b/go-sdk/airflow/if_test.go
index c61dae205e7..059e2fa21dc 100644
--- a/go-sdk/airflow/if_test.go
+++ b/go-sdk/airflow/if_test.go
@@ -199,9 +199,9 @@ func TestIfPanicsUnderItsOwnName(t *testing.T) {
want: "cannot take an input from task",
},
{
- name: "input that DagRef.Task did not return",
+ name: "input that DagRef.Task or TaskGroupRef.Task did
not return",
add: func(dag *DagRef) { dag.If(hasRows,
Inputs(&TaskRef{})) },
- want: "that DagRef.Task did not return",
+ want: "that DagRef.Task or TaskGroupRef.Task did not
return",
},
{
name: "missing input",
@@ -307,12 +307,12 @@ func TestThenAndElseRejectATaskOutsideTheDag(t
*testing.T) {
{
name: "zero TaskRef",
task: &TaskRef{},
- want: `got a *airflow.TaskRef that DagRef.Task did not
return`,
+ want: `got a *airflow.TaskRef that DagRef.Task or
TaskGroupRef.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`,
+ want: `got a *airflow.TaskRef that DagRef.Task or
TaskGroupRef.Task did not return`,
},
{
name: "task of another Dag",
@@ -416,11 +416,13 @@ func TestIfRefThatDagRefIfDidNotReturnPanics(t
*testing.T) {
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",
+ "airflow.IfRef.Then: DagRef.If or
TaskGroupRef.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",
+ "airflow.IfRef.Else: DagRef.If or
TaskGroupRef.If did not return "+
+ "the *airflow.IfRef",
func() { ref.Else(loaded) },
)
})
diff --git a/go-sdk/airflow/inputs.go b/go-sdk/airflow/inputs.go
index baba9d5f027..75775b3580c 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] 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()):
+// Inputs passes the results of tasks to the task that [DagRef.Task],
[DagRef.If], or one of the
+// methods of the same names on [TaskGroupRef] adds, and makes each of those
tasks an upstream task
+// of the new one. It is the Go form of a Python TaskFlow call such as
transform(extract()):
//
// extracted := dag.Task(extract)
// transformed := dag.Task(transform, airflow.Inputs(extracted))
@@ -58,9 +58,9 @@ 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] 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.
+// The method that adds the task panics unless each parameter after the
Context gets exactly one
+// task of the same Dag, inside a task group or not, 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
@@ -89,7 +89,7 @@ func (d *DagRef) checkInputs(
case d.tasksByID[upstream.taskID] != upstream:
panic(fmt.Sprintf(
"%s: task %q of Dag %q: airflow.Inputs got a
*airflow.TaskRef "+
- "at index %d that DagRef.Task did not
return",
+ "at index %d that DagRef.Task or
TaskGroupRef.Task did not return",
method, taskID, d.dagID, i,
))
}
diff --git a/go-sdk/airflow/inputs_test.go b/go-sdk/airflow/inputs_test.go
index af0c86c7372..5532dbd5a43 100644
--- a/go-sdk/airflow/inputs_test.go
+++ b/go-sdk/airflow/inputs_test.go
@@ -349,13 +349,13 @@ func TestTaskPanicsOnAnInputFromOutsideTheDag(t
*testing.T) {
name: "zero TaskRef",
inputs: Inputs(read, &TaskRef{}),
want: `task "mergeRows" of Dag "etl": airflow.Inputs
got a *airflow.TaskRef ` +
- `at index 1 that DagRef.Task did not return`,
+ `at index 1 that DagRef.Task or
TaskGroupRef.Task did not return`,
},
{
name: "copy of a task of the Dag",
inputs: Inputs(&readCopy),
want: `task "mergeRows" of Dag "etl": airflow.Inputs
got a *airflow.TaskRef ` +
- `at index 0 that DagRef.Task did not return`,
+ `at index 0 that DagRef.Task or
TaskGroupRef.Task did not return`,
},
{
name: "task of another Dag",
diff --git a/go-sdk/airflow/node.go b/go-sdk/airflow/node.go
index 6b0cefab74d..643c599d814 100644
--- a/go-sdk/airflow/node.go
+++ b/go-sdk/airflow/node.go
@@ -22,8 +22,9 @@ import (
"slices"
)
-// Node is what an edge between tasks connects. A [TaskRef] is one, so the
task that
-// [DagRef.Task] returns is an edge endpoint.
+// Node is what an edge connects: a task or a whole task group. A [TaskRef]
and a [TaskGroupRef]
+// are both Nodes, so the task that [DagRef.Task] returns and the group that
[DagRef.TaskGroup]
+// returns are both edge endpoints.
//
// Before and After declare the order-only edges that Python writes with >>
and <<, for tasks
// that have to run in an order but pass no data. [Inputs] declares the edge
that carries a
@@ -38,26 +39,62 @@ import (
//
// The second Before above starts from load, not from extract.
//
-// Node is the Go counterpart of Python's DAGNode, where set_upstream and
set_downstream live.
+// Node is the Go counterpart of Python's DAGNode, where set_upstream and
set_downstream live. In
+// Python, both an operator and a TaskGroup are DAGNodes.
// Its node method is unexported, so a type declared outside package airflow
can be a Node only by
// embedding one, and Before and After reject such a type as an argument.
type Node interface {
- // Before makes the receiver an upstream task of every node, as
Python's >> does.
+ // Before makes the receiver an upstream of every node, as Python's >>
does.
Before(nodes ...Node) Node
- // After makes the receiver a downstream task of every node, as
Python's << does.
+ // After makes the receiver a downstream of every node, as Python's <<
does.
After(nodes ...Node) Node
// node seals the interface. Only a type that package airflow declares
is a Node, so
// endpoints can read what one stands for.
node()
}
-// nodeEndpoint is one task that a Node stands for, with the label that
[Label] carries into the
-// edge verb the Node is passed to.
+// nodeEndpoint is one task or one task group that a Node stands for, with the
label that [Label]
+// carries into the edge verb the Node is passed to. Exactly one of task and
group is set, except
+// for the endpoint of a nil *TaskRef, which has neither.
type nodeEndpoint struct {
task *TaskRef
+ group *TaskGroupRef
label string
}
+// id returns the task_id or the group_id of e. Tasks and task groups share
one namespace of IDs,
+// so the ID names one node of the Dag.
+func (e nodeEndpoint) id() string {
+ if e.group != nil {
+ return e.group.groupID
+ }
+ return e.task.taskID
+}
+
+// owner returns the Dag that e belongs to, and nil for an endpoint that no
Dag returned.
+func (e nodeEndpoint) owner() *DagRef {
+ if e.group != nil {
+ return e.group.dag
+ }
+ return e.task.dag
+}
+
+// container returns the innermost task group that holds e, and nil when only
the Dag does.
+func (e nodeEndpoint) container() *TaskGroupRef {
+ if e.group != nil {
+ return e.group.parent
+ }
+ return e.task.group
+}
+
+// describe names e in a panic message, such as task "load" or task group
"transform".
+func (e nodeEndpoint) describe() string {
+ if e.group != nil {
+ return fmt.Sprintf("task group %q", e.group.groupID)
+ }
+ return fmt.Sprintf("task %q", e.task.taskID)
+}
+
// nodeSet is the argument set that Before and After return as one Node, and
the Node that
// [Label] returns. A label belongs to the edge of the verb the Node is passed
to, so a nodeSet
// carries its labels no further once it is the receiver of the next verb.
@@ -66,11 +103,11 @@ type nodeSet []nodeEndpoint
func (nodeSet) node() {}
// unlabelled returns the set with its labels dropped. A label belongs to the
one verb it was
-// passed to, so the Node a verb returns, which stands for the tasks it
pointed at, carries none.
+// passed to, so the Node a verb returns, which stands for the nodes it
pointed at, carries none.
func (s nodeSet) unlabelled() nodeSet {
plain := make(nodeSet, len(s))
for i, endpoint := range s {
- plain[i] = nodeEndpoint{task: endpoint.task}
+ plain[i] = nodeEndpoint{task: endpoint.task, group:
endpoint.group}
}
return plain
}
@@ -81,13 +118,19 @@ func (s nodeSet) After(nodes ...Node) Node { return
declareEdges(s, nodes, dirAf
func (*TaskRef) node() {}
-// endpoints returns the tasks that node stands for. The Node method is
unexported, so the types
-// of this package are the only Nodes, and a struct that embeds one is all
that reaches the
-// default case. where names the Node in a panic, such as
"airflow.Node.Before: nodes[0]".
+// endpoints returns the tasks and task groups that node stands for. The Node
method is
+// unexported, so the types of this package are the only Nodes, and a struct
that embeds one is all
+// that reaches the default case. where names the Node in a panic, such as
+// "airflow.Node.Before: nodes[0]".
func endpoints(where string, node Node) []nodeEndpoint {
switch node := node.(type) {
case *TaskRef:
return []nodeEndpoint{{task: node}}
+ case *TaskGroupRef:
+ if node == nil {
+ panic(where + " is a nil *airflow.TaskGroupRef")
+ }
+ return []nodeEndpoint{{group: node}}
case nodeSet:
return node
default:
@@ -111,10 +154,14 @@ func endpoints(where string, node Node) []nodeEndpoint {
//
// Declaring an edge that the Dag already has changes nothing, other than to
apply a [Label].
//
+// A node can also be a task group, which [TaskGroupRef.Before] describes.
+//
// Before panics if:
-// - a node is nil, or is a task that [DagRef.Task] did not return
+// - a node is nil, is a task that [DagRef.Task] or [TaskGroupRef.Task] did
not return, or is a
+// task group that [DagRef.TaskGroup] or [TaskGroupRef.TaskGroup] did not
return
// - a node belongs to another Dag
// - the edge would make a task depend on itself
+// - a node is a task group that holds the task
// - the Dag is already registered
//
// Every check runs before the call records any edge, so a fan-out that panics
leaves the Dag as
@@ -154,6 +201,11 @@ func (t *TaskRef) After(nodes ...Node) Node { return
declareEdges(t, nodes, dirA
// transformed := dag.Task(transform, Inputs(extracted))
// extracted.Before(Label(transformed, "rows"))
//
+// A label on an edge to or from a task group stays on that edge, as
[TaskGroupRef.Before]
+// describes. A label on an edge between two tasks stays on that edge too,
even when the tasks are
+// in different task groups. In that case Python can replace the receiver of
the verb with a group
+// that holds it, which makes the edge connect other tasks.
+//
// Label panics if node is nil or text is empty.
func Label(node Node, text string) Node {
if node == nil {
@@ -165,12 +217,13 @@ func Label(node Node, text string) Node {
labelling := endpoints("airflow.Label: node", node)
labelled := make(nodeSet, len(labelling))
for i, endpoint := range labelling {
- labelled[i] = nodeEndpoint{task: endpoint.task, label: text}
+ labelled[i] = nodeEndpoint{task: endpoint.task, group:
endpoint.group, label: text}
}
return labelled
}
-// edgeKey identifies an edge of a Dag. The task_ids of a Dag are unique, so
they name the ends.
+// edgeKey identifies an edge of a Dag. Tasks and task groups share one
namespace of IDs, so a
+// task_id or a group_id names each end.
type edgeKey struct{ upstream, downstream string }
// edgeDir is which way an edge verb points: [TaskRef.Before] from its
receiver, and
@@ -191,7 +244,7 @@ func (dir edgeDir) String() string {
// order returns the two ends of the edge between an endpoint of the receiver
and one of the nodes
// the verb was given.
-func (dir edgeDir) order(recv, arg *TaskRef) (upstream, downstream *TaskRef) {
+func (dir edgeDir) order(recv, arg nodeEndpoint) (upstream, downstream
nodeEndpoint) {
if dir == dirAfter {
return arg, recv
}
@@ -200,11 +253,16 @@ func (dir edgeDir) order(recv, arg *TaskRef) (upstream,
downstream *TaskRef) {
// pendingEdge is an edge that declareEdges has checked and is about to record.
type pendingEdge struct {
- upstream, downstream *TaskRef
+ upstream, downstream nodeEndpoint
key edgeKey
label string
}
+// betweenTasks reports whether both ends of e are tasks rather than task
groups.
+func (e pendingEdge) betweenTasks() bool {
+ return e.upstream.group == nil && e.downstream.group == nil
+}
+
// declareEdges records an edge from every endpoint of receiver to every node
it was given, or
// the other way round for After, and returns those nodes as one Node.
func declareEdges(receiver Node, nodes []Node, dir edgeDir) Node {
@@ -235,12 +293,14 @@ func declareEdges(receiver Node, nodes []Node, dir
edgeDir) Node {
}
for _, endpoint := range all {
// A zero TaskRef and a copy of a TaskRef get here.
- if dag.tasksByID[endpoint.task.taskID] != endpoint.task {
+ if endpoint.group == nil && dag.tasksByID[endpoint.task.taskID]
!= endpoint.task {
panic(fmt.Sprintf(
- "%s: Dag %q got a *airflow.TaskRef that
DagRef.Task did not return",
+ "%s: Dag %q got a *airflow.TaskRef that
DagRef.Task or TaskGroupRef.Task "+
+ "did not return",
where, dag.dagID,
))
}
+ dag.checkGroupLocked(where, endpoint.group)
}
// A verb with no node to point at, which a spread of an empty slice
reaches, declares no
// edge. So does one on the empty set that such a verb returned.
@@ -253,33 +313,55 @@ func declareEdges(receiver Node, nodes []Node, dir
edgeDir) Node {
at := make(map[edgeKey]int, len(ends)*len(args))
for _, end := range ends {
for _, arg := range args {
- upstream, downstream := dir.order(end.task, arg.task)
- if upstream == downstream {
+ upstream, downstream := dir.order(end, arg)
+ if upstream.task == downstream.task && upstream.group
== downstream.group {
panic(fmt.Sprintf(
- "%s: Dag %q: task %q cannot depend on
itself",
- where, dag.dagID, upstream.taskID,
+ "%s: Dag %q: %s cannot depend on
itself", where, dag.dagID, upstream.describe(),
))
}
- key := edgeKey{upstream: upstream.taskID, downstream:
downstream.taskID}
+ checkGroupEdgeEnds(where, dag, upstream, downstream)
+ key := edgeKey{upstream: upstream.id(), downstream:
downstream.id()}
i, declared := at[key]
if !declared {
i = len(pending)
at[key] = i
- pending = append(pending, pendingEdge{
- upstream: upstream, downstream:
downstream, key: key,
- label: dag.edgeLabels[key],
- })
+ edge := pendingEdge{upstream: upstream,
downstream: downstream, key: key}
+ if edge.betweenTasks() {
+ edge.label = dag.edgeLabels[key]
+ } else {
+ edge.label = dag.groupEdgeLabels[key]
+ }
+ pending = append(pending, edge)
}
// The label belongs to the node the verb was given, in
either direction.
pending[i].label = mergeLabel(pending[i].label,
arg.label)
}
}
for _, edge := range pending {
- dag.addEdgeLocked(edge.upstream, edge.downstream, edge.label)
+ if edge.betweenTasks() {
+ dag.addEdgeLocked(edge.upstream.task,
edge.downstream.task, edge.label)
+ } else {
+ dag.addGroupEdgeLocked(edge.upstream, edge.downstream,
edge.label)
+ }
}
return args.unlabelled()
}
+// checkGroupEdgeEnds panics if one end of an edge is a task group that holds
the other end. The
+// edge would order the group against part of itself: in Python, group >>
task_in_group makes the
+// last tasks of the group upstreams of a task that may be one of them.
+func checkGroupEdgeEnds(where string, dag *DagRef, upstream, downstream
nodeEndpoint) {
+ for _, pair := range [][2]nodeEndpoint{{upstream, downstream},
{downstream, upstream}} {
+ if group := pair[0].group; group != nil && group.holds(pair[1])
{
+ panic(fmt.Sprintf(
+ "%s: Dag %q: %s is inside task group %q, so an
edge cannot connect them; "+
+ "an edge connects a group to a task or
a group outside it",
+ where, dag.dagID, pair[1].describe(),
group.groupID,
+ ))
+ }
+ }
+}
+
// mergeLabel returns the label an edge carries once label is declared on it.
A declaration that
// carries no label leaves the edge's own label alone, which is what makes
redeclaring an edge
// idempotent, and one that carries a label overwrites it, as Python's
DAG.set_edge_info does.
@@ -291,29 +373,35 @@ func mergeLabel(declared, label string) string {
}
// edgeDag returns the Dag that every end of an edge belongs to. It panics
unless each end is a
-// task of that one Dag.
+// task or a task group of that one Dag.
func edgeDag(where string, ends []nodeEndpoint) *DagRef {
- var first *TaskRef
+ var first nodeEndpoint
+ var dag *DagRef
for _, end := range ends {
- task := end.task
switch {
- case task == nil:
+ case end.task == nil && end.group == nil:
panic(fmt.Sprintf("%s: got a nil *airflow.TaskRef",
where))
- case task.dag == nil:
+ case end.group != nil && end.group.dag == nil:
+ panic(fmt.Sprintf(
+ "%s: got a *airflow.TaskGroupRef that
DagRef.TaskGroup or "+
+ "TaskGroupRef.TaskGroup did not
return", where,
+ ))
+ case end.group == nil && end.task.dag == nil:
panic(fmt.Sprintf(
- "%s: got a *airflow.TaskRef that DagRef.Task
did not return", where,
+ "%s: got a *airflow.TaskRef that DagRef.Task or
TaskGroupRef.Task did not return",
+ where,
))
- case first == nil:
- first = task
- case task.dag != first.dag:
+ case dag == nil:
+ first, dag = end, end.owner()
+ case end.owner() != dag:
panic(fmt.Sprintf(
- "%s: cannot declare an edge between task %q of
Dag %q and task %q of Dag %q; "+
- "an edge connects tasks of one Dag",
- where, first.taskID, first.dag.dagID,
task.taskID, task.dag.dagID,
+ "%s: cannot declare an edge between %s of Dag
%q and %s of Dag %q; "+
+ "an edge connects the tasks and task
groups of one Dag",
+ where, first.describe(), dag.dagID,
end.describe(), end.owner().dagID,
))
}
}
- return first.dag
+ return dag
}
// addEdgeLocked records one edge of d, which the caller holds d.mu for. An
edge d already has is
diff --git a/go-sdk/airflow/node_test.go b/go-sdk/airflow/node_test.go
index e284a9f092d..474011af37c 100644
--- a/go-sdk/airflow/node_test.go
+++ b/go-sdk/airflow/node_test.go
@@ -354,11 +354,13 @@ func TestEdgeVerbsRejectATaskThatTaskDidNotReturn(t
*testing.T) {
copied := *loaded
assert.PanicsWithValue(t,
- "airflow.Node.Before: got a *airflow.TaskRef that DagRef.Task
did not return",
+ "airflow.Node.Before: got a *airflow.TaskRef that DagRef.Task
or TaskGroupRef.Task "+
+ "did not return",
func() { (&TaskRef{}).Before(loaded) },
)
assert.PanicsWithValue(t,
- `airflow.Node.Before: Dag "etl" got a *airflow.TaskRef that
DagRef.Task did not return`,
+ `airflow.Node.Before: Dag "etl" got a *airflow.TaskRef that
DagRef.Task or `+
+ `TaskGroupRef.Task did not return`,
func() { loaded.Before(&copied) },
)
}
@@ -369,7 +371,7 @@ func TestEdgeVerbsRejectATaskOfAnotherDag(t *testing.T) {
assert.PanicsWithValue(t,
`airflow.Node.Before: cannot declare an edge between task
"load" of Dag "etl" and `+
- `task "cleanup" of Dag "reporting"; an edge connects
tasks of one Dag`,
+ `task "cleanup" of Dag "reporting"; an edge connects
the tasks and task groups of one Dag`,
func() { loaded.Before(cleaned) },
)
}
@@ -449,7 +451,8 @@ func TestEdgeVerbsCheckTheTasksEvenWithNoNode(t *testing.T)
{
var none []Node
assert.PanicsWithValue(t,
- "airflow.Node.Before: got a *airflow.TaskRef that DagRef.Task
did not return",
+ "airflow.Node.Before: got a *airflow.TaskRef that DagRef.Task
or TaskGroupRef.Task "+
+ "did not return",
func() { (&TaskRef{}).Before(none...) },
)
assert.PanicsWithValue(t,
diff --git a/go-sdk/airflow/spec.gen.go b/go-sdk/airflow/spec.gen.go
index a4674daf848..d5eda56f3ad 100644
--- a/go-sdk/airflow/spec.gen.go
+++ b/go-sdk/airflow/spec.gen.go
@@ -75,7 +75,32 @@ type DagSpec struct {
Tags []string
}
-// TaskSpec holds the attributes of a task. DagRef.Task takes at most one per
task.
+// TaskGroupSpec holds the attributes of a task group other than its group_id.
+// DagRef.TaskGroup and TaskGroupRef.TaskGroup take at most one per group.
+type TaskGroupSpec struct {
+ // DocMD corresponds to the JSON schema field "doc_md".
+ DocMD string
+
+ // GroupDisplayName corresponds to the JSON schema field
"group_display_name".
+ GroupDisplayName string
+
+ // PrefixGroupID says whether the group_id prefixes the IDs of the
tasks and
+ // groups added through the group, as in "transform.cleanRows". When
PrefixGroupID
+ // is nil, the group_id prefixes them.
+ PrefixGroupID *bool
+
+ // Tooltip corresponds to the JSON schema field "tooltip".
+ Tooltip string
+
+ // UIColor corresponds to the JSON schema field "ui_color".
+ UIColor string
+
+ // UIFgColor corresponds to the JSON schema field "ui_fgcolor".
+ UIFgColor string
+}
+
+// TaskSpec holds the attributes of a task. DagRef.Task, DagRef.If and the
methods
+// of the same names on TaskGroupRef take at most one per task.
type TaskSpec struct {
// TaskDisplayName corresponds to the JSON schema field
"_task_display_name".
TaskDisplayName string
@@ -152,7 +177,9 @@ type TaskSpec struct {
// 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.
+ // Go function, so it needs a TaskID. A task added through a task group
takes the
+ // group_id as a prefix of its task_id, unless the TaskGroupSpec of the
group sets
+ // PrefixGroupID to false.
TaskID string
// TriggerRule corresponds to the JSON schema field "trigger_rule".
diff --git a/go-sdk/airflow/spec.go b/go-sdk/airflow/spec.go
index e36b0c70b40..61de3fa65f7 100644
--- a/go-sdk/airflow/spec.go
+++ b/go-sdk/airflow/spec.go
@@ -19,8 +19,8 @@ package airflow
import "reflect"
-// DagSpec and TaskSpec are generated from Airflow core's Dag serialization
schema,
-// which Python owns, so that neither struct drifts from it silently. genspec
+// DagSpec, TaskSpec and TaskGroupSpec are generated from Airflow core's Dag
serialization
+// schema, which Python owns, so that no struct drifts from it silently.
genspec
// rewrites the schema into the authoring shape, go-jsonschema writes the
structs,
// and genspec puts the license header back on what it wrote. The rewritten
schema
// is a build artifact under .build; spec.gen.go is committed.
@@ -36,7 +36,7 @@ import "reflect"
// json.Marshal(spec) to skip that step and produce a shape core misreads.
//go:generate go run ../internal/genspec -schema ../schema/dag-schema.json
-out ../../.build/go-sdk/spec.schema.json
-//go:generate go run github.com/atombender/[email protected] --only-models
--struct-name-from-title --tags "" --capitalization ID --capitalization JSON
--capitalization MD --capitalization XCom -p airflow -o spec.gen.go
../../.build/go-sdk/spec.schema.json
+//go:generate go run github.com/atombender/[email protected] --only-models
--struct-name-from-title --tags "" --capitalization ID --capitalization JSON
--capitalization FgColor --capitalization MD --capitalization UI
--capitalization XCom -p airflow -o spec.gen.go
../../.build/go-sdk/spec.schema.json
//go:generate go run ../internal/genspec -license spec.gen.go
// copySpec returns a copy of spec that shares no slice, map or pointer with
it, so that a
diff --git a/go-sdk/airflow/task_group.go b/go-sdk/airflow/task_group.go
new file mode 100644
index 00000000000..177e04467c5
--- /dev/null
+++ b/go-sdk/airflow/task_group.go
@@ -0,0 +1,553 @@
+// 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"
+ "slices"
+ "strings"
+ "unicode/utf8"
+)
+
+// groupIDMaxLength is the longest group_id that Python's validate_group_key
accepts, counted in
+// characters.
+const groupIDMaxLength = 200
+
+// The suffixes of the two join nodes that the Airflow UI draws for every task
group, as Python's
+// TaskGroup.upstream_join_id and TaskGroup.downstream_join_id name them. Both
IDs belong to the
+// namespace that task_ids and group_ids share, so no task or group can take
them.
+const (
+ upstreamJoinSuffix = ".upstream_join_id"
+ downstreamJoinSuffix = ".downstream_join_id"
+)
+
+// TaskGroupRef is a task group that [DagRef.TaskGroup] or
[TaskGroupRef.TaskGroup] added to a
+// Dag. It offers the methods of the Dag that add tasks and groups, and the
group_id prefixes the
+// ID of each task and group it adds, unless its [TaskGroupSpec] sets
PrefixGroupID to false. A
+// TaskGroupRef is a [Node], so [TaskGroupRef.Before] and [TaskGroupRef.After]
order the whole
+// group against a task or another group.
+type TaskGroupRef struct {
+ dag *DagRef
+ // parent is the group that the group was added through. It is nil for
a group that
+ // DagRef.TaskGroup added.
+ parent *TaskGroupRef
+ groupID string
+ spec TaskGroupSpec
+ // children holds the tasks and the groups added through the group, in
the order they were
+ // added. Each is a *TaskRef or a *TaskGroupRef. The groups of a Dag
form a tree, which a
+ // serialized Dag carries in task_group.children.
+ children []Node
+}
+
+// TaskGroup adds a task group to the Dag and returns it. Add the tasks of the
group through the
+// group, and order the group as a whole against a task or another group:
+//
+// group := dag.TaskGroup("transform")
+// cleaned := group.Task(cleanRows)
+// validated := group.Task(validateRows, airflow.Inputs(cleaned))
+//
+// extracted.Before(group) // extract >> transform
+//
+// The group_id prefixes the ID of each task and group added through the
group, as Python's
+// prefix_group_id does. So the tasks above are transform.cleanRows and
transform.validateRows.
+// [TaskGroupRef.TaskGroup] nests a group in another, and the prefixes nest
with it. Set
+// PrefixGroupID in a [TaskGroupSpec] to false to keep the IDs of the tasks
and groups added
+// through the group as they are written. Then they have to be unique across
the Dag.
+//
+// An optional TaskGroupSpec holds the rest of the group's attributes, and
TaskGroup panics if it
+// gets more than one TaskGroupSpec.
+//
+// Renaming a group renames each task whose task_id the group_id prefixes.
Airflow treats a
+// renamed task as a new task, as it does when the Go function of a task is
renamed.
+//
+// Task groups and tasks share one namespace of IDs, as they do in Python. The
IDs of the two join
+// nodes that the Airflow UI draws for a group, the group_id followed by
.upstream_join_id or
+// .downstream_join_id, are part of it too. Python reserves them only once the
group exists, so it
+// lets a task added before the group take one. TaskGroup panics in that case
too.
+//
+// TaskGroup panics if:
+// - groupID is empty, is longer than 200 characters, or holds a character
other than a letter,
+// a digit, an underscore or a dash, as Python's validate_group_key
requires
+// - spec holds more than one TaskGroupSpec
+// - a task, a task group or a join node already takes the ID that the group
would get, or the
+// ID of one of its join nodes
+// - the Dag is already registered
+func (d *DagRef) TaskGroup(groupID string, spec ...TaskGroupSpec)
*TaskGroupRef {
+ return d.addGroup("airflow.DagRef.TaskGroup", nil, groupID, spec)
+}
+
+// TaskGroup adds a task group inside g and returns it. The new group is
+// [DagRef.TaskGroup] one level down: its group_id takes the group_id of g as
a prefix, unless the
+// TaskGroupSpec of g sets PrefixGroupID to false.
+//
+// transform := dag.TaskGroup("transform")
+// checks := transform.TaskGroup("checks")
+// checks.Task(checkNulls) // the task transform.checks.checkNulls
+//
+// TaskGroup panics for the reasons that DagRef.TaskGroup lists. It also
panics if g is not the
+// TaskGroupRef that DagRef.TaskGroup or TaskGroupRef.TaskGroup returned, for
example a copy of it.
+func (g *TaskGroupRef) TaskGroup(groupID string, spec ...TaskGroupSpec)
*TaskGroupRef {
+ const method = "airflow.TaskGroupRef.TaskGroup"
+ return g.groupDag(method).addGroup(method, g, groupID, spec)
+}
+
+// Task adds a task that runs fn to the Dag of g, inside g, and returns the
new task. It is
+// [DagRef.Task] for a task of the group: fn, the options and the task_id
follow the same rules,
+// and the group_id of g then prefixes the task_id, unless the TaskGroupSpec
of g sets
+// PrefixGroupID to false. The group_id also prefixes a TaskID that a TaskSpec
sets:
+//
+// group := dag.TaskGroup("transform")
+// group.Task(cleanRows) //
transform.cleanRows
+// group.Task(validateRows, airflow.TaskSpec{TaskID: "v"}) // transform.v
+//
+// [Inputs] takes any task of the Dag, inside the group or not.
+//
+// Task panics for the reasons that DagRef.Task lists, and names the task by
the task_id that the
+// group gives it. It also panics if g is not the TaskGroupRef that
DagRef.TaskGroup or
+// TaskGroupRef.TaskGroup returned, for example a copy of it.
+func (g *TaskGroupRef) Task(fn any, opts ...TaskOption) *TaskRef {
+ const method = "airflow.TaskGroupRef.Task"
+ return g.groupDag(method).addTask(method, g, fn, opts, nil)
+}
+
+// If adds a condition to the Dag of g, inside g, and returns it. It is
[DagRef.If] for a
+// condition of the group, and the group_id of g prefixes the task_id of the
condition task as
+// [TaskGroupRef.Task] describes. [IfRef.Then] and [IfRef.Else] take any task
of the Dag, inside
+// the group or not.
+//
+// If panics for the reasons that DagRef.If and TaskGroupRef.Task list.
+func (g *TaskGroupRef) If(fn any, opts ...TaskOption) *IfRef {
+ const method = "airflow.TaskGroupRef.If"
+ return g.groupDag(method).addIf(method, g, fn, opts)
+}
+
+func (*TaskGroupRef) node() {}
+
+// Before makes the group an upstream of every node, which is Python's
+// transform >> [load, report]:
+//
+// transform.Before(loaded, reported)
+//
+// It is [TaskRef.Before] with the group at one end. A node can be a task or
another group, and
+// Before returns the nodes it was given as one Node, so a chain reads as it
does for a task:
+// extracted.Before(transform).Before(loaded) is extract >> transform >> load.
+//
+// An edge from a group stands for an edge from each of its last tasks, and an
edge to a group for
+// an edge to each of its first tasks. The first tasks of a group are the
tasks in it, nested
+// groups included, that no edge reaches from another task in the group, and
its last tasks are
+// those that no edge leaves for another task in the group, as Python's
TaskGroup.get_roots and
+// TaskGroup.get_leaves find them.
+//
+// [BundleRef.Register] expands the group edges one at a time, in the order
they were first
+// declared. Each expansion reads every task, every edge declared between two
tasks, and the task
+// edges that earlier group edges expanded into. So an author can add the
tasks of a group, and
+// the edges between them, after putting the group on an edge. The order of
two group edges
+// matters when one of them is inside a group that the other reaches: an edge
between two groups
+// inside transform changes the first and last tasks of transform only for the
edges to and from
+// transform declared after it. Python works out a group edge when >> runs, so
the two agree when
+// a Dag declares its group edges after its tasks and the edges between them,
every group at an end
+// of a group edge holds a task, and no [Label] sits on an edge whose receiver
is inside a task
+// group.
+//
+// Register steps over a group that holds no task: an edge to the group
continues along each edge
+// from it, and an edge from the group continues back along each edge to it,
whenever those were
+// declared. So extracted.Before(empty).Before(loaded) runs load after
extract, as Python's
+// extract >> empty >> load does. What Python does with an empty group depends
on how and when its
+// edges are declared: load << empty << extract adds no edge, and an empty
group with no task
+// before it falls back to the last tasks of the group that holds it, or of
the whole Dag.
+//
+// A [Label] on an edge to or from a group labels that edge, and none of the
edges between tasks
+// that it stands for. Depending on which end is the receiver and which groups
hold the ends,
+// Python labels those edges as well, as extract >> Label("rows") >> transform
does when no group
+// holds extract, or replaces the receiver with a group that holds it.
+//
+// Before panics for the reasons that TaskRef.Before lists. It also panics if:
+// - g or a node is a *TaskGroupRef that DagRef.TaskGroup or
TaskGroupRef.TaskGroup did not
+// return, such as a copy of one
+// - a node is g itself, a task or a group inside g, or a group that holds g
+func (g *TaskGroupRef) Before(nodes ...Node) Node { return declareEdges(g,
nodes, dirBefore) }
+
+// After makes the group a downstream of every node, which is Python's
transform << extract:
+//
+// transform.After(extracted)
+//
+// It is Before with the direction reversed, and it panics for the same
reasons.
+func (g *TaskGroupRef) After(nodes ...Node) Node { return declareEdges(g,
nodes, dirAfter) }
+
+// groupDag returns the Dag of g. It panics if g is nil or zero, which no
method of g can serve.
+// The Dag checks that g is the TaskGroupRef it returned once it holds its
lock.
+func (g *TaskGroupRef) groupDag(method string) *DagRef {
+ if g == nil || g.dag == nil {
+ panic(method + ": DagRef.TaskGroup or TaskGroupRef.TaskGroup
did not return " +
+ "the *airflow.TaskGroupRef")
+ }
+ return g.dag
+}
+
+// checkGroupLocked panics unless group is nil or a group that d added. A copy
of a TaskGroupRef
+// has the dag and group_id of the original, so only the identity of the
pointer tells them apart.
+// The caller holds d.mu.
+func (d *DagRef) checkGroupLocked(method string, group *TaskGroupRef) {
+ if group != nil && d.groupsByID[group.groupID] != group {
+ panic(fmt.Sprintf(
+ "%s: Dag %q got a *airflow.TaskGroupRef that
DagRef.TaskGroup or "+
+ "TaskGroupRef.TaskGroup did not return",
+ method, d.dagID,
+ ))
+ }
+}
+
+// childID returns the ID that a task or group added through g with the given
ID gets. g is nil
+// for one added to the Dag itself.
+func (g *TaskGroupRef) childID(id string) string {
+ if g == nil || (g.spec.PrefixGroupID != nil && !*g.spec.PrefixGroupID) {
+ return id
+ }
+ return g.groupID + "." + id
+}
+
+// holds reports whether node is inside g, at any depth. A group does not hold
itself.
+func (g *TaskGroupRef) holds(node nodeEndpoint) bool {
+ for parent := node.container(); parent != nil; parent = parent.parent {
+ if parent == g {
+ return true
+ }
+ }
+ return false
+}
+
+// addGroup adds a task group for DagRef.TaskGroup and TaskGroupRef.TaskGroup.
parent is the group
+// it is added through, and nil for DagRef.TaskGroup.
+func (d *DagRef) addGroup(
+ method string, parent *TaskGroupRef, groupID string, spec
[]TaskGroupSpec,
+) *TaskGroupRef {
+ if len(spec) > 1 {
+ panic(fmt.Sprintf(
+ "%s: task group %q of Dag %q got %d
airflow.TaskGroupSpec values; "+
+ "set all of the group's attributes in one
TaskGroupSpec",
+ method, parent.childID(groupID), d.dagID, len(spec),
+ ))
+ }
+ if err := checkGroupID(groupID); err != nil {
+ panic(fmt.Sprintf("%s: Dag %q: %v", method, d.dagID, err))
+ }
+
+ d.mu.Lock()
+ defer d.mu.Unlock()
+
+ if d.registered {
+ panic(fmt.Sprintf(
+ "%s: Dag %q has already been registered; add every task
group before Register",
+ method, d.dagID,
+ ))
+ }
+ d.checkGroupLocked(method, parent)
+ fullID := parent.childID(groupID)
+ for _, id := range []string{fullID, fullID + upstreamJoinSuffix, fullID
+ downstreamJoinSuffix} {
+ if taken := d.describeIDLocked(id); taken != "" {
+ panic(fmt.Sprintf(
+ "%s: Dag %q cannot add task group %q, because
%s already takes the ID %q; "+
+ "pass another group_id",
+ method, d.dagID, fullID, taken, id,
+ ))
+ }
+ }
+
+ group := &TaskGroupRef{dag: d, parent: parent, groupID: fullID}
+ if len(spec) == 1 {
+ group.spec = copySpec(spec[0])
+ }
+ if d.groupsByID == nil {
+ d.groupsByID = make(map[string]*TaskGroupRef)
+ }
+ d.groupsByID[fullID] = group
+ d.groups = append(d.groups, group)
+ if parent != nil {
+ parent.children = append(parent.children, group)
+ }
+ return group
+}
+
+// checkGroupID checks a group_id as Python's validate_group_key does. Python
matches the ID
+// against ^[\w-]+$, and its \w takes the characters for which str.isalnum is
true, and the
+// underscore. Python's $ also lets a trailing newline through, and
checkGroupID does not.
+func checkGroupID(groupID string) error {
+ if groupID == "" {
+ return fmt.Errorf("a task group needs a group_id, and got an
empty one")
+ }
+ if !utf8.ValidString(groupID) {
+ return fmt.Errorf("group_id %q is not valid UTF-8", groupID)
+ }
+ if n := utf8.RuneCountInString(groupID); n > groupIDMaxLength {
+ return fmt.Errorf(
+ "group_id %q has %d characters; a group_id has at most
%d",
+ groupID, n, groupIDMaxLength,
+ )
+ }
+ for _, r := range groupID {
+ if isWordRune(r) || r == '-' {
+ continue
+ }
+ var hint string
+ if r == '.' {
+ hint = "; nest a group in another with
TaskGroupRef.TaskGroup"
+ }
+ return fmt.Errorf(
+ "group_id %q holds %q; a group_id holds only letters,
digits, underscores and dashes%s",
+ groupID, r, hint,
+ )
+ }
+ return nil
+}
+
+// describeIDLocked names what takes id in the namespace that the tasks and
task groups of d
+// share, and returns the empty string when id is free. The caller holds d.mu.
+func (d *DagRef) describeIDLocked(id string) string {
+ if _, ok := d.tasksByID[id]; ok {
+ return fmt.Sprintf("task %q", id)
+ }
+ if _, ok := d.groupsByID[id]; ok {
+ return fmt.Sprintf("task group %q", id)
+ }
+ for _, suffix := range []string{upstreamJoinSuffix,
downstreamJoinSuffix} {
+ if groupID, ok := strings.CutSuffix(id, suffix); ok &&
d.groupsByID[groupID] != nil {
+ return fmt.Sprintf("a join node of task group %q",
groupID)
+ }
+ }
+ return ""
+}
+
+// groupEdge is an edge that has a task group at one end or both. It stands
for the edges between
+// the tasks of its ends, which [DagRef.expandGroupEdgesLocked] works out at
registration.
+type groupEdge struct {
+ upstream, downstream nodeEndpoint
+}
+
+// addGroupEdgeLocked records one edge of d that has a task group at one end
or both. The caller
+// holds d.mu, and settles the label of an edge that d already has with
[mergeLabel] first.
+func (d *DagRef) addGroupEdgeLocked(upstream, downstream nodeEndpoint, label
string) {
+ if d.groupEdgeLabels == nil {
+ d.groupEdgeLabels = make(map[edgeKey]string)
+ }
+ key := edgeKey{upstream: upstream.id(), downstream: downstream.id()}
+ if _, exists := d.groupEdgeLabels[key]; !exists {
+ d.groupEdges = append(d.groupEdges, groupEdge{
+ upstream: nodeEndpoint{task: upstream.task, group:
upstream.group},
+ downstream: nodeEndpoint{task: downstream.task, group:
downstream.group},
+ })
+ }
+ d.groupEdgeLabels[key] = label
+}
+
+// expandGroupEdgesLocked records the edges between tasks that the group edges
of d stand for, and
+// returns the keys of the edges it added. The caller holds d.mu.
+//
+// It expands the group edges in the order they were first declared, each from
the Dag as it stands
+// by then: every task, every edge between two tasks, and the edges that the
group edges before it
+// added. Stepping over a group with no task follows every group edge of that
group, whenever it
+// was declared. The TypeScript SDK expands the group edges of a Dag it
serializes by the same
+// rule, so that a Dag built alike in both SDKs ends up with the same task
edges.
+func (d *DagRef) expandGroupEdgesLocked() []expandedEdge {
+ expansion := newGroupExpansion(d.groupEdges)
+ var added []expandedEdge
+ for _, edge := range d.groupEdges {
+ from := edgeKey{upstream: edge.upstream.id(), downstream:
edge.downstream.id()}
+ upstreams := expansion.tasksAt(edge.upstream, false)
+ downstreams := expansion.tasksAt(edge.downstream, true)
+ for _, upstream := range upstreams {
+ for _, downstream := range downstreams {
+ key := edgeKey{upstream: upstream.taskID,
downstream: downstream.taskID}
+ if _, exists := d.edgeLabels[key]; exists {
+ continue
+ }
+ // The label of a group edge stays on the group
edge, as TaskGroupRef.Before says.
+ // Stepping over a group that holds no task can
lead back to the task it started
+ // from, as extract >> empty >> extract does,
and the cycle check reports that.
+ d.addEdgeLocked(upstream, downstream, "")
+ added = append(added, expandedEdge{key: key,
from: from})
+ }
+ }
+ }
+ return added
+}
+
+// expandedEdge is an edge between tasks that expandGroupEdgesLocked added,
with the group edge
+// that it stands for.
+type expandedEdge struct {
+ key, from edgeKey
+}
+
+// groupEdgesOn names the group edges that the expanded edges on cycle stand
for, in the order the
+// cycle runs through them. cycle is a list of task_ids that ends with the one
it starts from.
+func groupEdgesOn(cycle []string, expanded []expandedEdge) []string {
+ from := make(map[edgeKey]edgeKey, len(expanded))
+ for _, edge := range expanded {
+ from[edge.key] = edge.from
+ }
+ var names []string
+ named := make(map[edgeKey]bool)
+ for i := 0; i+1 < len(cycle); i++ {
+ group, ok := from[edgeKey{upstream: cycle[i], downstream:
cycle[i+1]}]
+ if ok && !named[group] {
+ named[group] = true
+ names = append(names, group.upstream+" ->
"+group.downstream)
+ }
+ }
+ return names
+}
+
+// removeEdgesLocked takes the edges that expandGroupEdgesLocked added back
out of d. Registration
+// calls it when it rejects the Dag, so that the Dag holds only the edges its
author declared. The
+// caller holds d.mu.
+func (d *DagRef) removeEdgesLocked(expanded []expandedEdge) {
+ for _, edge := range expanded {
+ key := edge.key
+ upstream, downstream := d.tasksByID[key.upstream],
d.tasksByID[key.downstream]
+ upstream.downstreams = slices.DeleteFunc(upstream.downstreams,
func(task *TaskRef) bool {
+ return task == downstream
+ })
+ downstream.upstreams = slices.DeleteFunc(downstream.upstreams,
func(task *TaskRef) bool {
+ return task == upstream
+ })
+ delete(d.edgeLabels, key)
+ }
+}
+
+// groupExpansion holds what expandGroupEdgesLocked needs besides the edges
between tasks: the
+// tasks in each group, which no expansion changes, and the group edges of
each group.
+type groupExpansion struct {
+ members map[*TaskGroupRef][]*TaskRef
+ inside map[*TaskGroupRef]map[*TaskRef]bool
+ // upstreams holds, for each group, the nodes whose group edges reach
it, and downstreams the
+ // nodes that its group edges reach, in the order the edges were first
declared.
+ upstreams, downstreams map[*TaskGroupRef][]nodeEndpoint
+}
+
+func newGroupExpansion(edges []groupEdge) *groupExpansion {
+ expansion := &groupExpansion{
+ members: make(map[*TaskGroupRef][]*TaskRef),
+ inside: make(map[*TaskGroupRef]map[*TaskRef]bool),
+ upstreams: make(map[*TaskGroupRef][]nodeEndpoint),
+ downstreams: make(map[*TaskGroupRef][]nodeEndpoint),
+ }
+ for _, edge := range edges {
+ if group := edge.downstream.group; group != nil {
+ expansion.upstreams[group] =
append(expansion.upstreams[group], edge.upstream)
+ }
+ if group := edge.upstream.group; group != nil {
+ expansion.downstreams[group] =
append(expansion.downstreams[group], edge.downstream)
+ }
+ }
+ return expansion
+}
+
+// tasksAt returns the tasks that an end of a group edge stands for. A task
stands for itself, and
+// a group for its first tasks when first is true, as the downstream end of an
edge, and for its
+// last tasks otherwise. A group that has none stands for the tasks beyond it.
+func (e *groupExpansion) tasksAt(end nodeEndpoint, first bool) []*TaskRef {
+ if end.group == nil {
+ return []*TaskRef{end.task}
+ }
+ if found := e.ends(end.group, first); len(found) > 0 {
+ return found
+ }
+ var found []*TaskRef
+ e.tasksBeyond(end.group, first, make(map[*TaskGroupRef]bool), &found)
+ return found
+}
+
+// tasksBeyond appends to found the tasks that an edge reaches through group
when group has no
+// first or last task to stop at. An edge into group continues along the group
edges from group
+// when first is true, and an edge out of group continues back along the group
edges into it
+// otherwise, as far as a task or a group with tasks to stop at. The
TypeScript SDK steps over
+// such a group the same way, and Python's extract >> empty >> load also runs
load after extract.
+//
+// A group has no first task when it holds no task, or when each of its tasks
has an upstream task
+// inside it, which only a cycle inside the group allows, and likewise for
last tasks with
+// downstream tasks. Registration rejects such a cycle anyway. seen holds the
groups already
+// stepped over, so that a loop of group edges between such groups ends.
+func (e *groupExpansion) tasksBeyond(
+ group *TaskGroupRef, first bool, seen map[*TaskGroupRef]bool, found
*[]*TaskRef,
+) {
+ if seen[group] {
+ return
+ }
+ seen[group] = true
+ next := e.upstreams[group]
+ if first {
+ next = e.downstreams[group]
+ }
+ for _, node := range next {
+ if node.group == nil {
+ *found = append(*found, node.task)
+ continue
+ }
+ if ends := e.ends(node.group, first); len(ends) > 0 {
+ *found = append(*found, ends...)
+ continue
+ }
+ e.tasksBeyond(node.group, first, seen, found)
+ }
+}
+
+// ends returns the first tasks of group when first is true, and its last
tasks otherwise, as the
+// Dag stands: the tasks in the group that no edge reaches from another task
in it, or that no edge
+// leaves for another task in it.
+func (e *groupExpansion) ends(group *TaskGroupRef, first bool) []*TaskRef {
+ members, inside := e.membersOf(group)
+ isInside := func(task *TaskRef) bool { return inside[task] }
+ var found []*TaskRef
+ for _, task := range members {
+ neighbours := task.downstreams
+ if first {
+ neighbours = task.upstreams
+ }
+ if !slices.ContainsFunc(neighbours, isInside) {
+ found = append(found, task)
+ }
+ }
+ return found
+}
+
+// membersOf returns the tasks in group, nested groups included, as a slice
and as a set. It lists
+// the tasks of a group before those of the groups nested in it, level by
level.
+func (e *groupExpansion) membersOf(group *TaskGroupRef) ([]*TaskRef,
map[*TaskRef]bool) {
+ if inside, ok := e.inside[group]; ok {
+ return e.members[group], inside
+ }
+ var members []*TaskRef
+ pending := []*TaskGroupRef{group}
+ for i := 0; i < len(pending); i++ {
+ for _, child := range pending[i].children {
+ if task, ok := child.(*TaskRef); ok {
+ members = append(members, task)
+ }
+ }
+ for _, child := range pending[i].children {
+ if nested, ok := child.(*TaskGroupRef); ok {
+ pending = append(pending, nested)
+ }
+ }
+ }
+ inside := make(map[*TaskRef]bool, len(members))
+ for _, task := range members {
+ inside[task] = true
+ }
+ e.members[group], e.inside[group] = members, inside
+ return members, inside
+}
diff --git a/go-sdk/airflow/task_group_test.go
b/go-sdk/airflow/task_group_test.go
new file mode 100644
index 00000000000..20a236545b3
--- /dev/null
+++ b/go-sdk/airflow/task_group_test.go
@@ -0,0 +1,1274 @@
+// 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"
+ "maps"
+ "slices"
+ "strings"
+ "sync"
+ "testing"
+
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+func cleanRows(Context) error { return nil }
+func validateRows(Context) error { return nil }
+
+// groupTask adds a task that passes no data to group, which is what an
order-only edge connects.
+func groupTask(t *testing.T, group *TaskGroupRef, taskID string) *TaskRef {
+ t.Helper()
+ return group.Task(ping, TaskSpec{TaskID: taskID})
+}
+
+func assertChildren(t *testing.T, group *TaskGroupRef, want ...Node) {
+ t.Helper()
+ require.Len(t, group.children, len(want))
+ for i := range want {
+ assert.Same(t, want[i], group.children[i], "child %d", i)
+ }
+}
+
+func assertGroupEdgeLabel(t *testing.T, dag *DagRef, upstream, downstream,
want string) {
+ t.Helper()
+ label, declared := dag.groupEdgeLabels[edgeKey{upstream: upstream,
downstream: downstream}]
+ require.True(t, declared, "the group edge %s -> %s was not declared",
upstream, downstream)
+ assert.Equal(t, want, label)
+}
+
+func TestTaskGroupPrefixesTheTaskIDs(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("transform")
+
+ read := group.Task(readRows)
+ counted := group.Task(countRows, Inputs(read), TaskSpec{TaskID:
"count"})
+
+ assert.Equal(t, "transform.readRows", read.taskID)
+ assert.Equal(t, "transform.count", counted.taskID)
+ assert.Same(t, group, read.group)
+ assertTasks(t, counted.inputs, read)
+ assertChildren(t, group, read, counted)
+ assertTasks(t, dag.tasks, read, counted)
+ assert.Same(t, read, dag.tasksByID["transform.readRows"])
+}
+
+func TestTaskGroupsNest(t *testing.T) {
+ dag := Dag("etl")
+ transform := dag.TaskGroup("transform")
+ cleaned := transform.Task(cleanRows)
+ checks := transform.TaskGroup("checks")
+ nulls := groupTask(t, checks, "nulls")
+
+ assert.Equal(t, "transform.checks", checks.groupID)
+ assert.Equal(t, "transform.checks.nulls", nulls.taskID)
+ assert.Same(t, transform, checks.parent)
+ assert.Nil(t, transform.parent)
+ // The groups form a tree, which is what a serialized Dag carries in
task_group.children.
+ assertChildren(t, transform, cleaned, checks)
+ assertChildren(t, checks, nulls)
+ assert.Equal(t, []*TaskGroupRef{transform, checks}, dag.groups)
+ assert.Same(t, checks, dag.groupsByID["transform.checks"])
+}
+
+// TestPrefixGroupIDFalseKeepsTheIDsOfWhatTheGroupHolds pins Python's rule: a
group decides
+// whether its own ID prefixes what is added through it, and its parent
decides whether the
+// parent's ID prefixes the group's.
+func TestPrefixGroupIDFalseKeepsTheIDsOfWhatTheGroupHolds(t *testing.T) {
+ dag := Dag("etl")
+ noPrefix := false
+ transform := dag.TaskGroup("transform")
+ checks := transform.TaskGroup("checks", TaskGroupSpec{PrefixGroupID:
&noPrefix})
+ nulls := groupTask(t, checks, "nulls")
+ inner := checks.TaskGroup("inner")
+
+ assert.Equal(t, "transform.checks", checks.groupID)
+ assert.Equal(t, "nulls", nulls.taskID)
+ assert.Equal(t, "inner", inner.groupID)
+}
+
+// TestPrefixesFollowTheParentOfEachGroup pins Python's rule across three
levels: mid keeps the
+// IDs of what it holds as written, so inner is not prefixed, while inner
prefixes its own task.
+func TestPrefixesFollowTheParentOfEachGroup(t *testing.T) {
+ dag := Dag("etl")
+ noPrefix := false
+ outer := dag.TaskGroup("outer")
+ mid := outer.TaskGroup("mid", TaskGroupSpec{PrefixGroupID: &noPrefix})
+ inner := mid.TaskGroup("inner")
+
+ assert.Equal(t, "outer.mid", mid.groupID)
+ assert.Equal(t, "inner", inner.groupID)
+ assert.Equal(t, "inner.rows", groupTask(t, inner, "rows").taskID)
+ assert.Equal(t, "mid_rows", groupTask(t, mid, "mid_rows").taskID)
+ assert.Equal(t, "outer.rows", groupTask(t, outer, "rows").taskID)
+}
+
+// TestTaskGroupNamesATaskByItsPrefixedTaskID pins that an error raised before
the task_id is
+// settled names the task by the task_id that the group gives it.
+func TestTaskGroupNamesATaskByItsPrefixedTaskID(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("transform")
+
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.Task: task "transform.v" of Dag "etl":
got more than one `+
+ `airflow.TaskSpec; set all of the task's attributes in
one TaskSpec`,
+ func() { group.Task(ping, TaskSpec{TaskID: "v"}, TaskSpec{}) },
+ )
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.Task: task "transform.ping" of Dag "etl":
got more than one `+
+ `airflow.Inputs; pass all of the task's inputs to one
airflow.Inputs`,
+ func() { group.Task(ping, Inputs(), Inputs()) },
+ )
+}
+
+// TestTaskChecksTheTaskIDWithItsPrefixes pins Python's validate_key, which
checks a task_id once
+// the group_ids of its task groups prefix it: each group_id can be valid
while the task_id they
+// make is too long.
+func TestTaskChecksTheTaskIDWithItsPrefixes(t *testing.T) {
+ long := strings.Repeat("g", 200)
+ for _, tc := range []struct {
+ name string
+ build func(dag *DagRef)
+ want string
+ }{
+ {
+ name: "too long once prefixed",
+ build: func(dag *DagRef) {
+ dag.TaskGroup(long).TaskGroup(long).Task(ping,
TaskSpec{TaskID: "t"})
+ },
+ want: fmt.Sprintf(
+ `airflow.TaskGroupRef.Task: Dag "etl": task_id
%q has 403 characters, `+
+ `counting the group_ids that prefix it,
and a task_id has at most 250; `+
+ `shorten a group_id, or set a shorter
task_id with airflow.TaskSpec{TaskID: ...}`,
+ long+"."+long+".t",
+ ),
+ },
+ {
+ name: "too long without a group",
+ build: func(dag *DagRef) { dag.Task(ping,
TaskSpec{TaskID: strings.Repeat("t", 251)}) },
+ want: fmt.Sprintf(
+ `airflow.DagRef.Task: Dag "etl": task_id %q has
251 characters, and a task_id `+
+ `has at most 250; set a shorter one
with airflow.TaskSpec{TaskID: ...}`,
+ strings.Repeat("t", 251),
+ ),
+ },
+ {
+ name: "a character Python rejects",
+ build: func(dag *DagRef) { dag.Task(ping,
TaskSpec{TaskID: "clean rows"}) },
+ want: `airflow.DagRef.Task: Dag "etl": task_id "clean
rows" holds ' ', and a task_id ` +
+ `holds only letters, digits, underscores,
dashes and dots; set another one with ` +
+ `airflow.TaskSpec{TaskID: ...}`,
+ },
+ {
+ name: "bytes that are not UTF-8",
+ build: func(dag *DagRef) { dag.Task(ping,
TaskSpec{TaskID: "rows\xff"}) },
+ want: `airflow.DagRef.Task: Dag "etl": task_id
"rows\xff" is not valid UTF-8; ` +
+ `set another one with airflow.TaskSpec{TaskID:
...}`,
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ dag := Dag("etl")
+
+ assert.PanicsWithValue(t, tc.want, func() {
tc.build(dag) })
+ assert.Empty(t, dag.tasks)
+ })
+ }
+}
+
+func TestTaskTakesTheTaskIDsPythonTakes(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup(strings.Repeat("g", 200))
+
+ // The last one makes a task_id of exactly 250 characters.
+ for _, taskID := range []string{"extract.rows", "clean-rows_2", "清理",
strings.Repeat("t", 49)} {
+ assert.NotPanics(t, func() { group.Task(ping, TaskSpec{TaskID:
taskID}) }, taskID)
+ }
+}
+
+func TestPrefixGroupIDTrueKeepsThePrefix(t *testing.T) {
+ dag := Dag("etl")
+ prefix := true
+ group := dag.TaskGroup("transform", TaskGroupSpec{PrefixGroupID:
&prefix})
+
+ assert.Equal(t, "transform.cleanRows", group.Task(cleanRows).taskID)
+}
+
+func TestTaskGroupAddsAConditionInsideTheGroup(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("gate")
+ loaded := dag.Task(load)
+
+ condition := group.If(isReady).Then(loaded)
+
+ assert.Equal(t, "gate.isReady", condition.task.taskID)
+ assert.Same(t, group, condition.task.group)
+ assertChildren(t, group, condition.task)
+ assertTasks(t, loaded.upstreams, condition.task)
+}
+
+func TestTaskGroupIfRejectsATriggerDagRun(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("gate")
+
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.If: Dag "etl": fn comes from
airflow.TriggerDagRun, `+
+ `but a condition function is a Go function that returns
(bool, error)`,
+ func() { group.If(TriggerDagRun(TriggerDagRunSpec{DagID:
"other"})) },
+ )
+}
+
+func TestTaskGroupTakesAtMostOneTaskGroupSpec(t *testing.T) {
+ dag := Dag("etl")
+
+ assert.PanicsWithValue(t,
+ `airflow.DagRef.TaskGroup: task group "transform" of Dag "etl"
got 2 `+
+ `airflow.TaskGroupSpec values; set all of the group's
attributes in one TaskGroupSpec`,
+ func() { dag.TaskGroup("transform", TaskGroupSpec{},
TaskGroupSpec{}) },
+ )
+ assert.Empty(t, dag.groups)
+
+ // The message names a nested group by the group_id it would get.
+ transform := dag.TaskGroup("transform")
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.TaskGroup: task group "transform.checks"
of Dag "etl" got 2 `+
+ `airflow.TaskGroupSpec values; set all of the group's
attributes in one TaskGroupSpec`,
+ func() { transform.TaskGroup("checks", TaskGroupSpec{},
TaskGroupSpec{}) },
+ )
+ assert.Empty(t, transform.children)
+}
+
+func TestTaskGroupCopiesItsSpec(t *testing.T) {
+ dag := Dag("etl")
+ noPrefix := false
+ group := dag.TaskGroup("transform", TaskGroupSpec{PrefixGroupID:
&noPrefix, Tooltip: "rows"})
+
+ noPrefix = true
+
+ assert.False(t, *group.spec.PrefixGroupID)
+ assert.Equal(t, "rows", group.spec.Tooltip)
+ assert.Equal(t, "cleanRows", group.Task(cleanRows).taskID)
+}
+
+func TestTaskGroupChecksTheGroupID(t *testing.T) {
+ for _, tc := range []struct {
+ name, groupID, want string
+ }{
+ {
+ name: "empty",
+ groupID: "",
+ want: "a task group needs a group_id, and got an
empty one",
+ },
+ {
+ name: "a dot",
+ groupID: "transform.checks",
+ want: `group_id "transform.checks" holds '.'; a
group_id holds only letters, ` +
+ `digits, underscores and dashes; nest a group
in another with TaskGroupRef.TaskGroup`,
+ },
+ {
+ name: "a space",
+ groupID: "clean rows",
+ want: `group_id "clean rows" holds ' '; a group_id
holds only letters, ` +
+ `digits, underscores and dashes`,
+ },
+ {
+ name: "too long",
+ groupID: strings.Repeat("g", 201),
+ want: fmt.Sprintf(
+ "group_id %q has 201 characters; a group_id has
at most 200",
+ strings.Repeat("g", 201),
+ ),
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ dag := Dag("etl")
+
+ assert.PanicsWithValue(t,
+ `airflow.DagRef.TaskGroup: Dag "etl": `+tc.want,
+ func() { dag.TaskGroup(tc.groupID) },
+ )
+ assert.Empty(t, dag.groups)
+ })
+ }
+}
+
+// TestTaskGroupTakesTheGroupIDsPythonTakes pins that a letter or a digit
outside ASCII is a word
+// character, as it is for Python's \w, and that the length is counted in
characters.
+func TestTaskGroupTakesTheGroupIDsPythonTakes(t *testing.T) {
+ dag := Dag("etl")
+
+ for _, groupID := range []string{"clean-rows_2", "清理", "étape٣",
strings.Repeat("清", 200)} {
+ assert.NotPanics(t, func() { dag.TaskGroup(groupID) }, groupID)
+ }
+}
+
+func TestTaskGroupsAndTasksShareOneNamespace(t *testing.T) {
+ for _, tc := range []struct {
+ name string
+ build func(dag *DagRef)
+ want string
+ }{
+ {
+ name: "a group with the ID of a group",
+ build: func(dag *DagRef) { dag.TaskGroup("transform");
dag.TaskGroup("transform") },
+ want: `airflow.DagRef.TaskGroup: Dag "etl" cannot add
task group "transform", ` +
+ `because task group "transform" already takes
the ID "transform"; ` +
+ `pass another group_id`,
+ },
+ {
+ name: "a group with the ID of a task",
+ build: func(dag *DagRef) { dag.Task(load);
dag.TaskGroup("load") },
+ want: `airflow.DagRef.TaskGroup: Dag "etl" cannot add
task group "load", ` +
+ `because task "load" already takes the ID
"load"; pass another group_id`,
+ },
+ {
+ name: "a task with the ID of a group",
+ build: func(dag *DagRef) { dag.TaskGroup("load");
dag.Task(load) },
+ want: `airflow.DagRef.Task: Dag "etl" cannot add task
"load", ` +
+ `because task group "load" already takes the
ID; ` +
+ `set another task_id with
airflow.TaskSpec{TaskID: ...}`,
+ },
+ {
+ name: "a nested group with the ID of a task its prefix
makes",
+ build: func(dag *DagRef) {
+ group := dag.TaskGroup("transform")
+ group.Task(cleanRows)
+ group.TaskGroup("cleanRows")
+ },
+ want: `airflow.TaskGroupRef.TaskGroup: Dag "etl" cannot
add task group ` +
+ `"transform.cleanRows", because task
"transform.cleanRows" already takes the ID ` +
+ `"transform.cleanRows"; pass another group_id`,
+ },
+ {
+ name: "a task with the ID of a join node",
+ build: func(dag *DagRef) {
+ dag.TaskGroup("transform").Task(ping,
TaskSpec{TaskID: "upstream_join_id"})
+ },
+ want: `airflow.TaskGroupRef.Task: Dag "etl" cannot add
task ` +
+ `"transform.upstream_join_id", because a join
node of task group "transform" ` +
+ `already takes the ID; set another task_id with
airflow.TaskSpec{TaskID: ...}`,
+ },
+ {
+ name: "a group whose join node has the ID of a task",
+ build: func(dag *DagRef) {
+ dag.Task(ping, TaskSpec{TaskID:
"transform.downstream_join_id"})
+ dag.TaskGroup("transform")
+ },
+ want: `airflow.DagRef.TaskGroup: Dag "etl" cannot add
task group "transform", ` +
+ `because task "transform.downstream_join_id"
already takes the ID ` +
+ `"transform.downstream_join_id"; pass another
group_id`,
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ assert.PanicsWithValue(t, tc.want, func() {
tc.build(Dag("etl")) })
+ })
+ }
+}
+
+// TestTaskGroupsWithoutAPrefixShareTheirNamespace pins that a group which
keeps the IDs of its
+// tasks as written puts them in the namespace of the whole Dag.
+func TestTaskGroupsWithoutAPrefixShareTheirNamespace(t *testing.T) {
+ dag := Dag("etl")
+ noPrefix := false
+ dag.Task(load)
+ group := dag.TaskGroup("transform", TaskGroupSpec{PrefixGroupID:
&noPrefix})
+
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.Task: Dag "etl" already has a task
"load"; `+
+ `set another task_id with airflow.TaskSpec{TaskID:
...}`,
+ func() { group.Task(load) },
+ )
+ assert.Empty(t, group.children)
+}
+
+func TestTaskGroupMethodsRejectAGroupTheDagDidNotReturn(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("transform")
+ copied := *group
+
+ for _, tc := range []struct {
+ name string
+ call func()
+ want string
+ }{
+ {
+ name: "Task on a copy",
+ call: func() { copied.Task(cleanRows) },
+ want: `airflow.TaskGroupRef.Task: Dag "etl" got a
*airflow.TaskGroupRef ` +
+ `that DagRef.TaskGroup or
TaskGroupRef.TaskGroup did not return`,
+ },
+ {
+ name: "TaskGroup on a copy",
+ call: func() { copied.TaskGroup("checks") },
+ want: `airflow.TaskGroupRef.TaskGroup: Dag "etl" got a
*airflow.TaskGroupRef ` +
+ `that DagRef.TaskGroup or
TaskGroupRef.TaskGroup did not return`,
+ },
+ {
+ name: "If on a copy",
+ call: func() { copied.If(isReady) },
+ want: `airflow.TaskGroupRef.If: Dag "etl" got a
*airflow.TaskGroupRef ` +
+ `that DagRef.TaskGroup or
TaskGroupRef.TaskGroup did not return`,
+ },
+ {
+ name: "Task on a zero group",
+ call: func() { (&TaskGroupRef{}).Task(cleanRows) },
+ want: "airflow.TaskGroupRef.Task: DagRef.TaskGroup or
TaskGroupRef.TaskGroup " +
+ "did not return the *airflow.TaskGroupRef",
+ },
+ {
+ name: "TaskGroup on a nil group",
+ call: func() { (*TaskGroupRef)(nil).TaskGroup("checks")
},
+ want: "airflow.TaskGroupRef.TaskGroup: DagRef.TaskGroup
or TaskGroupRef.TaskGroup " +
+ "did not return the *airflow.TaskGroupRef",
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ assert.PanicsWithValue(t, tc.want, tc.call)
+ })
+ }
+ assert.Empty(t, group.children)
+ assert.Empty(t, dag.tasks)
+ assert.Equal(t, []*TaskGroupRef{group}, dag.groups)
+}
+
+func TestTaskGroupMethodsAfterRegisterPanic(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("transform")
+ group.Task(cleanRows)
+ Bundle().Register(dag)
+ registered := snapshot(dag)
+
+ assert.PanicsWithValue(t,
+ `airflow.DagRef.TaskGroup: Dag "etl" has already been
registered; `+
+ `add every task group before Register`,
+ func() { dag.TaskGroup("load") },
+ )
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.TaskGroup: Dag "etl" has already been
registered; `+
+ `add every task group before Register`,
+ func() { group.TaskGroup("checks") },
+ )
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.Task: Dag "etl" has already been
registered; `+
+ `add every task before Register`,
+ func() { group.Task(validateRows) },
+ )
+ assert.PanicsWithValue(t,
+ `airflow.TaskGroupRef.If: Dag "etl" has already been
registered; `+
+ `add every task before Register`,
+ func() { group.If(isReady) },
+ )
+ assert.Equal(t, registered, snapshot(dag))
+}
+
+func TestTaskGroupIsSafeForConcurrentUse(t *testing.T) {
+ const workers = 8
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+
+ var wg sync.WaitGroup
+ for i := range workers {
+ wg.Go(func() {
+ task := groupTask(t, group, fmt.Sprintf("task_%d", i))
+ nested := group.TaskGroup(fmt.Sprintf("group_%d", i))
+ groupTask(t, nested, "rows")
+ extracted.Before(Label(nested, fmt.Sprintf("to %d", i)))
+ nested.After(task)
+ })
+ }
+ wg.Wait()
+ Bundle().Register(dag)
+
+ assert.Len(t, group.children, 2*workers)
+ assert.Len(t, dag.tasks, 1+2*workers)
+ assert.Len(t, dag.groups, 1+workers)
+ assert.Len(t, dag.groupEdges, 2*workers)
+ assert.Len(t, extracted.downstreams, workers)
+}
+
+func TestTaskGroupRefIsANode(t *testing.T) {
+ var _ Node = (*TaskGroupRef)(nil)
+}
+
+// TestGroupBeforeATaskReachesTheLastTasksOfTheGroup pins Python's transform
>> load: the edge
+// runs from each task of the group that no other task of the group runs after.
+func TestGroupBeforeATaskReachesTheLastTasksOfTheGroup(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("transform")
+ cleaned := groupTask(t, group, "clean")
+ validated := groupTask(t, group, "validate")
+ audited := groupTask(t, group, "audit")
+ cleaned.Before(validated)
+ loaded := orderedTask(t, dag, "load")
+
+ group.Before(loaded)
+
+ // The edge stays an edge to the group until Register.
+ assert.Empty(t, loaded.upstreams)
+ assertGroupEdgeLabel(t, dag, "transform", "load", "")
+
+ Bundle().Register(dag)
+
+ assertTasks(t, loaded.upstreams, validated, audited)
+ assertTasks(t, cleaned.downstreams, validated)
+ assertEdgeLabel(t, dag, "transform.validate", "load", "")
+ assertEdgeLabel(t, dag, "transform.audit", "load", "")
+}
+
+func TestTaskBeforeAGroupReachesTheFirstTasksOfTheGroup(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ cleaned := groupTask(t, group, "clean")
+ validated := groupTask(t, group, "validate")
+ audited := groupTask(t, group, "audit")
+ cleaned.Before(validated)
+
+ group.After(extracted)
+ Bundle().Register(dag)
+
+ assertTasks(t, extracted.downstreams, cleaned, audited)
+ assertTasks(t, validated.upstreams, cleaned)
+}
+
+func TestGroupBeforeAGroupConnectsTheLastTasksToTheFirstTasks(t *testing.T) {
+ dag := Dag("etl")
+ transform := dag.TaskGroup("transform")
+ cleaned := groupTask(t, transform, "clean")
+ validated := groupTask(t, transform, "validate")
+ cleaned.Before(validated)
+ publish := dag.TaskGroup("publish")
+ loaded := groupTask(t, publish, "load")
+ reported := groupTask(t, publish, "report")
+
+ transform.Before(publish)
+ Bundle().Register(dag)
+
+ assertTasks(t, validated.downstreams, loaded, reported)
+ assertTasks(t, loaded.upstreams, validated)
+ assertTasks(t, reported.upstreams, validated)
+ assertGroupEdgeLabel(t, dag, "transform", "publish", "")
+}
+
+// TestAGroupEdgeReachesTasksAddedAfterIt pins why Register, and not the edge
verb, works out the
+// tasks of a group edge. In Python, transform >> load reaches only the tasks
the group had when >>
+// ran.
+func TestAGroupEdgeReachesTasksAddedAfterIt(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("transform")
+ loaded := orderedTask(t, dag, "load")
+ group.Before(loaded)
+
+ cleaned := groupTask(t, group, "clean")
+ validated := groupTask(t, group, "validate")
+ cleaned.Before(validated)
+ Bundle().Register(dag)
+
+ assertTasks(t, loaded.upstreams, validated)
+}
+
+func TestGroupEdgeVerbsReturnTheirArgumentSet(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ cleaned := groupTask(t, group, "clean")
+ loaded := orderedTask(t, dag, "load")
+
+ extracted.Before(group).Before(loaded)
+ Bundle().Register(dag)
+
+ assertTasks(t, cleaned.upstreams, extracted)
+ assertTasks(t, cleaned.downstreams, loaded)
+}
+
+// TestGroupEdgesExpandInTheOrderTheyWereDeclared pins the rule that the
TypeScript SDK applies as
+// well: an edge between two groups inside transform, declared after the edges
to and from
+// transform, does not change which tasks of transform those edges reach.
Python's
+// extract >> transform >> load followed by clean >> check gives the same five
edges.
+func TestGroupEdgesExpandInTheOrderTheyWereDeclared(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ loaded := orderedTask(t, dag, "load")
+ transform := dag.TaskGroup("transform")
+ clean := transform.TaskGroup("clean")
+ cleaned := groupTask(t, clean, "rows")
+ check := transform.TaskGroup("check")
+ checked := groupTask(t, check, "rows")
+
+ extracted.Before(transform).Before(loaded)
+ clean.Before(check)
+ Bundle().Register(dag)
+
+ assertTasks(t, extracted.downstreams, cleaned, checked)
+ assertTasks(t, cleaned.downstreams, loaded, checked)
+ assertTasks(t, checked.downstreams, loaded)
+}
+
+// TestAGroupEdgeInsideAGroupDeclaredFirstShapesTheGroupsEnds is the same Dag
with the edge between
+// the two inner groups declared first. It makes clean.rows the only first
task of transform and
+// check.rows the only last one.
+func TestAGroupEdgeInsideAGroupDeclaredFirstShapesTheGroupsEnds(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ loaded := orderedTask(t, dag, "load")
+ transform := dag.TaskGroup("transform")
+ clean := transform.TaskGroup("clean")
+ cleaned := groupTask(t, clean, "rows")
+ check := transform.TaskGroup("check")
+ checked := groupTask(t, check, "rows")
+
+ clean.Before(check)
+ extracted.Before(transform).Before(loaded)
+ Bundle().Register(dag)
+
+ assertTasks(t, extracted.downstreams, cleaned)
+ assertTasks(t, cleaned.downstreams, checked)
+ assertTasks(t, loaded.upstreams, checked)
+}
+
+// TestAGroupEdgeLeavesTheLabelOfATaskEdgeAlone pins that a group edge labels
none of the edges it
+// stands for, and so keeps the label of one that was declared on its own.
Python puts the group
+// edge's label on extract >> transform.validate as well, over "rows".
+func TestAGroupEdgeLeavesTheLabelOfATaskEdgeAlone(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ cleaned := groupTask(t, group, "clean")
+ validated := groupTask(t, group, "validate")
+ extracted.Before(Label(validated, "rows"))
+
+ extracted.Before(Label(group, "to transform"))
+ Bundle().Register(dag)
+
+ assertGroupEdgeLabel(t, dag, "extract", "transform", "to transform")
+ assertEdgeLabel(t, dag, "extract", "transform.clean", "")
+ assertEdgeLabel(t, dag, "extract", "transform.validate", "rows")
+ assertTasks(t, extracted.downstreams, validated, cleaned)
+}
+
+// TestALabelOnAGroupEdgeStaysOnTheGroupEdge pins that a label on an edge to
or from a group labels
+// that edge and none of the edges between tasks that it stands for, whichever
end the label wraps.
+// Python also labels those when the receiver is a task that no group holds, as
+// extract >> Label("x") >> transform does.
+func TestALabelOnAGroupEdgeStaysOnTheGroupEdge(t *testing.T) {
+ for _, tc := range []struct {
+ name string
+ declare func(extracted *TaskRef, transform,
publish *TaskGroupRef)
+ upstream, downstream string
+ }{
+ {
+ name: "a task before a labelled group",
+ declare: func(extracted *TaskRef, transform, _
*TaskGroupRef) {
+ extracted.Before(Label(transform, "x"))
+ },
+ upstream: "extract", downstream: "transform",
+ },
+ {
+ name: "a task after a labelled group",
+ declare: func(extracted *TaskRef, transform, _
*TaskGroupRef) {
+ extracted.After(Label(transform, "x"))
+ },
+ upstream: "transform", downstream: "extract",
+ },
+ {
+ name: "a group before a labelled task",
+ declare: func(extracted *TaskRef, transform, _
*TaskGroupRef) {
+ transform.Before(Label(extracted, "x"))
+ },
+ upstream: "transform", downstream: "extract",
+ },
+ {
+ name: "a group after a labelled task",
+ declare: func(extracted *TaskRef, transform, _
*TaskGroupRef) {
+ transform.After(Label(extracted, "x"))
+ },
+ upstream: "extract", downstream: "transform",
+ },
+ {
+ name: "a group before a labelled group",
+ declare: func(_ *TaskRef, transform, publish
*TaskGroupRef) {
+ transform.Before(Label(publish, "x"))
+ },
+ upstream: "transform", downstream: "publish",
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ transform := dag.TaskGroup("transform")
+ groupTask(t, transform, "clean")
+ groupTask(t, transform, "validate")
+ publish := dag.TaskGroup("publish")
+ groupTask(t, publish, "push")
+
+ tc.declare(extracted, transform, publish)
+ Bundle().Register(dag)
+
+ assertGroupEdgeLabel(t, dag, tc.upstream,
tc.downstream, "x")
+ require.Len(t, dag.edgeLabels, 2)
+ for key, label := range dag.edgeLabels {
+ assert.Empty(t, label, "%s -> %s",
key.upstream, key.downstream)
+ }
+ })
+ }
+}
+
+// TestALabelOnAnEdgeBetweenTasksOfDifferentGroupsKeepsTheTaskEdge pins that a
labelled edge
+// between tasks in different task groups stays an edge between the two tasks.
Python replaces the
+// receiver with its group, which here would run load after transform.validate
instead of after
+// transform.clean.
+func TestALabelOnAnEdgeBetweenTasksOfDifferentGroupsKeepsTheTaskEdge(t
*testing.T) {
+ dag := Dag("etl")
+ transform := dag.TaskGroup("transform")
+ cleaned := groupTask(t, transform, "clean")
+ validated := groupTask(t, transform, "validate")
+ cleaned.Before(validated)
+ loaded := orderedTask(t, dag, "load")
+
+ cleaned.Before(Label(loaded, "rows"))
+ Bundle().Register(dag)
+
+ assertEdgeLabel(t, dag, "transform.clean", "load", "rows")
+ assertTasks(t, loaded.upstreams, cleaned)
+ assert.Empty(t, dag.groupEdges)
+}
+
+func TestRedeclaringAGroupEdgeIsIdempotent(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ groupTask(t, group, "clean")
+
+ extracted.Before(Label(group, "rows"))
+ group.After(extracted)
+
+ require.Len(t, dag.groupEdges, 1)
+ assertGroupEdgeLabel(t, dag, "extract", "transform", "rows")
+}
+
+func TestGroupEdgeVerbsRejectAnEdgeInsideTheGroup(t *testing.T) {
+ dag := Dag("etl")
+ transform := dag.TaskGroup("transform")
+ cleaned := groupTask(t, transform, "clean")
+ checks := transform.TaskGroup("checks")
+ nulls := groupTask(t, checks, "nulls")
+
+ for _, tc := range []struct {
+ name string
+ call func()
+ want string
+ }{
+ {
+ name: "a group before itself",
+ call: func() { transform.Before(transform) },
+ want: `airflow.Node.Before: Dag "etl": task group
"transform" cannot depend on itself`,
+ },
+ {
+ name: "a group before a task it holds",
+ call: func() { transform.Before(cleaned) },
+ want: `airflow.Node.Before: Dag "etl": task
"transform.clean" is inside task group ` +
+ `"transform", so an edge cannot connect them; `
+
+ `an edge connects a group to a task or a group
outside it`,
+ },
+ {
+ name: "a task before the group that holds it",
+ call: func() { cleaned.Before(transform) },
+ want: `airflow.Node.Before: Dag "etl": task
"transform.clean" is inside task group ` +
+ `"transform", so an edge cannot connect them; `
+
+ `an edge connects a group to a task or a group
outside it`,
+ },
+ {
+ name: "a group after a group it holds",
+ call: func() { transform.After(checks) },
+ want: `airflow.Node.After: Dag "etl": task group
"transform.checks" is inside task ` +
+ `group "transform", so an edge cannot connect
them; ` +
+ `an edge connects a group to a task or a group
outside it`,
+ },
+ {
+ name: "a task deeper inside the group",
+ call: func() { nulls.After(transform) },
+ want: `airflow.Node.After: Dag "etl": task
"transform.checks.nulls" is inside task ` +
+ `group "transform", so an edge cannot connect
them; ` +
+ `an edge connects a group to a task or a group
outside it`,
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ assert.PanicsWithValue(t, tc.want, tc.call)
+ })
+ }
+ assert.Empty(t, dag.groupEdges)
+ assert.Empty(t, dag.edgeLabels)
+}
+
+// TestGroupEdgeVerbsTakeAnEdgeBetweenGroupsOfOneParent pins that only holding
the other end
+// rules an edge out: two groups inside one group, and a group and a task
beside it, can be
+// ordered.
+func TestGroupEdgeVerbsTakeAnEdgeBetweenGroupsOfOneParent(t *testing.T) {
+ dag := Dag("etl")
+ transform := dag.TaskGroup("transform")
+ cleaned := groupTask(t, transform, "clean")
+ checks := transform.TaskGroup("checks")
+ groupTask(t, checks, "nulls")
+ audits := transform.TaskGroup("audits")
+ groupTask(t, audits, "rows")
+
+ assert.NotPanics(t, func() { cleaned.Before(checks) })
+ assert.NotPanics(t, func() { checks.Before(audits) })
+ assertGroupEdgeLabel(t, dag, "transform.clean", "transform.checks", "")
+ assertGroupEdgeLabel(t, dag, "transform.checks", "transform.audits", "")
+}
+
+func TestGroupEdgeVerbsRejectAGroupTheDagDidNotReturn(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ copied := *group
+ other := Dag("other").TaskGroup("transform")
+
+ for _, tc := range []struct {
+ name string
+ call func()
+ want string
+ }{
+ {
+ name: "a copy",
+ call: func() { extracted.Before(&copied) },
+ want: `airflow.Node.Before: Dag "etl" got a
*airflow.TaskGroupRef ` +
+ `that DagRef.TaskGroup or
TaskGroupRef.TaskGroup did not return`,
+ },
+ {
+ name: "a zero group",
+ call: func() { extracted.Before(&TaskGroupRef{}) },
+ want: "airflow.Node.Before: got a *airflow.TaskGroupRef
that DagRef.TaskGroup or " +
+ "TaskGroupRef.TaskGroup did not return",
+ },
+ {
+ name: "a nil group",
+ call: func() { extracted.After((*TaskGroupRef)(nil)) },
+ want: "airflow.Node.After: nodes[0] is a nil
*airflow.TaskGroupRef",
+ },
+ {
+ name: "a nil group as the receiver",
+ call: func() { (*TaskGroupRef)(nil).Before(extracted) },
+ want: "airflow.Node.Before: the receiver is a nil
*airflow.TaskGroupRef",
+ },
+ {
+ name: "a group of another Dag",
+ call: func() { extracted.Before(other) },
+ want: `airflow.Node.Before: cannot declare an edge
between task "extract" of Dag ` +
+ `"etl" and task group "transform" of Dag
"other"; ` +
+ `an edge connects the tasks and task groups of
one Dag`,
+ },
+ {
+ name: "a nil group in a label",
+ call: func() { Label((*TaskGroupRef)(nil), "rows") },
+ want: "airflow.Label: node is a nil
*airflow.TaskGroupRef",
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ assert.PanicsWithValue(t, tc.want, tc.call)
+ })
+ }
+ assert.Empty(t, dag.groupEdges)
+}
+
+func TestGroupEdgeVerbsAfterRegisterPanic(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ groupTask(t, group, "clean")
+ Bundle().Register(dag)
+
+ assert.PanicsWithValue(t,
+ `airflow.Node.After: Dag "etl" has already been registered; `+
+ `declare every edge before Register`,
+ func() { group.After(extracted) },
+ )
+ assert.Empty(t, dag.groupEdges)
+}
+
+// TestAnEdgeStepsOverAGroupWithNoTask pins Python's extract >> staging >>
load for a staging that
+// holds no task, which the TypeScript SDK serializes the same way: load runs
after extract.
+func TestAnEdgeStepsOverAGroupWithNoTask(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ loaded := orderedTask(t, dag, "load")
+ staging := dag.TaskGroup("staging")
+ // A group that holds only a group with no task holds no task either.
+ staging.TaskGroup("rows")
+
+ extracted.Before(staging).Before(loaded)
+ Bundle().Register(dag)
+
+ assertTasks(t, extracted.downstreams, loaded)
+ assertTasks(t, loaded.upstreams, extracted)
+}
+
+func TestAnEdgeStepsOverGroupsWithNoTaskInARow(t *testing.T) {
+ dag := Dag("etl")
+ transform := dag.TaskGroup("transform")
+ cleaned := groupTask(t, transform, "clean")
+ validated := groupTask(t, transform, "validate")
+ first := dag.TaskGroup("first")
+ second := dag.TaskGroup("second")
+ publish := dag.TaskGroup("publish")
+ pushed := groupTask(t, publish, "push")
+
+ transform.Before(first).Before(second).Before(publish)
+ Bundle().Register(dag)
+
+ assertTasks(t, pushed.upstreams, cleaned, validated)
+ assertTasks(t, cleaned.downstreams, pushed)
+}
+
+// TestAnEdgeStepsOverALoopOfGroupsWithNoTask pins that stepping over groups
with no task ends
+// when their group edges loop back. Without that, Register would recurse
until the stack runs out.
+func TestAnEdgeStepsOverALoopOfGroupsWithNoTask(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ loaded := orderedTask(t, dag, "load")
+ first := dag.TaskGroup("first")
+ second := dag.TaskGroup("second")
+
+ extracted.Before(first).Before(second).Before(first)
+ second.Before(loaded)
+ Bundle().Register(dag)
+
+ assertTasks(t, extracted.downstreams, loaded)
+ assertTasks(t, loaded.upstreams, extracted)
+}
+
+func TestAnEdgeToAGroupWithNoTaskAndNothingBeyondItAddsNoEdge(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ staging := dag.TaskGroup("staging")
+
+ extracted.Before(staging)
+ Bundle().Register(dag)
+
+ assert.Empty(t, extracted.downstreams)
+ assertGroupEdgeLabel(t, dag, "extract", "staging", "")
+}
+
+func TestRegisterRejectsACycleThroughAGroupWithNoTask(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ staging := dag.TaskGroup("staging")
+ extracted.Before(staging).Before(extracted)
+
+ assert.PanicsWithValue(t,
+ `airflow.BundleRef.Register: the task dependencies of Dag "etl"
contain a cycle: `+
+ `extract -> extract, through the group edge extract ->
staging`,
+ func() { Bundle().Register(dag) },
+ )
+ assert.Empty(t, extracted.downstreams)
+ assert.Empty(t, extracted.upstreams)
+ assert.Empty(t, dag.edgeLabels)
+}
+
+func TestRegisterTakesAGroupWithNoTaskThatNoEdgeReaches(t *testing.T) {
+ dag := Dag("etl")
+ orderedTask(t, dag, "extract")
+ dag.TaskGroup("staging")
+
+ assert.NotPanics(t, func() { Bundle().Register(dag) })
+}
+
+func TestRegisterRejectsACycleThroughAGroup(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ cleaned := groupTask(t, group, "clean")
+ loaded := orderedTask(t, dag, "load")
+ extracted.Before(group).Before(loaded)
+ loaded.Before(extracted)
+
+ assert.PanicsWithValue(t,
+ `airflow.BundleRef.Register: the task dependencies of Dag "etl"
contain a cycle: `+
+ `extract -> transform.clean -> load -> extract, `+
+ `through the group edges extract -> transform,
transform -> load`,
+ func() { Bundle().Register(dag) },
+ )
+ // Register takes back out the edges that the group edges stand for, so
the Dag holds only the
+ // edges its author declared.
+ assert.False(t, dag.registered)
+ assertTasks(t, extracted.downstreams)
+ assertTasks(t, cleaned.upstreams)
+ assertTasks(t, cleaned.downstreams)
+ assertTasks(t, loaded.upstreams)
+ assertTasks(t, loaded.downstreams, extracted)
+ assert.NotContains(
+ t,
+ dag.edgeLabels,
+ edgeKey{upstream: "extract", downstream: "transform.clean"},
+ )
+}
+
+// TestRegisterReportsADeclaredCycleAsDeclared pins that a cycle between edges
the author declared
+// is reported before any group edge adds edges to the graph, so the message
shows the cycle as
+// written rather than one that a group edge closes.
+func TestRegisterReportsADeclaredCycleAsDeclared(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ cleaned := groupTask(t, group, "clean")
+ validated := groupTask(t, group, "validate")
+ cleaned.Before(validated).Before(cleaned)
+ extracted.Before(group).Before(extracted)
+
+ assert.PanicsWithValue(t,
+ `airflow.BundleRef.Register: the task dependencies of Dag "etl"
contain a cycle: `+
+ `transform.clean -> transform.validate ->
transform.clean`,
+ func() { Bundle().Register(dag) },
+ )
+}
+
+// TestRegisterReportsACycleInsideAGroupOnAnEdge pins that a group whose tasks
all sit on a cycle
+// has no first or last task, so Register steps over it as it does a group
with no task, and then
+// reports the cycle.
+func TestRegisterReportsACycleInsideAGroupOnAnEdge(t *testing.T) {
+ dag := Dag("etl")
+ group := dag.TaskGroup("transform")
+ cleaned := groupTask(t, group, "clean")
+ validated := groupTask(t, group, "validate")
+ cleaned.Before(validated).Before(cleaned)
+ loaded := orderedTask(t, dag, "load")
+ group.Before(loaded)
+
+ assert.PanicsWithValue(t,
+ `airflow.BundleRef.Register: the task dependencies of Dag "etl"
contain a cycle: `+
+ `transform.clean -> transform.validate ->
transform.clean`,
+ func() { Bundle().Register(dag) },
+ )
+}
+
+// TestRegisterTakesADagThatAnotherBundleRegistered pins that registering a
Dag again, in another
+// bundle, changes nothing about it.
+func TestRegisterTakesADagThatAnotherBundleRegistered(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ group := dag.TaskGroup("transform")
+ groupTask(t, group, "clean")
+ extracted.Before(group)
+ Bundle().Register(dag)
+ registered := snapshot(dag)
+
+ assert.NotPanics(t, func() { Bundle().Register(dag) })
+ assert.Equal(t, registered, snapshot(dag))
+}
+
+// TestRegisterTakesAConditionOnceItHasATaskFromThen pins that a Dag that
Register rejects for a
+// condition without a task from Then can be completed and registered, with
its group edges
+// expanded then.
+func TestRegisterTakesAConditionOnceItHasATaskFromThen(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ loaded := orderedTask(t, dag, "load")
+ group := dag.TaskGroup("gate")
+ condition := group.If(isReady)
+ extracted.Before(group)
+
+ assert.PanicsWithValue(t,
+ `airflow.BundleRef.Register: condition "gate.isReady" of Dag
"etl" has no task from `+
+ `Then; name the task that runs when the condition is
true with IfRef.Then`,
+ func() { Bundle().Register(dag) },
+ )
+ assert.Empty(t, extracted.downstreams)
+
+ condition.Then(loaded)
+ Bundle().Register(dag)
+
+ assertTasks(t, extracted.downstreams, condition.task)
+ assertTasks(t, condition.task.downstreams, loaded)
+}
+
+func TestTasksOfAGroupTakeInputsAndConditionsFromOutsideIt(t *testing.T) {
+ dag := Dag("etl")
+ read := dag.Task(readRows)
+ group := dag.TaskGroup("transform")
+ counted := group.Task(countRows, Inputs(read))
+ reported := dag.Task(report, Inputs(counted))
+ gate := dag.If(isReady)
+ cleaned := group.Task(cleanRows)
+ gate.Then(cleaned)
+
+ Bundle().Register(dag)
+
+ assertTasks(t, counted.upstreams, read)
+ assertTasks(t, reported.upstreams, counted)
+ assertTasks(t, cleaned.upstreams, gate.task)
+}
+
+// dagSnapshot is what a Dag records, by ID, so that two snapshots are equal
when the Dag records
+// the same tasks, groups, inputs, conditions and edges in the same order.
+type dagSnapshot struct {
+ registered bool
+ tasks, groups []string
+ taskKeys, groupKeys []string
+ children map[string][]string
+ inputs map[string][]string
+ sides map[string][2]string
+ upstreams, downstreams map[string][]string
+ edgeLabels map[edgeKey]string
+ groupEdges []edgeKey
+ groupEdgeLabels map[edgeKey]string
+}
+
+func snapshot(dag *DagRef) dagSnapshot {
+ dag.mu.Lock()
+ defer dag.mu.Unlock()
+
+ s := dagSnapshot{
+ registered: dag.registered,
+ taskKeys: slices.Sorted(maps.Keys(dag.tasksByID)),
+ groupKeys: slices.Sorted(maps.Keys(dag.groupsByID)),
+ children: make(map[string][]string),
+ inputs: make(map[string][]string),
+ sides: make(map[string][2]string),
+ upstreams: make(map[string][]string),
+ downstreams: make(map[string][]string),
+ edgeLabels: maps.Clone(dag.edgeLabels),
+ groupEdgeLabels: maps.Clone(dag.groupEdgeLabels),
+ }
+ for _, task := range dag.tasks {
+ s.tasks = append(s.tasks, task.taskID)
+ s.inputs[task.taskID] = taskIDs(task.inputs)
+ s.upstreams[task.taskID] = taskIDs(task.upstreams)
+ s.downstreams[task.taskID] = taskIDs(task.downstreams)
+ if task.ifRef != nil {
+ var sides [2]string
+ if task.ifRef.thenTask != nil {
+ sides[0] = task.ifRef.thenTask.taskID
+ }
+ if task.ifRef.elseTask != nil {
+ sides[1] = task.ifRef.elseTask.taskID
+ }
+ s.sides[task.taskID] = sides
+ }
+ }
+ for _, group := range dag.groups {
+ s.groups = append(s.groups, group.groupID)
+ for _, child := range group.children {
+ s.children[group.groupID] = append(
+ s.children[group.groupID], endpoints("child",
child)[0].id(),
+ )
+ }
+ }
+ for _, edge := range dag.groupEdges {
+ s.groupEdges = append(s.groupEdges, edgeKey{edge.upstream.id(),
edge.downstream.id()})
+ }
+ if len(s.edgeLabels) == 0 {
+ s.edgeLabels = nil
+ }
+ return s
+}
+
+// TestAPanicLeavesTheDagAsItWas pins that every check runs before anything is
recorded, so a call
+// that panics leaves the Dag with what it recorded before the call, in the
same order. Each case
+// names the check it expects to fail, so that it cannot pass on a check that
fails earlier.
+func TestAPanicLeavesTheDagAsItWas(t *testing.T) {
+ type dagParts struct {
+ dag *DagRef
+ transform, checks *TaskGroupRef
+ extracted, cleaned *TaskRef
+ }
+ for _, tc := range []struct {
+ name string
+ prepare func(p dagParts)
+ call func(p dagParts)
+ want string
+ }{
+ {
+ name: "a second TaskGroupSpec",
+ call: func(p dagParts) { p.dag.TaskGroup("staging",
TaskGroupSpec{}, TaskGroupSpec{}) },
+ want: "got 2 airflow.TaskGroupSpec values",
+ },
+ {
+ name: "a group_id that Python rejects",
+ call: func(p dagParts) {
p.transform.TaskGroup("rows.checks") },
+ want: `group_id "rows.checks" holds '.'`,
+ },
+ {
+ name: "a group_id that is taken",
+ call: func(p dagParts) {
p.transform.TaskGroup("checks") },
+ want: `cannot add task group "transform.checks"`,
+ },
+ {
+ name: "a task_id that a group takes",
+ call: func(p dagParts) { p.dag.Task(ping,
TaskSpec{TaskID: "transform"}) },
+ want: `because task group "transform" already takes the
ID`,
+ },
+ {
+ name: "a task_id that Python rejects",
+ call: func(p dagParts) { p.transform.Task(ping,
TaskSpec{TaskID: "clean rows"}) },
+ want: `task_id "transform.clean rows" holds ' '`,
+ },
+ {
+ name: "Inputs that do not fit the function",
+ call: func(p dagParts) { p.transform.Task(countRows) },
+ want: "has 1 parameter(s) after airflow.Context",
+ },
+ {
+ name: "an edge from a group to a task it holds",
+ call: func(p dagParts) { p.transform.Before(p.cleaned)
},
+ want: `task "transform.clean" is inside task group
"transform"`,
+ },
+ {
+ name: "a fan-out that relabels a group edge before a
later pair fails",
+ call: func(p dagParts) {
+ p.extracted.Before(Label(p.transform,
"relabelled"), p.extracted)
+ },
+ want: `task "extract" cannot depend on itself`,
+ },
+ {
+ name: "a fan-out that declares a new group edge before
a later pair fails",
+ call: func(p dagParts) { p.extracted.Before(p.checks,
p.extracted) },
+ want: `task "extract" cannot depend on itself`,
+ },
+ {
+ name: "a cycle through a group at Register",
+ prepare: func(p dagParts) {
p.cleaned.Before(p.extracted) },
+ call: func(p dagParts) { Bundle().Register(p.dag) },
+ want: "contain a cycle",
+ },
+ {
+ name: "a condition without Then at Register",
+ prepare: func(p dagParts) { p.transform.If(isReady) },
+ call: func(p dagParts) { Bundle().Register(p.dag) },
+ want: "has no task from Then",
+ },
+ } {
+ t.Run(tc.name, func(t *testing.T) {
+ dag := Dag("etl")
+ extracted := orderedTask(t, dag, "extract")
+ transform := dag.TaskGroup("transform")
+ cleaned := groupTask(t, transform, "clean")
+ checks := transform.TaskGroup("checks")
+ extracted.Before(Label(transform, "rows"))
+ parts := dagParts{dag, transform, checks, extracted,
cleaned}
+ if tc.prepare != nil {
+ tc.prepare(parts)
+ }
+ before := snapshot(dag)
+
+ assert.Contains(t, panicMessage(t, func() {
tc.call(parts) }), tc.want)
+ assert.Equal(t, before, snapshot(dag))
+ })
+ }
+}
diff --git a/go-sdk/airflow/task_option.go b/go-sdk/airflow/task_option.go
index 248feec9641..f24e59487d3 100644
--- a/go-sdk/airflow/task_option.go
+++ b/go-sdk/airflow/task_option.go
@@ -19,13 +19,13 @@ package airflow
import "errors"
-// 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.
+// TaskOption is an option to [DagRef.Task], [DagRef.If] and the methods of
the same names on
+// [TaskGroupRef]. There are two kinds: a [TaskSpec] sets the attributes of
the task that the
+// method adds, and [Inputs] passes the results of other tasks to that task.
//
// Its only method is unexported, so a type outside this package cannot
declare it.
-// A struct that embeds a TaskSpec or a TaskOption still satisfies the
interface, but DagRef.Task
-// and DagRef.If panic when they get such a struct.
+// A struct that embeds a TaskSpec or a TaskOption still satisfies the
interface, but the methods
+// that take a TaskOption panic when they get such a struct.
type TaskOption interface{ applyTask(*taskConfig) error }
type taskConfig struct {
diff --git a/go-sdk/airflow/trigger_dag_run.go
b/go-sdk/airflow/trigger_dag_run.go
index d6ede6f724c..63fa87ce905 100644
--- a/go-sdk/airflow/trigger_dag_run.go
+++ b/go-sdk/airflow/trigger_dag_run.go
@@ -86,8 +86,8 @@ 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:
+// TriggerDagRun returns a value to pass to [DagRef.Task] or
[TaskGroupRef.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"}),
diff --git a/go-sdk/internal/genspec/authoring.go
b/go-sdk/internal/genspec/authoring.go
index d515065a5d5..85660dfef81 100644
--- a/go-sdk/internal/genspec/authoring.go
+++ b/go-sdk/internal/genspec/authoring.go
@@ -23,7 +23,7 @@ import (
"fmt"
)
-// authoringShape rewrites the two definitions the airflow package generates
from
+// authoringShape rewrites the definitions the airflow package generates from
// into the shape a Dag author writes, rather than the shape Airflow
serializes.
// Each definition drops the properties in exclude, rewrites the properties in
// override, and gains the properties in inject.
@@ -63,8 +63,9 @@ type propertyOverride struct {
}
var authoringShapes = map[string]authoringShape{
- "dag": dagShape,
- "operator": taskShape,
+ "dag": dagShape,
+ "operator": taskShape,
+ "task_group": taskGroupShape,
}
var dagShape = authoringShape{
@@ -115,7 +116,8 @@ var dagShape = authoringShape{
}
var taskShape = authoringShape{
- doc: "TaskSpec holds the attributes of a task. DagRef.Task takes at
most one per task.",
+ doc: "TaskSpec holds the attributes of a task. DagRef.Task, DagRef.If
and the methods of " +
+ "the same names on TaskGroupRef take at most one per task.",
exclude: map[string]string{
"task_type": "the operator class name,
which the SDK fills in",
"_task_module": "the operator's Python module,
which the SDK fills in",
@@ -180,7 +182,30 @@ 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. A task
from TriggerDagRun runs no Go function, so it needs a TaskID.",
+ 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. A task added
through a task group takes the group_id as a prefix of its task_id, unless the
TaskGroupSpec of the group sets PrefixGroupID to false.",
+ },
+ },
+}
+
+var taskGroupShape = authoringShape{
+ doc: "TaskGroupSpec holds the attributes of a task group other than its
group_id. " +
+ "DagRef.TaskGroup and TaskGroupRef.TaskGroup take at most one
per group.",
+ exclude: map[string]string{
+ "_group_id": "a positional parameter of
DagRef.TaskGroup and TaskGroupRef.TaskGroup",
+ "children": "the tasks and groups added through the
group",
+ "is_mapped": "derived from whether the group is
mapped, which the SDK does not model yet",
+ "upstream_group_ids": "the edges Before and After declare",
+ "downstream_group_ids": "the edges Before and After declare",
+ "upstream_task_ids": "the edges Before and After declare",
+ "downstream_task_ids": "the edges Before and After declare",
+ },
+ override: map[string]propertyOverride{
+ // The schema allows null for doc_md, as anyOf [string, null],
which rejectCombinators
+ // refuses. An empty DocMD means a group without docs, as null
does.
+ "doc_md": {goType: "string"},
+ "prefix_group_id": {
+ goType: "bool",
+ doc: "PrefixGroupID says whether the group_id
prefixes the IDs of the tasks and groups added through the group, as in
\"transform.cleanRows\". When PrefixGroupID is nil, the group_id prefixes
them.",
},
},
}
diff --git a/go-sdk/internal/genspec/authoring_test.go
b/go-sdk/internal/genspec/authoring_test.go
index bb5f9403faf..a56c96fa593 100644
--- a/go-sdk/internal/genspec/authoring_test.go
+++ b/go-sdk/internal/genspec/authoring_test.go
@@ -306,7 +306,7 @@ func TestShapeForAuthoringShapesTheCoreSchema(t *testing.T)
{
require.NoError(t, normalize(doc, specTitles))
definitions := doc["definitions"].(map[string]any)
- assert.Equal(t, []string{"dag", "operator"}, sortedKeys(definitions))
+ assert.Equal(t, []string{"dag", "operator", "task_group"},
sortedKeys(definitions))
dag :=
definitions["dag"].(map[string]any)["properties"].(map[string]any)
assert.Contains(t, dag, "schedule", "Schedule is injected over the
serialized timetable")
task :=
definitions["operator"].(map[string]any)["properties"].(map[string]any)
@@ -326,4 +326,24 @@ func TestShapeForAuthoringShapesTheCoreSchema(t
*testing.T) {
"the name of the Go function",
"the serialized property says nothing about how a task_id is
defaulted",
)
+ group :=
definitions["task_group"].(map[string]any)["properties"].(map[string]any)
+ assert.Equal(
+ t,
+ []string{
+ "doc_md",
+ "group_display_name",
+ "prefix_group_id",
+ "tooltip",
+ "ui_color",
+ "ui_fgcolor",
+ },
+ sortedKeys(group),
+ )
+ prefix :=
group["prefix_group_id"].(map[string]any)["goJSONSchema"].(map[string]any)
+ assert.Equal(
+ t,
+ true,
+ prefix["pointer"],
+ "a group_id prefixes by default, so false has to be
expressible",
+ )
}
diff --git a/go-sdk/internal/genspec/main.go b/go-sdk/internal/genspec/main.go
index 208af46833d..92bc6ea058c 100644
--- a/go-sdk/internal/genspec/main.go
+++ b/go-sdk/internal/genspec/main.go
@@ -16,8 +16,8 @@
// under the License.
// Command genspec rewrites Airflow core's Dag serialization schema (go-sdk's
vendored
-// copy of it, schema/dag-schema.json) into the schema the airflow package's
DagSpec
-// and TaskSpec generate from, so that the two structs are not hand-maintained.
+// copy of it, schema/dag-schema.json) into the schema the airflow package's
DagSpec,
+// TaskSpec and TaskGroupSpec generate from, so that the structs are not
hand-maintained.
//
// The schema is owned by Python and stays untouched; the rewritten copy is a
build
// artifact. genspec rewrites it in two passes: shapeForAuthoring turns the
diff --git a/go-sdk/internal/genspec/normalize.go
b/go-sdk/internal/genspec/normalize.go
index 307f12ba47c..575f22667f1 100644
--- a/go-sdk/internal/genspec/normalize.go
+++ b/go-sdk/internal/genspec/normalize.go
@@ -23,8 +23,9 @@ import (
)
var specTitles = map[string]string{
- "dag": "DagSpec",
- "operator": "TaskSpec",
+ "dag": "DagSpec",
+ "operator": "TaskSpec",
+ "task_group": "TaskGroupSpec",
}
// normalize rewrites doc in place so that go-jsonschema can read it, and
injects
@@ -108,10 +109,11 @@ func resolveNodeType(path string, node map[string]any)
error {
// --struct-name-from-title reads. It reports a definition that has gone
missing
// or already carries a title of its own, either of which means titles is
stale.
//
-// Neither definition carries a title, so the flag has nothing to read and
-// go-jsonschema falls back to capitalizing the definition keys dag and
operator.
-// Dag is already the name of the constructor and Operator is not the SDK's
-// vocabulary, which is why the titles are injected rather than left to the
tool.
+// No definition carries a title, so the flag has nothing to read and
go-jsonschema
+// falls back to capitalizing the definition keys dag, operator and
task_group. Dag
+// is already the name of the constructor, Operator is not the SDK's
vocabulary, and
+// TaskGroup is the name of the method that adds a group, which is why the
titles are
+// injected rather than left to the tool.
func injectTitles(doc map[string]any, titles map[string]string) error {
definitions, ok := doc["definitions"].(map[string]any)
if !ok {
diff --git a/go-sdk/internal/genspec/normalize_test.go
b/go-sdk/internal/genspec/normalize_test.go
index 263b4918a82..610b5c3339f 100644
--- a/go-sdk/internal/genspec/normalize_test.go
+++ b/go-sdk/internal/genspec/normalize_test.go
@@ -53,7 +53,8 @@ func TestNormalizeResolvesNullableType(t *testing.T) {
"type": "object",
"properties": {"rerun_with_latest_version":
{"type": ["boolean", "null"]}}
},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`)
@@ -72,7 +73,8 @@ func TestNormalizeKeepsASingleTypeAsItIs(t *testing.T) {
doc := schemaFrom(t, `{
"definitions": {
"dag": {"type": "object", "properties": {"dag_id":
{"type": "string"}}},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`)
@@ -87,6 +89,7 @@ func TestNormalizeDropsDependentRequired(t *testing.T) {
doc := schemaFrom(t, `{
"definitions": {
"dag": {"type": "object"},
+ "task_group": {"type": "object"},
"operator": {
"type": "object",
"dependencies": {
@@ -107,6 +110,7 @@ func TestNormalizeKeepsSchemaFormDependencies(t *testing.T)
{
doc := schemaFrom(t, `{
"definitions": {
"dag": {"type": "object"},
+ "task_group": {"type": "object"},
"operator": {
"type": "object",
"dependencies": {
@@ -128,7 +132,11 @@ func TestNormalizeKeepsSchemaFormDependencies(t
*testing.T) {
func TestNormalizeInjectsTheSpecTitles(t *testing.T) {
doc := schemaFrom(t, `{
- "definitions": {"dag": {"type": "object"}, "operator": {"type":
"object"}}
+ "definitions": {
+ "dag": {"type": "object"},
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
+ }
}`)
require.NoError(t, normalize(doc, specTitles))
@@ -136,6 +144,7 @@ func TestNormalizeInjectsTheSpecTitles(t *testing.T) {
definitions := doc["definitions"].(map[string]any)
assert.Equal(t, "DagSpec", definitions["dag"].(map[string]any)["title"])
assert.Equal(t, "TaskSpec",
definitions["operator"].(map[string]any)["title"])
+ assert.Equal(t, "TaskGroupSpec",
definitions["task_group"].(map[string]any)["title"])
}
func TestNormalizeRejects(t *testing.T) {
@@ -152,7 +161,8 @@ func TestNormalizeRejects(t *testing.T) {
"type": "object",
"properties": {"either":
{"type": ["boolean", "string"]}}
},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`,
wantErr: "/definitions/dag/properties/either",
@@ -162,7 +172,8 @@ func TestNormalizeRejects(t *testing.T) {
schema: `{
"definitions": {
"dag": {"type": "object", "properties":
{"nothing": {"type": ["null"]}}},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`,
wantErr: "allows nothing to generate from",
@@ -177,7 +188,8 @@ func TestNormalizeRejects(t *testing.T) {
schema: `{
"definitions": {
"dag": {"type": "object", "title":
"SerializedDag"},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`,
wantErr: `already has the title "SerializedDag"`,
@@ -196,7 +208,8 @@ func TestNormalizeAcceptsTheTitleItWouldInject(t
*testing.T) {
doc := schemaFrom(t, `{
"definitions": {
"dag": {"type": "object", "title": "DagSpec"},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`)
@@ -210,7 +223,8 @@ func TestNormalizeResolvesANullableTypeNestedInItems(t
*testing.T) {
"type": "object",
"properties": {"tags": {"type": "array",
"items": {"type": ["string", "null"]}}}
},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`)
@@ -232,7 +246,8 @@ func TestNormalizeReportsTheSameConstructOnEveryRun(t
*testing.T) {
"b": {"type": ["number", "string"]}
}
},
- "operator": {"type": "object"}
+ "operator": {"type": "object"},
+ "task_group": {"type": "object"}
}
}`
diff --git a/scripts/ci/prek/check_go_sdk_generated_drift.py
b/scripts/ci/prek/check_go_sdk_generated_drift.py
index b0f5b80191a..d902cf7014b 100755
--- a/scripts/ci/prek/check_go_sdk_generated_drift.py
+++ b/scripts/ci/prek/check_go_sdk_generated_drift.py
@@ -21,8 +21,9 @@ Keep the Go SDK's generated files in step with the schemas
they generate from.
Two of the Go SDK's surfaces are generated from schemas Python owns, and both
are
committed, so nothing regenerates them when the schema moves:
-* ``go-sdk/airflow/spec.gen.go`` — ``airflow.DagSpec`` and
``airflow.TaskSpec``, the
- structs a Dag author fills in, from ``go-sdk/schema/dag-schema.json``.
+* ``go-sdk/airflow/spec.gen.go`` — ``airflow.DagSpec``, ``airflow.TaskSpec``
and
+ ``airflow.TaskGroupSpec``, the structs a Dag author fills in, from
+ ``go-sdk/schema/dag-schema.json``.
* ``go-sdk/pkg/execution/genmodels/*.gen.go`` — the coordinator-protocol
messages,
from ``go-sdk/schema/supervisor-schema.json``.