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/,
+      );
+    });
+  });
+});

Reply via email to