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 874bc6c215b TS SDK: Add multi-way branching with dag.switch (#74070)
874bc6c215b is described below
commit 874bc6c215be59c0d31f5893284f74d8d6faec9a
Author: Guan-Ming Chiu <[email protected]>
AuthorDate: Wed Oct 7 13:01:38 2026 +0800
TS SDK: Add multi-way branching with dag.switch (#74070)
Co-authored-by: Jason(Zhe-You) Liu
<[email protected]>
---
.../language-sdks/typescript.rst | 24 ++
ts-sdk/adr/0002-native-dag-interface.md | 32 +++
ts-sdk/src/index.ts | 1 +
ts-sdk/src/sdk/dag.ts | 109 +++++++++-
ts-sdk/tests/sdk/branching.test.ts | 242 +++++++++++++++++++++
5 files changed, 397 insertions(+), 11 deletions(-)
diff --git
a/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst
b/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst
index a8ef96d5f01..99f568b679e 100644
--- a/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst
+++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst
@@ -419,6 +419,30 @@ The side not taken is skipped when the run reaches it, and
stays skipped if you
the branches named here are skipped, so a task that several branches converge
on still runs — unlike
Python's ``@task.branch``, which skips every immediate downstream it did not
follow.
+Multi-way branching
+~~~~~~~~~~~~~~~~~~~
+
+``dag.switch`` is the multi-way form: a handler that returns one of the cases
it is given.
+
+.. code-block:: typescript
+
+ async function pickPath({ rows }: { rows: number }): Promise<TaskRef> {
+ return rows > 1000 ? handleLong : handleShort;
+ }
+
+ dag.switch(pickPath, { rows: extracted
}).case(handleLong).case(handleShort);
+
+``dag.switch`` declares the decider the way ``dag.if`` does. A case is the
task reference itself, so
+the compiler checks the candidate exists and renaming a handler cannot
silently rewire a Dag. The
+task's own value is the chosen task's id, which a downstream task can read
from its XCom.
+
+There is no default case. A decider that returns anything outside its cases
fails the task, naming
+what it chose and what it could have chosen.
+
+Exactly one case is selected. Python's branch callable may return a list of
task ids, and no language
+SDK offers that yet: put the paths that run together behind one task, or gate
each with its own
+condition.
+
``new Dag`` and ``dag.task`` both take a trailing spec of Airflow options:
``{ schedule: "@daily", tags: ["etl"] }`` for the Dag, ``{ retries: 2,
retryDelay: 30 }`` for a task.
diff --git a/ts-sdk/adr/0002-native-dag-interface.md
b/ts-sdk/adr/0002-native-dag-interface.md
index f1c4faa8dcf..7b56849e6f6 100644
--- a/ts-sdk/adr/0002-native-dag-interface.md
+++ b/ts-sdk/adr/0002-native-dag-interface.md
@@ -178,6 +178,38 @@ A one-sided `if` is a branch with one candidate — it skips
`then` and follows
`ShortCircuitOperator`, which would also skip the whole downstream closure and
ignore trigger rules.
A guarded task takes no argument for the control edge: a condition's boolean
is a signal, not data.
+### Multi-way branching: `switch` and `case`
+
+`dag.switch(pickPath)` follows
+[`airflow-core/adr/lang-sdk/0008`](../../airflow-core/adr/lang-sdk/0008-control-flow-constructs.md)
+decision 2 without divergence: a case **is** the reference the SDK handed
back, not a label kept in
+step with one. The decider is a handler, declared and wired in one call as
`dag.if` declares a condition, as
+Go's `dag.Switch(pickPath)` does.
+
+```ts
+async function pickPath({ rows }: { rows: number }): Promise<TaskRef> {
+ return rows > 1000 ? handleLong : handleShort;
+}
+
+dag.switch(pickPath, { rows: extracted }).case(handleLong).case(handleShort);
+```
+
+A case is a task, never a condition or a branch: the decider returns the case
from an async handler,
+and a condition carries `then`, so returning one would be awaited as a
thenable instead.
+
+An earlier draft selected a case by a string label the author writes, on the
grounds that a handler's
+function name does not survive bundling. That concern does not apply: a
`TaskRef` carries the task's
+own id, which the SDK fixed when the task was declared and esbuild never
touches. Selecting by
+reference keeps the compiler checking that a candidate exists, which a label
cannot.
+
+The cases chain, as they do in Go. `case` reads the candidate list when the
task runs rather than
+when it is declared, which is what lets the chain follow the `dag.switch`
call; and unlike a
+condition's `then`, `case` is not a thenable trap, so nothing has to be
guarded here.
+
+**No default case**, per decision 3, and **exactly one case is selected**, the
limitation that ADR
+records for every Lang SDK. A branch with no case at all decides nothing, and
is rejected when the
+Dag is read.
+
## Consequences
- One authoring surface (`dag.task()` plus its factory) covers the graph and
each task's arguments,
diff --git a/ts-sdk/src/index.ts b/ts-sdk/src/index.ts
index 09eb66af399..8c0f90eb81f 100644
--- a/ts-sdk/src/index.ts
+++ b/ts-sdk/src/index.ts
@@ -27,6 +27,7 @@ export { SUPERVISOR_API_VERSION } from
"./coordinator/index.js";
export type { ArgNameMap } from "./sdk/arg-names.js";
export type { Registerable } from "./sdk/bundle.js";
export type {
+ Branch,
Condition,
ConditionElse,
DeciderArgs,
diff --git a/ts-sdk/src/sdk/dag.ts b/ts-sdk/src/sdk/dag.ts
index f2d1e66f3b3..c285a4c9e88 100644
--- a/ts-sdk/src/sdk/dag.ts
+++ b/ts-sdk/src/sdk/dag.ts
@@ -418,14 +418,14 @@ export type TaskOptions = TaskSpec;
export interface Condition extends Node {
/** Airflow task ID of the deciding task. */
readonly taskId: string;
- then(taskRef: TaskRef | Condition): ConditionElse;
+ then(taskRef: TaskRef | Condition | Branch): ConditionElse;
before(...downstream: readonly Node[]): Condition;
after(...upstream: readonly Node[]): Condition;
}
/** What `.then(...)` returns: the other side, which a one-sided condition
omits. */
export interface ConditionElse {
- else(taskRef: TaskRef | Condition): void;
+ else(taskRef: TaskRef | Condition | Branch): void;
}
/** The arguments after a decider's handler: its inputs, then its {@link
TaskSpec}. */
@@ -433,6 +433,21 @@ export type DeciderArgs<TArgs extends object | void> =
[TArgs] extends [void]
? [inputs?: undefined, spec?: TaskOptions]
: [inputs: TaskInputs<TArgs>, spec?: TaskOptions];
+/**
+ * A placed multi-way branch: name each task the decider chooses between.
+ *
+ * ```ts
+ * dag.switch(pickPath, { rows: extracted
}).case(handleLong).case(handleShort);
+ * ```
+ */
+export interface Branch extends Node {
+ /** Airflow task ID of the deciding task. */
+ readonly taskId: string;
+ case(taskRef: TaskRef): Branch;
+ before(...downstream: readonly Node[]): Branch;
+ after(...upstream: readonly Node[]): Branch;
+}
+
/** Per-task record a Dag retains: the reference, the handler, and its spec. */
export interface TaskRecord {
readonly task: TaskRef;
@@ -512,6 +527,7 @@ export class Dag {
readonly #orderEdges = new Map<string, OrderEdge>();
readonly #definedIn: string | undefined;
readonly #conditions = new Map<string, ConditionRecord>();
+ readonly #branches = new Map<string, readonly TaskRef[]>();
// Keyed by full group ID; a group's own record holds what it declares, so
// the tree is reconstructed by walking from the roots.
readonly #groups = new Map<string, MutableTaskGroupRecord>();
@@ -621,7 +637,7 @@ export class Dag {
});
const named = new Set<"then" | "else">();
- const name = (side: "then" | "else", target: TaskRef | Condition): void =>
{
+ const name = (side: "then" | "else", target: TaskRef | Condition |
Branch): void => {
// `await` calls `then(resolve, reject)`, so a function here means the
condition was awaited.
if (typeof target === "function") {
throw new Error(
@@ -665,14 +681,7 @@ export class Dag {
name("then", taskRef);
return elseStep;
},
- before: (...downstream) => {
- condition.before(...downstream);
- return placed;
- },
- after: (...upstream) => {
- condition.after(...upstream);
- return placed;
- },
+ ...this.#forwardEdges(condition, () => placed),
};
conditionTasks.set(placed, condition);
return Object.freeze(placed);
@@ -697,6 +706,76 @@ export class Dag {
this.#tasks.set(taskId, { ...record, canSkipDownstream: true, fn: wrapped
});
}
+ /**
+ * Declare a task that returns one of its cases; every other case is skipped.
+ *
+ * ```ts
+ * dag.switch(pickPath, { rows: extracted
}).case(handleLong).case(handleShort);
+ * ```
+ */
+ switch<TArgs extends object | void = void>(
+ handler: (args: TArgs) => TaskRef | Promise<TaskRef>,
+ ...args: DeciderArgs<NoInfer<TArgs>>
+ ): Branch {
+ return this.#placeBranch(this.#placeDecider(handler, args) as
TaskRef<TaskRef>);
+ }
+
+ #forwardEdges<T>(ref: TaskRef, self: () => T) {
+ return {
+ before: (...downstream: readonly Node[]) => {
+ ref.before(...downstream);
+ return self();
+ },
+ after: (...upstream: readonly Node[]) => {
+ ref.after(...upstream);
+ return self();
+ },
+ };
+ }
+
+ #placeBranch(decider: TaskRef<TaskRef>): Branch {
+ const taskId = decider.taskId;
+ const candidates: TaskRef[] = [];
+ this.#branches.set(taskId, candidates);
+ this.#wrapDecider(taskId, async (chosen: unknown) => {
+ const known = candidates.map((ref) => ref.taskId);
+ const picked = isTaskRef(chosen) && chosen.dagId === this.dagId ?
chosen.taskId : undefined;
+ if (picked === undefined || !known.includes(picked)) {
+ throw new Error(
+ `Task "${taskId}" of Dag "${this.dagId}" chose ` +
+ `${isTaskRef(chosen) ? `"${chosen.taskId}"` :
describeValue(chosen)}, ` +
+ `which is not one of its cases: ${known.join(", ")}`,
+ );
+ }
+ return { skip: known.filter((id) => id !== picked), result: picked };
+ });
+
+ const branch: Branch = {
+ dagId: this.dagId,
+ taskId,
+ case: (taskRef) => {
+ this.#validateOwnNode(taskRef, `a case of "${taskId}"`);
+ if (!isTaskRef(taskRef)) {
+ throw new Error(
+ `A case of Dag "${this.dagId}" branch "${taskId}" has to be a
task, not a task group`,
+ );
+ }
+ if (candidates.some((candidate) => candidate.taskId ===
taskRef.taskId)) {
+ throw new Error(
+ `Dag "${this.dagId}" branch "${taskId}" lists "${taskRef.taskId}"
twice; ` +
+ "each case names a different task",
+ );
+ }
+ candidates.push(taskRef);
+ decider.before(taskRef);
+ return branch;
+ },
+ ...this.#forwardEdges(decider, () => branch),
+ };
+ conditionTasks.set(branch, decider);
+ return Object.freeze(branch);
+ }
+
/**
* Declare a task group of this Dag.
*
@@ -1105,6 +1184,14 @@ export class Dag {
);
}
}
+ for (const [taskId, cases] of this.#branches) {
+ if (cases.length === 0) {
+ throw new Error(
+ `Branch "${taskId}" of Dag "${this.dagId}" has no cases, so it
decides nothing; ` +
+ "give it the tasks to choose between with dag.switch(handler,
inputs).case(task)",
+ );
+ }
+ }
for (const taskId of this.#tasks.keys()) {
if (!this.#inputs.has(taskId)) {
throw new Error(
diff --git a/ts-sdk/tests/sdk/branching.test.ts
b/ts-sdk/tests/sdk/branching.test.ts
new file mode 100644
index 00000000000..fb31c2022ee
--- /dev/null
+++ b/ts-sdk/tests/sdk/branching.test.ts
@@ -0,0 +1,242 @@
+/*!
+ * 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.
+ */
+
+import { describe, expect, it, vi } from "vitest";
+import {
+ Dag,
+ finalizeDag,
+ getDagOrderEdges,
+ getDagTaskInputs,
+ getDagTaskRecords,
+ type TaskRef,
+} from "../../src/sdk/dag.js";
+import { serializeDag } from "../../src/coordinator/serde.js";
+import { brand } from "../../src/sdk/brand.js";
+import type { TaskClient } from "../../src/sdk/client.js";
+import { runInTaskScope, type TaskContext } from "../../src/sdk/task.js";
+
+function branchedDag(dagId = "d", choose?: (refs: Record<string, TaskRef>) =>
unknown) {
+ const dag = new Dag(dagId);
+ const long = dag.task("handle_long", async () => undefined)();
+ const short = dag.task("handle_short", async () => undefined)();
+ const other = dag.task("handle_other", async () => undefined)();
+ const refs = { long, short, other };
+ const decider = dag.switch(async () => (choose ? (choose(refs) as TaskRef) :
long), undefined, {
+ taskId: "pick_path",
+ });
+ return { dag, long, short, other, decider };
+}
+
+function branded<T extends object>(value: T): T {
+ brand(value, "TaskRef");
+ return value;
+}
+
+async function runHandler(dag: Dag, taskId: string, args: unknown = {}) {
+ const skipDownstreamTasks = vi.fn(async () => undefined);
+ const setXCom = vi.fn(async () => undefined);
+ const client = { skipDownstreamTasks, setXCom } as unknown as TaskClient;
+ const ctx = { dagId: dag.dagId, taskId } as unknown as TaskContext;
+ const handler = getDagTaskRecords(dag).get(taskId)!.fn;
+ const returned = await runInTaskScope({ ctx, client }, () =>
+ (handler as (args: unknown) => Promise<unknown>)(args),
+ );
+ return { returned, skipDownstreamTasks, setXCom };
+}
+
+describe("dag.switch", () => {
+ it("marks the decider as deciding skips, which Airflow reads on a clear", ()
=> {
+ const { dag, long, short, decider } = branchedDag();
+ decider.case(long).case(short);
+
+ const tasks = (serializeDag(dag, "", ".") as { tasks: { __var:
Record<string, unknown> }[] })
+ .tasks;
+ const byId = new Map(tasks.map(({ __var }) => [__var["task_id"] as string,
__var]));
+ expect(byId.get("pick_path")?.["_can_skip_downstream"]).toBe(true);
+ expect(byId.get("handle_long")).not.toHaveProperty("_can_skip_downstream");
+ });
+
+ it("draws an order-only edge to each case, since no value flows", () => {
+ const { dag, long, short, decider } = branchedDag();
+ decider.case(long).case(short);
+
+ expect(getDagOrderEdges(dag)).toEqual([
+ { upstream: "pick_path", downstream: "handle_long" },
+ { upstream: "pick_path", downstream: "handle_short" },
+ ]);
+ });
+
+ describe("at run time", () => {
+ it("returns the chosen task's id, which is what a branch puts on the
wire", async () => {
+ const { dag, long, short, decider } = branchedDag("d", ({ short: s }) =>
s);
+ decider.case(long).case(short);
+
+ expect((await runHandler(dag,
"pick_path")).returned).toBe("handle_short");
+ });
+
+ it("skips every case it did not choose", async () => {
+ const { dag, long, short, other, decider } = branchedDag("d", ({ long: l
}) => l);
+ decider.case(long).case(short).case(other);
+
+ const { skipDownstreamTasks, setXCom } = await runHandler(dag,
"pick_path");
+
+ expect(skipDownstreamTasks).toHaveBeenCalledWith(["handle_short",
"handle_other"]);
+ expect(setXCom).toHaveBeenCalledWith({
+ key: "skipmixin_key",
+ value: { skipped: ["handle_short", "handle_other"] },
+ });
+ });
+
+ it("accepts a different reference to a case, matching it by task id",
async () => {
+ const { dag, long, short, decider } = branchedDag("d", () =>
+ branded({ dagId: "d", taskId: "handle_short" }),
+ );
+ decider.case(long).case(short);
+
+ const { returned, skipDownstreamTasks } = await runHandler(dag,
"pick_path");
+
+ expect(returned).toBe("handle_short");
+ expect(skipDownstreamTasks).toHaveBeenCalledWith(["handle_long"]);
+ });
+
+ it.each([
+ ["a string", "handle_long"],
+ ["a look-alike object", { dagId: "d", taskId: "handle_long" }],
+ ["another Dag's task of the same id", branded({ dagId: "other", taskId:
"handle_long" })],
+ ["null", null],
+ ])("fails the task when the decider returns %s", async (_label, value) => {
+ const { dag, long, short, decider } = branchedDag("d", () => value);
+ decider.case(long).case(short);
+
+ await expect(runHandler(dag, "pick_path")).rejects.toThrow(/which is not
one of its cases/);
+ });
+
+ it("skips nothing and fails when the decider chose wrongly", async () => {
+ const { dag, long, short, decider } = branchedDag("d", ({ other }) =>
other);
+ decider.case(long).case(short);
+
+ const skipDownstreamTasks = vi.fn(async () => undefined);
+ const client = { skipDownstreamTasks, setXCom: vi.fn() } as unknown as
TaskClient;
+ const ctx = { dagId: "d", taskId: "pick_path" } as unknown as
TaskContext;
+ const handler = getDagTaskRecords(dag).get("pick_path")!.fn;
+
+ await expect(
+ runInTaskScope({ ctx, client }, () => (handler as (args: unknown) =>
Promise<unknown>)({})),
+ ).rejects.toThrow(
+ /Task "pick_path" of Dag "d" chose "handle_other", which is not one of
its cases: handle_long, handle_short/,
+ );
+ expect(skipDownstreamTasks).not.toHaveBeenCalled();
+ });
+ });
+
+ it("carries edges of its own, like any task", () => {
+ const { dag, long, decider } = branchedDag();
+ const extracted = dag.task("extract", async () => undefined)();
+ decider.case(long);
+
+ decider.after(extracted);
+
+ expect(getDagOrderEdges(dag)).toContainEqual({
+ upstream: "extract",
+ downstream: "pick_path",
+ });
+ });
+
+ it("takes the decider's id from the handler's name, and its inputs from the
factory", () => {
+ const dag = new Dag("d");
+ const long = dag.task("handle_long", async () => undefined)();
+ const extracted = dag.task("extract", async () => 3)();
+ async function pickPath({ rows }: { rows: number }): Promise<TaskRef> {
+ return rows > 1 ? long : long;
+ }
+ const pick = dag.switch(pickPath, { rows: extracted });
+ pick.case(long);
+
+ expect(pick.taskId).toBe("pickPath");
+ expect(getDagTaskInputs(dag).get("pickPath")).toEqual({ rows: extracted });
+ });
+
+ it("stands at the other end of an edge, as Go's .After(pick) does", () => {
+ const { dag, long, decider } = branchedDag();
+ const notified = dag.task("notify", async () => undefined)();
+ decider.case(long);
+
+ notified.after(decider);
+
+ expect(getDagOrderEdges(dag)).toContainEqual({ upstream: "pick_path",
downstream: "notify" });
+ });
+
+ it("is a branch of a condition, so a switch can sit under dag.if", () => {
+ const { dag, long, short, decider } = branchedDag();
+ decider.case(long).case(short);
+ const gate = dag.if(async () => true, undefined, { taskId: "has_rows" });
+
+ gate.then(decider);
+
+ expect(getDagOrderEdges(dag)).toContainEqual({ upstream: "has_rows",
downstream: "pick_path" });
+ });
+
+ describe("rejects", () => {
+ it("a task group as a case", () => {
+ const { dag, decider } = branchedDag();
+ const group = dag.taskGroup("staging");
+
+ expect(() => decider.case(group as never)).toThrowError(
+ /A case of Dag "d" branch "pick_path" has to be a task, not a task
group/,
+ );
+ });
+
+ it("a case taken from another Dag", () => {
+ const { decider } = branchedDag("here");
+ const { long: foreign } = branchedDag("there");
+
+ expect(() => decider.case(foreign)).toThrowError(
+ /a case of "pick_path" cannot reach Dag "there" node "handle_long"/,
+ );
+ });
+
+ it.each([
+ ["a plain object", { dagId: "d", taskId: "handle_long" }],
+ ["a string", "handle_long"],
+ ["null", null],
+ ])("%s where a case belongs", (_label, value) => {
+ const { decider } = branchedDag();
+
+ expect(() => decider.case(value as TaskRef)).toThrowError(
+ /a case of "pick_path" on Dag "d" takes tasks and task groups/,
+ );
+ });
+
+ it("the same task listed twice", () => {
+ const { long, decider } = branchedDag();
+
+ expect(() => decider.case(long).case(long)).toThrowError(
+ /Dag "d" branch "pick_path" lists "handle_long" twice; each case names
a different task/,
+ );
+ });
+
+ it("a branch with no case, when the Dag is read", () => {
+ const { dag } = branchedDag();
+
+ expect(() => finalizeDag(dag)).toThrowError(
+ /Branch "pick_path" of Dag "d" has no cases, so it decides nothing/,
+ );
+ });
+ });
+});