jason810496 commented on code in PR #74075: URL: https://github.com/apache/airflow/pull/74075#discussion_r4165333509
########## go-sdk/airflow/if.go: ########## @@ -0,0 +1,225 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package airflow + +import ( + "fmt" + "reflect" + "strings" + + "github.com/apache/airflow/go-sdk/internal/bundle" +) + +// IfRef is a condition that [DagRef.If] added to a Dag. [IfRef.Then] names the task that runs +// when the condition is true, and [IfRef.Else] names the task that runs when it is false. +type IfRef struct { + // task is the task that runs the condition function. thenTask and elseTask run after task, + // but neither takes the result of task. + task *TaskRef + thenTask *TaskRef + elseTask *TaskRef +} + +// If adds a task that runs the condition function fn, and returns the condition. Use +// [IfRef.Then] to name the task that runs when fn returns true, and [IfRef.Else] to name the task +// that runs when fn returns false: +// +// extracted := dag.Task(extract) +// loaded := dag.Task(load) +// reportedEmpty := dag.Task(reportEmpty) +// +// gate := dag.If(hasRows, airflow.Inputs(extracted)) +// gate.Then(loaded) +// gate.Else(reportedEmpty) +// +// fn returns (bool, error). In every other way it follows the rules of a function passed to +// [DagRef.Task]: it takes a [Context] first, and [Inputs] fills the parameters after the Context. +// For the Dag above, the functions could be: +// +// func extract(actx airflow.Context) ([]string, error) +// func hasRows(actx airflow.Context, rows []string) (bool, error) +// func load(actx airflow.Context) error +// func reportEmpty(actx airflow.Context) error +// +// The task that If adds gets its task_id in the same way as a task from DagRef.Task, and a +// [TaskSpec] sets its attributes. +// +// When the task runs, the result of fn decides what it skips: +// - true skips the task from Else, or nothing if there is no task from Else +// - false skips the task from Then +// - an error fails the task, which then skips nothing +// +// The condition skips only that one task. In Python, a branch skips every task directly after it +// that it does not follow, and a short circuit skips every task after it. After the skip, the +// trigger rule of each task after the skipped task decides whether that task runs. A task that +// runs after both the task from Then and the task from Else needs a trigger rule that runs it when +// one of the two is skipped, such as [TriggerRuleNoneFailedMinOneSuccess]. +// +// [BundleRef.Register] panics if a condition has no task from Then. +// +// If panics for the same reasons as DagRef.Task does for a Go function. It also panics if fn +// comes from [TriggerDagRun] or does not return (bool, error). +func (d *DagRef) If(fn any, opts ...TaskOption) *IfRef { + if _, ok := fn.(TriggerDagRunTask); ok { + panic(fmt.Sprintf( + "airflow.DagRef.If: Dag %q: fn comes from airflow.TriggerDagRun, "+ + "but a condition function is a Go function that returns (bool, error)", + d.dagID, + )) + } + ifRef := &IfRef{} + d.addTask("airflow.DagRef.If", fn, opts, ifRef) + return ifRef +} + +// Then names the task that runs when the condition of g is true. When the condition is false, +// the task is skipped. Then returns g, so that a call to Else can follow. With the tasks from the +// example of [DagRef.If], the condition fits in one statement: +// +// dag.If(hasRows, airflow.Inputs(extracted)).Then(loaded).Else(reportedEmpty) +// +// The task runs after the condition, but its function does not take the result of the +// condition as a parameter. +// +// Then panics if: +// - g is not the IfRef that DagRef.If returned, for example a copy of that IfRef +// - task is nil, or DagRef.Task did not add task to the Dag of the condition +// - the condition already has a task from Then +// - task is the task from Else +// - the Dag is already registered +func (g *IfRef) Then(task *TaskRef) *IfRef { + g.setTask("Then", task) + return g +} + +// Else names the task that runs when the condition of g is false. When the condition is true, +// the task is skipped. Like Then, Else returns g. A condition does not need a task from Else. +// +// Else panics for the reasons that Then lists, with Then and Else swapped. +func (g *IfRef) Else(task *TaskRef) *IfRef { + g.setTask("Else", task) + return g +} + +// setTask makes task the task from Then or from Else of g. side is "Then" or "Else". +func (g *IfRef) setTask(side string, task *TaskRef) { + method := "airflow.IfRef." + side + // When the condition task runs, it reads the tasks from Then and Else from the IfRef that If + // returned. A task given to a copy of that IfRef would never reach the condition task. + if g == nil || g.task == nil || g.task.ifRef != g { + panic(method + ": DagRef.If did not return the *airflow.IfRef") + } + condition, d := g.task.taskID, g.task.dag + d.mu.Lock() + defer d.mu.Unlock() + + if d.registered { + panic(fmt.Sprintf( + "%s: Dag %q has already been registered; "+ + "name the tasks of every condition before Register", + method, d.dagID, + )) + } + switch { + case task == nil: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q got a nil *airflow.TaskRef", method, condition, d.dagID, + )) + case task.dag != nil && task.dag != d: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q cannot run task %q of another Dag, %q; "+ + "pass a task of the same Dag", + method, condition, d.dagID, task.taskID, task.dag.dagID, + )) + // A zero TaskRef and a copy of a TaskRef get here. + case d.tasksByID[task.taskID] != task: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q got a *airflow.TaskRef that DagRef.Task did not return", + method, condition, d.dagID, + )) + } Review Comment: `Then`/`Else` accept a task that is upstream of the condition, which creates a cycle that `Register` accepts: ```go gate := dag.If(hasRows, airflow.Inputs(extracted)) gate.Then(extracted) // extracted -> gate -> extracted ``` Suggest rejecting it here: ```suggestion case d.tasksByID[task.taskID] != task: panic(fmt.Sprintf( "%s: condition %q of Dag %q got a *airflow.TaskRef that DagRef.Task did not return", method, condition, d.dagID, )) case g.task.dependsOn(task): panic(fmt.Sprintf( "%s: condition %q of Dag %q takes task %q as an input, so it cannot run after the condition", method, condition, d.dagID, task.taskID, )) } ``` with: ```go func (t *TaskRef) dependsOn(other *TaskRef) bool { for _, in := range t.inputs { if in == other || in.dependsOn(other) { return true } } return false } ``` ########## go-sdk/airflow/if.go: ########## @@ -0,0 +1,225 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package airflow + +import ( + "fmt" + "reflect" + "strings" + + "github.com/apache/airflow/go-sdk/internal/bundle" +) + +// IfRef is a condition that [DagRef.If] added to a Dag. [IfRef.Then] names the task that runs +// when the condition is true, and [IfRef.Else] names the task that runs when it is false. +type IfRef struct { + // task is the task that runs the condition function. thenTask and elseTask run after task, + // but neither takes the result of task. + task *TaskRef + thenTask *TaskRef + elseTask *TaskRef +} + +// If adds a task that runs the condition function fn, and returns the condition. Use +// [IfRef.Then] to name the task that runs when fn returns true, and [IfRef.Else] to name the task +// that runs when fn returns false: +// +// extracted := dag.Task(extract) +// loaded := dag.Task(load) +// reportedEmpty := dag.Task(reportEmpty) +// +// gate := dag.If(hasRows, airflow.Inputs(extracted)) +// gate.Then(loaded) +// gate.Else(reportedEmpty) +// +// fn returns (bool, error). In every other way it follows the rules of a function passed to +// [DagRef.Task]: it takes a [Context] first, and [Inputs] fills the parameters after the Context. +// For the Dag above, the functions could be: +// +// func extract(actx airflow.Context) ([]string, error) +// func hasRows(actx airflow.Context, rows []string) (bool, error) +// func load(actx airflow.Context) error +// func reportEmpty(actx airflow.Context) error +// +// The task that If adds gets its task_id in the same way as a task from DagRef.Task, and a +// [TaskSpec] sets its attributes. +// +// When the task runs, the result of fn decides what it skips: +// - true skips the task from Else, or nothing if there is no task from Else +// - false skips the task from Then +// - an error fails the task, which then skips nothing +// +// The condition skips only that one task. In Python, a branch skips every task directly after it +// that it does not follow, and a short circuit skips every task after it. After the skip, the +// trigger rule of each task after the skipped task decides whether that task runs. A task that +// runs after both the task from Then and the task from Else needs a trigger rule that runs it when +// one of the two is skipped, such as [TriggerRuleNoneFailedMinOneSuccess]. +// +// [BundleRef.Register] panics if a condition has no task from Then. +// +// If panics for the same reasons as DagRef.Task does for a Go function. It also panics if fn +// comes from [TriggerDagRun] or does not return (bool, error). +func (d *DagRef) If(fn any, opts ...TaskOption) *IfRef { + if _, ok := fn.(TriggerDagRunTask); ok { + panic(fmt.Sprintf( + "airflow.DagRef.If: Dag %q: fn comes from airflow.TriggerDagRun, "+ + "but a condition function is a Go function that returns (bool, error)", + d.dagID, + )) + } + ifRef := &IfRef{} + d.addTask("airflow.DagRef.If", fn, opts, ifRef) + return ifRef +} + +// Then names the task that runs when the condition of g is true. When the condition is false, +// the task is skipped. Then returns g, so that a call to Else can follow. With the tasks from the +// example of [DagRef.If], the condition fits in one statement: +// +// dag.If(hasRows, airflow.Inputs(extracted)).Then(loaded).Else(reportedEmpty) +// +// The task runs after the condition, but its function does not take the result of the +// condition as a parameter. +// +// Then panics if: +// - g is not the IfRef that DagRef.If returned, for example a copy of that IfRef +// - task is nil, or DagRef.Task did not add task to the Dag of the condition +// - the condition already has a task from Then +// - task is the task from Else +// - the Dag is already registered +func (g *IfRef) Then(task *TaskRef) *IfRef { + g.setTask("Then", task) + return g +} + +// Else names the task that runs when the condition of g is false. When the condition is true, +// the task is skipped. Like Then, Else returns g. A condition does not need a task from Else. +// +// Else panics for the reasons that Then lists, with Then and Else swapped. +func (g *IfRef) Else(task *TaskRef) *IfRef { + g.setTask("Else", task) + return g +} + +// setTask makes task the task from Then or from Else of g. side is "Then" or "Else". +func (g *IfRef) setTask(side string, task *TaskRef) { + method := "airflow.IfRef." + side + // When the condition task runs, it reads the tasks from Then and Else from the IfRef that If + // returned. A task given to a copy of that IfRef would never reach the condition task. + if g == nil || g.task == nil || g.task.ifRef != g { + panic(method + ": DagRef.If did not return the *airflow.IfRef") + } + condition, d := g.task.taskID, g.task.dag + d.mu.Lock() + defer d.mu.Unlock() + + if d.registered { + panic(fmt.Sprintf( + "%s: Dag %q has already been registered; "+ + "name the tasks of every condition before Register", + method, d.dagID, + )) + } + switch { + case task == nil: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q got a nil *airflow.TaskRef", method, condition, d.dagID, + )) + case task.dag != nil && task.dag != d: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q cannot run task %q of another Dag, %q; "+ + "pass a task of the same Dag", + method, condition, d.dagID, task.taskID, task.dag.dagID, + )) + // A zero TaskRef and a copy of a TaskRef get here. + case d.tasksByID[task.taskID] != task: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q got a *airflow.TaskRef that DagRef.Task did not return", + method, condition, d.dagID, + )) + } + + slot, other, otherSide := &g.thenTask, g.elseTask, "Else" + if side == "Else" { + slot, other, otherSide = &g.elseTask, g.thenTask, "Then" + } + switch { + case *slot != nil: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q already has task %q from %s; call %s once", + method, condition, d.dagID, (*slot).taskID, side, side, + )) + case task == other: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q already has task %q from %s, "+ + "and a task cannot be on both sides of a condition", + method, condition, d.dagID, task.taskID, otherSide, + )) + } + *slot = task Review Comment: The condition -> Then/Else edge lives only on `IfRef`. The target's `inputs` never gets the condition, so anything that builds upstreams from `TaskRef.inputs` (serialization, ordering) would treat `loaded` from the doc example as a root task. Should the edge be recorded on the target task here, or is the follow-up serialization PR going to read it from `IfRef`? ########## go-sdk/airflow/dag.go: ########## @@ -183,57 +197,55 @@ func (d *DagRef) Task(fn any, opts ...TaskOption) *TaskRef { continue } panic(fmt.Sprintf( - "airflow.DagRef.Task: task %q of Dag %q: %v", findTaskName(fn, opts), d.dagID, err, + "%s: task %q of Dag %q: %v", method, findTaskName(fn, opts), d.dagID, err, )) } } taskID := cfg.spec.TaskID if taskID == "" && isTrigger { panic(fmt.Sprintf( - "airflow.DagRef.Task: Dag %q: a task from airflow.TriggerDagRun with DagID %q has no "+ + "%s: Dag %q: a task from airflow.TriggerDagRun with DagID %q has no "+ "Go function to take a task_id from; set one with airflow.TaskSpec{TaskID: ...}", - d.dagID, trigger.spec.DagID, + method, d.dagID, trigger.spec.DagID, )) } if taskID == "" { var ok bool if taskID, ok = taskIDFromFuncName(funcName(fn)); !ok { panic(fmt.Sprintf( - "airflow.DagRef.Task: Dag %q: %s has no name to use as the task_id; "+ + "%s: Dag %q: %s has no name to use as the task_id; "+ "set one with airflow.TaskSpec{TaskID: ...}", - d.dagID, funcName(fn), + method, d.dagID, funcName(fn), )) } } if triggerErr != nil { - panic(fmt.Sprintf( - "airflow.DagRef.Task: task %q of Dag %q: %v", taskID, d.dagID, triggerErr, - )) + panic(fmt.Sprintf("%s: task %q of Dag %q: %v", method, taskID, d.dagID, triggerErr)) } if err := checkTaskSpec(cfg.spec); err != nil { panic(fmt.Sprintf("airflow.DagRef.Task: task %q of Dag %q: %v", taskID, d.dagID, err)) Review Comment: This one still hard-codes the method, so `dag.If(cond, airflow.TaskSpec{TriggerRule: "bogus"})` reports `airflow.DagRef.Task`. ```suggestion panic(fmt.Sprintf("%s: task %q of Dag %q: %v", method, taskID, d.dagID, err)) ``` ########## go-sdk/airflow/if.go: ########## @@ -0,0 +1,225 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package airflow + +import ( + "fmt" + "reflect" + "strings" + + "github.com/apache/airflow/go-sdk/internal/bundle" +) + +// IfRef is a condition that [DagRef.If] added to a Dag. [IfRef.Then] names the task that runs +// when the condition is true, and [IfRef.Else] names the task that runs when it is false. +type IfRef struct { + // task is the task that runs the condition function. thenTask and elseTask run after task, + // but neither takes the result of task. + task *TaskRef + thenTask *TaskRef + elseTask *TaskRef +} + +// If adds a task that runs the condition function fn, and returns the condition. Use +// [IfRef.Then] to name the task that runs when fn returns true, and [IfRef.Else] to name the task +// that runs when fn returns false: +// +// extracted := dag.Task(extract) +// loaded := dag.Task(load) +// reportedEmpty := dag.Task(reportEmpty) +// +// gate := dag.If(hasRows, airflow.Inputs(extracted)) +// gate.Then(loaded) +// gate.Else(reportedEmpty) +// +// fn returns (bool, error). In every other way it follows the rules of a function passed to +// [DagRef.Task]: it takes a [Context] first, and [Inputs] fills the parameters after the Context. +// For the Dag above, the functions could be: +// +// func extract(actx airflow.Context) ([]string, error) +// func hasRows(actx airflow.Context, rows []string) (bool, error) +// func load(actx airflow.Context) error +// func reportEmpty(actx airflow.Context) error +// +// The task that If adds gets its task_id in the same way as a task from DagRef.Task, and a +// [TaskSpec] sets its attributes. +// +// When the task runs, the result of fn decides what it skips: +// - true skips the task from Else, or nothing if there is no task from Else +// - false skips the task from Then +// - an error fails the task, which then skips nothing +// +// The condition skips only that one task. In Python, a branch skips every task directly after it +// that it does not follow, and a short circuit skips every task after it. After the skip, the +// trigger rule of each task after the skipped task decides whether that task runs. A task that +// runs after both the task from Then and the task from Else needs a trigger rule that runs it when +// one of the two is skipped, such as [TriggerRuleNoneFailedMinOneSuccess]. +// +// [BundleRef.Register] panics if a condition has no task from Then. +// +// If panics for the same reasons as DagRef.Task does for a Go function. It also panics if fn +// comes from [TriggerDagRun] or does not return (bool, error). +func (d *DagRef) If(fn any, opts ...TaskOption) *IfRef { + if _, ok := fn.(TriggerDagRunTask); ok { + panic(fmt.Sprintf( + "airflow.DagRef.If: Dag %q: fn comes from airflow.TriggerDagRun, "+ + "but a condition function is a Go function that returns (bool, error)", + d.dagID, + )) + } + ifRef := &IfRef{} + d.addTask("airflow.DagRef.If", fn, opts, ifRef) + return ifRef +} + +// Then names the task that runs when the condition of g is true. When the condition is false, +// the task is skipped. Then returns g, so that a call to Else can follow. With the tasks from the +// example of [DagRef.If], the condition fits in one statement: +// +// dag.If(hasRows, airflow.Inputs(extracted)).Then(loaded).Else(reportedEmpty) +// +// The task runs after the condition, but its function does not take the result of the +// condition as a parameter. +// +// Then panics if: +// - g is not the IfRef that DagRef.If returned, for example a copy of that IfRef +// - task is nil, or DagRef.Task did not add task to the Dag of the condition +// - the condition already has a task from Then +// - task is the task from Else +// - the Dag is already registered +func (g *IfRef) Then(task *TaskRef) *IfRef { + g.setTask("Then", task) + return g +} + +// Else names the task that runs when the condition of g is false. When the condition is true, +// the task is skipped. Like Then, Else returns g. A condition does not need a task from Else. +// +// Else panics for the reasons that Then lists, with Then and Else swapped. +func (g *IfRef) Else(task *TaskRef) *IfRef { + g.setTask("Else", task) + return g +} + +// setTask makes task the task from Then or from Else of g. side is "Then" or "Else". +func (g *IfRef) setTask(side string, task *TaskRef) { + method := "airflow.IfRef." + side + // When the condition task runs, it reads the tasks from Then and Else from the IfRef that If + // returned. A task given to a copy of that IfRef would never reach the condition task. + if g == nil || g.task == nil || g.task.ifRef != g { + panic(method + ": DagRef.If did not return the *airflow.IfRef") + } + condition, d := g.task.taskID, g.task.dag + d.mu.Lock() + defer d.mu.Unlock() + + if d.registered { + panic(fmt.Sprintf( + "%s: Dag %q has already been registered; "+ + "name the tasks of every condition before Register", + method, d.dagID, + )) + } + switch { + case task == nil: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q got a nil *airflow.TaskRef", method, condition, d.dagID, + )) + case task.dag != nil && task.dag != d: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q cannot run task %q of another Dag, %q; "+ + "pass a task of the same Dag", + method, condition, d.dagID, task.taskID, task.dag.dagID, + )) + // A zero TaskRef and a copy of a TaskRef get here. + case d.tasksByID[task.taskID] != task: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q got a *airflow.TaskRef that DagRef.Task did not return", + method, condition, d.dagID, + )) + } + + slot, other, otherSide := &g.thenTask, g.elseTask, "Else" + if side == "Else" { + slot, other, otherSide = &g.elseTask, g.thenTask, "Then" + } + switch { + case *slot != nil: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q already has task %q from %s; call %s once", + method, condition, d.dagID, (*slot).taskID, side, side, + )) + case task == other: + panic(fmt.Sprintf( + "%s: condition %q of Dag %q already has task %q from %s, "+ + "and a task cannot be on both sides of a condition", + method, condition, d.dagID, task.taskID, otherSide, + )) + } + *slot = task +} + +// wrapCondition wraps fn as the task of g. The task skips the side of g that the result of fn +// does not take. +func (g *IfRef) wrapCondition(fn any) (bundle.Task, error) { + fnType := reflect.TypeOf(fn) + // DagRef.Task also takes a function whose last result has a concrete type that implements + // error. A nil value of that type becomes a non-nil error when the runtime reads it, so the + // condition task would always fail. + if fnType.NumOut() != 2 || + fnType.Out(0) != reflect.TypeFor[bool]() || + fnType.Out(1) != reflect.TypeFor[error]() { + return nil, fmt.Errorf( + "%s returns %s, but a condition function must return (bool, error)", + funcName(fn), describeResults(fnType), + ) + } + return bundle.NewPositionalBranchFunction(fn, g.findSkipped) +} + +// findSkipped returns the task_id of the task on the side of g that result does not take. It +// returns nil when that side has no task. +func (g *IfRef) findSkipped(result any) []string { + d := g.task.dag + d.mu.Lock() + defer d.mu.Unlock() + + notTaken := g.thenTask Review Comment: nit: `thenTask`/`elseTask` cannot change after `Register`, and only a registered Dag runs tasks, so the Dag-wide lock on every run isn't needed. ```suggestion func (g *IfRef) findSkipped(result any) []string { notTaken := g.thenTask ``` ########## go-sdk/internal/bundle/task.go: ########## @@ -202,6 +280,13 @@ func (f *taskFunction) validateFn( fnType.Out(fnType.NumOut()-1).Kind(), ) } + if f.findSkipped != nil && fnType.NumOut() != 2 { + return fmt.Errorf( + "task function %s returns only an error, but a task that skips downstream tasks "+ + "must return `<result>, error`", + f.fullName, + ) + } Review Comment: nit: `wrapCondition` already rejects anything but `(bool, error)` before calling `NewPositionalBranchFunction`, so this can't fail for the only caller. Suggest keeping the check in one place. ```suggestion ``` ########## go-sdk/internal/bundle/task.go: ########## @@ -139,9 +155,71 @@ func (f *taskFunction) call( res := retValues[0].Interface() f.sendXcom(ctx, res, sdkClient, logger) } + if err == nil && f.findSkipped != nil { Review Comment: The `skipmixin_key` XCom is only rewritten when `fn` succeeds. If try 1 returns false (skips `load`) and a cleared try 2 returns an error, the old list stays, and `NotPreviouslySkippedDep` skips an `all_done` `load` even though this try skipped nothing. The Go runtime ignores `XcomKeysToClear`, so the "always write" in `skipDownstream` doesn't cover the error path. Minimal fix, also writing the empty list on error: ```suggestion if f.findSkipped != nil { if err != nil { return errors.Join(err, f.skipDownstream(ctx, sdkClient, nil, logger)) } return f.skipDownstream(ctx, sdkClient, f.findSkipped(retValues[0].Interface()), logger) } ``` The cleaner fix is for `RunTask` to delete the `XcomKeysToClear` keys before running the task, as the Python task runner does. ########## go-sdk/internal/bundle/task.go: ########## @@ -139,9 +155,71 @@ func (f *taskFunction) call( res := retValues[0].Interface() f.sendXcom(ctx, res, sdkClient, logger) } + if err == nil && f.findSkipped != nil { + return f.skipDownstream(ctx, sdkClient, f.findSkipped(retValues[0].Interface()), logger) + } return err } +// skipMixinXComKey is the key of the XCom that lists the tasks a task skipped. When one of those +// tasks is cleared, NotPreviouslySkippedDep in Airflow core reads the XCom and skips the cleared +// task again instead of running it. SkipMixin in the standard provider writes the same key. +const skipMixinXComKey = "skipmixin_key" + +type skipDownstreamTasksKey struct{} + +// WithSkipDownstreamTasks returns a copy of ctx that carries skip, the function that a task from +// NewPositionalBranchFunction calls to skip tasks downstream of it. The runtime passes skip in the +// context and not in sdk.Client, so that task functions cannot call it. Airflow core reads the +// skipmixin_key XCom only from a task that has _can_skip_downstream set in the serialized Dag. A +// task skipped by a task without that flag would run when someone clears it. +func WithSkipDownstreamTasks( + ctx context.Context, + skip func(ctx context.Context, taskIDs []string) error, +) context.Context { + return context.WithValue(ctx, skipDownstreamTasksKey{}, skip) +} + +func (f *taskFunction) skipDownstream( + ctx context.Context, + client sdk.Client, + taskIDs []string, + logger *slog.Logger, +) error { + skip, ok := ctx.Value(skipDownstreamTasksKey{}).(func(context.Context, []string) error) Review Comment: This check runs after `fn` has already run and pushed `return_value`, so a runtime without `WithSkipDownstreamTasks` fails the task after its side effects, and a retry runs them again. Could `Execute` check it first? ```go if f.findSkipped != nil { if _, ok := ctx.Value(skipDownstreamTasksKey{}).(func(context.Context, []string) error); !ok { return errors.New("the task runtime cannot skip downstream tasks") } } ``` -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
