pub struct SpiralSearch {
x: i32,
y: i32,
dx: i32,
dy: i32,
count: usize,
max_distance: i32,
done: bool,
}
impl SpiralSearch {
#[must_use]
pub fn new(max_distance: i32) -> Self {
Self {
x: 0,
y: 0,
dx: 0,
dy: -1,
count: 0,
max_distance,
done: false,
}
}
}
impl Iterator for SpiralSearch {
type Item = (i32, i32);
fn next(&mut self) -> Option<Self::Item> {
if self.done {
return None;
}
if self.count != 0 {
if self.x == self.y
|| (self.x < 0 && self.x == -self.y)
|| (self.x > 0 && self.x == 1 - self.y)
{
let t = self.dx;
self.dx = -self.dy;
self.dy = t;
}
self.x += self.dx;
self.y += self.dy;
}
self.count += 1;
if self.x > self.max_distance {
self.done = true;
return None;
}
Some((self.x, self.y))
}
}
#[must_use]
pub fn spiral_len(max_distance: i32) -> u64 {
match u64::try_from(max_distance) {
Ok(m) => (2 * m + 1) * (2 * m + 1),
Err(_) => 0,
}
}
#[must_use]
pub fn spiral_position(index: u64) -> (i32, i32) {
if index == 0 {
return (0, 0);
}
let ku = index.isqrt().div_ceil(2);
let offset = index - (2 * ku - 1) * (2 * ku - 1);
let (side, t) = (offset / (2 * ku), (offset % (2 * ku)) as i64);
let k = ku as i64;
let (x, y) = match side {
0 => (k, 1 - k + t),
1 => (k - 1 - t, k),
2 => (-k, k - 1 - t),
_ => (1 - k + t, -k),
};
(x as i32, y as i32)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn closed_form_matches_the_walk() {
for max in -2..=40 {
let walked: Vec<(i32, i32)> = SpiralSearch::new(max).collect();
assert_eq!(walked.len() as u64, spiral_len(max), "max {max}");
for (i, &p) in walked.iter().enumerate() {
assert_eq!(spiral_position(i as u64), p, "index {i}");
}
}
}
#[test]
fn a_saturated_spiral_is_counted_not_built() {
let m = (1.0f64 / 1e-300 + 2.0) as i32;
assert_eq!(m, i32::MAX);
let n = spiral_len(m);
assert_eq!(n, 4_294_967_295u64 * 4_294_967_295);
assert_eq!(spiral_position(n - 1), (i32::MAX, -i32::MAX));
}
#[test]
fn first_few_positions() {
let v: Vec<_> = SpiralSearch::new(2).collect();
assert_eq!(v[0], (0, 0));
assert_eq!(v[1], (1, 0));
assert_eq!(v[2], (1, 1));
assert_eq!(v[3], (0, 1));
assert_eq!(v[4], (-1, 1));
assert_eq!(v[5], (-1, 0));
assert_eq!(v[6], (-1, -1));
assert_eq!(v[7], (0, -1));
assert_eq!(v[8], (1, -1));
}
#[test]
fn stops_after_max_distance() {
let max = 3;
let positions: Vec<_> = SpiralSearch::new(max).collect();
for &(x, _) in &positions {
assert!(x <= max, "x={x} exceeded max_distance={max}");
}
}
#[test]
fn max_distance_zero_gives_only_origin() {
let v: Vec<_> = SpiralSearch::new(0).collect();
assert_eq!(v, vec![(0, 0)]);
}
#[test]
fn covers_expected_area() {
let positions: Vec<_> = SpiralSearch::new(2).collect();
for &(x, y) in &positions {
assert!((-2..=2).contains(&x), "x={x}");
let _ = y; }
assert_eq!(positions[0], (0, 0));
}
}