use num_traits::{Num, Zero};
use std::ops::{Add, Sub};
pub struct Tour<'a, T, DFn> {
distance: &'a DFn,
points: &'a [T],
route: Vec<usize>,
}
impl<'a, T, DFn, D> Tour<'a, T, DFn>
where DFn: Fn(&T, &T) -> D,
D: Ord {
fn new(points: &'a [T], distance: &'a DFn) -> Self {
Self {
distance,
points,
route: (0..points.len()).collect(),
}
}
fn new_with_route(points: &'a [T], route: Vec<usize>, distance: &'a DFn) -> Self {
assert_eq!(points.len(), route.len());
Self {
distance,
points,
route,
}
}
fn len(&self) -> usize { self.points.len() }
fn next_index(&self, idx: usize) -> usize {
assert!(idx < self.len());
if idx + 1 < self.len() {
idx + 1
} else {
0
}
}
fn prev_index(&self, idx: usize) -> usize {
assert!(idx < self.len());
if idx == 0 {
self.len() - 1
} else {
idx - 1
}
}
fn element(&self, index: usize) -> &T {
&self.points[self.route[index]]
}
fn elements(&self, idx1: usize, idx2: usize) -> (&T, &T) {
(self.element(idx1), self.element(idx2))
}
fn edge_elements(&self, edge_start_idx: usize) -> (&T, &T) {
self.elements(edge_start_idx, self.next_index(edge_start_idx))
}
fn edge_length(&self, edge_start_idx: usize) -> D {
let (a, b) = self.edge_elements(edge_start_idx);
(self.distance)(a, b)
}
}
impl<'a, T, DFn, D> Tour<'a, T, DFn>
where DFn: Fn(&T, &T) -> D,
D: Ord + Add<Output=D> + Zero {
fn total_length_nowrap(&self) -> D {
(1..self.len())
.map(|i| self.edge_length(i-1))
.fold(Zero::zero(), |l, acc| l + acc)
}
fn length_of_adjacent_edges_nowrap(&self, idx: usize) -> D {
if idx == 0 {
self.edge_length(idx)
} else if idx == self.len() - 1 {
self.edge_length(self.prev_index(idx))
} else {
self.edge_length(idx) + self.edge_length(self.prev_index(idx))
}
}
}
pub fn solve_tsp<T, DFn, D>(points: &[T], distance_fn: &DFn) -> Vec<usize>
where DFn: Fn(&T, &T) -> D,
D: Ord + Add<Output=D> + Sub<Output=D> + Zero {
let route = solve_tsp_nearest_neighbour(points, distance_fn);
let mut tour = Tour::new_with_route(points, route, distance_fn);
let initial_cost = tour.total_length_nowrap();
let swap_improves_length = |tour: &Tour<T, DFn>, idx1: usize, idx2: usize| -> bool {
debug_assert!(idx1 < tour.len());
debug_assert!(idx2 < tour.len());
if idx1 == idx2 {
return false;
}
let (idx1, idx2) = if idx1 <= idx2 {
(idx1, idx2)
} else {
(idx2, idx1)
};
debug_assert_ne!(idx1 + 1, tour.len());
debug_assert_ne!(idx2, 0);
let p1 = tour.element(idx1);
let p2 = tour.element(idx2);
let p1_prev = tour.element(tour.prev_index(idx1));
let p2_prev = tour.element(tour.prev_index(idx2));
let p1_next = tour.element(tour.next_index(idx1));
let p2_next = tour.element(tour.next_index(idx2));
let d_now1 = if idx1 == 0 { Zero::zero() } else { distance_fn(p1_prev, p1) }
+ if idx1 + 1 == tour.len() || idx1 + 1 == idx2 { Zero::zero() } else { distance_fn(p1, p1_next) };
let d_now2 = if idx1 + 1 == idx2 { Zero::zero() } else { distance_fn(p2_prev, p2) }
+ if idx2 + 1 == tour.len() { Zero::zero() } else { distance_fn(p2, p2_next) };
let d_swapped1 = if idx1 == 0 { Zero::zero() } else { distance_fn(p1_prev, p2) }
+ if idx1 + 1 == tour.len() || idx1 + 1 == idx2 { Zero::zero() } else { distance_fn(p2, p1_next) };
let d_swapped2 = if idx1 + 1 == idx2 { Zero::zero() } else { distance_fn(p2_prev, p1) }
+ if idx2 + 1 == tour.len() { Zero::zero() } else { distance_fn(p1, p2_next) };
let d_now = d_now1 + d_now2;
let d_swapped = d_swapped1 + d_swapped2;
d_now > d_swapped
};
'outer: loop {
'inner: for i in 0..tour.len() {
for j in i..tour.len() {
if swap_improves_length(&tour, i, j) {
tour.route.swap(i, j);
break 'inner;
}
}
if i == tour.len() - 1 {
break 'outer;
}
}
}
let cost = tour.total_length_nowrap();
debug_assert!(cost <= initial_cost);
tour.route
}
fn solve_tsp_nearest_neighbour<T, DFn, D>(points: &[T], distance_fn: &DFn) -> Vec<usize>
where DFn: Fn(&T, &T) -> D,
D: Ord {
let mut tour = Tour::new(points, distance_fn);
if tour.len() <= 2 {
return tour.route;
}
for i in 0..tour.len() - 1 {
let nearest_neighbour_idx = tour.route[i + 1..].iter()
.copied()
.min_by_key(|&j| tour.edge_length(j))
.unwrap();
tour.route.swap(i + 1, nearest_neighbour_idx);
}
let start_index_of_longest_edge = (0..tour.len())
.max_by_key(|&a| tour.edge_length(a))
.unwrap();
let end_index_of_longest_edge = tour.next_index(start_index_of_longest_edge);
tour.route.rotate_left(end_index_of_longest_edge);
tour.route
}
#[test]
fn test_solve_tsp_nearest_neighbour() {
let points = vec![(2i32, 2i32), (0, 0), (1, 1)];
let tour = solve_tsp_nearest_neighbour(
&points,
&|(ax, ay), (bx, by)| (ax - bx).abs() + (ay - by).abs(),
);
assert_eq!(tour, [1, 2, 0]);
}
#[test]
fn test_solve_tsp() {
let points = vec![(2i32, 2i32), (0, 0), (1, 1)];
let tour = solve_tsp(
&points,
&|(ax, ay), (bx, by)| (ax - bx).abs() + (ay - by).abs(),
);
assert_eq!(tour, [1, 2, 0]);
}