use crate::decompose::{Embedding, EmbeddingOptions, MAX_EMBEDDING_DIM, embed};
use crate::error::VitriError;
use crate::tests::common::{make_formula, wide_component};
use crate::vtree::VarId;
fn at_dim(dim: usize) -> EmbeddingOptions {
EmbeddingOptions { dim }
}
#[test]
fn the_same_formula_and_options_give_the_same_points() {
let formula = wide_component();
let once = embed(&formula, &at_dim(2)).expect("the fixture embeds");
let twice = embed(&formula, &at_dim(2)).expect("the fixture embeds");
assert_eq!(once, twice);
}
#[test]
fn every_variable_has_a_row_of_its_own() {
let formula = wide_component();
for dim in 2..=MAX_EMBEDDING_DIM {
let e = embed(&formula, &at_dim(dim)).expect("the fixture embeds");
assert_eq!(e.dim, dim);
assert_eq!(e.num_vars(), formula.num_vars);
assert_eq!(e.coords.len(), formula.num_vars as usize * dim);
let v = VarId(3);
assert_eq!(e.position(v), &e.coords[v.idx() * dim..v.idx() * dim + dim]);
}
}
#[test]
fn a_dimension_outside_the_accepted_range_is_refused() {
let formula = wide_component();
for dim in [0, 1, MAX_EMBEDDING_DIM + 1] {
let err = embed(&formula, &at_dim(dim)).expect_err("out of range must be refused");
assert!(
matches!(err, VitriError::Input { .. }),
"an unusable dimension is bad input: {err:?}",
);
assert!(
err.to_string().contains(&MAX_EMBEDDING_DIM.to_string()),
"the refusal states the range it would have accepted: {err}",
);
}
}
#[test]
fn a_formula_with_no_variables_has_nothing_to_embed() {
let err = embed(&make_formula(0, Vec::new()), &at_dim(2))
.expect_err("an empty formula must be refused");
assert!(matches!(err, VitriError::Input { .. }), "{err:?}");
}
#[test]
fn variables_that_share_a_clause_are_placed_nearer_than_variables_that_do_not() {
let formula = wide_component();
let e = embed(&formula, &at_dim(2)).expect("the fixture embeds");
let mut together = Vec::new();
for clause in &formula.clauses {
for (i, a) in clause.literals.iter().enumerate() {
for b in &clause.literals[i + 1..] {
together.push(distance(&e, a.var.0, b.var.0));
}
}
}
let mut any = Vec::new();
for a in 0..formula.num_vars {
for b in (a + 1)..formula.num_vars {
any.push(distance(&e, a, b));
}
}
let mean = |d: &[f64]| d.iter().sum::<f64>() / d.len() as f64;
assert!(
mean(&together) < mean(&any),
"clause-sharing pairs average {} apart and arbitrary pairs {} — the \
embedding is not reading the formula",
mean(&together),
mean(&any),
);
}
fn distance(e: &Embedding, a: u32, b: u32) -> f64 {
e.position(VarId(a))
.iter()
.zip(e.position(VarId(b)))
.map(|(x, y)| (x - y) * (x - y))
.sum::<f64>()
.sqrt()
}