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 62917eab648 TS SDK: Add conditional branching with dag.if (#74069)
62917eab648 is described below

commit 62917eab648e7298d191b583b6b3e4ff30110ff0
Author: Guan-Ming Chiu <[email protected]>
AuthorDate: Wed Oct 7 10:44:32 2026 +0800

    TS SDK: Add conditional branching with dag.if (#74069)
    
    Co-authored-by: Jason(Zhe-You) Liu 
<[email protected]>
---
 .../language-sdks/typescript.rst                   |  28 ++
 ts-sdk/adr/0002-native-dag-interface.md            |  43 +++
 ts-sdk/src/coordinator/client.ts                   |  12 +
 ts-sdk/src/coordinator/protocol.ts                 |   1 +
 ts-sdk/src/coordinator/serde.ts                    |   4 +
 ts-sdk/src/index.ts                                |   3 +
 ts-sdk/src/sdk/dag.ts                              | 189 ++++++++++++-
 ts-sdk/tests/coordinator/client.test.ts            |  18 ++
 ts-sdk/tests/sdk/conditional.test.ts               | 308 +++++++++++++++++++++
 ts-sdk/tests/sdk/dag.test.ts                       |   4 +-
 10 files changed, 600 insertions(+), 10 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 8b33da1e640..a8ef96d5f01 100644
--- a/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst
+++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst
@@ -391,6 +391,34 @@ there. Set it once on the Dag and each task inherits it:
 ``queue`` on a task wins over the Dag's. See 
:ref:`typescript-sdk/coordinator-config` for the
 ``queue_to_coordinator`` entry that sends that queue to the coordinator.
 
+Conditional branching
+~~~~~~~~~~~~~~~~~~~~~
+
+``dag.if`` takes a handler that returns a boolean, and names the task each 
outcome runs:
+
+.. code-block:: typescript
+
+    async function hasRows({ rows }: { rows: number }): Promise<boolean> {
+      return rows > 0;
+    }
+
+    const gate = dag.if(hasRows, { rows: extracted });
+    gate.then(loaded).else(reportedEmpty);
+
+``dag.if`` declares the condition as a task: its id is the function's name, 
and a trailing spec sets
+its id or other options, as in ``dag.if(hasRows, { rows }, { taskId: 
"has_rows" })``. The compiler
+checks that the handler returns a boolean. ``else`` is optional: a one-sided 
condition skips its own
+branch when the condition fails and follows nothing. The condition is also a 
node, so
+``notified.after(gate)`` orders a task after it.
+
+A guarded task takes no argument for the control edge, because a condition's 
boolean decides whether
+the task runs rather than what it runs on. Read a value from the condition with
+``getClient().getXCom``.
+
+The side not taken is skipped when the run reaches it, and stays skipped if 
you clear it later. Only
+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.
+
 ``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 6614ee62091..f1c4faa8dcf 100644
--- a/ts-sdk/adr/0002-native-dag-interface.md
+++ b/ts-sdk/adr/0002-native-dag-interface.md
@@ -135,6 +135,49 @@ data is the wiring object itself — `summarize({ north: 
extractNorth(), south:
 does. This matches `Before`/`After` in the Go SDK's native Dag interface, 
spelled to TypeScript
 convention.
 
+### Conditional branching: `if` and `else`
+
+`dag.if(handler)` is TypeScript's spelling of the construct
+[`airflow-core/adr/lang-sdk/0008`](../../airflow-core/adr/lang-sdk/0008-control-flow-constructs.md)
+names after the host language's control flow:
+
+```ts
+async function hasRows({ rows }: { rows: number }): Promise<boolean> {
+  return rows > 0;
+}
+
+const gate = dag.if(hasRows, { rows: validated });
+gate.then(loadIfReady).else(loadFallback);
+```
+
+**The condition is a handler, declared and wired in one call.** That ADR writes
+`dag.If(hasRows, airflow.Inputs(validated), airflow.TaskSpec{...})` in Go, and
+`dag.if(hasRows, { rows: validated }, { ... })` reads the same way: inputs, 
then an optional
+`TaskSpec`. The deciding task takes its id from the function's name unless the 
spec's `taskId` sets
+it, and the spec carries its other options, such as `queue` and `retries`. The 
compiler checks that
+the handler returns a `boolean`, and that the inputs match its argument.
+
+**A condition is a node.** What `dag.if` returns carries `before` and `after`, 
and stands at the
+other end of an edge as the deciding task, so `notified.after(gate)` reads as 
Go's `.After(gate)`. It
+is also a valid branch of another condition.
+
+**A `then` chain is a thenable, and is guarded rather than avoided.** An 
object with a callable
+`then` is a *thenable*: were the condition `dag.if` returns to reach an 
`await`, the runtime would hand
+its `then` a resolve function where a task reference belongs. Two things 
contain that. `.then(...)`
+returns an object carrying only `.else`, so nothing past the first step is 
awaitable at all; and
+`.then` rejects a function argument by naming the cause, so an author who does 
await it reads "this
+builds a branch, drop the await" rather than a type error about references.
+
+**A branch is a real branch to Airflow.** The control edges serialize as 
ordinary order-only edges
+and carry no branch-candidate field, as that ADR's consequences require. The 
condition task is
+serialized with `_can_skip_downstream`, and it writes the `skipmixin_key` XCom 
alongside the skip, so
+clearing a skipped branch re-skips it the way a Python `@task.branch` does 
rather than running the
+side the condition rejected.
+
+A one-sided `if` is a branch with one candidate — it skips `then` and follows 
nothing — rather than a
+`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.
+
 ## Consequences
 
 - One authoring surface (`dag.task()` plus its factory) covers the graph and 
each task's arguments,
diff --git a/ts-sdk/src/coordinator/client.ts b/ts-sdk/src/coordinator/client.ts
index ba0be84faa2..9103d824d58 100644
--- a/ts-sdk/src/coordinator/client.ts
+++ b/ts-sdk/src/coordinator/client.ts
@@ -30,6 +30,7 @@ import type {
   GetXCom,
   SetXCom,
   GetConnection,
+  SkipDownstreamTasks,
   ConnectionResult as WireConnectionResult,
 } from "./protocol.js";
 
@@ -82,6 +83,9 @@ export interface CoordinatorClient extends TaskClient {
    * output fails the task, while one that pushed null binds null.
    */
   getXComEntry(opts: GetXComOpts): Promise<XComEntry>;
+
+  /** Mark direct downstream tasks of the running task as skipped; none is a 
no-op. */
+  skipDownstreamTasks(taskIds: readonly string[]): Promise<void>;
 }
 
 export function createCoordinatorClient(
@@ -202,6 +206,14 @@ export function createCoordinatorClient(
       await rpc("SetXCom", null, msg, () => undefined, "throw");
     },
 
+    // ---- Control flow ----
+
+    async skipDownstreamTasks(taskIds: readonly string[]): Promise<void> {
+      if (taskIds.length === 0) return;
+      const msg: SkipDownstreamTasks = { type: "SkipDownstreamTasks", tasks: 
[...taskIds] };
+      await rpc("SkipDownstreamTasks", null, msg, () => undefined, "throw");
+    },
+
     // ---- Connections ----
 
     async getConnection(connId: string): Promise<ConnectionResult | null> {
diff --git a/ts-sdk/src/coordinator/protocol.ts 
b/ts-sdk/src/coordinator/protocol.ts
index 323765036ea..1c95fee6f90 100644
--- a/ts-sdk/src/coordinator/protocol.ts
+++ b/ts-sdk/src/coordinator/protocol.ts
@@ -64,6 +64,7 @@ export type {
   GetXCom,
   SetXCom,
   GetConnection,
+  SkipDownstreamTasks,
 } from "../generated/supervisor.js";
 
 // -------- Frames from supervisor --------
diff --git a/ts-sdk/src/coordinator/serde.ts b/ts-sdk/src/coordinator/serde.ts
index 723a073b867..a26bf11837c 100644
--- a/ts-sdk/src/coordinator/serde.ts
+++ b/ts-sdk/src/coordinator/serde.ts
@@ -147,6 +147,7 @@ export function serializeDag(
         withDagQueue(record.spec, dag.spec.queue),
         graph.downstreamTaskIds.get(taskId),
         inputs.get(taskId),
+        record.canSkipDownstream === true,
       ),
     ),
     dag_dependencies: [],
@@ -184,6 +185,7 @@ function serializeTask(
   spec: object,
   downstream: ReadonlySet<string> | undefined,
   inputs: RecordedInputs | undefined,
+  canSkipDownstream: boolean,
 ): SerializedValue {
   const data: Record<string, SerializedValue> = {
     task_id: taskId,
@@ -202,6 +204,8 @@ function serializeTask(
   const label = `task "${taskId}" of Dag "${dagId}"`;
   const bindings = serializeArgBindings(inputs, label);
   if (bindings) data["_arg_bindings"] = bindings;
+  // Lets `NotPreviouslySkippedDep` re-skip a cleared downstream, as Python's 
SkipMixin does.
+  if (canSkipDownstream) data["_can_skip_downstream"] = true;
   applySchemaFields(data, spec, TASK_FIELD_RULES, label);
   if (downstream?.size) {
     data["downstream_task_ids"] = [...downstream].sort();
diff --git a/ts-sdk/src/index.ts b/ts-sdk/src/index.ts
index 4e7bee9a966..09eb66af399 100644
--- a/ts-sdk/src/index.ts
+++ b/ts-sdk/src/index.ts
@@ -27,6 +27,9 @@ 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 {
+  Condition,
+  ConditionElse,
+  DeciderArgs,
   DagSpec,
   Node,
   TaskFactory,
diff --git a/ts-sdk/src/sdk/dag.ts b/ts-sdk/src/sdk/dag.ts
index f643e290f45..f2d1e66f3b3 100644
--- a/ts-sdk/src/sdk/dag.ts
+++ b/ts-sdk/src/sdk/dag.ts
@@ -22,6 +22,7 @@
 // Dag and supplies its arguments, the way calling a TaskFlow function does in
 // Python.
 
+import type { CoordinatorClient } from "../coordinator/client.js";
 import {
   DAG_SCHEMA_FIELDS,
   TASK_SCHEMA_FIELDS,
@@ -31,7 +32,7 @@ import {
 import { brand, DUPLICATE_COPY_HINT, hasBrand } from "./brand.js";
 import type { JsonValue } from "./client-types.js";
 import { getCurrentModuleSource } from "./module-source.js";
-import type { TaskFunction } from "./task.js";
+import { getClient, type TaskFunction } from "./task.js";
 
 /** Internal: whether `value` is an object literal, not an array or a class 
instance. */
 export function isPlainRecord(value: unknown): value is Record<string, 
unknown> {
@@ -303,6 +304,12 @@ export function isTaskRef(value: unknown): value is 
TaskRef {
   return hasBrand(value, "TaskRef");
 }
 
+const conditionTasks = new WeakMap<object, TaskRef>();
+
+function resolveNode(node: Node): Node {
+  return (typeof node === "object" && node !== null && 
conditionTasks.get(node)) || node;
+}
+
 /**
  * The first reference reachable inside `value`, or `undefined` when there is
  * none.
@@ -401,13 +408,49 @@ export type TaskFactory<TArgs extends object | void = 
void, TReturn = unknown> =
  */
 export type TaskOptions = TaskSpec;
 
+/**
+ * A placed condition: name the task each outcome runs. `else` is optional.
+ *
+ * ```ts
+ * dag.if(hasRows, { rows: validated }).then(loadIfReady).else(loadFallback);
+ * ```
+ */
+export interface Condition extends Node {
+  /** Airflow task ID of the deciding task. */
+  readonly taskId: string;
+  then(taskRef: TaskRef | Condition): 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;
+}
+
+/** The arguments after a decider's handler: its inputs, then its {@link 
TaskSpec}. */
+export type DeciderArgs<TArgs extends object | void> = [TArgs] extends [void]
+  ? [inputs?: undefined, spec?: TaskOptions]
+  : [inputs: TaskInputs<TArgs>, spec?: TaskOptions];
+
 /** Per-task record a Dag retains: the reference, the handler, and its spec. */
 export interface TaskRecord {
   readonly task: TaskRef;
   readonly fn: TaskFunction;
   readonly spec: TaskSpec;
+  /** Whether this task decides which of its downstream tasks to skip. */
+  readonly canSkipDownstream?: boolean;
+}
+
+interface ConditionRecord {
+  whenTrue?: TaskRef;
+  whenFalse?: TaskRef;
 }
 
+// Mirrors `SkipMixin.skip` in `task-sdk/src/airflow/sdk/bases/skipmixin.py`.
+const SKIPMIXIN_XCOM_KEY = "skipmixin_key";
+const SKIPMIXIN_SKIPPED = "skipped";
+
 /**
  * Internal: what one call to a task factory recorded, by argument name.
  *
@@ -468,6 +511,7 @@ export class Dag {
   // insertion-ordered so the serialized Dag reads as written.
   readonly #orderEdges = new Map<string, OrderEdge>();
   readonly #definedIn: string | undefined;
+  readonly #conditions = new Map<string, ConditionRecord>();
   // 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>();
@@ -540,6 +584,119 @@ export class Dag {
     return this.#addTask(undefined, taskIdOrHandler, handlerOrOptions, 
maybeOptions);
   }
 
+  /**
+   * Declare a task whose boolean picks a branch; the side not taken is 
skipped.
+   *
+   * ```ts
+   * dag.if(hasRows, { rows: validated }).then(loadIfReady).else(loadFallback);
+   * dag.if(hasRows, { rows: validated }, { taskId: "has_rows", retries: 2 
}).then(loadIfReady);
+   * ```
+   */
+  if<TArgs extends object | void = void>(
+    handler: (args: TArgs) => boolean | Promise<boolean>,
+    ...args: DeciderArgs<NoInfer<TArgs>>
+  ): Condition {
+    return this.#placeCondition(this.#placeDecider(handler, args) as 
TaskRef<boolean>);
+  }
+
+  #placeDecider(handler: (args: never) => unknown, args: readonly unknown[]): 
TaskRef {
+    const [inputs, spec] = args as [unknown, TaskOptions | undefined];
+    const factory = this.#addTask(undefined, handler, spec) as (...inputs: 
unknown[]) => TaskRef;
+    return inputs === undefined ? factory() : factory(inputs);
+  }
+
+  #placeCondition(condition: TaskRef<boolean>): Condition {
+    const taskId = condition.taskId;
+    const branches: ConditionRecord = {};
+    this.#conditions.set(taskId, branches);
+    this.#wrapDecider(taskId, async (held: unknown) => {
+      if (typeof held !== "boolean") {
+        throw new Error(
+          `Condition "${taskId}" of Dag "${this.dagId}" returned 
${describeValue(held)} ` +
+            "rather than a boolean, so there is no branch to take",
+        );
+      }
+      const skipped = held ? branches.whenFalse : branches.whenTrue;
+      return { skip: skipped ? [skipped.taskId] : [], result: held };
+    });
+
+    const named = new Set<"then" | "else">();
+    const name = (side: "then" | "else", target: TaskRef | Condition): void => 
{
+      // `await` calls `then(resolve, reject)`, so a function here means the 
condition was awaited.
+      if (typeof target === "function") {
+        throw new Error(
+          `dag.if(...) of Dag "${this.dagId}" was awaited. It builds a branch 
rather than ` +
+            "doing work, so there is nothing to wait for; drop the await",
+        );
+      }
+      if (named.has(side)) {
+        throw new Error(
+          `Condition "${taskId}" of Dag "${this.dagId}" already has a 
"${side}" branch; ` +
+            "a condition names each side once",
+        );
+      }
+      const taskRef = resolveNode(target) as TaskRef;
+      this.#validateOwnNode(taskRef, `the "${side}" branch of "${taskId}"`);
+      if (!isTaskRef(taskRef)) {
+        throw new Error(
+          `The "${side}" branch of Dag "${this.dagId}" condition "${taskId}" 
has to be a task, ` +
+            "not a task group",
+        );
+      }
+      if (side === "else" && taskRef === branches.whenTrue) {
+        throw new Error(
+          `Both branches of Dag "${this.dagId}" condition "${taskId}" are ` +
+            `"${taskRef.taskId}", so the condition decides nothing; drop the 
else branch`,
+        );
+      }
+      named.add(side);
+      if (side === "then") branches.whenTrue = taskRef;
+      else branches.whenFalse = taskRef;
+      condition.before(taskRef);
+    };
+
+    const elseStep: ConditionElse = {
+      else: (taskRef) => name("else", taskRef),
+    };
+    const placed: Condition = {
+      dagId: this.dagId,
+      taskId,
+      then: (taskRef) => {
+        name("then", taskRef);
+        return elseStep;
+      },
+      before: (...downstream) => {
+        condition.before(...downstream);
+        return placed;
+      },
+      after: (...upstream) => {
+        condition.after(...upstream);
+        return placed;
+      },
+    };
+    conditionTasks.set(placed, condition);
+    return Object.freeze(placed);
+  }
+
+  #wrapDecider(
+    taskId: string,
+    decide: (returned: unknown) => Promise<{ skip: string[]; result: unknown 
}>,
+  ): void {
+    const record = this.#tasks.get(taskId)!;
+    const inner = record.fn;
+    const wrapped: TaskFunction = async (args) => {
+      const { skip, result } = await decide(await inner(args as never));
+      if (skip.length > 0) {
+        const client = getClient() as CoordinatorClient;
+        // Written before the skip so a cleared downstream is re-skipped, as 
SkipMixin does.
+        await client.setXCom({ key: SKIPMIXIN_XCOM_KEY, value: { 
[SKIPMIXIN_SKIPPED]: skip } });
+        await client.skipDownstreamTasks(skip);
+      }
+      return result;
+    };
+    this.#tasks.set(taskId, { ...record, canSkipDownstream: true, fn: wrapped 
});
+  }
+
   /**
    * Declare a task group of this Dag.
    *
@@ -804,7 +961,9 @@ export class Dag {
     return Object.freeze(task);
   }
 
-  #addOrderEdge(upstream: Node, downstream: Node, verb: "before" | "after"): 
void {
+  #addOrderEdge(upstreamNode: Node, downstreamNode: Node, verb: "before" | 
"after"): void {
+    const upstream = resolveNode(upstreamNode);
+    const downstream = resolveNode(downstreamNode);
     if (this.#finalized) {
       throw new Error(
         `An edge was drawn on Dag "${this.dagId}" after the Dag was read; ` +
@@ -814,7 +973,7 @@ export class Dag {
     // The argument is the one that can be foreign: the receiver is a node this
     // Dag handed out, since it is what carries the method.
     const other = verb === "before" ? downstream : upstream;
-    this.#validateOwnNode(other, verb);
+    this.#validateOwnNode(other, `${verb}()`);
     const upstreamId = nodeId(upstream);
     const downstreamId = nodeId(downstream);
     if (upstreamId === downstreamId) {
@@ -833,17 +992,17 @@ export class Dag {
     }
   }
 
-  #validateOwnNode(node: Node, verb: string): void {
+  #validateOwnNode(node: Node, label: string): void {
     const id = nodeId(node);
     if (id === undefined) {
       throw new Error(
-        `${verb}() on Dag "${this.dagId}" takes tasks and task groups this Dag 
handed out, ` +
+        `${label} on Dag "${this.dagId}" takes tasks and task groups this Dag 
handed out, ` +
           "not arbitrary values",
       );
     }
     if (node.dagId !== this.dagId) {
       throw new Error(
-        `${verb}() cannot draw an edge to Dag "${node.dagId}" node "${id}" 
from Dag ` +
+        `${label} cannot reach Dag "${node.dagId}" node "${id}" from Dag ` +
           `"${this.dagId}"; an edge joins two nodes of one Dag`,
       );
     }
@@ -853,7 +1012,7 @@ export class Dag {
     // the edge at this Dag's own group of that name.
     if (isTaskRef(node) ? this.#tasks.get(id)?.task !== node : 
this.#groupRefs.get(id) !== node) {
       throw new Error(
-        `${verb}() was given a reference to "${id}" that this Dag did not hand 
out; ` +
+        `${label} was given a reference to "${id}" that this Dag did not hand 
out; ` +
           `it comes from another Dag object with the same ID, or 
${DUPLICATE_COPY_HINT}`,
       );
     }
@@ -938,6 +1097,14 @@ export class Dag {
 
   #finalize(): void {
     if (this.#finalized) return;
+    for (const [taskId, branches] of this.#conditions) {
+      if (branches.whenTrue === undefined) {
+        throw new Error(
+          `Condition "${taskId}" of Dag "${this.dagId}" names no branch, so it 
decides nothing; ` +
+            "give it one with dag.if(handler, inputs).then(task)",
+        );
+      }
+    }
     for (const taskId of this.#tasks.keys()) {
       if (!this.#inputs.has(taskId)) {
         throw new Error(
@@ -961,6 +1128,14 @@ const TASK_ID_CHARACTERS = /^[\p{L}\p{N}_.-]+$/u;
 const GROUP_ID_CHARACTERS = /^[\p{L}\p{N}_-]+$/u;
 const GROUP_ID_MAX_LENGTH = 200;
 
+/** A value as an error message names it: its type, or the literal when short. 
*/
+function describeValue(value: unknown): string {
+  if (value === null) return "null";
+  if (value === undefined) return "undefined";
+  if (typeof value === "object") return Array.isArray(value) ? "an array" : 
"an object";
+  return `the ${typeof value} ${typeof value === "string" ? 
JSON.stringify(value) : String(value)}`;
+}
+
 /**
  * A handler's own function name, or undefined when it has none.
  *
diff --git a/ts-sdk/tests/coordinator/client.test.ts 
b/ts-sdk/tests/coordinator/client.test.ts
index 2934eec32f7..cba7154b516 100644
--- a/ts-sdk/tests/coordinator/client.test.ts
+++ b/ts-sdk/tests/coordinator/client.test.ts
@@ -134,6 +134,24 @@ describe("deleteVariable", () => {
   });
 });
 
+describe("skipDownstreamTasks", () => {
+  it("sends SkipDownstreamTasks with the task ids", async () => {
+    const { client: c, sent } = recordingClient({ type: "OKResponse", ok: true 
});
+
+    await c.skipDownstreamTasks(["load_fallback"]);
+
+    expect(sent).toEqual([{ type: "SkipDownstreamTasks", tasks: 
["load_fallback"] }]);
+  });
+
+  it("sends nothing for an empty list", async () => {
+    const { client: c, sent } = recordingClient();
+
+    await c.skipDownstreamTasks([]);
+
+    expect(sent).toEqual([]);
+  });
+});
+
 describe("writes do not read a supervisor 404 as absence", () => {
   it.each([
     ["setVariable", "PutVariable", (c: TaskClient) => c.setVariable("k", "v")],
diff --git a/ts-sdk/tests/sdk/conditional.test.ts 
b/ts-sdk/tests/sdk/conditional.test.ts
new file mode 100644
index 00000000000..f2487f2af6a
--- /dev/null
+++ b/ts-sdk/tests/sdk/conditional.test.ts
@@ -0,0 +1,308 @@
+/*!
+ * 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 type { TaskClient } from "../../src/sdk/client.js";
+import { runInTaskScope, type TaskContext } from "../../src/sdk/task.js";
+
+function gatedDag(dagId = "d", holds: boolean | unknown = true) {
+  const dag = new Dag(dagId);
+  const ifReady = dag.task("load_if_ready", async () => undefined)();
+  const fallback = dag.task("load_fallback", async () => undefined)();
+  const gated = dag.if(async () => holds as boolean, undefined, { taskId: 
"has_rows" });
+  return { dag, ifReady, fallback, gated };
+}
+
+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.if", () => {
+  it("serializes the condition as a task that decides skips", () => {
+    const { dag, ifReady, fallback, gated } = gatedDag();
+    gated.then(ifReady).else(fallback);
+
+    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("has_rows")?.["_can_skip_downstream"]).toBe(true);
+    
expect(byId.get("load_if_ready")).not.toHaveProperty("_can_skip_downstream");
+  });
+
+  it("serializes the spec given to the condition", () => {
+    const { dag, ifReady } = gatedDag();
+    dag
+      .if(async () => true, undefined, { taskId: "is_weekday", retries: 2, 
queue: "q" })
+      .then(ifReady);
+
+    const tasks = (serializeDag(dag, "", ".") as { tasks: { __var: 
Record<string, unknown> }[] })
+      .tasks;
+    expect(tasks.find(({ __var }) => __var["task_id"] === 
"is_weekday")?.__var).toMatchObject({
+      retries: 2,
+      queue: "q",
+    });
+  });
+
+  it("draws an order-only edge to each branch, since no value flows", () => {
+    const { dag, ifReady, fallback, gated } = gatedDag();
+    gated.then(ifReady).else(fallback);
+
+    expect(getDagOrderEdges(dag)).toEqual([
+      { upstream: "has_rows", downstream: "load_if_ready" },
+      { upstream: "has_rows", downstream: "load_fallback" },
+    ]);
+  });
+
+  it("keeps the condition's own arguments and its return value", async () => {
+    const dag = new Dag("d");
+    const loaded = dag.task("load", async () => undefined)();
+    const extracted = dag.task("extract", async () => 3)();
+    dag
+      .if(
+        async ({ rows }: { rows: number }) => rows > 0,
+        { rows: extracted },
+        { taskId: "has_rows" },
+      )
+      .then(loaded);
+
+    expect(getDagTaskInputs(dag).get("has_rows")).toEqual({ rows: extracted });
+    expect((await runHandler(dag, "has_rows", { rows: 2 
})).returned).toBe(true);
+    expect((await runHandler(dag, "has_rows", { rows: 0 
})).returned).toBe(false);
+  });
+
+  it("takes its inputs without a task id", () => {
+    const dag = new Dag("d");
+    const loaded = dag.task("load", async () => undefined)();
+    const extracted = dag.task("extract", async () => 3)();
+    async function hasRows({ rows }: { rows: number }): Promise<boolean> {
+      return rows > 0;
+    }
+    dag.if(hasRows, { rows: extracted }).then(loaded);
+
+    expect(getDagTaskInputs(dag).get("hasRows")).toEqual({ rows: extracted });
+  });
+
+  it("requires the handler's inputs at compile time", () => {
+    const dag = new Dag("d");
+    async function hasRows({ rows }: { rows: number }): Promise<boolean> {
+      return rows > 0;
+    }
+
+    // @ts-expect-error -- hasRows needs { rows }.
+    expect(() => dag.if(hasRows, undefined, { taskId: "has_rows" 
})).not.toThrow();
+  });
+
+  it("takes the condition's id from the handler's name", () => {
+    const dag = new Dag("d");
+    const loaded = dag.task("load", async () => undefined)();
+    async function hasRows(): Promise<boolean> {
+      return true;
+    }
+    const gate = dag.if(hasRows);
+    gate.then(loaded);
+
+    expect(gate.taskId).toBe("hasRows");
+    expect(dag.taskIds).toEqual(["load", "hasRows"]);
+  });
+
+  describe("at run time", () => {
+    it("skips the else branch when the condition holds", async () => {
+      const { dag, ifReady, fallback, gated } = gatedDag();
+      gated.then(ifReady).else(fallback);
+
+      const { skipDownstreamTasks, setXCom } = await runHandler(dag, 
"has_rows");
+
+      expect(skipDownstreamTasks).toHaveBeenCalledWith(["load_fallback"]);
+      expect(setXCom).toHaveBeenCalledWith({
+        key: "skipmixin_key",
+        value: { skipped: ["load_fallback"] },
+      });
+    });
+
+    it("skips the then branch when it does not", async () => {
+      const { dag, ifReady, fallback, gated } = gatedDag("d", false);
+      gated.then(ifReady).else(fallback);
+
+      const { skipDownstreamTasks } = await runHandler(dag, "has_rows");
+
+      expect(skipDownstreamTasks).toHaveBeenCalledWith(["load_if_ready"]);
+    });
+
+    it("skips nothing, and records nothing, when a one-sided condition holds", 
async () => {
+      const { dag, ifReady, gated } = gatedDag();
+      gated.then(ifReady);
+
+      const { skipDownstreamTasks, setXCom } = await runHandler(dag, 
"has_rows");
+
+      expect(skipDownstreamTasks).not.toHaveBeenCalled();
+      expect(setXCom).not.toHaveBeenCalled();
+    });
+
+    it("skips only its own branch when a one-sided condition fails", async () 
=> {
+      const { dag, ifReady, gated } = gatedDag("d", false);
+      gated.then(ifReady);
+
+      const { skipDownstreamTasks } = await runHandler(dag, "has_rows");
+
+      expect(skipDownstreamTasks).toHaveBeenCalledWith(["load_if_ready"]);
+    });
+
+    it.each([
+      ["a string", "yes"],
+      ["null", null],
+      ["a number", 1],
+      ["a bigint", 1n],
+    ])("fails the task when the condition returns %s", async (_label, value) 
=> {
+      const { dag, ifReady, gated } = gatedDag("d", value);
+      gated.then(ifReady);
+
+      await expect(runHandler(dag, "has_rows")).rejects.toThrow(
+        /Condition "has_rows" of Dag "d" returned .* rather than a boolean/,
+      );
+    });
+  });
+
+  it("carries edges of its own, like any task", () => {
+    const { dag, ifReady, gated } = gatedDag();
+    const extracted = dag.task("extract", async () => undefined)();
+    gated.then(ifReady);
+
+    gated.after(extracted);
+
+    expect(getDagOrderEdges(dag)).toContainEqual({
+      upstream: "extract",
+      downstream: "has_rows",
+    });
+  });
+
+  it("stands at the other end of an edge, as Go's .After(gate) does", () => {
+    const { dag, ifReady, gated } = gatedDag();
+    const notified = dag.task("notify", async () => undefined)();
+    gated.then(ifReady);
+
+    notified.after(gated);
+
+    expect(getDagOrderEdges(dag)).toContainEqual({ upstream: "has_rows", 
downstream: "notify" });
+  });
+
+  it("takes another condition as a branch", () => {
+    const { dag, ifReady, gated } = gatedDag();
+    const inner = dag.if(async () => true, undefined, { taskId: "is_weekday" 
});
+    inner.then(ifReady);
+
+    gated.then(inner);
+
+    expect(getDagOrderEdges(dag)).toContainEqual({
+      upstream: "has_rows",
+      downstream: "is_weekday",
+    });
+  });
+
+  describe("rejects", () => {
+    it("a task group as a branch", () => {
+      const { dag, gated } = gatedDag();
+      const group = dag.taskGroup("staging");
+
+      expect(() => gated.then(group as never)).toThrowError(
+        /The "then" branch of Dag "d" condition "has_rows" has to be a task, 
not a task group/,
+      );
+    });
+
+    it("a branch taken from another Dag", () => {
+      const { gated } = gatedDag("here");
+      const { ifReady: foreign } = gatedDag("there");
+
+      expect(() => gated.then(foreign)).toThrowError(
+        /the "then" branch of "has_rows" cannot reach Dag "there" node 
"load_if_ready"/,
+      );
+    });
+
+    it.each([
+      ["a plain object", { dagId: "d", taskId: "load_if_ready" }],
+      ["a string", "load_if_ready"],
+      ["null", null],
+    ])("%s where a branch belongs", (_label, value) => {
+      const { gated } = gatedDag();
+
+      expect(() => gated.then(value as TaskRef)).toThrowError(
+        /the "then" branch of "has_rows" on Dag "d" takes tasks and task 
groups/,
+      );
+    });
+
+    it("both branches naming the same task, which decides nothing", () => {
+      const { ifReady, gated } = gatedDag();
+
+      expect(() => gated.then(ifReady).else(ifReady)).toThrowError(
+        /Both branches of Dag "d" condition "has_rows" are "load_if_ready", so 
the condition decides nothing/,
+      );
+    });
+
+    it("a second else branch", () => {
+      const { ifReady, fallback, gated } = gatedDag();
+      const chain = gated.then(ifReady);
+      chain.else(fallback);
+
+      expect(() => chain.else(fallback)).toThrowError(
+        /Condition "has_rows" of Dag "d" already has an? "else" branch/,
+      );
+    });
+
+    it("a second then branch", () => {
+      const { ifReady, fallback, gated } = gatedDag();
+      gated.then(ifReady);
+
+      expect(() => gated.then(fallback)).toThrowError(
+        /Condition "has_rows" of Dag "d" already has a "then" branch/,
+      );
+    });
+
+    it("a condition that names no branch, when the Dag is read", () => {
+      const { dag } = gatedDag();
+
+      expect(() => finalizeDag(dag)).toThrowError(
+        /Condition "has_rows" of Dag "d" names no branch, so it decides 
nothing/,
+      );
+    });
+
+    it("being awaited, which would resolve the chain rather than build it", 
async () => {
+      const { gated } = gatedDag();
+
+      await expect(Promise.resolve(gated as unknown as 
Promise<void>)).rejects.toThrow(
+        /dag.if\(\.\.\.\) of Dag "d" was awaited/,
+      );
+    });
+  });
+});
diff --git a/ts-sdk/tests/sdk/dag.test.ts b/ts-sdk/tests/sdk/dag.test.ts
index 72c064a3df4..2a5061041c3 100644
--- a/ts-sdk/tests/sdk/dag.test.ts
+++ b/ts-sdk/tests/sdk/dag.test.ts
@@ -630,9 +630,7 @@ describe("Dag", () => {
       const { refs: there } = placedDag("there", "cleanup");
 
       expect(() => draw(here.load!, there.cleanup!)).toThrowError(
-        new RegExp(
-          `${verb}\\(\\) cannot draw an edge to Dag "there" node "cleanup" 
from Dag "here"`,
-        ),
+        new RegExp(`${verb}\\(\\) cannot reach Dag "there" node "cleanup" from 
Dag "here"`),
       );
     });
 

Reply via email to