1use crate::dag::graph::{ResourceRequirements, WorkflowDag};
4use crate::dag::topological_sort::create_execution_plan;
5use crate::error::Result;
6use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8
9#[derive(Debug, Clone, Serialize, Deserialize)]
11pub struct ResourcePool {
12 pub total_cpu_cores: f64,
14 pub total_memory_mb: u64,
16 pub total_gpus: u32,
18 pub total_disk_mb: u64,
20 pub custom_resources: HashMap<String, f64>,
22}
23
24impl Default for ResourcePool {
25 fn default() -> Self {
26 Self {
27 total_cpu_cores: num_cpus::get() as f64,
28 total_memory_mb: 8192,
29 total_gpus: 0,
30 total_disk_mb: 102400,
31 custom_resources: HashMap::new(),
32 }
33 }
34}
35
36#[derive(Debug, Clone)]
38pub struct AvailableResources {
39 pub cpu_cores: f64,
41 pub memory_mb: u64,
43 pub gpus: u32,
45 pub disk_mb: u64,
47 pub custom_resources: HashMap<String, f64>,
49}
50
51impl From<ResourcePool> for AvailableResources {
52 fn from(pool: ResourcePool) -> Self {
53 Self {
54 cpu_cores: pool.total_cpu_cores,
55 memory_mb: pool.total_memory_mb,
56 gpus: pool.total_gpus,
57 disk_mb: pool.total_disk_mb,
58 custom_resources: pool.custom_resources,
59 }
60 }
61}
62
63impl AvailableResources {
64 pub fn can_allocate(&self, requirements: &ResourceRequirements) -> bool {
66 if self.cpu_cores < requirements.cpu_cores {
67 return false;
68 }
69 if self.memory_mb < requirements.memory_mb {
70 return false;
71 }
72 if requirements.gpu && self.gpus == 0 {
73 return false;
74 }
75 if self.disk_mb < requirements.disk_mb {
76 return false;
77 }
78
79 for (key, &required_value) in &requirements.custom {
81 if let Some(&available_value) = self.custom_resources.get(key) {
82 if available_value < required_value {
83 return false;
84 }
85 } else {
86 return false;
87 }
88 }
89
90 true
91 }
92
93 pub fn allocate(&mut self, requirements: &ResourceRequirements) -> bool {
95 if !self.can_allocate(requirements) {
96 return false;
97 }
98
99 self.cpu_cores -= requirements.cpu_cores;
100 self.memory_mb -= requirements.memory_mb;
101 if requirements.gpu {
102 self.gpus -= 1;
103 }
104 self.disk_mb -= requirements.disk_mb;
105
106 for (key, &value) in &requirements.custom {
107 if let Some(available) = self.custom_resources.get_mut(key) {
108 *available -= value;
109 }
110 }
111
112 true
113 }
114
115 pub fn release(&mut self, requirements: &ResourceRequirements) {
117 self.cpu_cores += requirements.cpu_cores;
118 self.memory_mb += requirements.memory_mb;
119 if requirements.gpu {
120 self.gpus += 1;
121 }
122 self.disk_mb += requirements.disk_mb;
123
124 for (key, &value) in &requirements.custom {
125 *self.custom_resources.entry(key.clone()).or_insert(0.0) += value;
126 }
127 }
128}
129
130#[derive(Debug, Clone, Serialize, Deserialize)]
132pub struct ParallelSchedule {
133 pub waves: Vec<ExecutionWave>,
135 pub estimated_time_secs: u64,
137 pub max_parallelism: usize,
139}
140
141#[derive(Debug, Clone, Serialize, Deserialize)]
143pub struct ExecutionWave {
144 pub task_ids: Vec<String>,
146 pub estimated_time_secs: u64,
148}
149
150pub fn create_parallel_schedule(
152 dag: &WorkflowDag,
153 resource_pool: &ResourcePool,
154) -> Result<ParallelSchedule> {
155 let execution_plan = create_execution_plan(dag)?;
156 let mut waves = Vec::new();
157 let mut total_time = 0u64;
158 let mut max_parallelism = 0usize;
159
160 for level in execution_plan {
161 let mut available_resources = AvailableResources::from(resource_pool.clone());
162 let mut current_wave = Vec::new();
163 let mut waiting_tasks = level.clone();
164 let mut wave_time = 0u64;
165
166 let mut i = 0;
168 while i < waiting_tasks.len() {
169 let task_id = &waiting_tasks[i];
170 if let Some(task) = dag.get_task(task_id) {
171 if available_resources.can_allocate(&task.resources) {
172 available_resources.allocate(&task.resources);
173 current_wave.push(task_id.clone());
174 wave_time = wave_time.max(task.timeout_secs.unwrap_or(60));
175 waiting_tasks.remove(i);
176 } else {
177 i += 1;
178 }
179 } else {
180 i += 1;
181 }
182 }
183
184 if !current_wave.is_empty() {
185 max_parallelism = max_parallelism.max(current_wave.len());
186 waves.push(ExecutionWave {
187 task_ids: current_wave,
188 estimated_time_secs: wave_time,
189 });
190 total_time += wave_time;
191 }
192
193 while !waiting_tasks.is_empty() {
195 let mut available_resources = AvailableResources::from(resource_pool.clone());
196 let mut current_wave = Vec::new();
197 let mut wave_time = 0u64;
198 let mut i = 0;
199
200 while i < waiting_tasks.len() {
201 let task_id = &waiting_tasks[i];
202 if let Some(task) = dag.get_task(task_id) {
203 if available_resources.can_allocate(&task.resources) {
204 available_resources.allocate(&task.resources);
205 current_wave.push(task_id.clone());
206 wave_time = wave_time.max(task.timeout_secs.unwrap_or(60));
207 waiting_tasks.remove(i);
208 } else {
209 i += 1;
210 }
211 } else {
212 i += 1;
213 }
214 }
215
216 if !current_wave.is_empty() {
217 max_parallelism = max_parallelism.max(current_wave.len());
218 waves.push(ExecutionWave {
219 task_ids: current_wave,
220 estimated_time_secs: wave_time,
221 });
222 total_time += wave_time;
223 } else {
224 break;
226 }
227 }
228 }
229
230 Ok(ParallelSchedule {
231 waves,
232 estimated_time_secs: total_time,
233 max_parallelism,
234 })
235}
236
237pub fn calculate_resource_utilization(
239 dag: &WorkflowDag,
240 schedule: &ParallelSchedule,
241) -> Vec<ResourceUtilization> {
242 let mut utilization = Vec::new();
243 let mut current_time = 0u64;
244
245 for wave in &schedule.waves {
246 let mut cpu_used = 0.0;
247 let mut memory_used = 0u64;
248 let mut gpus_used = 0u32;
249
250 for task_id in &wave.task_ids {
251 if let Some(task) = dag.get_task(task_id) {
252 cpu_used += task.resources.cpu_cores;
253 memory_used += task.resources.memory_mb;
254 if task.resources.gpu {
255 gpus_used += 1;
256 }
257 }
258 }
259
260 utilization.push(ResourceUtilization {
261 time_secs: current_time,
262 cpu_cores_used: cpu_used,
263 memory_mb_used: memory_used,
264 gpus_used,
265 task_count: wave.task_ids.len(),
266 });
267
268 current_time += wave.estimated_time_secs;
269 }
270
271 utilization
272}
273
274#[derive(Debug, Clone, Serialize, Deserialize)]
276pub struct ResourceUtilization {
277 pub time_secs: u64,
279 pub cpu_cores_used: f64,
281 pub memory_mb_used: u64,
283 pub gpus_used: u32,
285 pub task_count: usize,
287}
288
289pub fn optimize_schedule(
291 dag: &WorkflowDag,
292 resource_pool: &ResourcePool,
293) -> Result<ParallelSchedule> {
294 let schedule = create_parallel_schedule(dag, resource_pool)?;
296
297 let mut waves = schedule.waves;
301
302 let mut optimized_waves: Vec<ExecutionWave> = Vec::new();
304 let mut i = 0;
305
306 while i < waves.len() {
307 let mut current_wave = waves[i].clone();
308 let mut current_resources = AvailableResources::from(resource_pool.clone());
309
310 for task_id in ¤t_wave.task_ids {
312 if let Some(task) = dag.get_task(task_id) {
313 current_resources.allocate(&task.resources);
314 }
315 }
316
317 if i + 1 < waves.len() {
320 let mut current_ids: std::collections::HashSet<String> =
323 current_wave.task_ids.iter().cloned().collect();
324 let next_ids: std::collections::HashSet<String> =
327 waves[i + 1].task_ids.iter().cloned().collect();
328
329 let candidates = waves[i + 1].task_ids.clone();
330 let mut pulled = Vec::new();
331
332 for task_id in &candidates {
333 if let Some(task) = dag.get_task(task_id) {
334 let deps = dag.get_dependencies(task_id);
339 let violates_ordering = deps
340 .iter()
341 .any(|d| current_ids.contains(d) || next_ids.contains(d));
342 if violates_ordering {
343 continue;
344 }
345
346 if current_resources.can_allocate(&task.resources) {
347 current_resources.allocate(&task.resources);
348 current_wave.estimated_time_secs = current_wave
349 .estimated_time_secs
350 .max(task.timeout_secs.unwrap_or(60));
351 current_ids.insert(task_id.clone());
352 current_wave.task_ids.push(task_id.clone());
353 pulled.push(task_id.clone());
354 }
355 }
356 }
357
358 if !pulled.is_empty() {
361 let pulled_set: std::collections::HashSet<&String> = pulled.iter().collect();
362 waves[i + 1].task_ids.retain(|id| !pulled_set.contains(id));
363 }
364 }
365
366 optimized_waves.push(current_wave);
367 i += 1;
368 }
369
370 optimized_waves.retain(|w| !w.task_ids.is_empty());
372
373 let total_time = optimized_waves.iter().map(|w| w.estimated_time_secs).sum();
375 let max_parallelism = optimized_waves
376 .iter()
377 .map(|w| w.task_ids.len())
378 .max()
379 .unwrap_or(0);
380
381 Ok(ParallelSchedule {
382 waves: optimized_waves,
383 estimated_time_secs: total_time,
384 max_parallelism,
385 })
386}
387
388#[cfg(test)]
389mod tests {
390 use super::*;
391 use crate::dag::graph::{ResourceRequirements, RetryPolicy, TaskEdge, TaskNode};
392 use std::collections::HashMap;
393
394 fn small_task(id: &str) -> TaskNode {
395 TaskNode {
396 id: id.to_string(),
397 name: id.to_string(),
398 description: None,
399 config: serde_json::json!({}),
400 retry: RetryPolicy::default(),
401 timeout_secs: Some(60),
402 resources: ResourceRequirements {
403 cpu_cores: 1.0,
404 memory_mb: 128,
405 gpu: false,
406 disk_mb: 0,
407 custom: HashMap::new(),
408 },
409 metadata: HashMap::new(),
410 }
411 }
412
413 #[test]
414 fn test_optimize_schedule_respects_dependencies() {
415 let mut dag = WorkflowDag::new();
419 dag.add_task(small_task("a")).expect("add a");
420 dag.add_task(small_task("b")).expect("add b");
421 dag.add_dependency("a", "b", TaskEdge::default())
422 .expect("add edge");
423
424 let pool = ResourcePool::default();
426 let schedule = optimize_schedule(&dag, &pool).expect("optimize");
427
428 let wave_of = |id: &str| {
429 schedule
430 .waves
431 .iter()
432 .position(|w| w.task_ids.iter().any(|t| t == id))
433 };
434 let a_wave = wave_of("a").expect("a scheduled");
435 let b_wave = wave_of("b").expect("b scheduled");
436
437 assert_ne!(a_wave, b_wave, "B must not share a wave with dependency A");
438 assert!(a_wave < b_wave, "A must run before B");
439
440 let b_count: usize = schedule
442 .waves
443 .iter()
444 .map(|w| w.task_ids.iter().filter(|t| *t == "b").count())
445 .sum();
446 assert_eq!(b_count, 1, "B must be scheduled exactly once");
447 }
448
449 #[test]
450 fn test_optimize_schedule_merges_independent_tasks() {
451 let mut dag = WorkflowDag::new();
455 dag.add_task(small_task("x")).expect("add x");
456 dag.add_task(small_task("y")).expect("add y");
457 let pool = ResourcePool::default();
460 let schedule = optimize_schedule(&dag, &pool).expect("optimize");
461
462 assert_eq!(schedule.waves.len(), 1);
464 assert_eq!(schedule.waves[0].task_ids.len(), 2);
465 }
466}