Skip to main content

proof_engine/pathfinding/
astar.rs

1// src/pathfinding/astar.rs
2// A* and pathfinding variants:
3//   - Generic A* over NodeId graph
4//   - Jump Point Search (JPS) for uniform-cost grid maps
5//   - Hierarchical A* with cluster-level precomputation
6//   - Flow fields for crowd simulation
7//   - Path caching with invalidation
8
9use std::collections::{BinaryHeap, HashMap, HashSet, VecDeque};
10use std::cmp::Ordering;
11
12// ── Vec2 (local, avoids cross-module dep) ────────────────────────────────────
13
14#[derive(Clone, Copy, Debug, PartialEq)]
15pub struct Vec2 {
16    pub x: f32,
17    pub y: f32,
18}
19
20impl Vec2 {
21    #[inline] pub fn new(x: f32, y: f32) -> Self { Self { x, y } }
22    #[inline] pub fn zero() -> Self { Self { x: 0.0, y: 0.0 } }
23    #[inline] pub fn dist(self, o: Self) -> f32 { ((self.x-o.x).powi(2)+(self.y-o.y).powi(2)).sqrt() }
24    #[inline] pub fn sub(self, o: Self) -> Self { Self::new(self.x-o.x, self.y-o.y) }
25    #[inline] pub fn add(self, o: Self) -> Self { Self::new(self.x+o.x, self.y+o.y) }
26    #[inline] pub fn scale(self, s: f32) -> Self { Self::new(self.x*s, self.y*s) }
27    #[inline] pub fn len(self) -> f32 { (self.x*self.x+self.y*self.y).sqrt() }
28    #[inline] pub fn norm(self) -> Self { let l=self.len(); if l<1e-9 {Self::zero()} else {self.scale(1.0/l)} }
29}
30
31// ── Node identifier ───────────────────────────────────────────────────────────
32
33#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
34pub struct NodeId(pub u32);
35
36// ── Generic A* graph trait ────────────────────────────────────────────────────
37
38/// Trait implemented by any graph that wants generic A*.
39pub trait AStarGraph {
40    type Cost: Copy + PartialOrd + std::ops::Add<Output = Self::Cost>;
41    fn zero_cost() -> Self::Cost;
42    fn max_cost() -> Self::Cost;
43    fn heuristic(&self, from: NodeId, to: NodeId) -> Self::Cost;
44    fn neighbors(&self, node: NodeId) -> Vec<(NodeId, Self::Cost)>;
45}
46
47/// Result of A* search.
48#[derive(Clone, Debug)]
49pub struct AStarResult {
50    pub path: Vec<NodeId>,
51    pub cost: f32,
52}
53
54pub struct AStarNode {
55    pub id:       NodeId,
56    pub position: Vec2,
57    pub walkable: bool,
58}
59
60// ── Priority entry ────────────────────────────────────────────────────────────
61
62#[derive(PartialEq)]
63struct PqEntry<C: PartialOrd> {
64    node: NodeId,
65    f:    C,
66}
67
68impl<C: PartialOrd> Eq for PqEntry<C> {}
69
70impl<C: PartialOrd> PartialOrd for PqEntry<C> {
71    fn partial_cmp(&self, other: &Self) -> Option<Ordering> { Some(self.cmp(other)) }
72}
73
74impl<C: PartialOrd> Ord for PqEntry<C> {
75    fn cmp(&self, other: &Self) -> Ordering {
76        other.f.partial_cmp(&self.f).unwrap_or(Ordering::Equal)
77    }
78}
79
80/// Run generic A* on any graph implementing AStarGraph.
81pub fn astar_search<G: AStarGraph>(
82    graph: &G,
83    start: NodeId,
84    goal: NodeId,
85) -> Option<AStarResult>
86where
87    G::Cost: std::fmt::Debug,
88    f32: From<G::Cost>,
89{
90    let mut open: BinaryHeap<PqEntry<G::Cost>> = BinaryHeap::new();
91    let mut came_from: HashMap<NodeId, NodeId> = HashMap::new();
92    let mut g_score: HashMap<NodeId, G::Cost> = HashMap::new();
93
94    g_score.insert(start, G::zero_cost());
95    open.push(PqEntry { node: start, f: graph.heuristic(start, goal) });
96
97    while let Some(PqEntry { node: current, .. }) = open.pop() {
98        if current == goal {
99            let path = reconstruct(start, goal, &came_from);
100            let cost = f32::from(*g_score.get(&goal).unwrap_or(&G::zero_cost()));
101            return Some(AStarResult { path, cost });
102        }
103        let cur_g = *g_score.get(&current).unwrap_or(&G::max_cost());
104        for (neighbor, edge_cost) in graph.neighbors(current) {
105            let tentative = cur_g + edge_cost;
106            if tentative < *g_score.get(&neighbor).unwrap_or(&G::max_cost()) {
107                came_from.insert(neighbor, current);
108                g_score.insert(neighbor, tentative);
109                let h = graph.heuristic(neighbor, goal);
110                open.push(PqEntry { node: neighbor, f: tentative + h });
111            }
112        }
113    }
114    None
115}
116
117fn reconstruct(start: NodeId, goal: NodeId, came_from: &HashMap<NodeId, NodeId>) -> Vec<NodeId> {
118    let mut path = Vec::new();
119    let mut cur = goal;
120    while cur != start {
121        path.push(cur);
122        match came_from.get(&cur) {
123            Some(&p) => cur = p,
124            None => break,
125        }
126    }
127    path.push(start);
128    path.reverse();
129    path
130}
131
132// ── Grid map ──────────────────────────────────────────────────────────────────
133
134/// A 2-D uniform grid map for JPS and flow fields.
135#[derive(Clone, Debug)]
136pub struct GridMap {
137    pub width:    usize,
138    pub height:   usize,
139    pub cells:    Vec<bool>,  // true = walkable
140    pub cell_size: f32,
141    pub origin:   Vec2,
142}
143
144impl GridMap {
145    pub fn new(width: usize, height: usize, cell_size: f32, origin: Vec2) -> Self {
146        Self {
147            width, height,
148            cells: vec![true; width * height],
149            cell_size, origin,
150        }
151    }
152
153    #[inline]
154    pub fn idx(&self, x: usize, y: usize) -> usize { y * self.width + x }
155
156    #[inline]
157    pub fn in_bounds(&self, x: i32, y: i32) -> bool {
158        x >= 0 && y >= 0 && (x as usize) < self.width && (y as usize) < self.height
159    }
160
161    #[inline]
162    pub fn walkable(&self, x: i32, y: i32) -> bool {
163        self.in_bounds(x, y) && self.cells[self.idx(x as usize, y as usize)]
164    }
165
166    pub fn set_walkable(&mut self, x: usize, y: usize, w: bool) {
167        let i = self.idx(x, y);
168        self.cells[i] = w;
169    }
170
171    /// Block a rectangular area.
172    pub fn block_rect(&mut self, x: usize, y: usize, w: usize, h: usize) {
173        for ry in y..((y+h).min(self.height)) {
174            for rx in x..((x+w).min(self.width)) {
175                let i = self.idx(rx, ry);
176                self.cells[i] = false;
177            }
178        }
179    }
180
181    pub fn node_id(&self, x: usize, y: usize) -> NodeId {
182        NodeId((y * self.width + x) as u32)
183    }
184
185    pub fn coords(&self, id: NodeId) -> (usize, usize) {
186        let i = id.0 as usize;
187        (i % self.width, i / self.width)
188    }
189
190    pub fn world_pos(&self, x: usize, y: usize) -> Vec2 {
191        Vec2::new(
192            self.origin.x + (x as f32 + 0.5) * self.cell_size,
193            self.origin.y + (y as f32 + 0.5) * self.cell_size,
194        )
195    }
196
197    pub fn grid_coords_for_world(&self, p: Vec2) -> Option<(usize, usize)> {
198        let gx = ((p.x - self.origin.x) / self.cell_size) as i32;
199        let gy = ((p.y - self.origin.y) / self.cell_size) as i32;
200        if self.in_bounds(gx, gy) {
201            Some((gx as usize, gy as usize))
202        } else {
203            None
204        }
205    }
206}
207
208// ── Jump Point Search ─────────────────────────────────────────────────────────
209
210/// JPS pathfinder for uniform-cost grid maps (8-directional movement).
211pub struct JpsPathfinder<'a> {
212    pub grid: &'a GridMap,
213}
214
215impl<'a> JpsPathfinder<'a> {
216    pub fn new(grid: &'a GridMap) -> Self { Self { grid } }
217
218    /// Find a path from `start` to `goal`, both as grid (x,y) coordinates.
219    pub fn find_path(&self, start: (usize, usize), goal: (usize, usize)) -> Option<Vec<(usize, usize)>> {
220        if !self.grid.walkable(start.0 as i32, start.1 as i32) { return None; }
221        if !self.grid.walkable(goal.0 as i32, goal.1 as i32) { return None; }
222        if start == goal { return Some(vec![start]); }
223
224        let mut open: BinaryHeap<JpsEntry> = BinaryHeap::new();
225        let mut came_from: HashMap<(usize,usize), (usize,usize)> = HashMap::new();
226        let mut g: HashMap<(usize,usize), f32> = HashMap::new();
227        let mut closed: HashSet<(usize,usize)> = HashSet::new();
228
229        g.insert(start, 0.0);
230        open.push(JpsEntry { pos: start, f: self.h(start, goal) });
231
232        while let Some(JpsEntry { pos: cur, .. }) = open.pop() {
233            if cur == goal {
234                return Some(self.reconstruct_path(start, goal, &came_from));
235            }
236            if closed.contains(&cur) { continue; }
237            closed.insert(cur);
238
239            let cur_g = *g.get(&cur).unwrap_or(&f32::MAX);
240            let successors = self.identify_successors(cur, goal, &came_from);
241
242            for succ in successors {
243                if closed.contains(&succ) { continue; }
244                let d = self.cost(cur, succ);
245                let ng = cur_g + d;
246                if ng < *g.get(&succ).unwrap_or(&f32::MAX) {
247                    g.insert(succ, ng);
248                    came_from.insert(succ, cur);
249                    open.push(JpsEntry { pos: succ, f: ng + self.h(succ, goal) });
250                }
251            }
252        }
253        None
254    }
255
256    fn h(&self, a: (usize,usize), b: (usize,usize)) -> f32 {
257        let dx = (a.0 as f32 - b.0 as f32).abs();
258        let dy = (a.1 as f32 - b.1 as f32).abs();
259        // Octile distance
260        let (mn, mx) = if dx < dy { (dx, dy) } else { (dy, dx) };
261        mx + mn * (std::f32::consts::SQRT_2 - 1.0)
262    }
263
264    fn cost(&self, a: (usize,usize), b: (usize,usize)) -> f32 {
265        let dx = (a.0 as i32 - b.0 as i32).abs();
266        let dy = (a.1 as i32 - b.1 as i32).abs();
267        if dx + dy == 2 { std::f32::consts::SQRT_2 } else { 1.0 }
268    }
269
270    fn identify_successors(
271        &self,
272        node: (usize,usize),
273        goal: (usize,usize),
274        came_from: &HashMap<(usize,usize),(usize,usize)>,
275    ) -> Vec<(usize,usize)> {
276        let neighbors = self.prune_neighbors(node, came_from);
277        let mut successors = Vec::new();
278        for nb in neighbors {
279            let dx = (nb.0 as i32 - node.0 as i32).signum();
280            let dy = (nb.1 as i32 - node.1 as i32).signum();
281            if let Some(jp) = self.jump(node, (dx, dy), goal) {
282                successors.push(jp);
283            }
284        }
285        successors
286    }
287
288    fn prune_neighbors(
289        &self,
290        node: (usize,usize),
291        came_from: &HashMap<(usize,usize),(usize,usize)>,
292    ) -> Vec<(usize,usize)> {
293        let parent = came_from.get(&node);
294        let (x, y) = (node.0 as i32, node.1 as i32);
295        if parent.is_none() {
296            // Start node: return all walkable neighbors
297            return self.all_neighbors(node);
298        }
299        let parent = parent.unwrap();
300        let dx = (x - parent.0 as i32).signum();
301        let dy = (y - parent.1 as i32).signum();
302        let mut neighbors = Vec::new();
303
304        if dx != 0 && dy != 0 {
305            // Diagonal
306            if self.grid.walkable(x, y + dy)     { neighbors.push((x as usize, (y+dy) as usize)); }
307            if self.grid.walkable(x + dx, y)     { neighbors.push(((x+dx) as usize, y as usize)); }
308            if self.grid.walkable(x + dx, y + dy) { neighbors.push(((x+dx) as usize, (y+dy) as usize)); }
309            if !self.grid.walkable(x - dx, y) && self.grid.walkable(x, y + dy) {
310                neighbors.push((x as usize, (y + dy) as usize));
311            }
312            if !self.grid.walkable(x, y - dy) && self.grid.walkable(x + dx, y) {
313                neighbors.push(((x + dx) as usize, y as usize));
314            }
315        } else if dx != 0 {
316            // Horizontal
317            if self.grid.walkable(x + dx, y) { neighbors.push(((x+dx) as usize, y as usize)); }
318            if !self.grid.walkable(x, y + 1) && self.grid.walkable(x + dx, y + 1) {
319                neighbors.push(((x+dx) as usize, (y+1) as usize));
320            }
321            if !self.grid.walkable(x, y - 1) && self.grid.walkable(x + dx, y - 1) {
322                neighbors.push(((x+dx) as usize, (y-1) as usize));
323            }
324        } else {
325            // Vertical
326            if self.grid.walkable(x, y + dy) { neighbors.push((x as usize, (y+dy) as usize)); }
327            if !self.grid.walkable(x + 1, y) && self.grid.walkable(x + 1, y + dy) {
328                neighbors.push(((x+1) as usize, (y+dy) as usize));
329            }
330            if !self.grid.walkable(x - 1, y) && self.grid.walkable(x - 1, y + dy) {
331                neighbors.push(((x-1) as usize, (y+dy) as usize));
332            }
333        }
334        neighbors.dedup();
335        neighbors
336    }
337
338    fn all_neighbors(&self, node: (usize,usize)) -> Vec<(usize,usize)> {
339        let (x, y) = (node.0 as i32, node.1 as i32);
340        let mut result = Vec::new();
341        for dy in -1i32..=1 {
342            for dx in -1i32..=1 {
343                if dx == 0 && dy == 0 { continue; }
344                if self.grid.walkable(x + dx, y + dy) {
345                    result.push(((x + dx) as usize, (y + dy) as usize));
346                }
347            }
348        }
349        result
350    }
351
352    fn jump(&self, node: (usize,usize), dir: (i32,i32), goal: (usize,usize)) -> Option<(usize,usize)> {
353        let (mut x, mut y) = (node.0 as i32, node.1 as i32);
354        let (dx, dy) = dir;
355        let max_steps = (self.grid.width + self.grid.height) * 2;
356        let mut steps = 0;
357
358        loop {
359            x += dx;
360            y += dy;
361            steps += 1;
362            if steps > max_steps { return None; }
363            if !self.grid.walkable(x, y) { return None; }
364            let cur = (x as usize, y as usize);
365            if cur == goal { return Some(cur); }
366
367            // Check for forced neighbors
368            if self.has_forced_neighbor(cur, dir) { return Some(cur); }
369
370            // Diagonal: recurse on both cardinal directions
371            if dx != 0 && dy != 0 {
372                if self.jump((x as usize, y as usize), (dx, 0), goal).is_some() { return Some(cur); }
373                if self.jump((x as usize, y as usize), (0, dy), goal).is_some() { return Some(cur); }
374            }
375        }
376    }
377
378    fn has_forced_neighbor(&self, node: (usize,usize), dir: (i32,i32)) -> bool {
379        let (x, y) = (node.0 as i32, node.1 as i32);
380        let (dx, dy) = dir;
381        if dx != 0 && dy != 0 {
382            // diagonal forced: blocked adjacent cardinal
383            (!self.grid.walkable(x - dx, y) && self.grid.walkable(x - dx, y + dy))
384            || (!self.grid.walkable(x, y - dy) && self.grid.walkable(x + dx, y - dy))
385        } else if dx != 0 {
386            (!self.grid.walkable(x, y + 1) && self.grid.walkable(x + dx, y + 1))
387            || (!self.grid.walkable(x, y - 1) && self.grid.walkable(x + dx, y - 1))
388        } else {
389            (!self.grid.walkable(x + 1, y) && self.grid.walkable(x + 1, y + dy))
390            || (!self.grid.walkable(x - 1, y) && self.grid.walkable(x - 1, y + dy))
391        }
392    }
393
394    fn reconstruct_path(
395        &self,
396        start: (usize,usize),
397        goal: (usize,usize),
398        came_from: &HashMap<(usize,usize),(usize,usize)>,
399    ) -> Vec<(usize,usize)> {
400        let mut path = Vec::new();
401        let mut cur = goal;
402        while cur != start {
403            path.push(cur);
404            match came_from.get(&cur) {
405                Some(&p) => cur = p,
406                None => break,
407            }
408        }
409        path.push(start);
410        path.reverse();
411        // Expand jump-point path into full grid steps
412        let mut expanded = Vec::new();
413        for i in 0..path.len().saturating_sub(1) {
414            expanded.push(path[i]);
415            let (ax, ay) = (path[i].0 as i32, path[i].1 as i32);
416            let (bx, by) = (path[i+1].0 as i32, path[i+1].1 as i32);
417            let sdx = (bx - ax).signum();
418            let sdy = (by - ay).signum();
419            let mut cx = ax + sdx;
420            let mut cy = ay + sdy;
421            while (cx, cy) != (bx, by) {
422                expanded.push((cx as usize, cy as usize));
423                cx += sdx;
424                cy += sdy;
425            }
426        }
427        if let Some(&last) = path.last() { expanded.push(last); }
428        expanded.dedup();
429        expanded
430    }
431}
432
433#[derive(PartialEq)]
434struct JpsEntry { pos: (usize,usize), f: f32 }
435impl Eq for JpsEntry {}
436impl PartialOrd for JpsEntry {
437    fn partial_cmp(&self, o: &Self) -> Option<Ordering> { Some(self.cmp(o)) }
438}
439impl Ord for JpsEntry {
440    fn cmp(&self, o: &Self) -> Ordering { o.f.partial_cmp(&self.f).unwrap_or(Ordering::Equal) }
441}
442
443// ── Hierarchical A* ───────────────────────────────────────────────────────────
444
445#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
446pub struct ClusterId(pub u32);
447
448/// A cluster groups nearby grid cells for hierarchical planning.
449#[derive(Clone, Debug)]
450pub struct Cluster {
451    pub id:            ClusterId,
452    pub x:             usize,   // grid cell offset
453    pub y:             usize,
454    pub width:         usize,
455    pub height:        usize,
456    pub entry_cells:   Vec<(usize, usize)>,   // border cells that connect to other clusters
457    pub neighbors:     Vec<(ClusterId, f32)>, // neighbor cluster + estimated cost
458}
459
460impl Cluster {
461    pub fn contains(&self, cx: usize, cy: usize) -> bool {
462        cx >= self.x && cx < self.x + self.width
463        && cy >= self.y && cy < self.y + self.height
464    }
465    pub fn center_cell(&self) -> (usize, usize) {
466        (self.x + self.width / 2, self.y + self.height / 2)
467    }
468}
469
470/// Hierarchical pathfinder: builds abstract cluster graph then refines.
471pub struct HierarchicalPathfinder {
472    pub clusters:      Vec<Cluster>,
473    pub cluster_map:   Vec<Option<ClusterId>>,  // per grid cell
474    pub grid_width:    usize,
475    pub grid_height:   usize,
476    pub cluster_size:  usize,
477}
478
479impl HierarchicalPathfinder {
480    /// Build cluster graph from a GridMap with given cluster size.
481    pub fn build(grid: &GridMap, cluster_size: usize) -> Self {
482        let cw = (grid.width  + cluster_size - 1) / cluster_size;
483        let ch = (grid.height + cluster_size - 1) / cluster_size;
484        let mut clusters = Vec::with_capacity(cw * ch);
485        let mut cluster_map = vec![None; grid.width * grid.height];
486        let mut id = 0u32;
487
488        for cy in 0..ch {
489            for cx in 0..cw {
490                let ox = cx * cluster_size;
491                let oy = cy * cluster_size;
492                let w  = cluster_size.min(grid.width  - ox);
493                let h  = cluster_size.min(grid.height - oy);
494
495                let mut entry_cells = Vec::new();
496                // Top/bottom border
497                for bx in ox..(ox+w) {
498                    if grid.walkable(bx as i32, oy as i32)        { entry_cells.push((bx, oy)); }
499                    let by = oy + h - 1;
500                    if grid.walkable(bx as i32, by as i32)        { entry_cells.push((bx, by)); }
501                }
502                // Left/right border
503                for by in oy..(oy+h) {
504                    if grid.walkable(ox as i32, by as i32)        { entry_cells.push((ox, by)); }
505                    let bx = ox + w - 1;
506                    if grid.walkable(bx as i32, by as i32)        { entry_cells.push((bx, by)); }
507                }
508                entry_cells.sort();
509                entry_cells.dedup();
510
511                let cid = ClusterId(id);
512                for gx in ox..(ox+w) {
513                    for gy in oy..(oy+h) {
514                        if grid.in_bounds(gx as i32, gy as i32) {
515                            let gi = gy * grid.width + gx;
516                            cluster_map[gi] = Some(cid);
517                        }
518                    }
519                }
520                clusters.push(Cluster {
521                    id: cid, x: ox, y: oy, width: w, height: h,
522                    entry_cells, neighbors: Vec::new(),
523                });
524                id += 1;
525            }
526        }
527
528        // Build neighbor edges between adjacent clusters
529        let mut hpf = HierarchicalPathfinder {
530            clusters, cluster_map,
531            grid_width: grid.width,
532            grid_height: grid.height,
533            cluster_size,
534        };
535        hpf.build_cluster_edges(grid);
536        hpf
537    }
538
539    fn build_cluster_edges(&mut self, grid: &GridMap) {
540        let nc = self.clusters.len();
541        for i in 0..nc {
542            let ci = &self.clusters[i];
543            // Check 4 adjacent cluster positions
544            let (cx, cy) = (ci.x, ci.y);
545            let cs = self.cluster_size;
546            let adj_offsets: [(i32,i32); 4] = [(1,0),(-1,0),(0,1),(0,-1)];
547            let mut neighbors = Vec::new();
548            for (aox, aoy) in adj_offsets {
549                let nx = cx as i32 + aox * cs as i32;
550                let ny = cy as i32 + aoy * cs as i32;
551                if nx < 0 || ny < 0 { continue; }
552                if let Some(j) = self.find_cluster_at(nx as usize, ny as usize) {
553                    if j != i {
554                        let cost = cs as f32; // approximate
555                        neighbors.push((self.clusters[j].id, cost));
556                    }
557                }
558            }
559            // Update neighbors (can't borrow mut + immut simultaneously, so rebuild)
560            let _ = neighbors; // will be set below
561        }
562        // Simplified: link adjacent grid clusters
563        let cw = (grid.width  + self.cluster_size - 1) / self.cluster_size;
564        let ch = (grid.height + self.cluster_size - 1) / self.cluster_size;
565        for cy in 0..ch {
566            for cx in 0..cw {
567                let idx = cy * cw + cx;
568                if idx >= self.clusters.len() { continue; }
569                let mut nbrs = Vec::new();
570                let pairs: [(i32,i32); 4] = [(1,0),(-1,0),(0,1),(0,-1)];
571                for (ddx, ddy) in pairs {
572                    let ncx = cx as i32 + ddx;
573                    let ncy = cy as i32 + ddy;
574                    if ncx < 0 || ncy < 0 || ncx >= cw as i32 || ncy >= ch as i32 { continue; }
575                    let nidx = (ncy as usize) * cw + (ncx as usize);
576                    if nidx < self.clusters.len() {
577                        let nid = self.clusters[nidx].id;
578                        nbrs.push((nid, self.cluster_size as f32));
579                    }
580                }
581                self.clusters[idx].neighbors = nbrs;
582            }
583        }
584    }
585
586    fn find_cluster_at(&self, x: usize, y: usize) -> Option<usize> {
587        self.clusters.iter().position(|c| c.contains(x, y))
588    }
589
590    pub fn cluster_for_cell(&self, x: usize, y: usize) -> Option<ClusterId> {
591        if x >= self.grid_width || y >= self.grid_height { return None; }
592        self.cluster_map[y * self.grid_width + x]
593    }
594
595    /// High-level path: returns sequence of ClusterIds.
596    pub fn abstract_path(&self, start_cell: (usize,usize), goal_cell: (usize,usize)) -> Vec<ClusterId> {
597        let sc = match self.cluster_for_cell(start_cell.0, start_cell.1) { Some(c) => c, None => return Vec::new() };
598        let gc = match self.cluster_for_cell(goal_cell.0, goal_cell.1) { Some(c) => c, None => return Vec::new() };
599        if sc == gc { return vec![sc]; }
600
601        let mut open: BinaryHeap<ClusterEntry> = BinaryHeap::new();
602        let mut came_from: HashMap<ClusterId, ClusterId> = HashMap::new();
603        let mut g: HashMap<ClusterId, f32> = HashMap::new();
604
605        g.insert(sc, 0.0);
606        open.push(ClusterEntry { id: sc, f: self.cluster_heuristic(sc, gc) });
607
608        while let Some(ClusterEntry { id: cur, .. }) = open.pop() {
609            if cur == gc {
610                let mut path = Vec::new();
611                let mut c = gc;
612                while c != sc {
613                    path.push(c);
614                    c = *came_from.get(&c).unwrap_or(&sc);
615                }
616                path.push(sc);
617                path.reverse();
618                return path;
619            }
620            let cur_g = *g.get(&cur).unwrap_or(&f32::MAX);
621            if let Some(cluster) = self.clusters.iter().find(|c| c.id == cur) {
622                for &(nid, edge_cost) in &cluster.neighbors {
623                    let ng = cur_g + edge_cost;
624                    if ng < *g.get(&nid).unwrap_or(&f32::MAX) {
625                        g.insert(nid, ng);
626                        came_from.insert(nid, cur);
627                        let h = self.cluster_heuristic(nid, gc);
628                        open.push(ClusterEntry { id: nid, f: ng + h });
629                    }
630                }
631            }
632        }
633        Vec::new()
634    }
635
636    fn cluster_heuristic(&self, a: ClusterId, b: ClusterId) -> f32 {
637        let ca = self.clusters.iter().find(|c| c.id == a).map(|c| c.center_cell()).unwrap_or((0,0));
638        let cb = self.clusters.iter().find(|c| c.id == b).map(|c| c.center_cell()).unwrap_or((0,0));
639        let dx = (ca.0 as f32 - cb.0 as f32).abs();
640        let dy = (ca.1 as f32 - cb.1 as f32).abs();
641        dx.max(dy)
642    }
643}
644
645#[derive(PartialEq)]
646struct ClusterEntry { id: ClusterId, f: f32 }
647impl Eq for ClusterEntry {}
648impl PartialOrd for ClusterEntry {
649    fn partial_cmp(&self, o: &Self) -> Option<Ordering> { Some(self.cmp(o)) }
650}
651impl Ord for ClusterEntry {
652    fn cmp(&self, o: &Self) -> Ordering { o.f.partial_cmp(&self.f).unwrap_or(Ordering::Equal) }
653}
654
655// ── Flow Field ────────────────────────────────────────────────────────────────
656
657/// Flow direction per cell: an 8-directional flow vector.
658#[derive(Clone, Copy, Debug, Default)]
659pub struct FlowVector {
660    pub dx: i8,   // -1, 0, +1
661    pub dy: i8,
662}
663
664impl FlowVector {
665    pub fn as_vec2(self) -> Vec2 {
666        Vec2::new(self.dx as f32, self.dy as f32).norm()
667    }
668    pub fn is_valid(self) -> bool { self.dx != 0 || self.dy != 0 }
669}
670
671/// A flow field: precomputed for a single goal, steers any number of agents.
672#[derive(Clone, Debug)]
673pub struct FlowField {
674    pub width:   usize,
675    pub height:  usize,
676    pub flow:    Vec<FlowVector>,
677    pub cost:    Vec<f32>,          // integration field (distance to goal)
678    pub goal:    (usize, usize),
679}
680
681/// Flow field grid: builds and stores flow fields.
682pub struct FlowFieldGrid<'a> {
683    pub grid: &'a GridMap,
684}
685
686impl<'a> FlowFieldGrid<'a> {
687    pub fn new(grid: &'a GridMap) -> Self { Self { grid } }
688
689    /// Build a flow field toward `goal` using Dijkstra integration.
690    pub fn build(&self, goal: (usize, usize)) -> FlowField {
691        let w = self.grid.width;
692        let h = self.grid.height;
693        let inf = f32::MAX / 2.0;
694        let mut cost = vec![inf; w * h];
695        let mut flow = vec![FlowVector::default(); w * h];
696
697        if !self.grid.walkable(goal.0 as i32, goal.1 as i32) {
698            return FlowField { width: w, height: h, flow, cost, goal };
699        }
700
701        let gi = goal.1 * w + goal.0;
702        cost[gi] = 0.0;
703
704        // BFS/Dijkstra integration field
705        let mut queue: VecDeque<(usize,usize)> = VecDeque::new();
706        queue.push_back(goal);
707
708        while let Some((cx, cy)) = queue.pop_front() {
709            let cur_cost = cost[cy * w + cx];
710            for (dx, dy) in DIRS_8 {
711                let nx = cx as i32 + dx;
712                let ny = cy as i32 + dy;
713                if !self.grid.walkable(nx, ny) { continue; }
714                let (nxi, nyi) = (nx as usize, ny as usize);
715                let ni = nyi * w + nxi;
716                let step_cost = if dx != 0 && dy != 0 { std::f32::consts::SQRT_2 } else { 1.0 };
717                let nc = cur_cost + step_cost;
718                if nc < cost[ni] {
719                    cost[ni] = nc;
720                    queue.push_back((nxi, nyi));
721                }
722            }
723        }
724
725        // Build flow vectors: each cell points toward the lowest-cost neighbor
726        for cy in 0..h {
727            for cx in 0..w {
728                if !self.grid.walkable(cx as i32, cy as i32) { continue; }
729                let ci = cy * w + cx;
730                if cost[ci] >= inf { continue; }
731
732                let mut best_cost = cost[ci];
733                let mut best_dx = 0i8;
734                let mut best_dy = 0i8;
735                for (dx, dy) in DIRS_8 {
736                    let nx = cx as i32 + dx;
737                    let ny = cy as i32 + dy;
738                    if !self.grid.walkable(nx, ny) { continue; }
739                    let ni = ny as usize * w + nx as usize;
740                    if cost[ni] < best_cost {
741                        best_cost = cost[ni];
742                        best_dx = dx as i8;
743                        best_dy = dy as i8;
744                    }
745                }
746                flow[ci] = FlowVector { dx: best_dx, dy: best_dy };
747            }
748        }
749
750        FlowField { width: w, height: h, flow, cost, goal }
751    }
752}
753
754const DIRS_8: [(i32,i32); 8] = [
755    (1,0),(-1,0),(0,1),(0,-1),(1,1),(1,-1),(-1,1),(-1,-1)
756];
757
758impl FlowField {
759    /// Get the flow vector at grid cell (x, y).
760    pub fn get_flow(&self, x: usize, y: usize) -> FlowVector {
761        if x < self.width && y < self.height {
762            self.flow[y * self.width + x]
763        } else {
764            FlowVector::default()
765        }
766    }
767
768    /// Get the integration cost at grid cell (x, y).
769    pub fn get_cost(&self, x: usize, y: usize) -> f32 {
770        if x < self.width && y < self.height {
771            self.cost[y * self.width + x]
772        } else {
773            f32::MAX
774        }
775    }
776
777    /// Sample the flow direction at world position `p` (grid-space integer lookup).
778    pub fn sample(&self, gx: usize, gy: usize) -> Vec2 {
779        self.get_flow(gx, gy).as_vec2()
780    }
781}
782
783// ── Path cache with invalidation ──────────────────────────────────────────────
784
785/// A cached path entry.
786#[derive(Clone, Debug)]
787pub struct CachedPath {
788    pub start:   (usize, usize),
789    pub goal:    (usize, usize),
790    pub path:    Vec<(usize, usize)>,
791    pub version: u64,
792}
793
794/// Cache of computed paths, invalidated when the grid changes.
795pub struct PathCache {
796    entries:       HashMap<((usize,usize),(usize,usize)), CachedPath>,
797    pub version:   u64,
798    capacity:      usize,
799    // LRU tracking via insertion order
800    order:         VecDeque<((usize,usize),(usize,usize))>,
801}
802
803impl PathCache {
804    pub fn new(capacity: usize) -> Self {
805        Self {
806            entries: HashMap::new(),
807            version: 0,
808            capacity,
809            order: VecDeque::new(),
810        }
811    }
812
813    /// Increment version, invalidating all stale cache entries.
814    pub fn invalidate(&mut self) {
815        self.version += 1;
816    }
817
818    /// Clear all entries.
819    pub fn clear(&mut self) {
820        self.entries.clear();
821        self.order.clear();
822    }
823
824    /// Look up a cached path; returns None if not present or stale.
825    pub fn get(&self, start: (usize,usize), goal: (usize,usize)) -> Option<&Vec<(usize,usize)>> {
826        let key = (start, goal);
827        let entry = self.entries.get(&key)?;
828        if entry.version == self.version {
829            Some(&entry.path)
830        } else {
831            None
832        }
833    }
834
835    /// Store a path in the cache, evicting LRU entry if over capacity.
836    pub fn insert(&mut self, start: (usize,usize), goal: (usize,usize), path: Vec<(usize,usize)>) {
837        let key = (start, goal);
838        if self.entries.contains_key(&key) {
839            self.entries.get_mut(&key).unwrap().path = path;
840            self.entries.get_mut(&key).unwrap().version = self.version;
841        } else {
842            if self.entries.len() >= self.capacity {
843                if let Some(evict_key) = self.order.pop_front() {
844                    self.entries.remove(&evict_key);
845                }
846            }
847            self.entries.insert(key, CachedPath { start, goal, path, version: self.version });
848            self.order.push_back(key);
849        }
850    }
851
852    /// Get or compute a path, using JPS if not cached.
853    pub fn get_or_compute<'g>(&mut self, grid: &'g GridMap, start: (usize,usize), goal: (usize,usize)) -> Vec<(usize,usize)> {
854        if let Some(cached) = self.get(start, goal) {
855            return cached.clone();
856        }
857        let jps = JpsPathfinder::new(grid);
858        let path = jps.find_path(start, goal).unwrap_or_default();
859        self.insert(start, goal, path.clone());
860        path
861    }
862
863    pub fn entry_count(&self) -> usize { self.entries.len() }
864}
865
866// ── Simple concrete graph for generic A* ─────────────────────────────────────
867
868/// Simple flat graph with node positions and weighted edges.
869pub struct SimpleGraph {
870    pub nodes:     Vec<Vec2>,
871    pub edges:     Vec<Vec<(NodeId, f32)>>,
872}
873
874impl SimpleGraph {
875    pub fn new() -> Self { Self { nodes: Vec::new(), edges: Vec::new() } }
876
877    pub fn add_node(&mut self, pos: Vec2) -> NodeId {
878        let id = NodeId(self.nodes.len() as u32);
879        self.nodes.push(pos);
880        self.edges.push(Vec::new());
881        id
882    }
883
884    pub fn add_edge(&mut self, a: NodeId, b: NodeId, cost: f32) {
885        let ai = a.0 as usize;
886        let bi = b.0 as usize;
887        if ai < self.edges.len() { self.edges[ai].push((b, cost)); }
888        if bi < self.edges.len() { self.edges[bi].push((a, cost)); }
889    }
890}
891
892impl AStarGraph for SimpleGraph {
893    type Cost = f32;
894    fn zero_cost() -> f32 { 0.0 }
895    fn max_cost() -> f32  { f32::MAX / 2.0 }
896    fn heuristic(&self, from: NodeId, to: NodeId) -> f32 {
897        let a = self.nodes.get(from.0 as usize).copied().unwrap_or(Vec2::zero());
898        let b = self.nodes.get(to.0  as usize).copied().unwrap_or(Vec2::zero());
899        a.dist(b)
900    }
901    fn neighbors(&self, node: NodeId) -> Vec<(NodeId, f32)> {
902        self.edges.get(node.0 as usize).cloned().unwrap_or_default()
903    }
904}
905
906// ── Tests ─────────────────────────────────────────────────────────────────────
907
908#[cfg(test)]
909mod tests {
910    use super::*;
911
912    #[test]
913    fn test_astar_simple() {
914        let mut g = SimpleGraph::new();
915        let a = g.add_node(Vec2::new(0.0, 0.0));
916        let b = g.add_node(Vec2::new(1.0, 0.0));
917        let c = g.add_node(Vec2::new(2.0, 0.0));
918        g.add_edge(a, b, 1.0);
919        g.add_edge(b, c, 1.0);
920        let res = astar_search(&g, a, c).unwrap();
921        assert_eq!(res.path, vec![a, b, c]);
922        assert!((res.cost - 2.0).abs() < 1e-4);
923    }
924
925    #[test]
926    fn test_jps_straight() {
927        let mut grid = GridMap::new(10, 10, 1.0, Vec2::zero());
928        let jps = JpsPathfinder::new(&grid);
929        let path = jps.find_path((0,0), (5,0)).unwrap();
930        assert!(!path.is_empty());
931        assert_eq!(path[0], (0,0));
932        assert_eq!(*path.last().unwrap(), (5,0));
933    }
934
935    #[test]
936    fn test_jps_with_obstacle() {
937        let mut grid = GridMap::new(10, 10, 1.0, Vec2::zero());
938        // Wall in the middle
939        for y in 0..8 { grid.set_walkable(5, y, false); }
940        let jps = JpsPathfinder::new(&grid);
941        let path = jps.find_path((0,5), (9,5));
942        // Should find path around wall
943        assert!(path.is_some());
944    }
945
946    #[test]
947    fn test_flow_field() {
948        let grid = GridMap::new(8, 8, 1.0, Vec2::zero());
949        let ffg = FlowFieldGrid::new(&grid);
950        let ff = ffg.build((7, 7));
951        // Cell (0,0) should have valid flow toward goal
952        let fv = ff.get_flow(0, 0);
953        assert!(fv.is_valid());
954    }
955
956    #[test]
957    fn test_path_cache() {
958        let grid = GridMap::new(10, 10, 1.0, Vec2::zero());
959        let mut cache = PathCache::new(16);
960        let path = cache.get_or_compute(&grid, (0,0), (9,9));
961        assert!(!path.is_empty());
962        // Second call should hit cache
963        let path2 = cache.get_or_compute(&grid, (0,0), (9,9));
964        assert_eq!(path, path2);
965        // After invalidation, cache entry is stale
966        cache.invalidate();
967        let cached = cache.get((0,0), (9,9));
968        assert!(cached.is_none());
969    }
970
971    #[test]
972    fn test_hierarchical_abstract_path() {
973        let grid = GridMap::new(16, 16, 1.0, Vec2::zero());
974        let hpf = HierarchicalPathfinder::build(&grid, 4);
975        let abstract_path = hpf.abstract_path((0,0), (15,15));
976        assert!(!abstract_path.is_empty());
977    }
978}