use yo_common::num::DIGITS_MAX;
use yo_common::{Code, Error, Result};
use crate::db::Db;
use crate::elem::Elements;
use crate::geo::{self, Kind, Shape, Unit};
use crate::keyspace::Keyspace;
use crate::strings;
use crate::zset::{Bound, Zset};
use crate::zsets::{ZAdd, member_bytes};
const NO_MEMBER: &str = "could not decode requested zset member";
#[must_use]
pub fn out_of_range(lon: f64, lat: f64) -> Error {
yo_alloc::allow(|| {
Error::fmt(
Code::Invalid,
format_args!("invalid longitude,latitude pair {lon:.6},{lat:.6}"),
)
})
}
#[must_use]
pub fn no_member() -> Error {
Error::new(Code::Invalid, NO_MEMBER)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Sort {
Near,
Far,
}
#[derive(Debug, Clone, Copy, Default)]
pub struct Limit {
pub sort: Option<Sort>,
pub count: Option<usize>,
pub any: bool,
}
impl Limit {
#[must_use]
fn ordering(&self) -> Option<Sort> {
match self.sort {
Some(s) => Some(s),
None if self.count.is_some() && !self.any => Some(Sort::Near),
None => None,
}
}
#[must_use]
fn cap(&self) -> Option<usize> {
self.any.then_some(self.count).flatten()
}
}
#[derive(Debug, Clone, Copy)]
pub struct Hit {
at: usize,
len: usize,
pub score: u64,
pub lon: f64,
pub lat: f64,
pub metres: f64,
}
#[derive(Debug, Default)]
pub struct Scratch {
hits: Vec<Hit>,
names: Vec<u8>,
}
impl Scratch {
pub fn iter(&self) -> impl Iterator<Item = (&[u8], &Hit)> {
self.hits
.iter()
.map(|h| (&self.names[h.at..h.at + h.len], h))
}
#[must_use]
pub fn len(&self) -> usize {
self.hits.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.hits.is_empty()
}
fn clear(&mut self) {
self.hits.clear();
self.names.clear();
}
fn push(&mut self, name: &[u8], score: u64, lon: f64, lat: f64, metres: f64) {
let at = self.names.len();
self.names.extend_from_slice(name);
self.hits.push(Hit {
at,
len: name.len(),
score,
lon,
lat,
metres,
});
}
fn order(&mut self, limit: Limit) {
let Some(sort) = limit.ordering() else {
self.hits.truncate(limit.count.unwrap_or(usize::MAX));
return;
};
let near = |a: &Hit, b: &Hit| a.metres.total_cmp(&b.metres);
let far = |a: &Hit, b: &Hit| b.metres.total_cmp(&a.metres);
let want = limit.count.unwrap_or(self.hits.len()).min(self.hits.len());
if want < self.hits.len() {
match sort {
Sort::Near => self.hits.select_nth_unstable_by(want, near),
Sort::Far => self.hits.select_nth_unstable_by(want, far),
};
self.hits.truncate(want);
}
match sort {
Sort::Near => self.hits.sort_unstable_by(near),
Sort::Far => self.hits.sort_unstable_by(far),
}
}
}
impl Keyspace {
pub fn geoadd<'m, I>(&mut self, key: &[u8], points: I, opts: ZAdd) -> Result<usize>
where
I: Iterator<Item = (f64, f64, &'m [u8])> + Clone,
{
for (lon, lat, member) in points.clone() {
strings::check_len(key, member.len())?;
if geo::score(lon, lat).is_none() {
return Err(out_of_range(lon, lat));
}
}
self.zadd(
key,
points.map(|(lon, lat, m)| {
(geo::score(lon, lat).expect("checked") as f64, m)
}),
opts,
)
}
pub fn geopos<'m, F>(
&mut self,
key: &[u8],
members: impl Iterator<Item = &'m [u8]>,
mut f: F,
) -> Result<()>
where
F: FnMut(Option<(f64, f64)>),
{
let Some(at) = self.zset_slot(key)? else {
members.for_each(|_| f(None));
return Ok(());
};
let z = self.zset_at(at);
for m in members {
f(z.score(m).and_then(geo::decode));
}
Ok(())
}
pub fn geohash<'m, F>(
&mut self,
key: &[u8],
members: impl Iterator<Item = &'m [u8]>,
mut f: F,
) -> Result<()>
where
F: FnMut(Option<&[u8]>),
{
let Some(at) = self.zset_slot(key)? else {
members.for_each(|_| f(None));
return Ok(());
};
let z = self.zset_at(at);
for m in members {
let text = z
.score(m)
.and_then(geo::decode)
.and_then(|(lon, lat)| geo::geohash(lon, lat));
match text {
Some(bytes) => f(Some(&bytes)),
None => f(None),
}
}
Ok(())
}
pub fn geodist(&mut self, key: &[u8], a: &[u8], b: &[u8]) -> Result<Option<f64>> {
let Some(at) = self.zset_slot(key)? else {
return Ok(None);
};
let z = self.zset_at(at);
let (Some(sa), Some(sb)) = (z.score(a), z.score(b)) else {
return Ok(None);
};
let (Some(pa), Some(pb)) = (geo::decode(sa), geo::decode(sb)) else {
return Ok(None);
};
Ok(Some(geo::distance(pa.0, pa.1, pb.0, pb.1)))
}
pub fn geocentre(&mut self, key: &[u8], member: &[u8]) -> Result<Option<(f64, f64)>> {
let Some(at) = self.zset_slot(key)? else {
return Ok(None);
};
match self.zset_at(at).score(member).and_then(geo::decode) {
Some(xy) => Ok(Some(xy)),
None => Err(no_member()),
}
}
pub fn geosearch(&mut self, key: &[u8], shape: &Shape, limit: Limit) -> Result<usize> {
let mut found = std::mem::take(&mut self.geo);
found.clear();
let outcome = match self.zset_slot(key) {
Err(e) => Err(e),
Ok(None) => Ok(()),
Ok(Some(at)) => {
collect(self.zset_at(at), shape, limit, &mut found);
Ok(())
}
};
found.order(limit);
let n = found.hits.len();
self.geo = found;
outcome.map(|()| n)
}
#[must_use]
pub fn geohits(&self) -> &Scratch {
&self.geo
}
pub fn geosearchstore(
&mut self,
dest: &[u8],
src: &[u8],
shape: &Shape,
limit: Limit,
dist: bool,
) -> Result<usize> {
let n = self.geosearch(src, shape, limit)?;
let found = std::mem::take(&mut self.geo);
let mut got = Elements::with_capacity(n.max(16));
for (name, hit) in found.iter() {
let score = if dist {
hit.metres / shape.unit.metres()
} else {
hit.score as f64
};
let _ = got.insert(name, score);
}
self.geo = found;
let limits = self.zset_limits;
let built = Zset::from_elements(got, &limits);
Ok(self.put_zset(dest, built))
}
}
impl Db {
pub fn geosearchstore(
&mut self,
dest: &[u8],
src: &[u8],
shape: &Shape,
limit: Limit,
dist: bool,
) -> Result<usize> {
let (home, onto) = (self.stripe_of(src), self.stripe_of(dest));
if home == onto {
return self
.stripe_mut(home)
.geosearchstore(dest, src, shape, limit, dist);
}
let n = self.stripe_mut(home).geosearch(src, shape, limit)?;
let mut got = Elements::with_capacity(n.max(16));
for (name, hit) in self.stripe(home).geohits().iter() {
let score = if dist {
hit.metres / shape.unit.metres()
} else {
hit.score as f64
};
let _ = got.insert(name, score);
}
let limits = self.stripe(onto).zset_limits;
let built = Zset::from_elements(got, &limits);
Ok(self.stripe_mut(onto).put_zset(dest, built))
}
}
fn collect(z: &Zset, shape: &Shape, limit: Limit, out: &mut Scratch) {
let search = geo::areas(shape);
let cap = limit.cap();
let mut digits = [0u8; DIGITS_MAX];
let mut last = 0usize;
for i in 0..search.boxes.len() {
let hash = search.boxes[i];
if hash.bits == 0 && hash.step == 0 {
continue;
}
if last != 0 && hash == search.boxes[last] {
continue;
}
if cap.is_some_and(|n| out.hits.len() >= n) {
break;
}
let (low, high) = geo::range(hash);
let window = z.window_by_score(Bound::closed(low as f64), Bound::open(high as f64));
z.walk(window.start, window.len(), false, |m, raw| {
if cap.is_some_and(|n| out.hits.len() >= n) {
return;
}
let Some((lon, lat)) = geo::decode(raw) else {
return;
};
let Some(metres) = shape.covers(lon, lat) else {
return;
};
out.push(member_bytes(m, &mut digits), raw as u64, lon, lat, metres);
});
last = i;
}
}
#[must_use]
pub fn circle(lon: f64, lat: f64, radius: f64, unit: Unit) -> Shape {
Shape {
lon,
lat,
kind: Kind::Circle { radius },
unit,
}
}
#[cfg(test)]
mod tests {
use super::*;
const PLACES: [(f64, f64, &[u8]); 4] = [
(13.361389, 38.115556, b"Palermo"),
(15.087269, 37.502669, b"Catania"),
(12.758489, 38.788135, b"edge"),
(2.352222, 48.856613, b"Paris"),
];
fn ks() -> Keyspace {
let mut db = Keyspace::new();
let opts = ZAdd::default();
db.geoadd(b"g", PLACES.iter().copied(), opts)
.expect("the places are all in range");
db
}
fn names(db: &Keyspace) -> Vec<Vec<u8>> {
db.geohits().iter().map(|(n, _)| n.to_vec()).collect()
}
#[test]
fn a_geo_key_is_a_sorted_set_of_hashes() {
let mut db = ks();
assert_eq!(db.zcard(b"g").expect("a zset"), 4);
assert_eq!(
db.zscore(b"g", b"Palermo").expect("a zset"),
Some(3_479_099_956_230_698.0)
);
assert_eq!(
db.zscore(b"g", b"Catania").expect("a zset"),
Some(3_479_447_370_796_909.0)
);
}
#[test]
fn nothing_is_stored_when_one_pair_is_out_of_range() {
let mut db = Keyspace::new();
let bad: [(f64, f64, &[u8]); 2] = [(13.0, 38.0, b"good"), (13.0, 86.0, b"bad")];
let err = db
.geoadd(b"g", bad.iter().copied(), ZAdd::default())
.expect_err("86 is past the projection");
assert_eq!(
err.message(),
"invalid longitude,latitude pair 13.000000,86.000000"
);
assert_eq!(db.zcard(b"g").expect("a zset"), 0);
}
#[test]
fn a_position_comes_back_where_it_went_in_give_or_take_two_metres() {
let mut db = ks();
let mut got = Vec::new();
db.geopos(b"g", [&b"Palermo"[..], b"nope"].into_iter(), |p| {
got.push(p)
})
.expect("a zset");
let (lon, lat) = got[0].expect("Palermo is there");
assert_eq!(format!("{lon}"), "13.361389338970184");
assert_eq!(format!("{lat}"), "38.1155563954963");
assert_eq!(got[1], None);
}
#[test]
fn the_distance_between_two_members_is_the_one_a_real_server_answers() {
let mut db = ks();
let d = db
.geodist(b"g", b"Palermo", b"Catania")
.expect("a zset")
.expect("both are there");
assert_eq!(format!("{d:.4}"), "166274.1516");
assert_eq!(format!("{:.4}", d / Unit::Km.metres()), "166.2742");
assert_eq!(db.geodist(b"g", b"Palermo", b"nope").expect("a zset"), None);
assert_eq!(db.geodist(b"nope", b"a", b"b").expect("no key"), None);
}
#[test]
fn a_hash_string_is_eleven_characters_and_ends_in_a_zero() {
let mut db = ks();
let mut got: Vec<Option<Vec<u8>>> = Vec::new();
db.geohash(
b"g",
[&b"Palermo"[..], b"Catania", b"nope"].into_iter(),
|h| {
got.push(h.map(<[u8]>::to_vec));
},
)
.expect("a zset");
assert_eq!(got[0].as_deref(), Some(&b"sqc8b49rny0"[..]));
assert_eq!(got[1].as_deref(), Some(&b"sqdtr74hyu0"[..]));
assert_eq!(got[2], None);
}
#[test]
fn a_radius_search_finds_what_is_inside_it_nearest_first() {
let mut db = ks();
let shape = circle(15.0, 37.0, 200.0, Unit::Km);
let limit = Limit {
sort: Some(Sort::Near),
..Limit::default()
};
assert_eq!(db.geosearch(b"g", &shape, limit).expect("a zset"), 2);
assert_eq!(names(&db), [b"Catania".to_vec(), b"Palermo".to_vec()]);
let hits: Vec<f64> = db.geohits().iter().map(|(_, h)| h.metres).collect();
assert_eq!(format!("{:.4}", hits[0] / 1000.0), "56.4413");
assert_eq!(format!("{:.4}", hits[1] / 1000.0), "190.4424");
}
#[test]
fn a_box_search_reaches_the_corners_a_circle_does_not() {
let mut db = ks();
let shape = Shape {
lon: 13.361389,
lat: 38.115556,
kind: Kind::Rect {
width: 400.0,
height: 400.0,
},
unit: Unit::Km,
};
let limit = Limit {
sort: Some(Sort::Near),
..Limit::default()
};
assert_eq!(db.geosearch(b"g", &shape, limit).expect("a zset"), 3);
assert_eq!(
names(&db),
[b"Palermo".to_vec(), b"edge".to_vec(), b"Catania".to_vec()]
);
}
#[test]
fn a_count_takes_the_nearest_and_desc_takes_the_furthest() {
let mut db = ks();
let shape = circle(15.0, 37.0, 200.0, Unit::Km);
let limit = Limit {
count: Some(1),
..Limit::default()
};
assert_eq!(db.geosearch(b"g", &shape, limit).expect("a zset"), 1);
assert_eq!(names(&db), [b"Catania".to_vec()]);
let limit = Limit {
sort: Some(Sort::Far),
count: Some(1),
..Limit::default()
};
assert_eq!(db.geosearch(b"g", &shape, limit).expect("a zset"), 1);
assert_eq!(names(&db), [b"Palermo".to_vec()]);
}
#[test]
fn any_stops_at_the_count_rather_than_finding_the_nearest() {
let mut db = ks();
let shape = circle(15.0, 37.0, 200.0, Unit::Km);
let limit = Limit {
count: Some(1),
any: true,
..Limit::default()
};
assert_eq!(db.geosearch(b"g", &shape, limit).expect("a zset"), 1);
assert_eq!(db.geohits().len(), 1);
}
#[test]
fn a_search_that_finds_nothing_is_not_an_error() {
let mut db = ks();
let shape = circle(0.0, 0.0, 1.0, Unit::M);
assert_eq!(
db.geosearch(b"g", &shape, Limit::default())
.expect("a zset"),
0
);
assert!(db.geohits().is_empty());
assert_eq!(
db.geosearch(b"nope", &shape, Limit::default())
.expect("no key"),
0
);
}
#[test]
fn a_search_around_a_member_is_a_search_around_where_it_is() {
let mut db = ks();
let (lon, lat) = db
.geocentre(b"g", b"Palermo")
.expect("a zset")
.expect("Palermo is there");
let shape = circle(lon, lat, 200.0, Unit::Km);
let limit = Limit {
sort: Some(Sort::Near),
..Limit::default()
};
assert_eq!(db.geosearch(b"g", &shape, limit).expect("a zset"), 3);
assert_eq!(names(&db)[0], b"Palermo".to_vec());
assert!(db.geocentre(b"g", b"nope").is_err());
assert_eq!(db.geocentre(b"nope", b"nope").expect("no key"), None);
}
#[test]
fn a_store_keeps_the_hashes_and_a_storedist_keeps_the_distances() {
let mut db = ks();
let shape = circle(15.0, 37.0, 200.0, Unit::Km);
let limit = Limit {
sort: Some(Sort::Near),
..Limit::default()
};
assert_eq!(
db.geosearchstore(b"d", b"g", &shape, limit, false)
.expect("a zset"),
2
);
assert_eq!(
db.zscore(b"d", b"Catania").expect("a zset"),
Some(3_479_447_370_796_909.0)
);
assert_eq!(
db.geodist(b"d", b"Palermo", b"Catania")
.expect("a zset")
.map(|d| format!("{d:.4}")),
Some("166274.1516".to_string())
);
assert_eq!(
db.geosearchstore(b"e", b"g", &shape, limit, true)
.expect("a zset"),
2
);
let d = db
.zscore(b"e", b"Catania")
.expect("a zset")
.expect("stored");
assert_eq!(format!("{d:.4}"), "56.4413");
}
#[test]
fn a_store_that_finds_nothing_deletes_what_was_there() {
let mut db = ks();
let shape = circle(15.0, 37.0, 200.0, Unit::Km);
assert_eq!(
db.geosearchstore(b"d", b"g", &shape, Limit::default(), false)
.expect("a zset"),
2
);
let empty = circle(0.0, 0.0, 1.0, Unit::M);
assert_eq!(
db.geosearchstore(b"d", b"g", &empty, Limit::default(), false)
.expect("a zset"),
0
);
assert_eq!(db.zcard(b"d").expect("no key"), 0);
}
#[test]
fn a_search_across_the_date_line_finds_both_sides_of_it() {
let mut db = Keyspace::new();
let pair: [(f64, f64, &[u8]); 2] = [(179.9, 0.0, b"west"), (-179.9, 0.0, b"east")];
db.geoadd(b"d", pair.iter().copied(), ZAdd::default())
.expect("both are in range");
let d = db
.geodist(b"d", b"west", b"east")
.expect("a zset")
.expect("both are there");
assert_eq!(format!("{:.4}", d / Unit::Km.metres()), "22.2454");
let shape = circle(179.95, 0.0, 50.0, Unit::Km);
assert_eq!(
db.geosearch(b"d", &shape, Limit::default())
.expect("a zset"),
2
);
let shape = circle(180.0, 0.0, 50.0, Unit::Km);
assert_eq!(
db.geosearch(b"d", &shape, Limit::default())
.expect("a zset"),
1
);
assert_eq!(names(&db), [b"west".to_vec()]);
}
#[test]
fn a_key_holding_something_else_is_refused_everywhere() {
let mut db = Keyspace::new();
db.set(b"s", b"v", strings::SetOptions::default())
.expect("a fresh key");
let shape = circle(0.0, 0.0, 1.0, Unit::Km);
assert_eq!(
db.geoadd(b"s", PLACES.iter().copied(), ZAdd::default())
.expect_err("a string")
.code(),
Code::WrongType
);
assert!(db.geopos(b"s", [&b"x"[..]].into_iter(), |_| {}).is_err());
assert!(db.geohash(b"s", [&b"x"[..]].into_iter(), |_| {}).is_err());
assert!(db.geodist(b"s", b"a", b"b").is_err());
assert!(db.geosearch(b"s", &shape, Limit::default()).is_err());
assert!(
db.geosearchstore(b"d", b"s", &shape, Limit::default(), false)
.is_err()
);
}
}