dotloom_constraints/
graph.rs1use std::collections::BTreeMap;
4
5use crate::{Backend, Problem, VarId};
6
7#[derive(Debug, Clone, PartialEq)]
9pub struct Component {
10 pub vars: Vec<VarId>,
12 pub rules: Vec<usize>,
14 pub backend: Backend,
16}
17
18struct Dsu {
19 parent: Vec<usize>,
20}
21
22impl Dsu {
23 fn new(n: usize) -> Self {
24 Self { parent: (0..n).collect() }
25 }
26
27 fn find(&mut self, mut a: usize) -> usize {
28 while let Some(&p) = self.parent.get(a) {
29 if p == a {
30 break;
31 }
32 let gp = self.parent.get(p).copied().unwrap_or(p);
33 if let Some(slot) = self.parent.get_mut(a) {
34 *slot = gp;
35 }
36 a = p;
37 }
38 a
39 }
40
41 fn union(&mut self, a: usize, b: usize) {
42 let (ra, rb) = (self.find(a), self.find(b));
43 if ra != rb {
44 let (lo, hi) = if ra < rb { (ra, rb) } else { (rb, ra) };
46 if let Some(slot) = self.parent.get_mut(hi) {
47 *slot = lo;
48 }
49 }
50 }
51}
52
53#[must_use]
62pub fn components(p: &Problem) -> Vec<Component> {
63 let n = p.vars.len();
64 let mut dsu = Dsu::new(n);
65 let fixed = |v: VarId| p.is_fixed(v);
66 let mut rule_vars: Vec<Vec<VarId>> = Vec::with_capacity(p.rules.len());
67 for r in &p.rules {
68 let vs: Vec<VarId> = r.vars().into_iter().filter(|v| !fixed(*v)).collect();
69 for w in vs.windows(2) {
70 dsu.union(w[0].index(), w[1].index());
71 }
72 rule_vars.push(vs);
73 }
74 let mut by_root: BTreeMap<usize, Component> = BTreeMap::new();
75 let mut constant_rules = Vec::new();
76 for (ri, vs) in rule_vars.iter().enumerate() {
77 match vs.first() {
78 None => constant_rules.push(ri),
79 Some(v) => {
80 let root = dsu.find(v.index());
81 by_root
82 .entry(root)
83 .or_insert_with(|| Component { vars: Vec::new(), rules: Vec::new(), backend: Backend::Linear })
84 .rules
85 .push(ri);
86 }
87 }
88 }
89 let x = p.values();
90 for v in 0..n {
91 let id = VarId(u32::try_from(v).unwrap_or(u32::MAX));
92 if fixed(id) {
93 continue;
94 }
95 let root = dsu.find(v);
96 by_root
97 .entry(root)
98 .or_insert_with(|| Component { vars: Vec::new(), rules: Vec::new(), backend: Backend::Trivial })
99 .vars
100 .push(id);
101 }
102 let mut out: Vec<Component> = by_root.into_values().collect();
103 for c in &mut out {
104 if c.rules.is_empty() {
105 c.backend = Backend::Trivial;
106 continue;
107 }
108 let linear = c.rules.iter().all(|ri| {
109 p.rules.get(*ri).is_some_and(|r| {
110 r.unsupported.is_none() && r.rows.iter().all(|row| row.expr.linear_form(&x, &fixed).is_some())
111 })
112 });
113 c.backend = if linear { Backend::Linear } else { Backend::Numeric };
114 }
115 for ri in constant_rules {
116 out.push(Component { vars: Vec::new(), rules: vec![ri], backend: Backend::Trivial });
117 }
118 out
119}
120
121#[cfg(test)]
122mod tests {
123 use super::*;
124 use crate::{Expr, Rule, Strength, Variable, rules};
125
126 #[test]
127 fn components_split_and_classify() {
128 let mut p = Problem::default();
129 let a = p.add_var(Variable::new(1.0));
130 let b = p.add_var(Variable::new(2.0));
131 let c = p.add_var(Variable::new(3.0));
132 let d = p.add_var(Variable::new(4.0));
133 let f = p.add_var(Variable::new(5.0).fixed(true));
134 p.rules.push(Rule::new(1, rules::equal(Expr::Var(a), Expr::Var(b), 1.0), Strength::Required));
136 p.rules.push(Rule::new(2, rules::equal(Expr::Var(a), Expr::Var(f), 1.0), Strength::Required));
137 p.rules.push(Rule::new(3, rules::fix(Expr::mul(Expr::Var(c), Expr::Var(d)), 12.0, 1.0), Strength::Required));
138 let comps = components(&p);
139 assert_eq!(comps.len(), 2);
140 assert_eq!(comps[0].vars, vec![a, b]);
141 assert_eq!(comps[0].backend, Backend::Linear);
142 assert_eq!(comps[0].rules, vec![0, 1]);
143 assert_eq!(comps[1].vars, vec![c, d]);
144 assert_eq!(comps[1].backend, Backend::Numeric);
145 }
146
147 #[test]
148 fn constant_rule_is_its_own_component() {
149 let mut p = Problem::default();
150 let f = p.add_var(Variable::new(5.0).fixed(true));
151 p.rules.push(Rule::new(1, rules::fix(Expr::Var(f), 6.0, 1.0), Strength::Required));
152 let comps = components(&p);
153 assert_eq!(comps.len(), 1);
154 assert!(comps[0].vars.is_empty());
155 assert_eq!(comps[0].backend, Backend::Trivial);
156 }
157}