1use kasuari::{Constraint, Expression, RelationalOperator, Solver, Strength as KStrength, Term, Variable as KVar};
27
28use crate::{
29 Backend, ComponentReport, Problem, Relation, Solution, SolveOptions, Status, Strength, VarId,
30 analysis::Columns,
31 graph::{Component, components},
32 job::solve_trivial,
33 linear::{STAY_BASE, kstrength, read_back, verify_linear},
34};
35
36#[derive(Debug, Clone, PartialEq)]
38struct Spec {
39 terms: Vec<(usize, f64)>,
40 constant: f64,
41 relation: Relation,
42 strength: f64,
43}
44
45#[derive(Debug)]
47struct Slot {
48 spec: Spec,
49 constraint: Constraint,
50}
51
52#[derive(Debug)]
54struct Edit {
55 aux: KVar,
56 terms: Vec<(usize, f64)>,
57 link: Constraint,
58 value: f64,
59}
60
61#[derive(Debug, Clone, PartialEq)]
63struct EditSpec {
64 terms: Vec<(usize, f64)>,
65 value: f64,
66 strength: f64,
67}
68
69struct LinearPart {
70 comp: Component,
71 cols: Columns,
72 kvars: Vec<KVar>,
73 solver: Solver,
74 slots: Vec<Slot>,
75 edits: Vec<Edit>,
76}
77
78enum Part {
79 Trivial(Component),
80 Linear(Box<LinearPart>),
81}
82
83impl core::fmt::Debug for Part {
84 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
85 match self {
86 Self::Trivial(c) => f.debug_tuple("Trivial").field(c).finish(),
87 Self::Linear(p) => f
88 .debug_struct("Linear")
89 .field("component", &p.comp)
90 .field("constraints", &p.slots.len())
91 .field("edits", &p.edits.len())
92 .finish_non_exhaustive(),
93 }
94 }
95}
96
97#[derive(Debug, PartialEq)]
99struct Signature {
100 vars: Vec<(bool, u64, u64)>,
101 rules: Vec<(Strength, Vec<VarId>, Vec<Relation>, bool)>,
102 targets: Vec<(VarId, Strength)>,
103}
104
105fn signature(p: &Problem) -> Signature {
106 Signature {
107 vars: p.vars.iter().map(|v| (v.fixed, v.scale.to_bits(), crate::stay_factor(v.stay).to_bits())).collect(),
108 rules: p
109 .rules
110 .iter()
111 .map(|r| {
112 let free: Vec<VarId> = r.vars().into_iter().filter(|v| !p.is_fixed(*v)).collect();
113 (r.strength, free, r.rows.iter().map(|row| row.relation).collect(), r.unsupported.is_some())
114 })
115 .collect(),
116 targets: p.targets.iter().map(|t| (t.var, t.strength)).collect(),
117 }
118}
119
120fn row_form(p: &Problem, rule: usize, row: usize, cols: &Columns, x: &[f64]) -> Option<(Vec<(usize, f64)>, f64)> {
122 let r = p.rules.get(rule)?.rows.get(row)?;
123 let lf = r.expr.linear_form(x, &|v| p.is_fixed(v))?;
124 let sigma = if r.scale.is_finite() && r.scale > 0.0 { r.scale } else { 1.0 };
125 let mut terms = Vec::with_capacity(lf.terms.len());
126 for (v, c) in lf.terms {
127 let col = cols.col.get(v.index()).copied().flatten()?;
128 let s = cols.scales.get(col).copied().unwrap_or(1.0);
129 terms.push((col, c * s / sigma));
130 }
131 Some((terms, lf.constant / sigma))
132}
133
134fn constraint(spec: &Spec, kvars: &[KVar]) -> Option<Constraint> {
135 let mut terms = Vec::with_capacity(spec.terms.len());
136 for &(col, c) in &spec.terms {
137 terms.push(Term::new(*kvars.get(col)?, c));
138 }
139 let op = match spec.relation {
140 Relation::Eq => RelationalOperator::Equal,
141 Relation::Le => RelationalOperator::LessOrEqual,
142 };
143 Some(Constraint::new(Expression::new(terms, spec.constant), op, KStrength::new(spec.strength)))
144}
145
146fn link(terms: &[(usize, f64)], aux: KVar, kvars: &[KVar]) -> Option<Constraint> {
147 let mut t = Vec::with_capacity(terms.len() + 1);
148 for &(col, c) in terms {
149 t.push(Term::new(*kvars.get(col)?, c));
150 }
151 t.push(Term::new(aux, -1.0));
152 Some(Constraint::new(Expression::new(t, 0.0), RelationalOperator::Equal, KStrength::REQUIRED))
153}
154
155impl LinearPart {
156 fn specs(&self, p: &Problem, moving: &[usize], x: &[f64]) -> Option<(Vec<Spec>, Vec<EditSpec>)> {
159 let mut slots = Vec::new();
160 let mut edits = Vec::new();
161 for hard_pass in [true, false] {
162 for &ri in &self.comp.rules {
163 let rule = p.rules.get(ri)?;
164 if rule.is_hard() != hard_pass {
165 continue;
166 }
167 let is_moving = moving.binary_search(&ri).is_ok();
168 for (k, row) in rule.rows.iter().enumerate() {
169 let (terms, constant) = row_form(p, ri, k, &self.cols, x)?;
170 if is_moving {
171 if row.relation != Relation::Eq {
172 return None;
173 }
174 let strength = if rule.is_hard() { Strength::Strong } else { rule.strength };
175 edits.push(EditSpec { terms, value: -constant, strength: kstrength(strength).value() });
176 } else {
177 let strength = if hard_pass { KStrength::REQUIRED } else { kstrength(rule.strength) };
178 slots.push(Spec { terms, constant, relation: row.relation, strength: strength.value() });
179 }
180 }
181 }
182 }
183 for t in &p.targets {
184 let Some(col) = self.cols.col.get(t.var.index()).copied().flatten() else { continue };
185 let s = self.cols.scales.get(col).copied().unwrap_or(1.0);
186 let strength = if t.strength == Strength::Required { Strength::Strong } else { t.strength };
187 slots.push(Spec {
188 terms: vec![(col, 1.0)],
189 constant: -t.value / s,
190 relation: Relation::Eq,
191 strength: kstrength(strength).value(),
192 });
193 }
194 let n = self.cols.n();
195 for (col, var) in self.cols.vars.iter().enumerate() {
196 let s = self.cols.scales.get(col).copied().unwrap_or(1.0);
197 let reference = x.get(var.index()).copied().unwrap_or(0.0);
198 let mult = crate::stay_factor(p.vars.get(var.index()).map_or(1.0, |v| v.stay));
199 let w = STAY_BASE * mult * (1.0 + 0.5 * col as f64 / n.max(1) as f64);
200 slots.push(Spec { terms: vec![(col, 1.0)], constant: -reference / s, relation: Relation::Eq, strength: w });
201 }
202 Some((slots, edits))
203 }
204
205 fn new(p: &Problem, comp: Component, moving: &[usize], x: &[f64]) -> Option<Self> {
206 let cols = Columns::new(p, &comp.vars);
207 let kvars: Vec<KVar> = (0..cols.n()).map(|_| KVar::new()).collect();
208 let mut part = Self { comp, cols, kvars, solver: Solver::new(), slots: Vec::new(), edits: Vec::new() };
209 let (slots, edits) = part.specs(p, moving, x)?;
210 for spec in slots {
211 let c = constraint(&spec, &part.kvars)?;
212 let required = spec.strength >= KStrength::REQUIRED.value();
213 if part.solver.add_constraint(c.clone()).is_err() && required {
214 return None;
215 }
216 part.slots.push(Slot { spec, constraint: c });
217 }
218 for e in edits {
219 let aux = KVar::new();
220 let l = link(&e.terms, aux, &part.kvars)?;
221 part.solver.add_constraint(l.clone()).ok()?;
222 part.solver.add_edit_variable(aux, KStrength::new(e.strength)).ok()?;
223 part.solver.suggest_value(aux, e.value).ok()?;
224 part.edits.push(Edit { aux, terms: e.terms, link: l, value: e.value });
225 }
226 Some(part)
227 }
228
229 fn update(&mut self, p: &Problem, moving: &[usize], x: &[f64]) -> Option<()> {
232 let (slots, edits) = self.specs(p, moving, x)?;
233 if slots.len() != self.slots.len() || edits.len() != self.edits.len() {
234 return None;
235 }
236 for (cur, want) in self.slots.iter_mut().zip(slots) {
237 if cur.spec == want {
238 continue;
239 }
240 let c = constraint(&want, &self.kvars)?;
241 self.solver.remove_constraint(&cur.constraint).ok()?;
242 self.solver.add_constraint(c.clone()).ok()?;
243 *cur = Slot { spec: want, constraint: c };
244 }
245 for (cur, want) in self.edits.iter_mut().zip(edits) {
246 if cur.terms != want.terms {
247 let l = link(&want.terms, cur.aux, &self.kvars)?;
248 self.solver.remove_constraint(&cur.link).ok()?;
249 self.solver.add_constraint(l.clone()).ok()?;
250 cur.link = l;
251 cur.terms = want.terms;
252 }
253 if cur.value.to_bits() != want.value.to_bits() {
254 self.solver.suggest_value(cur.aux, want.value).ok()?;
255 cur.value = want.value;
256 }
257 }
258 Some(())
259 }
260}
261
262#[derive(Debug)]
265pub struct LinearSession {
266 signature: Signature,
267 moving: Vec<usize>,
268 parts: Vec<Part>,
269 opts: SolveOptions,
270 broken: bool,
271}
272
273impl LinearSession {
274 #[must_use]
280 pub fn new(problem: &Problem, moving: &[usize], opts: SolveOptions) -> Option<Self> {
281 if problem.vars.iter().any(|v| !v.value.is_finite()) || problem.rules.iter().any(|r| r.unsupported.is_some()) {
282 return None;
283 }
284 let mut moving = moving.to_vec();
285 moving.sort_unstable();
286 moving.dedup();
287 let x = problem.values();
288 let mut parts = Vec::new();
289 for comp in components(problem) {
290 match comp.backend {
291 Backend::Numeric => return None,
292 Backend::Trivial => parts.push(Part::Trivial(comp)),
293 Backend::Linear => parts.push(Part::Linear(Box::new(LinearPart::new(problem, comp, &moving, &x)?))),
294 }
295 }
296 Some(Self { signature: signature(problem), moving, parts, opts, broken: false })
297 }
298
299 #[must_use]
302 pub fn is_broken(&self) -> bool {
303 self.broken
304 }
305
306 pub fn resolve(&mut self, problem: &Problem) -> Option<Solution> {
311 if self.broken || signature(problem) != self.signature {
312 return None;
313 }
314 let x_in = problem.values();
315 if x_in.iter().any(|v| !v.is_finite()) {
316 return None;
317 }
318 let mut x = x_in.clone();
319 let mut reports = Vec::with_capacity(self.parts.len());
320 let mut diagnostics = Vec::new();
321 let mut iterations = 0;
322 for part in &mut self.parts {
323 match part {
324 Part::Trivial(comp) => {
325 let (status, d, max_hard) = solve_trivial(problem, comp, &mut x, &self.opts);
326 if !status.is_acceptable() {
327 return None;
328 }
329 diagnostics.extend(d);
330 reports.push(report(problem, comp, Backend::Trivial, status, 0, max_hard));
331 }
332 Part::Linear(lp) => {
333 if lp.update(problem, &self.moving, &x_in).is_none() {
334 self.broken = true;
335 return None;
336 }
337 read_back(&lp.solver, &lp.cols, &lp.kvars, &mut x);
338 let out = verify_linear(problem, &lp.comp, &lp.cols, &mut x, &x_in, &self.opts);
339 if !out.status.is_acceptable() {
340 return None;
341 }
342 iterations += 1;
343 diagnostics.extend(out.diagnostics);
344 reports.push(report(problem, &lp.comp, Backend::Linear, out.status, 1, out.max_hard));
345 }
346 }
347 }
348 let status = reports.iter().fold(Status::Solved, |acc, r| acc.combine(r.status));
349 Some(Solution { values: x, status, components: reports, diagnostics, iterations })
350 }
351}
352
353fn report(
354 p: &Problem,
355 comp: &Component,
356 backend: Backend,
357 status: Status,
358 iterations: u32,
359 max_hard: f64,
360) -> ComponentReport {
361 ComponentReport {
362 vars: comp.vars.clone(),
363 rules: comp.rules.iter().filter_map(|ri| p.rules.get(*ri).map(|r| r.id)).collect(),
364 backend,
365 status,
366 iterations,
367 max_hard_residual: max_hard,
368 }
369}