1use std::collections::HashMap;
8use std::sync::Arc;
9
10use serde_json::Value;
11
12use crate::error::ConfigError;
13use crate::state::State;
14
15use super::contract::ComputedContract;
16
17pub type ComputeFn = Arc<dyn Fn(&State) -> Option<Value> + Send + Sync>;
19
20pub struct ComputedVar {
27 pub key: String,
29 pub dependencies: Vec<String>,
31 pub compute: ComputeFn,
33}
34
35pub struct ComputedRegistry {
41 vars: Vec<ComputedVar>,
43 dep_index: HashMap<String, Vec<usize>>,
46}
47
48impl Default for ComputedRegistry {
49 fn default() -> Self {
50 Self::new()
51 }
52}
53
54impl ComputedRegistry {
55 pub fn new() -> Self {
57 Self {
58 vars: Vec::new(),
59 dep_index: HashMap::new(),
60 }
61 }
62
63 pub fn register(&mut self, var: ComputedVar) -> Result<(), ConfigError> {
70 let previous = match self.vars.iter().position(|v| v.key == var.key) {
71 Some(pos) => Some((pos, std::mem::replace(&mut self.vars[pos], var))),
72 None => {
73 self.vars.push(var);
74 None
75 }
76 };
77 if let Err(err) = self.topo_sort() {
78 match previous {
80 Some((pos, old)) => self.vars[pos] = old,
81 None => {
82 self.vars.pop();
83 }
84 }
85 return Err(err);
86 }
87 self.rebuild_dep_index();
88 Ok(())
89 }
90
91 pub fn recompute(&self, state: &State) -> Vec<String> {
94 let mut changed = Vec::new();
95 for var in &self.vars {
96 if let Some(new_val) = (var.compute)(state) {
97 let derived_key = format!("derived:{}", var.key);
98 let old_val = state.get_raw(&derived_key);
99 let did_change = old_val.as_ref() != Some(&new_val);
100 let _ = state.set(&derived_key, new_val);
101 if did_change {
102 changed.push(var.key.clone());
103 }
104 }
105 }
106 changed
107 }
108
109 pub fn recompute_affected(&self, state: &State, changed_keys: &[String]) -> Vec<String> {
115 let mut visited = vec![false; self.vars.len()];
117 let mut affected_set = Vec::new();
118
119 let mut work_keys: Vec<String> = changed_keys.to_vec();
121
122 while let Some(key) = work_keys.pop() {
123 if let Some(indices) = self.dep_index.get(&key) {
125 for &idx in indices {
126 if !visited[idx] {
127 visited[idx] = true;
128 affected_set.push(idx);
129 work_keys.push(self.vars[idx].key.clone());
132 }
133 }
134 }
135 }
136
137 affected_set.sort_unstable();
139
140 let mut changed = Vec::new();
141 for idx in affected_set {
142 let var = &self.vars[idx];
143 if let Some(new_val) = (var.compute)(state) {
144 let derived_key = format!("derived:{}", var.key);
145 let old_val = state.get_raw(&derived_key);
146 let did_change = old_val.as_ref() != Some(&new_val);
147 let _ = state.set(&derived_key, new_val);
148 if did_change {
149 changed.push(var.key.clone());
150 }
151 }
152 }
153 changed
154 }
155
156 pub fn validate(&self) -> Result<(), ConfigError> {
159 let n = self.vars.len();
161 if n == 0 {
162 return Ok(());
163 }
164
165 let key_to_idx: HashMap<&str, usize> = self
166 .vars
167 .iter()
168 .enumerate()
169 .map(|(i, v)| (v.key.as_str(), i))
170 .collect();
171
172 let mut in_degree = vec![0usize; n];
173 let mut adj: Vec<Vec<usize>> = vec![Vec::new(); n];
174
175 for (i, var) in self.vars.iter().enumerate() {
176 for dep in &var.dependencies {
177 if let Some(&dep_idx) = key_to_idx.get(dep.as_str()) {
178 adj[dep_idx].push(i);
179 in_degree[i] += 1;
180 }
181 }
183 }
184
185 let mut queue: Vec<usize> = (0..n).filter(|&i| in_degree[i] == 0).collect();
186 let mut visited = 0usize;
187
188 while let Some(node) = queue.pop() {
189 visited += 1;
190 for &neighbor in &adj[node] {
191 in_degree[neighbor] -= 1;
192 if in_degree[neighbor] == 0 {
193 queue.push(neighbor);
194 }
195 }
196 }
197
198 if visited == n {
199 Ok(())
200 } else {
201 let cycle_vars: Vec<&str> = (0..n)
203 .filter(|&i| in_degree[i] > 0)
204 .map(|i| self.vars[i].key.as_str())
205 .collect();
206 Err(ConfigError::new(format!(
207 "Cycle detected among computed variables: {cycle_vars:?}"
208 )))
209 }
210 }
211
212 pub fn len(&self) -> usize {
214 self.vars.len()
215 }
216
217 pub fn is_empty(&self) -> bool {
219 self.vars.is_empty()
220 }
221
222 pub fn describe(&self) -> Vec<ComputedContract> {
224 self.vars
225 .iter()
226 .map(|var| ComputedContract {
227 key: var.key.clone(),
228 dependencies: var.dependencies.clone(),
229 })
230 .collect()
231 }
232
233 fn topo_sort(&mut self) -> Result<(), ConfigError> {
239 let n = self.vars.len();
240
241 for var in &self.vars {
243 if var.dependencies.contains(&var.key) {
244 return Err(ConfigError::new(format!(
245 "Cycle detected among computed variables: {:?}",
246 vec![var.key.as_str()]
247 )));
248 }
249 }
250
251 if n <= 1 {
252 return Ok(());
253 }
254
255 let key_to_idx: HashMap<&str, usize> = self
257 .vars
258 .iter()
259 .enumerate()
260 .map(|(i, v)| (v.key.as_str(), i))
261 .collect();
262
263 let mut in_degree = vec![0usize; n];
266 let mut adj: Vec<Vec<usize>> = vec![Vec::new(); n];
267
268 for (i, var) in self.vars.iter().enumerate() {
269 for dep in &var.dependencies {
270 if let Some(&dep_idx) = key_to_idx.get(dep.as_str()) {
271 adj[dep_idx].push(i);
272 in_degree[i] += 1;
273 }
274 }
275 }
276
277 let mut queue: Vec<usize> = (0..n).filter(|&i| in_degree[i] == 0).collect();
279 let mut order: Vec<usize> = Vec::with_capacity(n);
280
281 while let Some(node) = queue.pop() {
282 order.push(node);
283 for &neighbor in &adj[node] {
284 in_degree[neighbor] -= 1;
285 if in_degree[neighbor] == 0 {
286 queue.push(neighbor);
287 }
288 }
289 }
290
291 if order.len() != n {
292 let cycle_vars: Vec<&str> = (0..n)
293 .filter(|&i| in_degree[i] > 0)
294 .map(|i| self.vars[i].key.as_str())
295 .collect();
296 return Err(ConfigError::new(format!(
297 "Cycle detected among computed variables: {cycle_vars:?}"
298 )));
299 }
300
301 let mut slots: Vec<Option<ComputedVar>> = self.vars.drain(..).map(Some).collect();
304 for &idx in &order {
305 if let Some(var) = slots[idx].take() {
306 self.vars.push(var);
307 }
308 }
309 Ok(())
310 }
311
312 fn rebuild_dep_index(&mut self) {
314 self.dep_index.clear();
315 for (i, var) in self.vars.iter().enumerate() {
316 for dep in &var.dependencies {
317 self.dep_index.entry(dep.clone()).or_default().push(i);
318 }
319 }
320 }
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326 use serde_json::json;
327
328 #[test]
331 fn single_var_register_and_recompute() {
332 let mut registry = ComputedRegistry::new();
333 registry
334 .register(ComputedVar {
335 key: "doubled".into(),
336 dependencies: vec!["app:count".into()],
337 compute: Arc::new(|state| {
338 let count: i64 = state.get("app:count")?;
339 Some(json!(count * 2))
340 }),
341 })
342 .unwrap();
343
344 let state = State::new();
345 let _ = state.set("app:count", 5);
346
347 let changed = registry.recompute(&state);
348 assert_eq!(changed, vec!["doubled"]);
349 assert_eq!(state.get::<i64>("derived:doubled"), Some(10));
350 }
351
352 #[test]
355 fn dependency_ordering() {
356 let mut registry = ComputedRegistry::new();
357
358 registry
360 .register(ComputedVar {
361 key: "derived_from_base".into(),
362 dependencies: vec!["base".into()],
363 compute: Arc::new(|state| {
364 let base: i64 = state.get("derived:base")?;
365 Some(json!(base + 100))
366 }),
367 })
368 .unwrap();
369
370 registry
372 .register(ComputedVar {
373 key: "base".into(),
374 dependencies: vec!["app:input".into()],
375 compute: Arc::new(|state| {
376 let input: i64 = state.get("app:input")?;
377 Some(json!(input * 2))
378 }),
379 })
380 .unwrap();
381
382 let state = State::new();
383 let _ = state.set("app:input", 3);
384
385 let changed = registry.recompute(&state);
386 assert_eq!(state.get::<i64>("derived:base"), Some(6));
388 assert_eq!(state.get::<i64>("derived:derived_from_base"), Some(106));
389 assert!(changed.contains(&"base".to_string()));
390 assert!(changed.contains(&"derived_from_base".to_string()));
391 }
392
393 #[test]
396 fn cycle_detection_is_an_error_and_rolls_back() {
397 let mut registry = ComputedRegistry::new();
398 registry
399 .register(ComputedVar {
400 key: "a".into(),
401 dependencies: vec!["b".into()],
402 compute: Arc::new(|_| Some(json!(1))),
403 })
404 .unwrap();
405 let result = registry.register(ComputedVar {
406 key: "b".into(),
407 dependencies: vec!["a".into()],
408 compute: Arc::new(|_| Some(json!(2))),
409 });
410 let err = result.expect_err("cycle must be rejected");
411 assert!(err.to_string().contains("Cycle detected"), "{err}");
412 assert_eq!(registry.len(), 1);
414 assert!(registry.validate().is_ok());
415 }
416
417 #[test]
420 fn recompute_returns_only_changed_keys() {
421 let mut registry = ComputedRegistry::new();
422 registry
423 .register(ComputedVar {
424 key: "level".into(),
425 dependencies: vec!["app:score".into()],
426 compute: Arc::new(|state| {
427 let score: f64 = state.get("app:score")?;
428 if score > 0.5 {
429 Some(json!("high"))
430 } else {
431 Some(json!("low"))
432 }
433 }),
434 })
435 .unwrap();
436
437 let state = State::new();
438 let _ = state.set("app:score", 0.8);
439
440 let changed = registry.recompute(&state);
442 assert_eq!(changed, vec!["level"]);
443
444 let changed = registry.recompute(&state);
446 assert!(changed.is_empty());
447
448 let _ = state.set("app:score", 0.2);
450 let changed = registry.recompute(&state);
451 assert_eq!(changed, vec!["level"]);
452 assert_eq!(
453 state.get::<String>("derived:level"),
454 Some("low".to_string())
455 );
456 }
457
458 #[test]
461 fn recompute_affected_only_recomputes_affected() {
462 let call_count_a = Arc::new(std::sync::atomic::AtomicUsize::new(0));
463 let call_count_b = Arc::new(std::sync::atomic::AtomicUsize::new(0));
464
465 let cc_a = call_count_a.clone();
466 let cc_b = call_count_b.clone();
467
468 let mut registry = ComputedRegistry::new();
469 registry
470 .register(ComputedVar {
471 key: "from_x".into(),
472 dependencies: vec!["app:x".into()],
473 compute: Arc::new(move |state| {
474 cc_a.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
475 let x: i64 = state.get("app:x")?;
476 Some(json!(x + 1))
477 }),
478 })
479 .unwrap();
480 registry
481 .register(ComputedVar {
482 key: "from_y".into(),
483 dependencies: vec!["app:y".into()],
484 compute: Arc::new(move |state| {
485 cc_b.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
486 let y: i64 = state.get("app:y")?;
487 Some(json!(y + 1))
488 }),
489 })
490 .unwrap();
491
492 let state = State::new();
493 let _ = state.set("app:x", 10);
494 let _ = state.set("app:y", 20);
495
496 let changed = registry.recompute_affected(&state, &["app:x".into()]);
498 assert_eq!(changed, vec!["from_x"]);
499 assert_eq!(call_count_a.load(std::sync::atomic::Ordering::SeqCst), 1);
500 assert_eq!(call_count_b.load(std::sync::atomic::Ordering::SeqCst), 0);
501
502 assert_eq!(state.get::<i64>("derived:from_x"), Some(11));
503 assert_eq!(state.get_raw("derived:from_y"), None);
505 }
506
507 #[test]
510 fn validate_catches_cycles() {
511 let mut registry = ComputedRegistry::new();
512 registry.vars.push(ComputedVar {
514 key: "x".into(),
515 dependencies: vec!["y".into()],
516 compute: Arc::new(|_| Some(json!(1))),
517 });
518 registry.vars.push(ComputedVar {
519 key: "y".into(),
520 dependencies: vec!["x".into()],
521 compute: Arc::new(|_| Some(json!(2))),
522 });
523
524 let result = registry.validate();
525 assert!(result.is_err());
526 let msg = result.unwrap_err().to_string();
527 assert!(msg.contains("Cycle detected"));
528 }
529
530 #[test]
533 fn validate_succeeds_on_valid_graph() {
534 let mut registry = ComputedRegistry::new();
535 registry
536 .register(ComputedVar {
537 key: "a".into(),
538 dependencies: vec!["app:input".into()],
539 compute: Arc::new(|_| Some(json!(1))),
540 })
541 .unwrap();
542 registry
543 .register(ComputedVar {
544 key: "b".into(),
545 dependencies: vec!["a".into()],
546 compute: Arc::new(|_| Some(json!(2))),
547 })
548 .unwrap();
549
550 assert!(registry.validate().is_ok());
551 }
552
553 #[test]
556 fn compute_returning_none_skips_write() {
557 let mut registry = ComputedRegistry::new();
558 registry
559 .register(ComputedVar {
560 key: "maybe".into(),
561 dependencies: vec!["app:flag".into()],
562 compute: Arc::new(|state| {
563 let flag: bool = state.get("app:flag")?;
564 if flag { Some(json!("yes")) } else { None }
565 }),
566 })
567 .unwrap();
568
569 let state = State::new();
570 let changed = registry.recompute(&state);
572 assert!(changed.is_empty());
573 assert_eq!(state.get_raw("derived:maybe"), None);
574
575 let _ = state.set("app:flag", false);
577 let changed = registry.recompute(&state);
578 assert!(changed.is_empty());
579 assert_eq!(state.get_raw("derived:maybe"), None);
580
581 let _ = state.set("app:flag", true);
583 let changed = registry.recompute(&state);
584 assert_eq!(changed, vec!["maybe"]);
585 assert_eq!(
586 state.get::<String>("derived:maybe"),
587 Some("yes".to_string())
588 );
589 }
590
591 #[test]
594 fn diamond_dependency() {
595 let mut registry = ComputedRegistry::new();
603
604 registry
605 .register(ComputedVar {
606 key: "d".into(),
607 dependencies: vec!["app:root".into()],
608 compute: Arc::new(|state| {
609 let root: i64 = state.get("app:root")?;
610 Some(json!(root))
611 }),
612 })
613 .unwrap();
614
615 registry
616 .register(ComputedVar {
617 key: "a".into(),
618 dependencies: vec!["d".into()],
619 compute: Arc::new(|state| {
620 let d: i64 = state.get("derived:d")?;
621 Some(json!(d + 10))
622 }),
623 })
624 .unwrap();
625
626 registry
627 .register(ComputedVar {
628 key: "b".into(),
629 dependencies: vec!["d".into()],
630 compute: Arc::new(|state| {
631 let d: i64 = state.get("derived:d")?;
632 Some(json!(d + 20))
633 }),
634 })
635 .unwrap();
636
637 registry
638 .register(ComputedVar {
639 key: "c".into(),
640 dependencies: vec!["a".into(), "b".into()],
641 compute: Arc::new(|state| {
642 let a: i64 = state.get("derived:a")?;
643 let b: i64 = state.get("derived:b")?;
644 Some(json!(a + b))
645 }),
646 })
647 .unwrap();
648
649 let state = State::new();
650 let _ = state.set("app:root", 1);
651
652 let changed = registry.recompute(&state);
653 assert_eq!(state.get::<i64>("derived:d"), Some(1));
654 assert_eq!(state.get::<i64>("derived:a"), Some(11));
655 assert_eq!(state.get::<i64>("derived:b"), Some(21));
656 assert_eq!(state.get::<i64>("derived:c"), Some(32));
657 assert_eq!(changed.len(), 4);
658 }
659
660 #[test]
663 fn empty_registry_recompute_returns_empty() {
664 let registry = ComputedRegistry::new();
665 let state = State::new();
666 let changed = registry.recompute(&state);
667 assert!(changed.is_empty());
668 }
669
670 #[test]
673 fn len_and_is_empty() {
674 let mut registry = ComputedRegistry::new();
675 assert!(registry.is_empty());
676 assert_eq!(registry.len(), 0);
677
678 registry
679 .register(ComputedVar {
680 key: "x".into(),
681 dependencies: vec![],
682 compute: Arc::new(|_| Some(json!(1))),
683 })
684 .unwrap();
685 assert!(!registry.is_empty());
686 assert_eq!(registry.len(), 1);
687 }
688
689 #[test]
692 fn recompute_affected_diamond() {
693 let mut registry = ComputedRegistry::new();
694
695 registry
696 .register(ComputedVar {
697 key: "root_derived".into(),
698 dependencies: vec!["app:root".into()],
699 compute: Arc::new(|state| {
700 let r: i64 = state.get("app:root")?;
701 Some(json!(r * 10))
702 }),
703 })
704 .unwrap();
705
706 registry
707 .register(ComputedVar {
708 key: "leaf".into(),
709 dependencies: vec!["root_derived".into()],
710 compute: Arc::new(|state| {
711 let rd: i64 = state.get("derived:root_derived")?;
712 Some(json!(rd + 5))
713 }),
714 })
715 .unwrap();
716
717 let state = State::new();
718 let _ = state.set("app:root", 2);
719
720 registry.recompute(&state);
722 assert_eq!(state.get::<i64>("derived:root_derived"), Some(20));
723 assert_eq!(state.get::<i64>("derived:leaf"), Some(25));
724
725 let _ = state.set("app:root", 3);
727 let changed = registry.recompute_affected(&state, &["app:root".into()]);
728 assert!(changed.contains(&"root_derived".to_string()));
730 assert_eq!(state.get::<i64>("derived:root_derived"), Some(30));
731 assert!(changed.contains(&"leaf".to_string()));
734 assert_eq!(state.get::<i64>("derived:leaf"), Some(35));
735 }
736
737 #[test]
740 fn validate_empty_registry() {
741 let registry = ComputedRegistry::new();
742 assert!(registry.validate().is_ok());
743 }
744
745 #[test]
748 fn self_cycle_is_an_error() {
749 let mut registry = ComputedRegistry::new();
750 let err = registry
751 .register(ComputedVar {
752 key: "self_ref".into(),
753 dependencies: vec!["self_ref".into()],
754 compute: Arc::new(|_| Some(json!(1))),
755 })
756 .expect_err("self-cycle must be rejected");
757 assert!(err.to_string().contains("Cycle detected"), "{err}");
758 assert!(registry.is_empty());
759 }
760}