uranusjr commented on code in PR #74230:
URL: https://github.com/apache/airflow/pull/74230#discussion_r4203711485


##########
java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt:
##########
@@ -133,12 +143,122 @@ class DagDef(
     task.owner?.let { owner ->
       throw IllegalArgumentException("Task '${task.id}' already belongs to Dag 
'${owner.id}'")
     }
+    require(task.id !in groups) { "Dag '$id' already has a task group with ID: 
${task.id}" }
     require(tasks.putIfAbsent(task.id, task) == null) {
       "Tasks in Dag have duplicate ID: ${task.id}"
     }
     task.owner = this
     return this
   }
+
+  /**
+   * Declares a task group of this Dag.
+   *
+   * ```java
+   * var staging = dag.taskGroup("staging");
+   * var stage = staging.task("stage", Stage.class); // task "staging.stage"
+   * extract.before(staging);
+   * ```
+   *
+   * @param id Group ID. Must contain only ASCII letters, digits, underscores,
+   *    or dashes, and differ from every task and group ID in this Dag.
+   * @return The group, to declare tasks in and to wire edges with.
+   * @throws IllegalArgumentException if [id] is not a valid group ID, or the
+   *    Dag already has a task or task group with that ID.
+   */
+  fun taskGroup(id: String): TaskGroupRef = addGroup(null, id)
+
+  internal fun addGroup(
+    parent: TaskGroupRef?,
+    localId: String,
+  ): TaskGroupRef {
+    require(GROUP_ID.matches(localId)) {
+      "Task group ID '$localId' must contain only ASCII letters, digits, 
underscores, or dashes"
+    }
+    val groupId = parent?.qualify(localId) ?: localId
+    require(groupId !in tasks && groupId !in groups) {
+      "Dag '$id' already has a task or task group with ID: $groupId"
+    }
+    return TaskGroupRef(this, groupId, parent).also {
+      groups[groupId] = it
+      parent?.children?.add(it)
+    }
+  }
+
+  /**
+   * Turns every edge drawn to or from a task group into edges between tasks,
+   * and records it on the group as Python's `TaskGroup` does.
+   *
+   * A group upstream stands for its leaves and a group downstream for its
+   * roots. Edges expand in the order they were drawn, each reading the task
+   * edges the ones before it left behind, which is how Python resolves a
+   * group's endpoints at every `>>`. Runs once: a group edge drawn after
+   * that is rejected where it is drawn rather than silently left out.
+   */
+  internal fun expandGroupEdges() {
+    if (groupEdgesExpanded) return
+    groupEdgesExpanded = true
+    for ((upstream, downstream) in groupEdges) {
+      val from = leavesOf(upstream)
+      rootsOf(downstream).forEach { it.dependsOn(*from.toTypedArray()) }
+      if (downstream is TaskGroupRef) {
+        downstream.upstreamTaskIds += from.map { it.id }
+        if (upstream is TaskGroupRef) downstream.upstreamGroupIds += 
upstream.id
+      }
+      // When both ends are groups, the upstream records the downstream group
+      // only, not its tasks, which is how Python leaves it.
+      if (upstream is TaskGroupRef) {
+        if (downstream is TaskGroupRef) {
+          upstream.downstreamGroupIds += downstream.id
+        } else {
+          upstream.downstreamTaskIds += (downstream as TaskDef).id
+        }
+      }
+    }
+  }
+
+  /**
+   * The tasks an edge out of [endpoint] starts from: Python's `find_leaves`.
+   *
+   * A group stands for its own leaves. One holding none stands for whatever
+   * already runs before it, then for the group edges drawn into it, and
+   * failing both for the group it is nested in, so an edge out of an empty
+   * group still reaches the tasks around it.
+   */
+  private fun leavesOf(
+    endpoint: Any,
+    seen: MutableSet<TaskGroupRef> = mutableSetOf(),
+  ): List<TaskDef> {
+    if (endpoint is TaskDef) return listOf(endpoint)
+    var group: TaskGroupRef? = endpoint as TaskGroupRef
+    while (group != null && seen.add(group)) {
+      group.leaves().takeIf { it.isNotEmpty() }?.let { return it }
+      group.upstreamTaskIds.takeIf { it.isNotEmpty() }?.let { ids -> return 
ids.map { tasks.getValue(it) } }
+      val current = group
+      groupEdges.filter { it.second === current }.takeIf { it.isNotEmpty() 
}?.let { drawn ->
+        return drawn.flatMap { leavesOf(it.first, seen) }
+      }

Review Comment:
   Python does not do this, if I understand correctly. It only falls back to 
`upstream_task_ids` and then the parent.



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