guan404ming commented on code in PR #74070:
URL: https://github.com/apache/airflow/pull/74070#discussion_r4197883402


##########
ts-sdk/src/sdk/dag.ts:
##########
@@ -540,6 +600,188 @@ 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, "has_rows", { rows: validated }).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 [taskId, inputs] =
+      typeof args[0] === "string" ? [args[0], args[1]] : [undefined, args[0]];
+    const factory = (
+      taskId === undefined
+        ? this.#addTask(undefined, handler)
+        : this.#addTask(undefined, taskId, handler)
+    ) 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 | 
Branch): 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 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>(

Review Comment:
   Sounds good, I'll do it in a follow-up once this series lands.



##########
ts-sdk/src/sdk/dag.ts:
##########
@@ -540,6 +600,188 @@ 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, "has_rows", { rows: validated }).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 [taskId, inputs] =
+      typeof args[0] === "string" ? [args[0], args[1]] : [undefined, args[0]];
+    const factory = (
+      taskId === undefined
+        ? this.#addTask(undefined, handler)
+        : this.#addTask(undefined, taskId, handler)
+    ) 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 | 
Branch): 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 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>);
+  }
+
+  #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);
+      if (!candidates.includes(chosen as TaskRef)) {

Review Comment:
   Good catch, it now matches by task id.



##########
ts-sdk/src/sdk/dag.ts:
##########
@@ -540,6 +600,188 @@ 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, "has_rows", { rows: validated }).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 [taskId, inputs] =
+      typeof args[0] === "string" ? [args[0], args[1]] : [undefined, args[0]];
+    const factory = (
+      taskId === undefined
+        ? this.#addTask(undefined, handler)
+        : this.#addTask(undefined, taskId, handler)
+    ) 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 | 
Branch): 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) => {

Review Comment:
   Done, both now share one `#forwardEdges` helper.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to