1use std::cmp::Ordering;
2use std::collections::{BinaryHeap, HashSet};
3
4use crate::scheduler::policy::SchedulerPolicyConfig;
5use crate::types::error::{DeepStrikeError, Result};
6use crate::types::result::LoopResult;
7use crate::types::task::RuntimeTask;
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub enum TaskStatus {
11 Pending,
12 Ready,
13 Running,
14 Completed,
15 CompletedPartial,
16 Failed,
17 SkippedUpstreamFailed,
18}
19
20impl TaskStatus {
21 pub fn is_terminal(self) -> bool {
22 matches!(
23 self,
24 Self::Completed | Self::CompletedPartial | Self::Failed | Self::SkippedUpstreamFailed
25 )
26 }
27}
28
29#[derive(Debug, Clone)]
30pub struct TaskNode {
31 pub id: usize,
32 pub task: RuntimeTask,
33 pub status: TaskStatus,
34 pub result: Option<LoopResult>,
35 pub dependencies: Vec<usize>,
36}
37
38pub struct TaskGraph {
42 nodes: Vec<TaskNode>,
43 in_degree: Vec<usize>,
46 reverse_adjacency: Vec<Vec<usize>>,
48 ready_heap: BinaryHeap<ReadyEntry>,
49 ready_generation: Vec<u64>,
50 enqueued_round: Vec<u64>,
51 enqueue_sequence: u64,
52 ready_round: u64,
53 scheduling: Vec<SchedulingMetadata>,
54 scheduler_policy: SchedulerPolicyConfig,
55}
56
57#[derive(Debug, Clone, Copy, Default)]
58struct SchedulingMetadata {
59 critical_path_remaining: u64,
60 downstream_fanout: u64,
61 token_cost: u64,
62}
63
64#[derive(Debug, Clone, Copy, PartialEq, Eq)]
65struct ReadyEntry {
66 priority: i128,
67 enqueue_sequence: u64,
68 node_id: usize,
69 generation: u64,
70}
71
72impl Ord for ReadyEntry {
73 fn cmp(&self, other: &Self) -> Ordering {
74 self.priority
75 .cmp(&other.priority)
76 .then_with(|| other.enqueue_sequence.cmp(&self.enqueue_sequence))
77 .then_with(|| other.node_id.cmp(&self.node_id))
78 }
79}
80
81impl PartialOrd for ReadyEntry {
82 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
83 Some(self.cmp(other))
84 }
85}
86
87impl TaskGraph {
88 pub fn new() -> Self {
89 Self {
90 nodes: Vec::new(),
91 in_degree: Vec::new(),
92 reverse_adjacency: Vec::new(),
93 ready_heap: BinaryHeap::new(),
94 ready_generation: Vec::new(),
95 enqueued_round: Vec::new(),
96 enqueue_sequence: 0,
97 ready_round: 0,
98 scheduling: Vec::new(),
99 scheduler_policy: SchedulerPolicyConfig::default(),
100 }
101 }
102
103 pub fn add(&mut self, task: RuntimeTask, mut dependencies: Vec<usize>) -> usize {
107 let mut seen = std::collections::HashSet::new();
108 dependencies.retain(|d| seen.insert(*d));
109 let id = self.nodes.len();
110 let deg = dependencies.len();
111 let max_index = dependencies.iter().copied().max().unwrap_or(id).max(id);
112 self.reverse_adjacency.resize_with(max_index + 1, Vec::new);
113 for &dependency in &dependencies {
114 self.reverse_adjacency[dependency].push(id);
115 }
116 self.nodes.push(TaskNode {
117 id,
118 task,
119 status: if deg == 0 {
120 TaskStatus::Ready
121 } else {
122 TaskStatus::Pending
123 },
124 result: None,
125 dependencies,
126 });
127 self.in_degree.push(deg);
128 self.ready_generation.push(0);
129 self.enqueued_round.push(self.ready_round);
130 self.scheduling.push(SchedulingMetadata::default());
131 if deg == 0 {
132 self.enqueue_ready(id);
133 }
134 id
135 }
136
137 pub fn topological_sort(&self) -> Result<Vec<usize>> {
139 let n = self.nodes.len();
140 let mut in_deg: Vec<usize> = self
144 .nodes
145 .iter()
146 .map(|node| node.dependencies.len())
147 .collect();
148
149 let mut queue: Vec<usize> = (0..n).filter(|&i| in_deg[i] == 0).collect();
150 let mut order = Vec::with_capacity(n);
151
152 while let Some(id) = queue.pop() {
153 order.push(id);
154 for &next in self.reverse_adjacency.get(id).into_iter().flatten() {
155 in_deg[next] -= 1;
156 if in_deg[next] == 0 {
157 queue.push(next);
158 }
159 }
160 }
161
162 if order.len() != n {
163 return Err(DeepStrikeError::OrchestrationCycle);
164 }
165 Ok(order)
166 }
167
168 pub fn ready_tasks(&mut self) -> Vec<usize> {
170 let mut valid_entries = Vec::new();
174 let mut ready = Vec::new();
175 while let Some(entry) = self.ready_heap.pop() {
176 if self.nodes.get(entry.node_id).map(|node| node.status) == Some(TaskStatus::Ready)
177 && self.ready_generation[entry.node_id] == entry.generation
178 {
179 ready.push(entry.node_id);
180 valid_entries.push(entry);
181 }
182 }
183 self.ready_heap.extend(valid_entries);
184 self.ready_round = self.ready_round.saturating_add(1);
185 ready
186 }
187
188 pub fn start(&mut self, task_id: usize) {
190 if let Some(node) = self.nodes.get_mut(task_id) {
191 node.status = TaskStatus::Running;
192 }
193 }
194
195 pub fn set_ready(&mut self, task_id: usize) {
199 if let Some(node) = self.nodes.get_mut(task_id) {
200 if node.status != TaskStatus::Ready {
201 node.status = TaskStatus::Ready;
202 self.enqueue_ready(task_id);
203 }
204 }
205 }
206
207 pub fn complete(&mut self, task_id: usize, result: LoopResult) {
213 {
214 let Some(node) = self.nodes.get_mut(task_id) else {
215 return;
216 };
217 if node.status.is_terminal() {
218 return;
219 }
220 node.status = TaskStatus::Completed;
221 node.result = Some(result);
222 }
223 let dependents = self
224 .reverse_adjacency
225 .get(task_id)
226 .cloned()
227 .unwrap_or_default();
228 for dep_id in dependents {
229 self.in_degree[dep_id] -= 1;
230 if self.in_degree[dep_id] == 0 {
231 let should_enqueue =
232 self.nodes.get(dep_id).map(|n| n.status) == Some(TaskStatus::Pending);
233 if should_enqueue {
234 self.nodes[dep_id].status = TaskStatus::Ready;
235 self.enqueue_ready(dep_id);
236 }
237 }
238 }
239 }
240
241 pub fn complete_partial(&mut self, task_id: usize, result: LoopResult) {
242 if let Some(node) = self.nodes.get_mut(task_id) {
243 if !node.status.is_terminal() {
244 node.status = TaskStatus::CompletedPartial;
245 node.result = Some(result);
246 }
247 }
248 }
249
250 pub fn fail(&mut self, task_id: usize) {
254 if let Some(node) = self.nodes.get_mut(task_id) {
255 if !node.status.is_terminal() {
256 node.status = TaskStatus::Failed;
257 }
258 }
259 }
260
261 pub fn fail_with_result(&mut self, task_id: usize, result: LoopResult) {
262 if let Some(node) = self.nodes.get_mut(task_id) {
263 if !node.status.is_terminal() {
264 node.status = TaskStatus::Failed;
265 node.result = Some(result);
266 }
267 }
268 }
269
270 pub fn skip_upstream_failed(&mut self, task_id: usize) {
271 if let Some(node) = self.nodes.get_mut(task_id) {
272 if !node.status.is_terminal() {
273 node.status = TaskStatus::SkippedUpstreamFailed;
274 }
275 }
276 }
277
278 pub(crate) fn restore_runtime_state(
284 &mut self,
285 states: &[(TaskStatus, Option<LoopResult>)],
286 ) -> std::result::Result<(), String> {
287 if states.len() != self.nodes.len() {
288 return Err(format!(
289 "workflow checkpoint carries {} node states for a {} node graph",
290 states.len(),
291 self.nodes.len()
292 ));
293 }
294
295 for (node, (status, result)) in self.nodes.iter_mut().zip(states) {
296 node.status = *status;
297 node.result = result.clone();
298 }
299 self.in_degree = self
300 .nodes
301 .iter()
302 .map(|node| {
303 node.dependencies
304 .iter()
305 .filter(|&&dependency| {
306 self.nodes.get(dependency).map(|node| node.status)
307 != Some(TaskStatus::Completed)
308 })
309 .count()
310 })
311 .collect();
312 self.ready_heap.clear();
313 self.ready_generation.fill(0);
314 self.enqueued_round.fill(0);
315 self.enqueue_sequence = 0;
316 self.ready_round = 0;
317 for node in 0..self.nodes.len() {
318 if self.nodes[node].status == TaskStatus::Ready {
319 self.enqueue_ready(node);
320 }
321 }
322 Ok(())
323 }
324
325 pub fn get(&self, task_id: usize) -> Option<&TaskNode> {
326 self.nodes.get(task_id)
327 }
328
329 pub fn len(&self) -> usize {
330 self.nodes.len()
331 }
332
333 pub fn is_empty(&self) -> bool {
334 self.nodes.is_empty()
335 }
336
337 pub fn all_done(&self) -> bool {
338 self.nodes.iter().all(|n| n.status.is_terminal())
339 }
340
341 pub fn configure_scheduling(&mut self, policy: SchedulerPolicyConfig, token_costs: &[u64]) {
342 self.scheduler_policy = policy;
343 let order = self
344 .topological_sort()
345 .unwrap_or_else(|_| (0..self.nodes.len()).collect());
346 let mut reachable: Vec<HashSet<usize>> = vec![HashSet::new(); self.nodes.len()];
347 for &node in order.iter().rev() {
348 let mut critical = 1u64;
349 let children = self
350 .reverse_adjacency
351 .get(node)
352 .cloned()
353 .unwrap_or_default();
354 for child in children {
355 critical = critical.max(1 + self.scheduling[child].critical_path_remaining);
356 reachable[node].insert(child);
357 let descendants: Vec<usize> = reachable[child].iter().copied().collect();
358 reachable[node].extend(descendants);
359 }
360 self.scheduling[node] = SchedulingMetadata {
361 critical_path_remaining: critical,
362 downstream_fanout: reachable[node].len() as u64,
363 token_cost: token_costs.get(node).copied().unwrap_or(0),
364 };
365 }
366 self.rebuild_ready_heap();
367 }
368
369 fn rebuild_ready_heap(&mut self) {
370 self.ready_heap.clear();
371 for node_id in 0..self.nodes.len() {
372 if self.nodes[node_id].status == TaskStatus::Ready {
373 self.push_ready_entry(node_id);
374 }
375 }
376 }
377
378 fn enqueue_ready(&mut self, task_id: usize) {
379 self.ready_generation[task_id] = self.ready_generation[task_id].saturating_add(1);
380 self.enqueued_round[task_id] = self.ready_round;
381 self.enqueue_sequence = self.enqueue_sequence.saturating_add(1);
382 self.push_ready_entry(task_id);
383 }
384
385 fn push_ready_entry(&mut self, task_id: usize) {
386 let metadata = self.scheduling[task_id];
387 let policy = self.scheduler_policy;
388 let priority = i128::from(policy.critical_path_weight)
389 * i128::from(metadata.critical_path_remaining)
390 + i128::from(policy.fanout_weight) * i128::from(metadata.downstream_fanout)
391 - i128::from(policy.age_weight) * i128::from(self.enqueued_round[task_id])
392 - i128::from(policy.token_cost_weight) * i128::from(metadata.token_cost);
393 self.ready_heap.push(ReadyEntry {
394 priority,
395 enqueue_sequence: self.enqueue_sequence,
396 node_id: task_id,
397 generation: self.ready_generation[task_id],
398 });
399 }
400}
401
402impl Default for TaskGraph {
403 fn default() -> Self {
404 Self::new()
405 }
406}
407
408#[cfg(test)]
409mod tests {
410 use super::*;
411
412 #[test]
413 fn topological_sort_linear() {
414 let mut g = TaskGraph::new();
415 let a = g.add(RuntimeTask::new("A"), vec![]);
416 let b = g.add(RuntimeTask::new("B"), vec![a]);
417 let c = g.add(RuntimeTask::new("C"), vec![b]);
418
419 let order = g.topological_sort().unwrap();
420 assert_eq!(order, vec![0, 1, 2]);
421 let _ = (a, c);
422 }
423
424 #[test]
425 fn detects_cycle() {
426 let mut g = TaskGraph::new();
427 g.nodes.push(TaskNode {
428 id: 0,
429 task: RuntimeTask::new("A"),
430 status: TaskStatus::Pending,
431 result: None,
432 dependencies: vec![1],
433 });
434 g.nodes.push(TaskNode {
435 id: 1,
436 task: RuntimeTask::new("B"),
437 status: TaskStatus::Pending,
438 result: None,
439 dependencies: vec![0],
440 });
441 g.in_degree.push(1);
442 g.in_degree.push(1);
443
444 assert!(g.topological_sort().is_err());
445 }
446
447 #[test]
448 fn ready_tasks_respects_deps() {
449 let mut g = TaskGraph::new();
450 let a = g.add(RuntimeTask::new("A"), vec![]);
451 let _b = g.add(RuntimeTask::new("B"), vec![a]);
452
453 assert_eq!(g.ready_tasks(), vec![0]); }
455
456 #[test]
457 fn set_ready_rearms_without_promoting_dependents() {
458 let mut g = TaskGraph::new();
459 let a = g.add(RuntimeTask::new("A"), vec![]); let b = g.add(RuntimeTask::new("B"), vec![a]); g.start(a);
462 g.set_ready(a);
464 assert_eq!(g.nodes[a].status, TaskStatus::Ready);
465 assert_eq!(g.nodes[b].status, TaskStatus::Pending);
466 assert_eq!(g.ready_tasks(), vec![a]);
467 }
468
469 #[test]
470 fn complete_promotes_dependent() {
471 use crate::types::result::{LoopResult, TerminationReason};
472 let mut g = TaskGraph::new();
473 let a = g.add(RuntimeTask::new("A"), vec![]);
474 let b = g.add(RuntimeTask::new("B"), vec![a]);
475
476 assert_eq!(g.nodes[b].status, TaskStatus::Pending);
477 g.complete(
478 a,
479 LoopResult {
480 termination: TerminationReason::Completed,
481 final_message: None,
482 turns_used: 1,
483 total_tokens_used: 0,
484 loop_continue: None,
485 classify_branch: None,
486 tournament_winner: None,
487 pace_decision: None,
488 },
489 );
490 assert_eq!(g.nodes[b].status, TaskStatus::Ready);
491 }
492
493 #[test]
494 fn duplicate_complete_is_idempotent() {
495 use crate::types::result::{LoopResult, TerminationReason};
496 let result = || LoopResult {
497 termination: TerminationReason::Completed,
498 final_message: None,
499 turns_used: 1,
500 total_tokens_used: 0,
501 loop_continue: None,
502 classify_branch: None,
503 tournament_winner: None,
504 pace_decision: None,
505 };
506 let mut g = TaskGraph::new();
508 let a = g.add(RuntimeTask::new("A"), vec![]);
509 let c = g.add(RuntimeTask::new("C"), vec![]);
510 let b = g.add(RuntimeTask::new("B"), vec![a, c]);
511
512 g.complete(a, result());
513 g.complete(a, result()); assert_eq!(g.nodes[b].status, TaskStatus::Pending);
515 g.complete(c, result());
516 assert_eq!(g.nodes[b].status, TaskStatus::Ready);
517 g.fail(a);
519 assert_eq!(g.nodes[a].status, TaskStatus::Completed);
520 }
521
522 #[test]
523 fn critical_path_priority_beats_lower_node_id() {
524 let mut g = TaskGraph::new();
525 let wide = g.add(RuntimeTask::new("wide"), vec![]);
526 let chain = g.add(RuntimeTask::new("chain"), vec![]);
527 g.add(RuntimeTask::new("wide-child-a"), vec![wide]);
528 g.add(RuntimeTask::new("wide-child-b"), vec![wide]);
529 let chain_2 = g.add(RuntimeTask::new("chain-2"), vec![chain]);
530 let chain_3 = g.add(RuntimeTask::new("chain-3"), vec![chain_2]);
531 g.add(RuntimeTask::new("chain-4"), vec![chain_3]);
532
533 g.configure_scheduling(SchedulerPolicyConfig::default(), &[]);
534
535 assert_eq!(g.ready_tasks(), vec![chain, wide]);
536 }
537
538 #[test]
539 fn zero_weights_use_fifo_and_loop_rearm_yields() {
540 let mut g = TaskGraph::new();
541 let loop_node = g.add(RuntimeTask::new("loop"), vec![]);
542 let peer = g.add(RuntimeTask::new("peer"), vec![]);
543 let policy = SchedulerPolicyConfig {
544 critical_path_weight: 0,
545 fanout_weight: 0,
546 age_weight: 0,
547 token_cost_weight: 0,
548 ..SchedulerPolicyConfig::default()
549 };
550 g.configure_scheduling(policy, &[]);
551 assert_eq!(g.ready_tasks(), vec![loop_node, peer]);
552
553 g.start(loop_node);
554 g.set_ready(loop_node);
555 assert_eq!(g.ready_tasks(), vec![peer, loop_node]);
556 assert_eq!(
557 g.ready_heap.len(),
558 2,
559 "stale loop generations must be collected"
560 );
561 }
562
563 #[test]
564 fn reverse_adjacency_tracks_only_outgoing_dependents() {
565 let mut g = TaskGraph::new();
566 let root = g.add(RuntimeTask::new("root"), vec![]);
567 let unrelated = g.add(RuntimeTask::new("unrelated"), vec![]);
568 let child = g.add(RuntimeTask::new("child"), vec![root]);
569 g.add(RuntimeTask::new("grandchild"), vec![child]);
570
571 assert_eq!(g.reverse_adjacency[root], vec![child]);
572 assert!(g.reverse_adjacency[unrelated].is_empty());
573 }
574}