1use crate::{
10 Backend, Certainty, ComponentReport, Diagnostic, DiagnosticKind, Problem, Solution, SolveOptions, Status, Strength,
11 analysis::{Columns, eval_rows, max_violation},
12 graph::{Component, components},
13 linear::{conflict_diagnostic, solve_linear},
14 numeric::NumericSolver,
15};
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
19pub enum Progress {
20 Running {
22 completed: usize,
24 total: usize,
26 },
27 Finished,
29}
30
31#[derive(Debug)]
33pub struct SolveJob {
34 problem: Problem,
35 opts: SolveOptions,
36 x: Vec<f64>,
37 x_in: Vec<f64>,
38 comps: Vec<Component>,
39 next: usize,
40 current: Option<NumericSolverBox>,
41 reports: Vec<ComponentReport>,
42 diagnostics: Vec<Diagnostic>,
43 cancelled: bool,
44 invalid: Option<String>,
45 iterations: u32,
46}
47
48struct NumericSolverBox(NumericSolver);
50
51impl core::fmt::Debug for NumericSolverBox {
52 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
53 f.write_str("NumericSolver")
54 }
55}
56
57fn validate(p: &Problem) -> Result<(), String> {
58 let n = p.vars.len();
59 if u32::try_from(n).is_err() {
60 return Err("too many variables".into());
61 }
62 for (i, v) in p.vars.iter().enumerate() {
63 if !v.value.is_finite() {
64 return Err(format!("variable {i} ({}) is not finite", v.label));
65 }
66 }
67 for r in &p.rules {
68 for v in r.vars() {
69 if v.index() >= n {
70 return Err(format!("rule {} references unknown variable {}", r.id, v.0));
71 }
72 }
73 for row in &r.rows {
74 if row.expr.size() > 100_000 {
75 return Err(format!("rule {} expression too large", r.id));
76 }
77 }
78 }
79 for t in &p.targets {
80 if t.var.index() >= n || !t.value.is_finite() {
81 return Err("invalid target".into());
82 }
83 }
84 Ok(())
85}
86
87impl SolveJob {
88 #[must_use]
91 pub fn new(problem: Problem, opts: SolveOptions) -> Self {
92 let x = problem.values();
93 let invalid = validate(&problem).err();
94 let comps = if invalid.is_some() { Vec::new() } else { components(&problem) };
95 Self {
96 x_in: x.clone(),
97 x,
98 comps,
99 next: 0,
100 current: None,
101 reports: Vec::new(),
102 diagnostics: Vec::new(),
103 cancelled: false,
104 invalid,
105 iterations: 0,
106 problem,
107 opts,
108 }
109 }
110
111 pub fn cancel(&mut self) {
113 self.cancelled = true;
114 }
115
116 #[must_use]
118 pub fn is_finished(&self) -> bool {
119 self.cancelled || self.invalid.is_some() || self.next >= self.comps.len()
120 }
121
122 #[must_use]
125 pub fn peek_status(&self) -> Status {
126 if self.cancelled {
127 return Status::Cancelled;
128 }
129 if self.invalid.is_some() {
130 return Status::Unsupported;
131 }
132 self.reports.iter().fold(Status::Solved, |acc, r| acc.combine(r.status))
133 }
134
135 pub fn step(&mut self, budget: u32) -> Progress {
137 let mut left = budget.max(1);
138 while left > 0 && !self.is_finished() {
139 let Some(comp) = self.comps.get(self.next).cloned() else { break };
140 match comp.backend {
141 Backend::Trivial => {
142 if let Some(unsupported) = self.unsupported(&comp) {
144 self.push_unsupported(&comp, unsupported);
145 } else {
146 self.solve_trivial(&comp);
147 }
148 self.next += 1;
149 }
150 Backend::Linear => {
151 if let Some(unsupported) = self.unsupported(&comp) {
152 self.push_unsupported(&comp, unsupported);
153 } else {
154 let out = solve_linear(&self.problem, &comp, &mut self.x, &self.x_in, &self.opts);
155 self.iterations += 1;
156 self.reports.push(self.report(&comp, Backend::Linear, out.status, 1, out.max_hard));
157 self.diagnostics.extend(out.diagnostics);
158 }
159 self.next += 1;
160 left -= 1;
161 continue;
162 }
163 Backend::Numeric => {
164 if let Some(unsupported) = self.unsupported(&comp) {
165 self.push_unsupported(&comp, unsupported);
166 self.next += 1;
167 continue;
168 }
169 let solver = self
170 .current
171 .get_or_insert_with(|| NumericSolverBox(NumericSolver::new(&self.problem, &comp, &self.x_in)));
172 let finished = solver.0.iterate(&self.problem, &mut self.x, &self.opts);
173 self.iterations += 1;
174 left -= 1;
175 if finished && let Some(NumericSolverBox(s)) = self.current.take() {
176 self.finish_numeric(&comp, &s);
177 self.next += 1;
178 }
179 continue;
180 }
181 }
182 }
183 if self.is_finished() {
184 Progress::Finished
185 } else {
186 Progress::Running { completed: self.next, total: self.comps.len() }
187 }
188 }
189
190 fn unsupported(&self, comp: &Component) -> Option<(usize, String)> {
191 comp.rules
192 .iter()
193 .find_map(|ri| self.problem.rules.get(*ri).and_then(|r| r.unsupported.clone().map(|u| (*ri, u))))
194 }
195
196 fn push_unsupported(&mut self, comp: &Component, (ri, why): (usize, String)) {
197 let mut d = conflict_diagnostic(&self.problem, &[ri], Certainty::Certain, None, why);
198 d.kind = DiagnosticKind::Unsupported;
199 self.diagnostics.push(d);
200 self.reports.push(self.report(comp, comp.backend, Status::Unsupported, 0, f64::NAN));
201 }
202
203 fn report(
204 &self,
205 comp: &Component,
206 backend: Backend,
207 status: Status,
208 iterations: u32,
209 max_hard: f64,
210 ) -> ComponentReport {
211 ComponentReport {
212 vars: comp.vars.clone(),
213 rules: comp.rules.iter().filter_map(|ri| self.problem.rules.get(*ri).map(|r| r.id)).collect(),
214 backend,
215 status,
216 iterations,
217 max_hard_residual: max_hard,
218 }
219 }
220
221 fn solve_trivial(&mut self, comp: &Component) {
222 let (status, diagnostics, max_hard) = solve_trivial(&self.problem, comp, &mut self.x, &self.opts);
223 self.diagnostics.extend(diagnostics);
224 self.reports.push(self.report(comp, Backend::Trivial, status, 0, max_hard));
225 }
226
227 fn finish_numeric(&mut self, comp: &Component, s: &NumericSolver) {
228 let mut status = s.done.unwrap_or(Status::NotConverged { suspected_conflict: false });
229 if status == Status::Conflicting {
230 let evidence = s.inconsistent_rules(&self.problem, &self.x, &self.opts);
231 self.diagnostics.push(conflict_diagnostic(
232 &self.problem,
233 &evidence,
234 Certainty::Certain,
235 None,
236 "hard rules fix the same value to different numbers".into(),
237 ));
238 } else if let Status::NotConverged { suspected_conflict } = status {
239 let evidence = s.inconsistent_rules(&self.problem, &self.x, &self.opts);
240 let fixed = |v| self.problem.is_fixed(v);
241 let all_linear = !evidence.is_empty()
242 && evidence.iter().all(|ri| {
243 self.problem
244 .rules
245 .get(*ri)
246 .is_some_and(|r| r.rows.iter().all(|row| row.expr.linear_form(&self.x, &fixed).is_some()))
247 });
248 let residual = Some(s.max_hard);
249 if all_linear {
250 status = Status::Conflicting;
251 self.diagnostics.push(conflict_diagnostic(
252 &self.problem,
253 &evidence,
254 Certainty::Certain,
255 residual,
256 "linear rules in a mixed component contradict each other".into(),
257 ));
258 } else if suspected_conflict || !evidence.is_empty() {
259 status = Status::NotConverged { suspected_conflict: true };
260 let rules = if evidence.is_empty() { comp.rules.clone() } else { evidence };
261 self.diagnostics.push(conflict_diagnostic(
262 &self.problem,
263 &rules,
264 Certainty::Suspected,
265 residual,
266 "hard rules could not be satisfied; the residual stopped decreasing".into(),
267 ));
268 } else {
269 let mut d = conflict_diagnostic(
270 &self.problem,
271 &comp.rules,
272 Certainty::Suspected,
273 residual,
274 format!("iteration budget of {} exhausted", self.opts.max_iterations),
275 );
276 d.kind = DiagnosticKind::NotConverged;
277 self.diagnostics.push(d);
278 }
279 } else {
280 let redundant = s.redundant_rules(&self.problem, &self.x, &self.opts);
281 if !redundant.is_empty() {
282 let mut d = conflict_diagnostic(
283 &self.problem,
284 &redundant,
285 Certainty::Certain,
286 None,
287 "redundant rules (consistent)".into(),
288 );
289 d.kind = DiagnosticKind::Redundant;
290 self.diagnostics.push(d);
291 }
292 }
293 self.reports.push(self.report(comp, Backend::Numeric, status, s.iterations, s.max_hard));
294 }
295
296 #[must_use]
298 pub fn into_solution(mut self) -> Solution {
299 while !self.is_finished() {
300 self.step(u32::MAX);
301 }
302 if let Some(why) = self.invalid.take() {
303 return Solution {
304 values: self.x_in,
305 status: Status::Unsupported,
306 components: Vec::new(),
307 diagnostics: vec![Diagnostic {
308 kind: DiagnosticKind::Unsupported,
309 certainty: Certainty::Certain,
310 rules: Vec::new(),
311 labels: Vec::new(),
312 sources: Vec::new(),
313 entities: Vec::new(),
314 residual: None,
315 message: why,
316 }],
317 iterations: 0,
318 };
319 }
320 if self.cancelled {
321 let mut diagnostics = self.diagnostics;
322 diagnostics.push(Diagnostic {
323 kind: DiagnosticKind::Cancelled,
324 certainty: Certainty::Certain,
325 rules: Vec::new(),
326 labels: Vec::new(),
327 sources: Vec::new(),
328 entities: Vec::new(),
329 residual: None,
330 message: "solve cancelled".into(),
331 });
332 return Solution {
333 values: self.x_in,
334 status: Status::Cancelled,
335 components: self.reports,
336 diagnostics,
337 iterations: self.iterations,
338 };
339 }
340 let status = self.reports.iter().fold(Status::Solved, |acc, r| acc.combine(r.status));
341 let values = if status.is_acceptable() { self.x } else { self.x_in };
342 Solution {
343 values,
344 status,
345 components: self.reports,
346 diagnostics: self.diagnostics,
347 iterations: self.iterations,
348 }
349 }
350}
351
352#[must_use]
354pub fn solve(problem: &Problem, opts: &SolveOptions) -> Solution {
355 SolveJob::new(problem.clone(), *opts).into_solution()
356}
357
358pub(crate) fn solve_trivial(
361 p: &Problem,
362 comp: &Component,
363 x: &mut [f64],
364 opts: &SolveOptions,
365) -> (Status, Vec<Diagnostic>, f64) {
366 if comp.rules.is_empty() {
367 for v in &comp.vars {
368 let best = p
369 .targets
370 .iter()
371 .filter(|t| t.var == *v)
372 .max_by_key(|t| if t.strength == Strength::Required { Strength::Strong } else { t.strength });
373 if let (Some(t), Some(slot)) = (best, x.get_mut(v.index())) {
374 *slot = t.value;
375 }
376 }
377 let dof = comp.vars.len();
378 let status = if dof == 0 { Status::Solved } else { Status::Underconstrained { dof } };
379 return (status, Vec::new(), 0.0);
380 }
381 let cols = Columns::new(p, &[]);
382 let rows = eval_rows(p, &comp.rules, &cols, x);
383 let hard: Vec<_> = rows.iter().filter(|r| p.rules.get(r.rule).is_some_and(crate::Rule::is_hard)).cloned().collect();
384 let max_hard = max_violation(&hard);
385 if max_hard > opts.tolerance {
386 let msg = "rule depends only on fixed values and is violated".to_owned();
387 let residual = hard.iter().map(|r| r.value.abs()).fold(0.0, f64::max);
388 let d = conflict_diagnostic(p, &comp.rules, Certainty::Certain, Some(residual), msg);
389 return (Status::Conflicting, vec![d], max_hard);
390 }
391 let soft_unmet =
392 rows.iter().any(|r| p.rules.get(r.rule).is_some_and(|x| !x.is_hard()) && r.violation() > opts.tolerance);
393 let mut diagnostics = Vec::new();
394 if soft_unmet {
395 let mut d = conflict_diagnostic(
396 p,
397 &comp.rules,
398 Certainty::Certain,
399 None,
400 "preference cannot be met: all its variables are fixed".into(),
401 );
402 d.kind = DiagnosticKind::PreferenceUnmet;
403 diagnostics.push(d);
404 }
405 (Status::Solved, diagnostics, max_hard)
406}