pub mod codec;
pub mod conf;
pub use codec::{
BASE32_ALPHABET, BASE32_DECODE_TABLE, D_R, EARTH_RADIUS_METERS, GEO_LAT_MAX, GEO_LAT_MIN,
GEO_LAT_RANGE, GEO_LAT_RANGE_STANDARD, GEO_LON_MAX, GEO_LON_MIN, GEO_LON_RANGE,
GEO_LON_RANGE_STANDARD, GEO_STEP_MAX, MERCATOR_MAX, align_52bits, base32_to_coords,
bounding_box, convert_meters_to_unit, convert_unit_to_meters, coords_to_base32, decode_geohash,
deinterleave64, encode_geohash, encode_geohash_string, estimate_steps_by_radius,
geohash_decode, geohash_decode_area_to_long_lat, geohash_decode_wgs84, geohash_encode,
geohash_encode_wgs84, geohash_move_x, geohash_move_y, geohash_neighbors, geohash_to_base32,
get_areas_by_shape_wgs84, haversine_distance, interleave64, scores_of_geohash_box,
validate_long_lat,
};
pub use conf::{
DistanceSort, DistanceUnit, GeoHashArea, GeoHashBits, GeoHashNeighbors, GeoHashRadius,
GeoHashRange, GeoPoint, GeoRadiusOption, GeoSearch, GeoSearchStoreOption, GeoShape,
GeoShapeType, OriginPoint,
};
use crate::db::WeDb;
use crate::error::{Error, Result};
use crate::zset::conf::{RangeScoreSpec, ZAdd};
use rapidhash::RapidHashSet;
impl WeDb {
pub fn geoadd<K: AsRef<[u8]>, M: AsRef<[u8]>>(
&self,
key: K,
items: &[(f64, f64, M)],
) -> Result<usize> {
self.geoadd_opts(key, items, [])
}
pub fn geoadd_opts<K: AsRef<[u8]>, M: AsRef<[u8]>>(
&self,
key: K,
items: &[(f64, f64, M)],
options: impl AsRef<[ZAdd]>,
) -> Result<usize> {
let mut score_members = Vec::with_capacity(items.len());
for (lon, lat, member) in items {
validate_long_lat(*lon, *lat)?;
let bits = geohash_encode_wgs84(*lon, *lat, GEO_STEP_MAX)
.ok_or_else(|| Error::invalid_data("ERR invalid longitude/latitude coordinates"))?;
let score = align_52bits(bits) as f64;
score_members.push((score, member));
}
self.zadd(key, &score_members, options)
}
pub fn geodist<K: AsRef<[u8]>, M1: AsRef<[u8]>, M2: AsRef<[u8]>>(
&self,
key: K,
member1: M1,
member2: M2,
unit: Option<&str>,
) -> Result<Option<f64>> {
let u = match unit {
Some(s) => DistanceUnit::parse(s).ok_or_else(|| {
Error::invalid_data("ERR unsupported unit provided. please use M, KM, FT, MI")
})?,
None => DistanceUnit::Meters,
};
let pos = self.geopos(&key, &[member1.as_ref(), member2.as_ref()])?;
if let (Some(Some((lon1, lat1))), Some(Some((lon2, lat2)))) =
(pos.first().copied(), pos.get(1).copied())
{
let d_meters = haversine_distance(lon1, lat1, lon2, lat2);
Ok(Some(u.from_meters(d_meters)))
} else {
Ok(None)
}
}
pub fn geopos<K: AsRef<[u8]>, M: AsRef<[u8]>>(
&self,
key: K,
members: &[M],
) -> Result<Vec<Option<(f64, f64)>>> {
let mut results = Vec::with_capacity(members.len());
for m in members {
if let Some(score) = self.zscore(&key, m)? {
let bits = GeoHashBits {
bits: score as u64,
step: GEO_STEP_MAX,
};
let (lon, lat) = geohash_decode_wgs84(bits);
results.push(Some((lon, lat)));
} else {
results.push(None);
}
}
Ok(results)
}
pub fn geohash<K: AsRef<[u8]>, M: AsRef<[u8]>>(
&self,
key: K,
members: &[M],
) -> Result<Vec<Option<String>>> {
let mut results = Vec::with_capacity(members.len());
for m in members {
if let Some(score) = self.zscore(&key, m)? {
let bits = GeoHashBits {
bits: score as u64,
step: GEO_STEP_MAX,
};
let (lon, lat) = geohash_decode_wgs84(bits);
results.push(Some(encode_geohash_string(lon, lat)));
} else {
results.push(None);
}
}
Ok(results)
}
pub fn georadius<K: AsRef<[u8]>>(
&self,
key: K,
longitude: f64,
latitude: f64,
radius: f64,
opt: &GeoRadiusOption,
) -> Result<Vec<GeoPoint>> {
validate_long_lat(longitude, latitude)?;
if radius.is_nan() || radius.is_infinite() || radius < 0.0 {
return Err(Error::invalid_data(
"ERR radius must be greater than or equal to 0",
));
}
if self.zcard(key.as_ref())? == 0 {
if let Some(ref store_k) = opt.store_key {
self.del(&[store_k])?;
}
if let Some(ref store_dist_k) = opt.store_dist_key {
self.del(&[store_dist_k])?;
}
return Ok(Vec::new());
}
let mut shape = GeoShape::new_circular_with_unit(longitude, latitude, radius, opt.unit);
let mut points = self.search_shape(&key, &mut shape, opt.unit)?;
let mut sort_order = opt.sort;
if opt.count.is_some() && !opt.any && sort_order == DistanceSort::None {
sort_order = DistanceSort::Asc;
}
match sort_order {
DistanceSort::Asc => {
points.sort_unstable_by(|a, b| a.dist.total_cmp(&b.dist));
}
DistanceSort::Desc => {
points.sort_unstable_by(|a, b| b.dist.total_cmp(&a.dist));
}
DistanceSort::None => {}
}
if let Some(limit) = opt.count {
points.truncate(limit);
}
if let Some(ref store_k) = opt.store_key {
self.del(&[store_k])?;
if !points.is_empty() {
let items: Vec<(f64, &str)> = points
.iter()
.map(|p| (p.score, p.member.as_str()))
.collect();
self.zadd(store_k, &items, [])?;
}
}
if let Some(ref store_dist_k) = opt.store_dist_key {
self.del(&[store_dist_k])?;
if !points.is_empty() {
let items: Vec<(f64, &str)> =
points.iter().map(|p| (p.dist, p.member.as_str())).collect();
self.zadd(store_dist_k, &items, [])?;
}
}
Ok(points)
}
pub fn georadiusbymember<K: AsRef<[u8]>, M: AsRef<[u8]>>(
&self,
key: K,
member: M,
radius: f64,
opt: &GeoRadiusOption,
) -> Result<Vec<GeoPoint>> {
let k_ref = key.as_ref();
let m_ref = member.as_ref();
if self.zcard(k_ref)? == 0 {
if let Some(ref store_k) = opt.store_key {
self.del(&[store_k])?;
}
if let Some(ref store_dist_k) = opt.store_dist_key {
self.del(&[store_dist_k])?;
}
return Ok(Vec::new());
}
let pos = self.geopos(k_ref, &[m_ref])?;
if let Some(Some((lon, lat))) = pos.into_iter().next() {
self.georadius(k_ref, lon, lat, radius, opt)
} else {
Err(Error::invalid_data(
"ERR could not decode requested zset member",
))
}
}
pub fn geosearch<K: AsRef<[u8]>>(
&self,
key: K,
origin: &OriginPoint,
shape: &mut GeoShape,
opt: &GeoSearch,
) -> Result<Vec<GeoPoint>> {
let k_ref = key.as_ref();
let (lon, lat) = match origin {
OriginPoint::Coord { lon, lat } => {
validate_long_lat(*lon, *lat)?;
(*lon, *lat)
}
OriginPoint::Member(m) => {
if self.zcard(k_ref)? == 0 {
return Ok(Vec::new());
}
let pos = self.geopos(k_ref, &[m.as_bytes()])?;
match pos.into_iter().next().flatten() {
Some((l, t)) => (l, t),
None => {
return Err(Error::invalid_data(
"ERR could not decode requested zset member",
));
}
}
}
};
shape.center_lon = lon;
shape.center_lat = lat;
bounding_box(shape);
let mut points = self.search_shape(k_ref, shape, opt.unit)?;
let mut sort_order = if opt.asc || opt.sort == DistanceSort::Asc {
DistanceSort::Asc
} else {
opt.sort
};
if opt.count.is_some() && !opt.any && sort_order == DistanceSort::None {
sort_order = DistanceSort::Asc;
}
match sort_order {
DistanceSort::Asc => {
points.sort_unstable_by(|a, b| a.dist.total_cmp(&b.dist));
}
DistanceSort::Desc => {
points.sort_unstable_by(|a, b| b.dist.total_cmp(&a.dist));
}
DistanceSort::None => {}
}
if let Some(limit) = opt.count {
points.truncate(limit);
}
Ok(points)
}
pub fn geosearchstore<K: AsRef<[u8]>, D: AsRef<[u8]>>(
&self,
destination: D,
source: K,
origin: &OriginPoint,
shape: &mut GeoShape,
opt: &GeoSearchStoreOption,
) -> Result<usize> {
let dest_ref = destination.as_ref();
let src_ref = source.as_ref();
let (lon, lat) = match origin {
OriginPoint::Coord { lon, lat } => {
validate_long_lat(*lon, *lat)?;
if self.zcard(src_ref)? == 0 {
self.del(&[dest_ref])?;
return Ok(0);
}
(*lon, *lat)
}
OriginPoint::Member(m) => {
if self.zcard(src_ref)? == 0 {
self.del(&[dest_ref])?;
return Ok(0);
}
let pos = self.geopos(src_ref, &[m.as_bytes()])?;
match pos.into_iter().next().flatten() {
Some((l, t)) => (l, t),
None => {
return Err(Error::invalid_data(
"ERR could not decode requested zset member",
));
}
}
}
};
shape.center_lon = lon;
shape.center_lat = lat;
bounding_box(shape);
let mut points = self.search_shape(src_ref, shape, opt.unit)?;
let mut sort_order = opt.sort;
if opt.count.is_some() && !opt.any && sort_order == DistanceSort::None {
sort_order = DistanceSort::Asc;
}
match sort_order {
DistanceSort::Asc => {
points.sort_unstable_by(|a, b| a.dist.total_cmp(&b.dist));
}
DistanceSort::Desc => {
points.sort_unstable_by(|a, b| b.dist.total_cmp(&a.dist));
}
DistanceSort::None => {}
}
if let Some(limit) = opt.count {
points.truncate(limit);
}
self.del(&[dest_ref])?;
if points.is_empty() {
return Ok(0);
}
let items: Vec<(f64, &str)> = points
.iter()
.map(|p| {
let score = if opt.store_dist { p.dist } else { p.score };
(score, p.member.as_str())
})
.collect();
self.zadd(dest_ref, &items, [])?;
Ok(items.len())
}
fn search_shape<K: AsRef<[u8]>>(
&self,
key: K,
geo_shape: &mut GeoShape,
unit: DistanceUnit,
) -> Result<Vec<GeoPoint>> {
let georadius = get_areas_by_shape_wgs84(geo_shape);
let raw_neighbors = [
georadius.hash,
georadius.neighbors.north,
georadius.neighbors.south,
georadius.neighbors.east,
georadius.neighbors.west,
georadius.neighbors.north_east,
georadius.neighbors.north_west,
georadius.neighbors.south_east,
georadius.neighbors.south_west,
];
let mut unique_hashes = [GeoHashBits::default(); 9];
let mut unique_count = 0usize;
for hash in raw_neighbors {
if hash.is_zero() {
continue;
}
let mut duplicate = false;
for prev in &unique_hashes[..unique_count] {
if prev.bits == hash.bits && prev.step == hash.step {
duplicate = true;
break;
}
}
if !duplicate {
unique_hashes[unique_count] = hash;
unique_count += 1;
}
}
let mut points = Vec::new();
let mut seen_members = RapidHashSet::default();
for &hash in &unique_hashes[..unique_count] {
let (min_bits, max_bits) = scores_of_geohash_box(hash);
let spec = RangeScoreSpec {
min: min_bits as f64,
max: max_bits as f64,
minex: false,
maxex: true, offset: 0,
count: None,
};
let range_members = self.zrangebyscore(&key, &spec)?;
for (member_bytes, score) in range_members {
if !seen_members.insert(member_bytes.clone()) {
continue;
}
let bits = GeoHashBits {
bits: score as u64,
step: GEO_STEP_MAX,
};
let (pt_lon, pt_lat) = geohash_decode_wgs84(bits);
let d_meters =
haversine_distance(geo_shape.center_lon, geo_shape.center_lat, pt_lon, pt_lat);
let is_inside = match geo_shape.shape_type {
GeoShapeType::Circular => d_meters <= geo_shape.radius * geo_shape.conversion,
GeoShapeType::Rectangular => {
pt_lon >= geo_shape.bounds[0]
&& pt_lon <= geo_shape.bounds[2]
&& pt_lat >= geo_shape.bounds[1]
&& pt_lat <= geo_shape.bounds[3]
}
GeoShapeType::None => false,
};
if is_inside {
let member_str = String::from_utf8(member_bytes)
.unwrap_or_else(|e| String::from_utf8_lossy(e.as_bytes()).into_owned());
let dist = unit.from_meters(d_meters);
points.push(GeoPoint {
longitude: pt_lon,
latitude: pt_lat,
member: member_str,
dist,
score,
});
}
}
}
Ok(points)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_interleave_deinterleave_roundtrip() {
let x: u32 = 0x12345678;
let y: u32 = 0x87654321;
let interleaved = interleave64(x, y);
let (rx, ry) = deinterleave64(interleaved);
assert_eq!(x, rx);
assert_eq!(y, ry);
}
#[test]
fn test_validate_long_lat_limits() {
assert!(validate_long_lat(0.0, 0.0).is_ok());
assert!(validate_long_lat(GEO_LON_MIN, GEO_LAT_MIN).is_ok());
assert!(validate_long_lat(GEO_LON_MAX, GEO_LAT_MAX).is_ok());
assert!(validate_long_lat(GEO_LON_MIN - 0.0001, 0.0).is_err());
assert!(validate_long_lat(GEO_LON_MAX + 0.0001, 0.0).is_err());
assert!(validate_long_lat(0.0, GEO_LAT_MIN - 0.0001).is_err());
assert!(validate_long_lat(0.0, GEO_LAT_MAX + 0.0001).is_err());
assert!(validate_long_lat(f64::NAN, 0.0).is_err());
assert!(validate_long_lat(0.0, f64::INFINITY).is_err());
}
#[test]
fn test_geohash_encode_decode_roundtrip() {
let lon = 116.4074;
let lat = 39.9042;
let hash = encode_geohash(lon, lat);
let (d_lon, d_lat) = decode_geohash(hash);
assert!((lon - d_lon).abs() < 1e-4);
assert!((lat - d_lat).abs() < 1e-4);
let hash_str = encode_geohash_string(lon, lat);
assert_eq!(hash_str.len(), 11);
assert_eq!(hash_str, geohash_to_base32(hash));
let (b_lon, b_lat) = base32_to_coords(&hash_str).unwrap();
assert!((lon - b_lon).abs() < 1e-4);
assert!((lat - b_lat).abs() < 1e-4);
}
#[test]
fn test_haversine_distance_known_points() {
let d = haversine_distance(2.3522, 48.8566, -0.1278, 51.5074);
assert!((d - 343_556.0).abs() < 5000.0);
assert_eq!(haversine_distance(12.34, 56.78, 12.34, 56.78), 0.0);
}
#[test]
fn test_geoshape_contains_point() {
let circle = GeoShape::new_circular(0.0, 0.0, 1000.0);
assert!(circle.contains_point(0.0, 0.0));
assert!(circle.contains_point(0.001, 0.001));
assert!(!circle.contains_point(1.0, 1.0));
let rect = GeoShape::new_rectangular(0.0, 0.0, 2000.0, 2000.0);
assert!(rect.contains_point(0.0, 0.0));
assert!(!rect.contains_point(1.0, 1.0));
}
}