use crate::navmesh::points_equal;
use crate::{
Navmesh, NavmeshPathfinder, NavmeshQuery, NavmeshQueryResult, NavmeshSearchResult, Point2,
PolygonPath,
};
use super::channel_search::{portal_midpoint, search_cell_corridor};
const EPSILON: f64 = 1e-9;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct TAStar;
impl NavmeshPathfinder for TAStar {
fn name(&self) -> &'static str {
"ta-star"
}
fn search(&self, navmesh: &Navmesh, query: NavmeshQuery) -> NavmeshSearchResult {
let (start_cell, goal_cell) = match navmesh.query(query) {
NavmeshQueryResult::Connected {
start_cell,
goal_cell,
} => (start_cell, goal_cell),
NavmeshQueryResult::InvalidStart => {
return Err(crate::NavmeshSearchError::InvalidStart { point: query.start });
}
NavmeshQueryResult::InvalidGoal => {
return Err(crate::NavmeshSearchError::InvalidGoal { point: query.goal });
}
NavmeshQueryResult::NoPath { .. } => return crate::navmesh::search_not_found(0),
};
if points_equal(query.start, query.goal) {
return crate::navmesh::search_found(
PolygonPath::from_points(vec![query.start])
.expect("polygon path contains at least one point"),
1,
);
}
let (Some(cells), visited_nodes) =
search_cell_corridor(navmesh, start_cell, goal_cell, query.budget)?
else {
return crate::navmesh::search_not_found(0);
};
let Some(corridor) = crate::navmesh::corridor::NavmeshCorridor::from_cells(
navmesh,
query.start,
query.goal,
&cells,
) else {
return crate::navmesh::search_not_found(visited_nodes);
};
let midpoint_seed = corridor
.portals
.iter()
.map(portal_midpoint)
.collect::<Vec<_>>();
let baseline_points =
crate::navmesh::funnel::pull_string(navmesh, &corridor, midpoint_seed.clone());
if baseline_points.len() >= 2 && !navmesh.path_is_walkable(&baseline_points) {
return crate::navmesh::search_not_found(visited_nodes);
}
let refined_seed = refine_query_locally(navmesh, &corridor, midpoint_seed);
let refined_points = crate::navmesh::funnel::pull_string(navmesh, &corridor, refined_seed);
let chosen_points = if refined_points.len() >= 2
&& navmesh.path_is_walkable(&refined_points)
&& path_cost(&refined_points) + EPSILON < path_cost(&baseline_points)
{
refined_points
} else {
baseline_points
};
crate::navmesh::search_found(
PolygonPath::from_points(chosen_points)
.expect("polygon path contains at least one point"),
visited_nodes,
)
}
}
fn refine_query_locally(
navmesh: &Navmesh,
corridor: &crate::navmesh::corridor::NavmeshCorridor,
seed_points: Vec<Point2>,
) -> Vec<Point2> {
if seed_points.is_empty() {
return seed_points;
}
let mut refined = seed_points;
let max_passes = corridor.portals.len().max(1);
for _ in 0..max_passes {
let mut improved = false;
for index in 0..corridor.portals.len() {
let portal = corridor.portals[index];
let prev = if index == 0 {
corridor.start
} else {
refined[index - 1]
};
let next = if index + 1 == refined.len() {
corridor.goal
} else {
refined[index + 1]
};
let current = refined[index];
let current_cost = local_turn_cost(prev, current, next);
let mut best_point = current;
let mut best_cost = current_cost;
for candidate in [portal.start, portal.end] {
if !navmesh.segment_is_walkable(prev, candidate)
|| !navmesh.segment_is_walkable(candidate, next)
{
continue;
}
let candidate_cost = local_turn_cost(prev, candidate, next);
if candidate_cost + EPSILON < best_cost {
best_point = candidate;
best_cost = candidate_cost;
}
}
if !points_equal(best_point, current) {
refined[index] = best_point;
improved = true;
}
}
if !improved {
break;
}
}
refined
}
fn local_turn_cost(prev: Point2, current: Point2, next: Point2) -> f64 {
segment_cost(prev, current) + segment_cost(current, next)
}
fn path_cost(points: &[Point2]) -> f64 {
points
.windows(2)
.map(|segment| segment_cost(segment[0], segment[1]))
.sum()
}
fn segment_cost(a: Point2, b: Point2) -> f64 {
((a.x - b.x).powi(2) + (a.y - b.y).powi(2)).sqrt()
}