use rusqlite::Connection;
const EARTH_RADIUS_KM: f64 = 6371.0;
pub const CONTINENTS: &[(&str, f64, f64, f64, f64)] = &[
("Antarctica", -90.0, -180.0, -60.0, 180.0),
("Oceania", -50.0, 110.0, 0.0, 180.0),
("Europe", 34.0, -25.0, 72.0, 45.0),
("Asia", -10.0, 25.0, 80.0, 180.0),
("Africa", -37.0, -20.0, 35.0, 52.0),
("North America", 7.0, -170.0, 85.0, -50.0),
("South America", -56.0, -82.0, 13.0, -34.0),
];
pub fn continent_of(lat: f64, lon: f64) -> &'static str {
for (name, south, west, north, east) in CONTINENTS {
if lat >= *south && lat <= *north && lon >= *west && lon <= *east {
return name;
}
}
"Other"
}
pub fn haversine_km(lat1: f64, lon1: f64, lat2: f64, lon2: f64) -> f64 {
let d_lat = (lat2 - lat1).to_radians();
let d_lon = (lon2 - lon1).to_radians();
let lat1r = lat1.to_radians();
let lat2r = lat2.to_radians();
let a = (d_lat / 2.0).sin().powi(2) + lat1r.cos() * lat2r.cos() * (d_lon / 2.0).sin().powi(2);
let c = 2.0 * a.sqrt().atan2((1.0 - a).sqrt());
EARTH_RADIUS_KM * c
}
pub fn register_haversine_sql_function(conn: &Connection) -> rusqlite::Result<()> {
use rusqlite::functions::FunctionFlags;
conn.create_scalar_function(
"haversine_km",
4,
FunctionFlags::SQLITE_UTF8 | FunctionFlags::SQLITE_DETERMINISTIC,
|ctx| {
Ok(haversine_km(
ctx.get(0)?,
ctx.get(1)?,
ctx.get(2)?,
ctx.get(3)?,
))
},
)
}
use std::collections::{BinaryHeap, HashMap};
struct HeapEntry {
dist: f64,
i: usize,
j: usize,
}
impl Eq for HeapEntry {}
impl PartialEq for HeapEntry {
fn eq(&self, other: &Self) -> bool {
self.dist == other.dist
}
}
impl Ord for HeapEntry {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other.dist.total_cmp(&self.dist)
}
}
impl PartialOrd for HeapEntry {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
const CELL_DIVISOR: f64 = 10.0;
pub fn cluster_by_distance(points: &[(f64, f64)], radius_km: f64) -> Vec<Vec<usize>> {
if points.is_empty() {
return Vec::new();
}
let cells = snap_cells(points, radius_km);
let weights: Vec<f64> = cells.members.iter().map(|m| m.len() as f64).collect();
agglomerate(&cells.centroids, &weights, radius_km)
.into_iter()
.map(|group| {
group
.into_iter()
.flat_map(|cell| cells.members[cell].iter().copied())
.collect()
})
.collect()
}
struct Cells {
centroids: Vec<(f64, f64)>,
members: Vec<Vec<usize>>,
}
fn snap_cells(points: &[(f64, f64)], radius_km: f64) -> Cells {
let side_km = radius_km / CELL_DIVISOR;
let height = (side_km / EARTH_RADIUS_KM).to_degrees();
let snap = side_km > 0.0 && height.is_finite();
let mut index: HashMap<(i64, i64), usize> = HashMap::new();
let mut members: Vec<Vec<usize>> = Vec::new();
for (p, &(lat, lon)) in points.iter().enumerate() {
let key = if snap {
let band = (lat / height).floor();
let edge = (band * height).abs().max(((band + 1.0) * height).abs());
let cos = edge.min(90.0).to_radians().cos();
let column = if cos < 1e-6 {
0.0
} else {
(lon / (height / cos)).floor()
};
(band as i64, column as i64)
} else {
(lat.to_bits() as i64, lon.to_bits() as i64)
};
let cell = *index.entry(key).or_insert_with(|| {
members.push(Vec::new());
members.len() - 1
});
members[cell].push(p);
}
let centroids = members.iter().map(|m| centroid(points, m)).collect();
Cells { centroids, members }
}
fn within_radius_pairs(points: &[(f64, f64)], radius_km: f64) -> Vec<(usize, usize, f64)> {
const SLACK_DEG: f64 = 1e-9;
let mut order: Vec<usize> = (0..points.len()).collect();
order.sort_by(|&a, &b| points[a].0.total_cmp(&points[b].0));
let max_dlat = (radius_km / EARTH_RADIUS_KM).to_degrees() + SLACK_DEG;
let whole_circle = radius_km >= std::f64::consts::PI * EARTH_RADIUS_KM;
let half_angle_sin = (radius_km / (2.0 * EARTH_RADIUS_KM)).sin();
let mut pairs = Vec::new();
for (a, &i) in order.iter().enumerate() {
let (lat_i, lon_i) = points[i];
for &j in &order[a + 1..] {
let (lat_j, lon_j) = points[j];
if lat_j - lat_i > max_dlat {
break;
}
let cos = lat_i.abs().max(lat_j.abs()).min(90.0).to_radians().cos();
let max_dlon = if whole_circle || half_angle_sin >= cos {
180.0
} else {
(2.0 * (half_angle_sin / cos).asin()).to_degrees() + SLACK_DEG
};
let mut dlon = (lon_j - lon_i).abs() % 360.0;
if dlon > 180.0 {
dlon = 360.0 - dlon;
}
if dlon > max_dlon {
continue;
}
let d = haversine_km(lat_i, lon_i, lat_j, lon_j);
if d <= radius_km {
pairs.push((i.min(j), i.max(j), d));
}
}
}
pairs
}
fn mean_distance(points: &[(f64, f64)], weights: &[f64], a: &[usize], b: &[usize]) -> f64 {
let mut sum = 0.0;
let mut weight = 0.0;
for &p in a {
for &q in b {
let w = weights[p] * weights[q];
sum += w * haversine_km(points[p].0, points[p].1, points[q].0, points[q].1);
weight += w;
}
}
sum / weight
}
fn agglomerate(points: &[(f64, f64)], weights: &[f64], radius_km: f64) -> Vec<Vec<usize>> {
let n = points.len();
let mut rows: Vec<Option<HashMap<usize, f64>>> = (0..n).map(|_| Some(HashMap::new())).collect();
let mut heap: BinaryHeap<HeapEntry> = BinaryHeap::new();
for (i, j, d) in within_radius_pairs(points, radius_km) {
rows[i].as_mut().expect("alive").insert(j, d);
rows[j].as_mut().expect("alive").insert(i, d);
heap.push(HeapEntry { dist: d, i, j });
}
let mut members: Vec<Vec<usize>> = (0..n).map(|p| vec![p]).collect();
let mut sizes: Vec<f64> = weights.to_vec();
while let Some(HeapEntry { dist: d, i, j }) = heap.pop() {
if rows[i].as_ref().and_then(|row| row.get(&j)) != Some(&d) {
continue;
}
let row_i = rows[i].take().expect("alive");
let row_j = rows[j].take().expect("alive");
let (size_i, size_j) = (sizes[i], sizes[j]);
let mut others: Vec<usize> = row_i
.keys()
.chain(row_j.keys())
.copied()
.filter(|&k| k != i && k != j)
.collect();
others.sort_unstable();
others.dedup();
let mut merged: HashMap<usize, f64> = HashMap::new();
for k in others {
let d_ik = row_i
.get(&k)
.copied()
.unwrap_or_else(|| mean_distance(points, weights, &members[i], &members[k]));
let d_jk = row_j
.get(&k)
.copied()
.unwrap_or_else(|| mean_distance(points, weights, &members[j], &members[k]));
let new_d = (size_i * d_ik + size_j * d_jk) / (size_i + size_j);
let row_k = rows[k].as_mut().expect("a neighbour is alive");
row_k.remove(&i);
row_k.remove(&j);
if new_d <= radius_km {
row_k.insert(i, new_d);
merged.insert(k, new_d);
heap.push(HeapEntry {
dist: new_d,
i: i.min(k),
j: i.max(k),
});
}
}
let moved = std::mem::take(&mut members[j]);
members[i].extend(moved);
sizes[i] = size_i + size_j;
rows[i] = Some(merged);
}
(0..n)
.filter(|&r| rows[r].is_some())
.map(|r| std::mem::take(&mut members[r]))
.collect()
}
pub fn centroid(points: &[(f64, f64)], member_idxs: &[usize]) -> (f64, f64) {
let n = member_idxs.len() as f64;
let sum_lat: f64 = member_idxs.iter().map(|&i| points[i].0).sum();
let sum_lon: f64 = member_idxs.iter().map(|&i| points[i].1).sum();
(sum_lat / n, sum_lon / n)
}
pub fn ensure_location_clusters_table(conn: &Connection) -> rusqlite::Result<()> {
conn.execute_batch(
"CREATE TABLE IF NOT EXISTS location_clusters (
id INTEGER PRIMARY KEY,
centroid_lat REAL NOT NULL,
centroid_lon REAL NOT NULL,
name TEXT,
photo_count INTEGER NOT NULL,
radius_km REAL NOT NULL,
created_at TEXT NOT NULL
);",
)
}
pub fn ensure_gps_index(conn: &Connection) {
let _ = conn.execute_batch(
"CREATE INDEX IF NOT EXISTS idx_file_hashes_gps ON file_hashes(gps_lat, gps_lon)",
);
}
pub const LOCATIONS_GPS_FINGERPRINT: &str = "locations_gps_fingerprint";
pub const LOCATIONS_RADIUS: &str = "locations_radius";
pub const DEFAULT_CLUSTER_RADIUS_KM: f64 = 15.0;
pub fn gps_fingerprint(conn: &Connection) -> rusqlite::Result<String> {
let rows: i64 = conn.query_row(
"SELECT COUNT(*) FROM file_hashes
WHERE gps_lat IS NOT NULL AND gps_lon IS NOT NULL",
[],
|r| r.get(0),
)?;
let mut stmt = conn.prepare(
"SELECT gps_lat, gps_lon FROM file_hashes
WHERE gps_lat IS NOT NULL AND gps_lon IS NOT NULL
ORDER BY gps_lat, gps_lon",
)?;
let mut distinct = 0i64;
let mut sum_lat = 0.0f64;
let mut sum_lon = 0.0f64;
let mut last: Option<(f64, f64)> = None;
let mut iter = stmt.query_map([], |r| Ok((r.get::<_, f64>(0)?, r.get::<_, f64>(1)?)))?;
while let Some((lat, lon)) = iter.next().transpose()? {
if last != Some((lat, lon)) {
distinct += 1;
}
sum_lat += lat;
sum_lon += lon;
last = Some((lat, lon));
}
Ok(format!("v1:{rows}:{distinct}:{sum_lat}:{sum_lon}"))
}
pub struct RecomputedCluster {
pub id: i64,
pub name: Option<String>,
pub centroid_lat: f64,
pub centroid_lon: f64,
pub photo_count: i64,
}
pub fn recompute_all(
conn: &Connection,
cache: &crate::library::CachePaths,
radius_km: f64,
quiet: bool,
) -> anyhow::Result<Vec<RecomputedCluster>> {
let tx = conn.unchecked_transaction()?;
ensure_location_clusters_table(&tx)?;
ensure_gps_index(&tx);
let coords: Vec<(f64, f64)> = {
let mut stmt = tx.prepare(
"SELECT DISTINCT gps_lat, gps_lon FROM file_hashes \
WHERE gps_lat IS NOT NULL AND gps_lon IS NOT NULL",
)?;
let rows = stmt
.query_map([], |r| Ok((r.get::<_, f64>(0)?, r.get::<_, f64>(1)?)))?
.collect::<rusqlite::Result<Vec<_>>>()?;
rows
};
tx.execute(
"UPDATE file_hashes SET location_cluster_id = NULL WHERE location_cluster_id IS NOT NULL",
[],
)?;
tx.execute("DELETE FROM location_clusters", [])?;
if coords.is_empty() {
tx.commit()?;
store_recompute_state(conn, radius_km)?;
return Ok(Vec::new());
}
if !quiet {
tracing::info!(
"Clustering {} distinct coordinate(s) at radius {}km...",
coords.len(),
radius_km
);
}
let member_groups = cluster_by_distance(&coords, radius_km);
if !quiet {
tracing::info!(
"{} cluster(s); naming them and assigning photos",
member_groups.len()
);
}
let progress =
crate::progress::Progress::new_counting(coords.len() as u64, quiet, "coordinates");
let mut clusters = Vec::with_capacity(member_groups.len());
for members in &member_groups {
let (centroid_lat, centroid_lon) = centroid(&coords, members);
let name = crate::location::location_name_in(cache, centroid_lat, centroid_lon)?;
tx.execute(
"INSERT INTO location_clusters \
(centroid_lat, centroid_lon, name, photo_count, radius_km, created_at) \
VALUES (?1, ?2, ?3, ?4, ?5, datetime('now'))",
rusqlite::params![centroid_lat, centroid_lon, name, 0i64, radius_km],
)?;
let id = tx.last_insert_rowid();
let mut photo_count = 0i64;
for &idx in members {
let (lat, lon) = coords[idx];
let affected = tx.execute(
"UPDATE file_hashes SET location_cluster_id = ?1 \
WHERE gps_lat = ?2 AND gps_lon = ?3",
rusqlite::params![id, lat, lon],
)?;
photo_count += affected as i64;
progress.tick();
}
tx.execute(
"UPDATE location_clusters SET photo_count = ?1 WHERE id = ?2",
rusqlite::params![photo_count, id],
)?;
clusters.push(RecomputedCluster {
id,
name,
centroid_lat,
centroid_lon,
photo_count,
});
}
progress.finish();
clusters.sort_by_key(|c| std::cmp::Reverse(c.photo_count));
tx.commit()?;
store_recompute_state(conn, radius_km)?;
Ok(clusters)
}
fn store_recompute_state(conn: &Connection, radius_km: f64) -> anyhow::Result<()> {
let fingerprint = gps_fingerprint(conn)?;
crate::library_state::set_string(conn, LOCATIONS_GPS_FINGERPRINT, &fingerprint)?;
crate::library_state::set_string(conn, LOCATIONS_RADIUS, &format!("{radius_km}"))?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gps_fingerprint_is_stable_for_identical_data() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (path TEXT PRIMARY KEY, gps_lat REAL, gps_lon REAL);
INSERT INTO file_hashes VALUES ('a', 52.5, 13.4), ('b', 52.5, 13.4), ('c', 48.8, 2.35);",
)
.unwrap();
let first = gps_fingerprint(&conn).unwrap();
let second = gps_fingerprint(&conn).unwrap();
assert_eq!(first, second, "identical data, identical fingerprint");
}
#[test]
fn gps_fingerprint_changes_when_gps_data_changes() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (path TEXT PRIMARY KEY, gps_lat REAL, gps_lon REAL);
INSERT INTO file_hashes VALUES ('a', 52.5, 13.4);",
)
.unwrap();
let before = gps_fingerprint(&conn).unwrap();
conn.execute("UPDATE file_hashes SET gps_lon = 14.0", [])
.unwrap();
let moved = gps_fingerprint(&conn).unwrap();
assert_ne!(before, moved, "a moved coordinate is a change");
conn.execute("INSERT INTO file_hashes VALUES ('b', 48.8, 2.35)", [])
.unwrap();
let added = gps_fingerprint(&conn).unwrap();
assert_ne!(moved, added, "a new coordinate is a change");
conn.execute("DELETE FROM file_hashes WHERE path = 'b'", [])
.unwrap();
conn.execute("UPDATE file_hashes SET gps_lon = 13.4", [])
.unwrap();
let shrunk = gps_fingerprint(&conn).unwrap();
assert_ne!(added, shrunk, "losing a row is a change");
}
#[test]
fn gps_fingerprint_on_a_library_without_gps_is_a_known_value() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (path TEXT PRIMARY KEY, gps_lat REAL, gps_lon REAL);",
)
.unwrap();
let fp = gps_fingerprint(&conn).unwrap();
assert_eq!(
fp, "v1:0:0:0:0",
"document the empty shape: rows:distinct:sumlat:sumlon"
);
}
fn temp_cache() -> crate::library::CachePaths {
let dir = tempfile::tempdir().unwrap();
let base = dir.path().to_path_buf();
std::mem::forget(dir);
crate::library::CachePaths {
thumbnails: base.join("thumbnails"),
geo: base.join("geo"),
base,
}
}
#[test]
fn recompute_all_clusters_assigns_and_counts() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (path TEXT PRIMARY KEY, gps_lat REAL, gps_lon REAL,
location_cluster_id INTEGER);
INSERT INTO file_hashes (path, gps_lat, gps_lon) VALUES
('a', 52.5, 13.4), ('b', 52.5001, 13.4001), ('c', -33.87, 151.21);",
)
.unwrap();
let cache = temp_cache();
let clusters = recompute_all(&conn, &cache, 15.0, true).unwrap();
assert_eq!(clusters.len(), 2, "{:?}", clusters.len());
let total: i64 = clusters.iter().map(|c| c.photo_count).sum();
assert_eq!(total, 3, "every photo lands in exactly one cluster");
let assigned: i64 = conn
.query_row(
"SELECT COUNT(*) FROM file_hashes WHERE location_cluster_id IS NOT NULL",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(assigned, 3);
}
#[test]
fn recompute_all_on_empty_gps_wipes_cleanly() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (path TEXT PRIMARY KEY, gps_lat REAL, gps_lon REAL,
location_cluster_id INTEGER);",
)
.unwrap();
let cache = temp_cache();
ensure_location_clusters_table(&conn).unwrap();
conn.execute(
"INSERT INTO location_clusters \
(centroid_lat, centroid_lon, name, photo_count, radius_km, created_at) \
VALUES (52.5, 13.4, 'ghost', 1, 15.0, datetime('now'))",
[],
)
.unwrap();
recompute_all(&conn, &cache, 15.0, true).unwrap();
let clusters: i64 = conn
.query_row("SELECT COUNT(*) FROM location_clusters", [], |r| r.get(0))
.unwrap();
assert_eq!(clusters, 0, "the stale cluster must be wiped");
}
#[test]
fn the_gps_index_exists_and_serves_an_exact_match() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (
path TEXT PRIMARY KEY, hash TEXT NOT NULL, gps_lat REAL, gps_lon REAL);",
)
.unwrap();
ensure_gps_index(&conn);
ensure_gps_index(&conn);
let exact: String = conn
.query_row(
"EXPLAIN QUERY PLAN SELECT 1 FROM file_hashes WHERE gps_lat = 1.0 AND gps_lon = 2.0",
[],
|r| r.get(3),
)
.unwrap();
assert!(
exact.contains("idx_file_hashes_gps"),
"an exact match must use the index, got: {exact}"
);
let rounded: String = conn
.query_row(
"EXPLAIN QUERY PLAN SELECT 1 FROM file_hashes \
WHERE ROUND(gps_lat, 6) = ROUND(1.0, 6) AND ROUND(gps_lon, 6) = ROUND(2.0, 6)",
[],
|r| r.get(3),
)
.unwrap();
assert!(
rounded.contains("SCAN"),
"ROUND() on the column must still scan - this is why it was removed: {rounded}"
);
}
#[test]
fn two_coordinates_that_round_alike_stay_separate() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (
path TEXT PRIMARY KEY, hash TEXT NOT NULL, gps_lat REAL, gps_lon REAL);
INSERT INTO file_hashes VALUES
('/a.jpg','a', 52.55360000001, 13.43),
('/b.jpg','b', 52.55360000002, 13.43);",
)
.unwrap();
let exact: i64 = conn
.query_row(
"SELECT COUNT(*) FROM file_hashes WHERE gps_lat = 52.55360000001 AND gps_lon = 13.43",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(exact, 1, "exact matching claims only its own row");
let rounded: i64 = conn
.query_row(
"SELECT COUNT(*) FROM file_hashes \
WHERE ROUND(gps_lat, 6) = ROUND(52.55360000001, 6) AND ROUND(gps_lon, 6) = ROUND(13.43, 6)",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(rounded, 2, "ROUND claims both - the double-count");
}
#[test]
fn haversine_zero_distance_for_identical_points() {
let d = haversine_km(48.8566, 2.3522, 48.8566, 2.3522);
assert!(d.abs() < 1e-9, "expected ~0, got {d}");
}
#[test]
fn haversine_one_degree_of_latitude_is_about_111_km() {
let d = haversine_km(0.0, 0.0, 1.0, 0.0);
assert!((d - 111.19).abs() < 0.5, "expected ~111.19km, got {d}");
}
#[test]
fn haversine_paris_to_london_is_about_343_km() {
let d = haversine_km(48.8566, 2.3522, 51.5074, -0.1278);
assert!((d - 343.0).abs() < 5.0, "expected ~343km, got {d}");
}
#[test]
fn registered_haversine_sql_function_matches_rust_distance() {
let conn = Connection::open_in_memory().unwrap();
register_haversine_sql_function(&conn).unwrap();
let sql_distance: f64 = conn
.query_row(
"SELECT haversine_km(?1, ?2, ?3, ?4)",
rusqlite::params![52.52, 13.405, 48.8566, 2.3522],
|row| row.get(0),
)
.unwrap();
let rust_distance = haversine_km(52.52, 13.405, 48.8566, 2.3522);
assert!((sql_distance - rust_distance).abs() < 1e-9);
}
#[test]
fn cluster_by_distance_empty_input_returns_empty() {
assert!(cluster_by_distance(&[], 15.0).is_empty());
}
#[test]
fn cluster_by_distance_single_point_is_its_own_cluster() {
let clusters = cluster_by_distance(&[(48.8566, 2.3522)], 15.0);
assert_eq!(clusters, vec![vec![0]]);
}
#[test]
fn cluster_by_distance_groups_nearby_points_and_isolates_far_ones() {
let points = vec![
(48.8566, 2.3522), (48.8606, 2.3376), (48.8530, 2.3499), (51.5074, -0.1278), ];
let clusters = cluster_by_distance(&points, 50.0);
assert_eq!(clusters.len(), 2, "expected 2 clusters, got {clusters:?}");
let mut sizes: Vec<usize> = clusters.iter().map(|c| c.len()).collect();
sizes.sort();
assert_eq!(sizes, vec![1, 3]);
}
#[test]
fn cluster_by_distance_all_points_merge_when_radius_is_huge() {
let points = vec![(48.8566, 2.3522), (51.5074, -0.1278)];
let clusters = cluster_by_distance(&points, 10_000.0);
assert_eq!(clusters.len(), 1);
}
fn dense_reference(points: &[(f64, f64)], radius_km: f64) -> Vec<Vec<usize>> {
let n = points.len();
let mut dist: Vec<Vec<f64>> = vec![vec![0.0; n]; n];
let mut heap: BinaryHeap<HeapEntry> = BinaryHeap::new();
for i in 0..n {
for j in (i + 1)..n {
let d = haversine_km(points[i].0, points[i].1, points[j].0, points[j].1);
dist[i][j] = d;
dist[j][i] = d;
if d <= radius_km {
heap.push(HeapEntry { dist: d, i, j });
}
}
}
let mut members: Vec<Vec<usize>> = (0..n).map(|i| vec![i]).collect();
let mut alive = vec![true; n];
while let Some(HeapEntry { dist: d, i, j }) = heap.pop() {
if !alive[i] || !alive[j] {
continue;
}
if dist[i][j] != d {
continue;
}
if d > radius_km {
break;
}
let size_i = members[i].len() as f64;
let size_j = members[j].len() as f64;
let moved = std::mem::take(&mut members[j]);
members[i].extend(moved);
alive[j] = false;
for k in 0..n {
if k == i || k == j || !alive[k] {
continue;
}
let new_d = (size_i * dist[i][k] + size_j * dist[j][k]) / (size_i + size_j);
if new_d != dist[i][k] {
dist[i][k] = new_d;
dist[k][i] = new_d;
heap.push(HeapEntry {
dist: new_d,
i: i.min(k),
j: i.max(k),
});
}
}
}
(0..n)
.filter(|&r| alive[r])
.map(|r| std::mem::take(&mut members[r]))
.collect()
}
fn canonical(mut partition: Vec<Vec<usize>>) -> Vec<Vec<usize>> {
for group in &mut partition {
group.sort_unstable();
}
partition.sort();
partition
}
fn noise(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*state >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0
}
fn gps_fixture(seed: u64) -> Vec<(f64, f64)> {
let mut state = seed;
let mut points = Vec::new();
for (lat, lon, count) in [(41.01, 28.98, 60), (52.52, 13.405, 40), (69.65, 18.96, 20)] {
for _ in 0..count {
points.push((lat + noise(&mut state) * 0.2, lon + noise(&mut state) * 0.3));
}
}
for step in 0..12 {
points.push((38.0 + noise(&mut state) * 0.001, 27.0 + step as f64 * 0.103));
}
points.push((-16.5, 179.99));
points.push((-16.5, -179.99));
points.extend([(84.0, 10.0), (84.0, 11.5), (89.9, 0.0), (89.9, 180.0)]);
points
}
#[test]
fn sparse_linkage_matches_the_dense_matrix() {
for seed in [1, 2, 3] {
let points = gps_fixture(seed);
let weights = vec![1.0; points.len()];
for radius in [1.0, 15.0, 50.0, 10_000.0] {
assert_eq!(
canonical(agglomerate(&points, &weights, radius)),
canonical(dense_reference(&points, radius)),
"seed {seed}, radius {radius}km"
);
}
}
}
#[test]
fn a_chain_is_not_strung_into_one_cluster() {
let chain: Vec<(f64, f64)> = (0..12).map(|s| (38.0, 27.0 + s as f64 * 0.103)).collect();
assert!(cluster_by_distance(&chain, 15.0).len() > 1);
}
#[test]
fn within_radius_pairs_finds_exactly_the_brute_force_pairs() {
let points = gps_fixture(7);
for radius in [1.0, 15.0, 500.0] {
let mut expected = Vec::new();
for i in 0..points.len() {
for j in (i + 1)..points.len() {
let d = haversine_km(points[i].0, points[i].1, points[j].0, points[j].1);
if d <= radius {
expected.push((i, j));
}
}
}
let mut found: Vec<(usize, usize)> = within_radius_pairs(&points, radius)
.into_iter()
.map(|(i, j, _)| (i, j))
.collect();
found.sort_unstable();
assert_eq!(found, expected, "radius {radius}km");
}
}
#[test]
fn nearby_coordinates_share_a_cell_and_distant_ones_do_not() {
let cells = snap_cells(&[(41.0, 29.0), (41.00009, 29.0), (41.018, 29.0)], 15.0);
let cell_of = |p: usize| cells.members.iter().position(|m| m.contains(&p));
assert_eq!(cell_of(0), cell_of(1));
assert_ne!(cell_of(0), cell_of(2));
}
#[test]
fn a_cell_stays_within_about_its_side() {
for lat in [0.0, 45.0, 70.0] {
let points: Vec<(f64, f64)> = (0..200)
.map(|k| {
(
lat + (k % 20) as f64 * 0.0007,
10.0 + (k / 20) as f64 * 0.0011,
)
})
.collect();
let cells = snap_cells(&points, 15.0);
for (cell, members) in cells.members.iter().enumerate() {
let (clat, clon) = cells.centroids[cell];
for &p in members {
let d = haversine_km(points[p].0, points[p].1, clat, clon);
let side = 15.0 / CELL_DIVISOR;
assert!(d <= side * 1.5, "{d}km from its cell at latitude {lat}");
}
}
}
}
#[test]
fn a_zero_radius_snaps_nothing() {
let cells = snap_cells(&[(1.0, 1.0), (1.0, 1.0), (1.0, 1.0000001)], 0.0);
assert_eq!(cells.members, vec![vec![0, 1], vec![2]]);
assert_eq!(
cluster_by_distance(&[(1.0, 1.0), (1.0, 1.0000001)], 0.0).len(),
2
);
}
#[test]
#[ignore]
fn clusters_a_large_library_quickly() {
let mut state = 42;
let mut points = Vec::new();
for (lat, lon) in [(41.01, 28.98), (52.52, 13.405), (40.71, -74.0)] {
for _ in 0..5_833 {
points.push((
lat + noise(&mut state) * 0.18,
lon + noise(&mut state) * 0.24,
));
}
}
for _ in 0..7_500 {
points.push((noise(&mut state) * 70.0, noise(&mut state) * 180.0));
}
let started = std::time::Instant::now();
let clusters = cluster_by_distance(&points, DEFAULT_CLUSTER_RADIUS_KM);
eprintln!(
"{} coordinates -> {} clusters in {:?}",
points.len(),
clusters.len(),
started.elapsed()
);
}
#[test]
fn centroid_is_unweighted_mean() {
let points = vec![(0.0, 0.0), (2.0, 4.0)];
let (lat, lon) = centroid(&points, &[0, 1]);
assert!((lat - 1.0).abs() < 1e-9);
assert!((lon - 2.0).abs() < 1e-9);
}
#[test]
fn continent_of_assigns_representative_points() {
let cases = [
(52.52, 13.405, "Europe"),
(35.68, 139.69, "Asia"),
(-33.87, 151.21, "Oceania"),
(40.71, -74.01, "North America"),
(-23.55, -46.63, "South America"),
(6.45, 3.39, "Africa"),
(-75.0, 0.0, "Antarctica"),
(55.75, 37.62, "Europe"),
(-41.29, 174.78, "Oceania"),
];
for (lat, lon, expected) in cases {
assert_eq!(continent_of(lat, lon), expected, "at {lat},{lon}");
}
}
#[test]
fn continent_of_names_the_ocean_fallback_for_unmatched_points() {
assert_eq!(continent_of(0.0, -30.0), "Other");
}
#[test]
fn ensure_location_clusters_table_is_idempotent() {
let conn = Connection::open_in_memory().unwrap();
ensure_location_clusters_table(&conn).unwrap();
ensure_location_clusters_table(&conn).unwrap(); }
#[test]
fn recompute_clears_foreign_key_children_first() {
let conn = Connection::open_in_memory().unwrap();
conn.execute_batch("PRAGMA foreign_keys = ON").unwrap();
ensure_location_clusters_table(&conn).unwrap();
conn.execute_batch(
"CREATE TABLE file_hashes (
path TEXT PRIMARY KEY, hash TEXT NOT NULL, gps_lat REAL, gps_lon REAL,
location_cluster_id INTEGER REFERENCES location_clusters(id)
ON DELETE RESTRICT ON UPDATE RESTRICT
);
INSERT INTO location_clusters (id, centroid_lat, centroid_lon, name, photo_count, radius_km, created_at)
VALUES (1, 52.52, 13.40, 'berlin', 1, 25.0, datetime('now'));
INSERT INTO file_hashes (path, hash, gps_lat, gps_lon, location_cluster_id)
VALUES ('/p/x.jpg', 'x', 52.51, 13.39, 1);",
)
.unwrap();
let cache = temp_cache();
recompute_all(&conn, &cache, 15.0, true)
.expect("recompute must clear child references before deleting the parent clusters");
let clusters: i64 = conn
.query_row("SELECT COUNT(*) FROM location_clusters", [], |r| r.get(0))
.unwrap();
let refs: i64 = conn
.query_row(
"SELECT COUNT(*) FROM file_hashes WHERE location_cluster_id IS NOT NULL",
[],
|r| r.get(0),
)
.unwrap();
assert_eq!(clusters, 1, "the located shot regains a fresh cluster");
assert_eq!(refs, 1, "its reference points at the fresh cluster");
let violations: i64 = conn
.query_row("SELECT COUNT(*) FROM pragma_foreign_key_check", [], |r| {
r.get(0)
})
.unwrap();
assert_eq!(violations, 0);
}
}