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"`),
);
});