use rand::{Rng, RngExt};
use std::fmt;
use std::hash::{Hash, Hasher};
use crate::state::State;
pub const fn neighborhood_size<const R: usize>() -> usize {
2 * R + 1
}
#[derive(Clone)]
pub struct CAState<const N: usize, const R: usize = 1> {
cells: [u8; N],
}
impl<const N: usize, const R: usize> CAState<N, R> {
pub fn new(cells: [u8; N]) -> Self {
for &c in &cells {
assert!(c <= 1, "Cell values must be 0 or 1, got {}", c);
}
Self { cells }
}
pub fn random(rng: &mut impl Rng) -> Self {
let mut cells = [0u8; N];
for c in cells.iter_mut() {
*c = rng.random_range(0..=1);
}
Self { cells }
}
pub fn cell(&self, index: i32) -> u8 {
let i = index.rem_euclid(N as i32) as usize;
self.cells[i]
}
pub fn cells(&self) -> &[u8; N] {
&self.cells
}
pub fn neighborhood(&self, i: usize) -> usize {
let mut pattern = 0usize;
for r in 0..neighborhood_size::<R>() {
let offset = (r as i32) - (R as i32);
let bit = self.cell(i as i32 + offset) as usize;
pattern |= bit << (neighborhood_size::<R>() - 1 - r);
}
pattern
}
pub fn with_cells(&self, cells: [u8; N]) -> Self {
Self { cells }
}
}
impl<const N: usize, const R: usize> State for CAState<N, R> {
type Encoding = Vec<u8>;
fn canonical_encoding(&self) -> Self::Encoding {
self.cells.to_vec()
}
fn distance(&self, other: &Self) -> u32 {
self.cells
.iter()
.zip(other.cells.iter())
.map(|(a, b)| if a != b { 1 } else { 0 })
.sum()
}
}
impl<const N: usize, const R: usize> PartialEq for CAState<N, R> {
fn eq(&self, other: &Self) -> bool {
self.cells == other.cells
}
}
impl<const N: usize, const R: usize> Eq for CAState<N, R> {}
impl<const N: usize, const R: usize> Hash for CAState<N, R> {
fn hash<H: Hasher>(&self, state: &mut H) {
self.cells.hash(state);
}
}
impl<const N: usize, const R: usize> fmt::Debug for CAState<N, R> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "CAState<{}, {}>({:?})", N, R, self.cells)
}
}
impl<const N: usize, const R: usize> fmt::Display for CAState<N, R> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for &c in &self.cells {
write!(f, "{}", if c == 1 { '■' } else { '□' })?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_neighborhood_r1() {
let state = CAState::<8, 1>::new([1, 0, 1, 0, 0, 0, 0, 0]);
assert_eq!(state.neighborhood(0), 2);
assert_eq!(state.neighborhood(1), 5);
}
#[test]
fn test_periodic_boundary() {
let state = CAState::<4, 1>::new([1, 0, 0, 0]);
assert_eq!(state.cell(-1), 0); assert_eq!(state.cell(4), 1); }
#[test]
fn test_distance() {
let s1 = CAState::<4, 1>::new([0, 0, 0, 0]);
let s2 = CAState::<4, 1>::new([1, 0, 0, 0]);
assert_eq!(s1.distance(&s2), 1);
assert_eq!(s1.distance(&s1), 0);
}
}