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


##########
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:
   Not a blocking item from being merged, can be follow-up after this series of 
PRs:
   
   We should validate the incoming `fn` for the `if` and `switch` if possible. 
(e.g. we shouldn't accept the `triggerDagRun` in the condition)



##########
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:
   This validates the decider's return value by object identity 
(`candidates.includes(chosen)`), not by `taskId`, unlike `dag.if`'s 
branch-naming path which resolves through `resolveNode`. A handler that returns 
a `TaskRef` re-fetched some other way (same `taskId`, different JS object) 
would fail here with "not one of its cases" even though the chosen id is a 
legitimate case.
   



##########
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:
   `#placeCondition`'s `before`/`after` here and `#placeBranch`'s own copy 
below (around line 772) hand-write the identical forwarding closure 
independently, four near-identical closures total across the two.
   
   Not sure would it be possible to introduce a help to reuse them.



-- 
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