use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
use std::fmt;
use crate::solver::{LinExpr, SVar};
use crate::Number;
use super::project::{project, Projected};
use super::store::{Addr, Cell, PendingDif, Store};
use super::symbol::Symbols;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Answer {
pub equations: Vec<String>,
pub disequations: Vec<String>,
pub constraints: Vec<String>,
}
impl Answer {
pub fn is_true(&self) -> bool {
self.equations.is_empty() && self.disequations.is_empty() && self.constraints.is_empty()
}
pub fn parts(&self) -> impl Iterator<Item = &str> {
self.equations
.iter()
.chain(self.disequations.iter())
.chain(self.constraints.iter())
.map(String::as_str)
}
}
impl fmt::Display for Answer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.is_true() {
return f.write_str("true");
}
let mut first = true;
for part in self
.equations
.iter()
.chain(self.disequations.iter())
.chain(self.constraints.iter())
{
if !first {
f.write_str(", ")?;
}
first = false;
f.write_str(part)?;
}
Ok(())
}
}
struct Printer<'a> {
symbols: &'a Symbols,
store: &'a Store,
names: HashMap<Addr, String>,
cyclic: HashSet<Addr>,
pending_nodes: Vec<Addr>,
next_internal: usize,
}
impl<'a> Printer<'a> {
fn new(symbols: &'a Symbols, store: &'a Store) -> Self {
Printer {
symbols,
store,
names: HashMap::new(),
cyclic: HashSet::new(),
pending_nodes: Vec::new(),
next_internal: 1,
}
}
fn fresh_name(&mut self) -> String {
let n = format!("_{}", self.next_internal);
self.next_internal += 1;
n
}
fn name_for(&mut self, a: Addr) -> String {
if let Some(n) = self.names.get(&a) {
return n.clone();
}
let n = self.fresh_name();
self.names.insert(a, n.clone());
n
}
fn find_cycles(&mut self, root: Addr) {
enum Step {
Enter(Addr),
Leave(Addr),
}
let mut on_path: HashSet<Addr> = HashSet::new();
let mut done: HashSet<Addr> = HashSet::new();
let mut stack = vec![Step::Enter(root)];
while let Some(step) = stack.pop() {
match step {
Step::Enter(a) => {
let a = self.store.deref(a);
let Cell::Struct(_, args) = self.store.cell(a) else { continue };
if on_path.contains(&a) {
self.cyclic.insert(a);
continue;
}
if done.contains(&a) {
continue;
}
on_path.insert(a);
stack.push(Step::Leave(a));
for &arg in args.iter().rev() {
stack.push(Step::Enter(arg));
}
}
Step::Leave(a) => {
on_path.remove(&a);
done.insert(a);
}
}
}
}
fn render(&mut self, root: Addr) -> String {
enum Item {
Term(Addr),
Text(&'static str),
}
let root = self.store.deref(root);
let mut out = String::new();
let mut stack = vec![Item::Term(root)];
let mut first = true;
while let Some(item) = stack.pop() {
match item {
Item::Text(t) => out.push_str(t),
Item::Term(a) => {
let a = self.store.deref(a);
let is_root = first && a == root;
first = false;
match self.store.cell(a).clone() {
Cell::Var(_) => {
let fixed = self.store.numvar.get(&a).and_then(|sv| self.store.class_value(*sv));
match fixed {
Some(c) => out.push_str(&c.to_string()),
None => {
let n = self.name_for(a);
out.push_str(&n);
}
}
}
Cell::Const(c) => out.push_str(self.symbols.name(c)),
Cell::Num(n) => out.push_str(&n.to_string()),
Cell::Struct(f, args) => {
if self.cyclic.contains(&a) && !is_root {
if !self.names.contains_key(&a) {
let n = self.fresh_name();
self.names.insert(a, n);
self.pending_nodes.push(a);
}
out.push_str(&self.names[&a]);
continue;
}
out.push_str(self.symbols.name(f));
out.push('(');
stack.push(Item::Text(")"));
for (i, &arg) in args.iter().enumerate().rev() {
if i + 1 < args.len() {
stack.push(Item::Text(", "));
}
stack.push(Item::Term(arg));
}
}
}
}
}
}
out
}
}
pub(crate) fn render_answer(
symbols: &Symbols,
store: &Store,
query_vars: &[(String, Addr)],
) -> Answer {
let mut p = Printer::new(symbols, store);
let mut equations = Vec::new();
let reachable = reachable_vars(store, query_vars.iter().map(|(_, a)| *a));
let difs: Vec<PendingDif> = store
.pending_difs()
.into_iter()
.filter(|(a, b, _)| reachable_vars(store, [*a, *b]).iter().any(|v| reachable.contains(v)))
.collect();
for (_, addr) in query_vars {
p.find_cycles(*addr);
}
for (a, b, _) in &difs {
p.find_cycles(*a);
p.find_cycles(*b);
}
let mut aliases: Vec<Option<String>> = Vec::with_capacity(query_vars.len());
for (name, addr) in query_vars {
let d = store.deref(*addr);
match store.cell(d) {
Cell::Var(_) => {
if let Some(existing) = p.names.get(&d) {
aliases.push(Some(format!("{existing} = {name}")));
} else {
p.names.insert(d, name.clone());
aliases.push(None);
}
}
_ => {
if p.cyclic.contains(&d) && !p.names.contains_key(&d) {
p.names.insert(d, name.clone());
}
aliases.push(None);
}
}
}
for ((name, addr), alias) in query_vars.iter().zip(aliases) {
if let Some(alias) = alias {
equations.push(alias);
continue;
}
let d = store.deref(*addr);
if !matches!(store.cell(d), Cell::Var(_)) {
let rendered = p.render(d);
equations.push(format!("{name} = {rendered}"));
} else if let Some(c) = store.numvar.get(&d).and_then(|sv| store.class_value(*sv)) {
equations.push(format!("{name} = {c}"));
}
}
let mut disequations = Vec::new();
for (a, b, reduced) in difs {
let (l, r) = if reduced.len() == 1 {
let (v, t) = reduced[0];
let t_is_var = matches!(store.cell(store.deref(t)), Cell::Var(_));
let (v, t) = if t_is_var && store.deref(t) < store.deref(v) { (t, v) } else { (v, t) };
(p.render(v), p.render(t))
} else {
(p.render(a), p.render(b))
};
disequations.push(format!("{l} != {r}"));
}
while let Some(node) = p.pending_nodes.pop() {
let name = p.names[&node].clone();
let rendered = p.render(node);
equations.push(format!("{name} = {rendered}"));
}
let constraints = render_numeric(&mut p, store, &reachable);
Answer { equations, disequations, constraints }
}
fn public_names(p: &mut Printer<'_>, store: &Store, reachable: &HashSet<Addr>) -> BTreeMap<SVar, String> {
let mut names: BTreeMap<SVar, String> = BTreeMap::new();
let mut heap_named: Vec<(SVar, Addr)> = reachable
.iter()
.filter_map(|v| store.numvar.get(v).map(|sv| (store.root(*sv), *v)))
.collect();
heap_named.sort();
for (r, v) in heap_named {
let n = p.name_for(v);
names.entry(r).or_insert(n);
}
let mut attr_named: Vec<(SVar, Addr)> = store
.attrs
.iter()
.enumerate()
.filter(|(i, _)| store.attr_visible(*i))
.filter(|(_, e)| reachable_vars(store, [e.term]).iter().all(|v| reachable.contains(v)))
.map(|(_, e)| (store.root(e.svar), e.term))
.collect();
attr_named.sort();
attr_named.dedup_by_key(|(r, _)| *r);
for (r, term) in attr_named {
if let std::collections::btree_map::Entry::Vacant(slot) = names.entry(r) {
let rendered = p.render(term);
slot.insert(rendered);
}
}
names
}
fn render_numeric(p: &mut Printer<'_>, store: &Store, reachable: &HashSet<Addr>) -> Vec<String> {
let mut names = public_names(p, store, reachable);
if names.is_empty() {
return Vec::new();
}
let public: BTreeSet<SVar> = names.keys().copied().collect();
let Projected { eqs, ineqs, difs, survivors } = project(store, &public);
let mut attr_render: BTreeMap<SVar, Addr> = BTreeMap::new();
for e in &store.attrs {
attr_render.entry(store.root(e.svar)).or_insert(e.term);
}
for v in &survivors {
if !names.contains_key(v) {
let n = match attr_render.get(v) {
Some(term) => p.render(*term),
None => p.fresh_name(),
};
names.insert(*v, n);
}
}
let render_expr = |e: &LinExpr| -> Option<String> {
let mut out = String::new();
let mut first = true;
let lead_const = e.constant.is_positive()
&& e.terms.iter().next().is_some_and(|(_, a)| a.is_negative());
if lead_const {
out.push_str(&e.constant.to_string());
first = false;
}
for (v, a) in &e.terms {
let name = names.get(v)?;
if first {
if a.is_negative() {
out.push('-');
}
} else if a.is_negative() {
out.push_str(" - ");
} else {
out.push_str(" + ");
}
first = false;
let m = a.abs();
if m != Number::one() {
out.push_str(&format!("{m}*"));
}
out.push_str(name);
}
if first {
out.push_str(&e.constant.to_string());
} else if lead_const {
} else if e.constant.is_positive() {
out.push_str(&format!(" + {}", e.constant));
} else if e.constant.is_negative() {
out.push_str(&format!(" - {}", e.constant.abs()));
}
Some(out)
};
let render_rel = |e: &LinExpr, op: &str| -> Option<String> {
let c = -&e.constant;
let mut terms = e.clone();
terms.constant = Number::zero();
if terms.terms.len() == 2 && c.is_zero() {
let mut it = terms.terms.iter();
let (v1, a1) = it.next().unwrap();
let (v2, a2) = it.next().unwrap();
if *a1 == Number::one() && *a2 == -Number::one() {
return Some(format!("{} {op} {}", names.get(v1)?, names.get(v2)?));
}
if *a2 == Number::one() && *a1 == -Number::one() {
return Some(format!("{} {op} {}", names.get(v2)?, names.get(v1)?));
}
}
Some(format!("{} {op} {c}", render_expr(&terms)?))
};
let first_var = |e: &LinExpr| e.terms.keys().next().copied().unwrap_or(SVar(usize::MAX));
let mut out = Vec::new();
for (r, name) in &names {
if attr_render.contains_key(r) {
if let Some(c) = store.class_value(*r) {
if !reachable.iter().any(|v| store.numvar.get(v).map(|sv| store.root(*sv)) == Some(*r)) {
out.push(format!("{name} = {c}"));
}
}
}
}
for (subject, rhs) in &eqs {
if let (Some(name), Some(r)) = (names.get(subject), render_expr(rhs)) {
out.push(format!("{name} = {r}"));
}
}
let mut lines: Vec<(SVar, bool, String)> = ineqs
.iter()
.filter_map(|(e, strict)| {
let leading_negative = e.terms.iter().next().is_some_and(|(_, a)| a.is_negative());
let (expr, op) = if leading_negative {
let mut n = e.clone();
n.negate();
(n, if *strict { "<" } else { "<=" })
} else {
(e.clone(), if *strict { ">" } else { ">=" })
};
render_rel(&expr, op).map(|s| (first_var(&expr), leading_negative, s))
})
.collect();
lines.sort();
out.extend(lines.into_iter().map(|(_, _, s)| s));
let mut dlines: Vec<(SVar, String)> = difs
.iter()
.filter_map(|e| render_rel(e, "!=").map(|s| (first_var(e), s)))
.collect();
dlines.sort();
out.extend(dlines.into_iter().map(|(_, s)| s));
out
}
fn reachable_vars(store: &Store, roots: impl IntoIterator<Item = Addr>) -> HashSet<Addr> {
let mut vars = HashSet::new();
let mut seen = HashSet::new();
let mut stack: Vec<Addr> = roots.into_iter().collect();
while let Some(a) = stack.pop() {
let a = store.deref(a);
if !seen.insert(a) {
continue;
}
match store.cell(a) {
Cell::Var(_) => {
vars.insert(a);
}
Cell::Struct(_, args) => stack.extend(args.iter().copied()),
_ => {}
}
}
vars
}
pub(crate) fn render_terms(
symbols: &Symbols,
store: &Store,
query_vars: &[(String, Addr)],
addrs: &[Addr],
) -> Vec<String> {
let mut p = Printer::new(symbols, store);
for (name, addr) in query_vars {
let d = store.deref(*addr);
if matches!(store.cell(d), Cell::Var(_)) {
p.names.entry(d).or_insert_with(|| name.clone());
}
}
for &a in addrs {
p.find_cycles(a);
}
addrs.iter().map(|a| p.render(*a)).collect()
}