Skip to content

Commit 1e367df

Browse files
committed
TS SDK: add conditional branching with dag.if
1 parent dc0d191 commit 1e367df

10 files changed

Lines changed: 592 additions & 9 deletions

File tree

‎airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst‎

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -359,6 +359,30 @@ there. Set it once on the Dag and each task inherits it:
359359
``queue`` on a task wins over the Dag's. See :ref:`typescript-sdk/coordinator-config` for the
360360
``queue_to_coordinator`` entry that sends that queue to the coordinator.
361361

362+
Conditional branching
363+
~~~~~~~~~~~~~~~~~~~~~
364+
365+
``dag.if`` takes a task whose handler returns a boolean, and names the task each outcome runs:
366+
367+
.. code-block:: typescript
368+
369+
const condition = dag.task("has_rows", async ({ rows }: { rows: number }) => rows > 0);
370+
const gated = condition({ rows: extracted });
371+
372+
dag.if(gated).then(loaded).else(reportedEmpty);
373+
374+
The condition is an ordinary task, so it is declared, typed and wired like any other, and the
375+
compiler checks that its handler really returns a boolean. ``else`` is optional: a one-sided
376+
condition skips its own branch when the condition fails and follows nothing.
377+
378+
A guarded task takes no argument for the control edge, because a condition's boolean decides whether
379+
the task runs rather than what it runs on. Read a value from the condition with
380+
``getClient().getXCom``.
381+
382+
The side not taken is skipped when the run reaches it, and stays skipped if you clear it later. Only
383+
the branches named here are skipped, so a task that several branches converge on still runs — unlike
384+
Python's ``@task.branch``, which skips every immediate downstream it did not follow.
385+
362386
``new Dag`` and ``dag.task`` both take a trailing spec of Airflow options:
363387
``{ schedule: "@daily", tags: ["etl"] }`` for the Dag, ``{ retries: 2, retryDelay: 30 }`` for a task.
364388

‎ts-sdk/adr/0002-native-dag-interface.md‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -135,6 +135,43 @@ data is the wiring object itself — `summarize({ north: extractNorth(), south:
135135
does. This matches `Before`/`After` in the Go SDK's native Dag interface, spelled to TypeScript
136136
convention.
137137

138+
### Conditional branching: `if` and `else`
139+
140+
`dag.if(condition)` is TypeScript's spelling of the construct
141+
[`airflow-core/adr/lang-sdk/0008`](../../airflow-core/adr/lang-sdk/0008-control-flow-constructs.md)
142+
names after the host language's control flow:
143+
144+
```ts
145+
const gated = dag.task("has_rows", async ({ rows }: { rows: number }) => rows > 0)({
146+
rows: extracted,
147+
});
148+
149+
dag.if(gated).then(loadIfReady).else(loadFallback);
150+
```
151+
152+
**The condition is a task reference, not a task id and a function.** `dag.if` takes a `TaskRef` the
153+
Dag already handed back, whose handler's return type the compiler checks is `boolean`. So nothing
154+
depends on the task's id or on its function name, and a condition is declared, named and typed the
155+
same way every other task is. The same reasoning that makes a *case* a reference in that ADR's
156+
decision 2 applies to the condition itself.
157+
158+
**A `then` chain is a thenable, and is guarded rather than avoided.** An object with a callable
159+
`then` is a *thenable*: were the object `dag.if` returns to reach an `await`, the runtime would hand
160+
its `then` a resolve function where a task reference belongs. Two things contain that. `.then(...)`
161+
returns an object carrying only `.else`, so nothing past the first step is awaitable at all; and
162+
`.then` rejects a function argument by naming the cause, so an author who does await it reads "this
163+
builds a branch, drop the await" rather than a type error about references.
164+
165+
**A branch is a real branch to Airflow.** The control edges serialize as ordinary order-only edges
166+
and carry no branch-candidate field, as that ADR's consequences require. The condition task is
167+
serialized with `_can_skip_downstream`, and it writes the `skipmixin_key` XCom alongside the skip, so
168+
clearing a skipped branch re-skips it the way a Python `@task.branch` does rather than running the
169+
side the condition rejected.
170+
171+
A one-sided `if` is a branch with one candidate — it skips `then` and follows nothing — rather than a
172+
`ShortCircuitOperator`, which would also skip the whole downstream closure and ignore trigger rules.
173+
A guarded task takes no argument for the control edge: a condition's boolean is a signal, not data.
174+
138175
## Consequences
139176

140177
- One authoring surface (`dag.task()` plus its factory) covers the graph and each task's arguments,

‎ts-sdk/src/coordinator/client.ts‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ import type {
3030
GetXCom,
3131
SetXCom,
3232
GetConnection,
33+
SkipDownstreamTasks,
3334
ConnectionResult as WireConnectionResult,
3435
} from "./protocol.js";
3536

@@ -202,6 +203,16 @@ export function createCoordinatorClient(
202203
await rpc("SetXCom", null, msg, () => undefined, "throw");
203204
},
204205

206+
// ---- Control flow ----
207+
208+
async skipDownstreamTasks(taskIds: readonly string[]): Promise<void> {
209+
// The supervisor treats an empty list as "skip nothing", and sending it
210+
// would cost a round trip to say so.
211+
if (taskIds.length === 0) return;
212+
const msg: SkipDownstreamTasks = { type: "SkipDownstreamTasks", tasks: [...taskIds] };
213+
await rpc("SkipDownstreamTasks", null, msg, () => undefined, "throw");
214+
},
215+
205216
// ---- Connections ----
206217

207218
async getConnection(connId: string): Promise<ConnectionResult | null> {

‎ts-sdk/src/coordinator/protocol.ts‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,7 @@ export type {
6464
GetXCom,
6565
SetXCom,
6666
GetConnection,
67+
SkipDownstreamTasks,
6768
} from "../generated/supervisor.js";
6869

6970
// -------- Frames from supervisor --------

‎ts-sdk/src/coordinator/serde.ts‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -154,6 +154,7 @@ export function serializeDag(
154154
withDagQueue(record.spec, dag.spec.queue),
155155
graph.downstreamTaskIds.get(taskId),
156156
inputs.get(taskId),
157+
record.canSkipDownstream === true,
157158
),
158159
),
159160
dag_dependencies: [],
@@ -191,6 +192,7 @@ function serializeTask(
191192
spec: object,
192193
downstream: ReadonlySet<string> | undefined,
193194
inputs: RecordedInputs | undefined,
195+
canSkipDownstream: boolean,
194196
): SerializedValue {
195197
const data: Record<string, SerializedValue> = {
196198
task_id: taskId,
@@ -208,6 +210,10 @@ function serializeTask(
208210
};
209211
const bindings = serializeArgBindings(inputs);
210212
if (bindings) data["_arg_bindings"] = bindings;
213+
// What Python writes for a SkipMixin operator, and what makes
214+
// `NotPreviouslySkippedDep` consult this task's `skipmixin_key` XCom when one
215+
// of its skipped downstream tasks is cleared.
216+
if (canSkipDownstream) data["_can_skip_downstream"] = true;
211217
applySchemaFields(data, spec, TASK_FIELD_RULES, `task "${taskId}" of Dag "${dagId}"`);
212218
if (downstream?.size) {
213219
data["downstream_task_ids"] = [...downstream].sort();

‎ts-sdk/src/index.ts‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,8 @@ export { SUPERVISOR_API_VERSION } from "./coordinator/index.js";
2727
export type { ArgNameMap } from "./sdk/arg-names.js";
2828
export type { Registerable } from "./sdk/bundle.js";
2929
export type {
30+
Condition,
31+
ConditionElse,
3032
DagSpec,
3133
Node,
3234
TaskFactory,

‎ts-sdk/src/sdk/client.ts‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -107,6 +107,19 @@ export interface TaskClient {
107107
* @throws {@link Exceptions!ConnectionNotFoundError | ConnectionNotFoundError} when the connection does not exist.
108108
*/
109109
getConnectionOrThrow(connId: string): Promise<ConnectionResult>;
110+
111+
/**
112+
* Mark downstream tasks of the running task as skipped.
113+
*
114+
* The mechanism behind `dag.if(...)`: a condition's boolean is a run-time
115+
* signal, so which branch a Dag run takes is decided here rather than
116+
* recorded in the Dag. Python reaches the same supervisor message through
117+
* `skip()` and `skip_all_except()`.
118+
*
119+
* Each id must name a direct downstream of the running task. Passing none is
120+
* a no-op, so a branch that skips nothing needs no guard at the call site.
121+
*/
122+
skipDownstreamTasks(taskIds: readonly string[]): Promise<void>;
110123
}
111124

112125
/** Error thrown by {@link TaskClient.getVariableOrThrow}. */

0 commit comments

Comments
 (0)