From 038cc36b781ac243df199162baefcff92d41833c Mon Sep 17 00:00:00 2001 From: lashman Date: Sat, 24 May 2025 22:58:00 +0300 Subject: [PATCH] fill constraint propagation --- cruciverb-core/src/solver/csp.rs | 530 +++++++++++++++++++++++++++++++ 1 file changed, 530 insertions(+) create mode 100644 cruciverb-core/src/solver/csp.rs diff --git a/cruciverb-core/src/solver/csp.rs b/cruciverb-core/src/solver/csp.rs new file mode 100644 index 0000000..c8cb374 --- /dev/null +++ b/cruciverb-core/src/solver/csp.rs @@ -0,0 +1,530 @@ +use std::collections::VecDeque; +use std::sync::atomic::Ordering; +use std::time::Instant; + +use rand::rngs::StdRng; +use rand::Rng; +use rand::SeedableRng; + +use crate::config::GenerationConfig; +use crate::dict::Dictionary; +use crate::error::{CrucivError, Result}; +use crate::grid::slot::Slot; +use crate::grid::Grid; + +use super::state::SolverState; + +pub struct SolveResult { + pub grid: Grid, + pub word_assignments: Vec<(usize, usize)>, +} + +pub fn solve( + grid: &Grid, + slots: &[Slot], + dict: &Dictionary, + config: &GenerationConfig, +) -> Result { + let start = Instant::now(); + let timeout = std::time::Duration::from_secs(config.timeout_seconds); + + let mut rng = match config.seed { + Some(s) => StdRng::seed_from_u64(s), + None => StdRng::from_os_rng(), + }; + + let total_restarts = config.max_restarts.max(slots.len()); + let bt_budget = (config.max_backtrack as u64).max(slots.len() as u64 * 2000); + + for restart in 0..total_restarts { + if start.elapsed() > timeout { + return Err(CrucivError::Timeout(config.timeout_seconds)); + } + if let Some(ref flag) = config.cancelled { + if flag.load(Ordering::Relaxed) { + return Err(CrucivError::Generation("cancelled".into())); + } + } + + let mut state = SolverState::new(slots, grid.width, grid.height, dict.word_count()); + state.init_domains(slots, &dict.index); + + // run initial domain propagation to detect infeasible cells early + if !propagate_all_domains(&mut state, slots, dict) { + tracing::debug!( + "restart {}: initial domain propagation found infeasible cell", + restart + 1 + ); + continue; + } + + let mut bt = 0u64; + let rand_top = config.rand_top_base + restart.min(5); + + let result = backtrack( + &mut state, slots, dict, config, &start, timeout, &mut rng, &mut bt, rand_top, + bt_budget, + ); + + if let Some(assignments) = result { + let mut filled = grid.clone(); + for &(si, wi) in &assignments { + let word = &dict.get(wi).text; + for (pos, (r, c)) in slots[si].cells().iter().enumerate() { + filled.set( + *r, + *c, + crate::grid::Cell::Letter(word.as_bytes()[pos] as char), + ); + } + } + return Ok(SolveResult { + grid: filled, + word_assignments: assignments, + }); + } + + tracing::debug!("restart {} failed ({} backtracks)", restart + 1, bt); + } + + Err(CrucivError::Generation(format!( + "failed to fill grid after {} restarts", + total_restarts + ))) +} + +fn backtrack( + state: &mut SolverState, + slots: &[Slot], + dict: &Dictionary, + config: &GenerationConfig, + start: &Instant, + timeout: std::time::Duration, + rng: &mut StdRng, + bt: &mut u64, + rand_top: usize, + bt_budget: u64, +) -> Option> { + if *bt > bt_budget { + return None; + } + if *bt % 512 == 0 { + if start.elapsed() > timeout { + return None; + } + if let Some(ref flag) = config.cancelled { + if flag.load(Ordering::Relaxed) { + return None; + } + } + } + + let slot_idx = match pick_slot(state, slots, dict) { + Some(i) => i, + None => { + return Some( + state + .assignments + .iter() + .enumerate() + .filter_map(|(i, a)| a.map(|w| (i, w))) + .collect(), + ); + } + }; + + let slot = &slots[slot_idx]; + let constraints = state.current_constraints(slot); + + let all_candidates = + dict.candidates_filtered(slot.length, &constraints, config.frequency_threshold); + + // score and filter candidates + let mut scored: Vec<(usize, f64)> = all_candidates + .iter() + .filter(|&&w| !state.is_used(w)) + .map(|&w| { + let word = dict.get(w); + let mut s = word.frequency; + let bytes = word.text.as_bytes(); + let mut seen = [false; 26]; + let mut unique = 0; + for &b in bytes { + let idx = (b - b'A') as usize; + if !seen[idx] { + seen[idx] = true; + unique += 1; + } + } + s += (unique as f64 / bytes.len() as f64) * 0.1; + // penalize recently used words so the solver picks fresh ones + if let Some(ref staleness) = config.staleness { + if let Some(&penalty) = staleness.get(&word.text) { + s *= penalty; + } + } + (w, s) + }) + .collect(); + + if scored.is_empty() { + return None; + } + + scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + + // randomize top candidates + let top_n = scored.len().min(rand_top); + if top_n > 1 { + for i in 0..top_n - 1 { + let j = rng.random_range(i..top_n); + scored.swap(i, j); + } + } + + scored.truncate(24); + + // filter out words related to already-placed words + let placed_texts: Vec<&str> = state + .assignments + .iter() + .filter_map(|a| a.map(|wi| dict.get(wi).text.as_str())) + .collect(); + if !placed_texts.is_empty() { + scored.retain(|(w, _)| { + let text = &dict.get(*w).text; + !placed_texts + .iter() + .any(|p| super::candidate::is_related(text, p)) + }); + } + + // LCV ordering + if scored.len() > 1 && scored.len() <= 20 && slots[slot_idx].intersections.len() >= 2 { + lcv_sort(&mut scored, slot_idx, state, slots, dict); + } + + let max_try = scored.len().min(16); + + for i in 0..max_try { + let (word_idx, _) = scored[i]; + let word_bytes: Vec = dict.get(word_idx).text.bytes().collect(); + + // save state for undo + let cells = slot.cells(); + let mut prev_letters: Vec<(usize, usize, Option)> = Vec::with_capacity(cells.len()); + let mut prev_domains: Vec<(usize, usize, u32)> = Vec::new(); + + for (pos, &(r, c)) in cells.iter().enumerate() { + prev_letters.push((r, c, state.get_letter(r, c))); + state.set_letter(r, c, word_bytes[pos]); + } + state.assignments[slot_idx] = Some(word_idx); + state.mark_used(word_idx); + + // AC-3 propagation with domain narrowing + let ac3_ok = ac3_propagate(state, slots, dict, slot_idx, &mut prev_domains); + + if ac3_ok { + if let Some(result) = backtrack( + state, slots, dict, config, start, timeout, rng, bt, rand_top, bt_budget, + ) { + return Some(result); + } + } + + *bt += 1; + + // undo + state.assignments[slot_idx] = None; + state.unmark_used(word_idx); + for &(r, c, p) in &prev_letters { + match p { + Some(l) => state.set_letter(r, c, l), + None => state.clear_letter(r, c), + } + } + for &(r, c, d) in &prev_domains { + state.set_domain(r, c, d); + } + } + + None +} + +/// Compute the set of letters possible at a given position in a slot, +/// given the current constraints. Returns a 26-bit mask. +fn compute_slot_domain_at(state: &SolverState, slot: &Slot, pos: usize, dict: &Dictionary) -> u32 { + let constraints = state.current_constraints(slot); + let len = slot.length; + let chunks = match dict.index.chunk_count(len) { + Some(c) => c, + None => return 0, + }; + let n = dict.index.count_of_length(len); + if n == 0 { + return 0; + } + + // build intersection bitset for all existing constraints + let mut result = vec![u64::MAX; chunks]; + let last_bits = n % 64; + if last_bits > 0 { + result[chunks - 1] = (1u64 << last_bits) - 1; + } + for &(cpos, letter) in &constraints { + if cpos == pos { + continue; + } // skip the position we're computing domain for + if let Some(bits) = dict.index.get_letter_bits(len, cpos, letter) { + for (i, chunk) in result.iter_mut().enumerate() { + *chunk &= bits[i]; + } + } else { + return 0; + } + } + + // now union all letters that appear at `pos` across remaining candidates + let mut domain = 0u32; + for letter in 0u8..26 { + if let Some(bits) = dict.index.get_letter_bits(len, pos, b'A' + letter) { + for (i, &chunk) in result.iter().enumerate() { + if chunk & bits[i] != 0 { + domain |= 1u32 << letter; + break; + } + } + } + } + domain +} + +/// Full domain propagation for all unassigned cells. +/// Returns false if any cell has an empty domain. +fn propagate_all_domains(state: &mut SolverState, slots: &[Slot], dict: &Dictionary) -> bool { + let mut changed = true; + while changed { + changed = false; + for slot in slots { + if state.assignments[slot.index].is_some() { + continue; + } + for (pos, &(r, c)) in slot.cells().iter().enumerate() { + if state.get_letter(r, c).is_some() { + continue; + } + let new_domain = compute_slot_domain_at(state, slot, pos, dict); + let old_domain = state.get_domain(r, c); + let narrowed = old_domain & new_domain; + if narrowed != old_domain { + if narrowed == 0 { + return false; + } + state.set_domain(r, c, narrowed); + changed = true; + } + } + } + } + true +} + +/// AC-3 propagation after placing a word. +/// Checks crossing slots and narrows cell domains. +/// Saves changed domains for undo. +fn ac3_propagate( + state: &mut SolverState, + slots: &[Slot], + dict: &Dictionary, + placed: usize, + prev_domains: &mut Vec<(usize, usize, u32)>, +) -> bool { + let mut queue: VecDeque = VecDeque::new(); + let mut checked = vec![false; slots.len()]; + checked[placed] = true; + + for ix in &slots[placed].intersections { + if state.assignments[ix.other_slot].is_none() { + queue.push_back(ix.other_slot); + } + } + + while let Some(si) = queue.pop_front() { + if checked[si] || state.assignments[si].is_some() { + continue; + } + checked[si] = true; + + let slot = &slots[si]; + let constraints = state.current_constraints(slot); + let count = dict.index.candidate_count(slot.length, &constraints); + + if count == 0 { + return false; + } + + // narrow cell domains for unconstrained positions in this slot + let mut domain_changed = false; + for (pos, &(r, c)) in slot.cells().iter().enumerate() { + if state.get_letter(r, c).is_some() { + continue; + } + let new_domain = compute_slot_domain_at(state, slot, pos, dict); + let old_domain = state.get_domain(r, c); + let narrowed = old_domain & new_domain; + if narrowed == 0 { + return false; + } + if narrowed != old_domain { + prev_domains.push((r, c, old_domain)); + state.set_domain(r, c, narrowed); + domain_changed = true; + } + } + + // if domains changed, propagate to crossings + if domain_changed || count <= 20 { + for ix in &slot.intersections { + if !checked[ix.other_slot] && state.assignments[ix.other_slot].is_none() { + queue.push_back(ix.other_slot); + } + } + } + } + + true +} + +// Full MRV: check candidate_count for ALL unassigned slots. +// Tiebreak by degree (number of unassigned crossings). +fn pick_slot(state: &SolverState, slots: &[Slot], dict: &Dictionary) -> Option { + let mut best_idx = None; + let mut best_count = usize::MAX; + let mut best_degree = 0usize; + + for (i, a) in state.assignments.iter().enumerate() { + if a.is_some() { + continue; + } + let slot = &slots[i]; + let constraints = state.current_constraints(slot); + let count = dict.index.candidate_count(slot.length, &constraints); + + if count == 0 { + return Some(i); // force failure fast + } + + let degree = slot + .intersections + .iter() + .filter(|ix| state.assignments[ix.other_slot].is_none()) + .count(); + + if count < best_count || (count == best_count && degree > best_degree) { + best_count = count; + best_degree = degree; + best_idx = Some(i); + } + } + + best_idx +} + +// LCV: sort candidates by how much freedom they leave in crossing slots. +fn lcv_sort( + scored: &mut Vec<(usize, f64)>, + slot_idx: usize, + state: &SolverState, + slots: &[Slot], + dict: &Dictionary, +) { + let slot = &slots[slot_idx]; + let crossings: Vec<_> = slot + .intersections + .iter() + .filter(|ix| state.assignments[ix.other_slot].is_none()) + .collect(); + + if crossings.is_empty() { + return; + } + + let mut freedom: Vec<(usize, f64, i64)> = scored + .iter() + .map(|&(wi, score)| { + let word = dict.get(wi); + let bytes = word.text.as_bytes(); + let mut total_freedom = 0i64; + + for ix in &crossings { + let letter = bytes[ix.self_pos]; + let other = &slots[ix.other_slot]; + let mut other_constraints = state.current_constraints(other); + other_constraints.push((ix.other_pos, letter)); + let count = dict.index.candidate_count(other.length, &other_constraints); + total_freedom += count as i64; + } + + (wi, score, total_freedom) + }) + .collect(); + + freedom.sort_by(|a, b| b.2.cmp(&a.2).then(b.1.partial_cmp(&a.1).unwrap())); + + *scored = freedom.into_iter().map(|(wi, s, _)| (wi, s)).collect(); +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::grid::slot::extract_slots; + use crate::grid::Cell; + + #[test] + fn solve_5x5_real_dict() { + use std::path::Path; + let path = Path::new("../data/wordlists/enable.txt"); + if !path.exists() { + return; + } + let dict = Dictionary::load(path, None, 3).unwrap(); + let grid = Grid::new(5, 5); + let slots = extract_slots(&grid, 3); + let config = GenerationConfig { + width: 5, + height: 5, + timeout_seconds: 10, + frequency_threshold: 0.0, + max_restarts: 5, + seed: Some(42), + ..Default::default() + }; + let result = solve(&grid, &slots, &dict, &config); + assert!(result.is_ok()); + } + + #[test] + fn solve_with_real_dict() { + use std::path::Path; + let path = Path::new("../data/wordlists/enable.txt"); + if !path.exists() { + return; + } + let dict = Dictionary::load(path, None, 3).unwrap(); + let mut grid = Grid::new(5, 5); + grid.set(2, 2, Cell::Black); + let slots = extract_slots(&grid, 3); + let config = GenerationConfig { + width: 5, + height: 5, + timeout_seconds: 10, + frequency_threshold: 0.0, + max_restarts: 5, + seed: Some(42), + ..Default::default() + }; + let result = solve(&grid, &slots, &dict, &config); + assert!(result.is_ok()); + } +}