use super::sorted_set_object::{SortedSetEntry, SortedSetObject};
use crate::{
geo::{
GeoAddOptions, GeoOrder, GeoOriginType, GeoSearchOptions, GeoSearchType,
geo_hash::{GeoDistanceUnitType, GeoHash},
},
parse_utils::{try_get_geo_distance_unit, try_get_geo_lon_lat},
resp::{ObjectInput, ObjectOutput},
};
#[derive(Debug, Clone)]
pub struct GeoSearchData {
pub member: Vec<u8>,
pub distance: f64,
pub geo_hash: i64,
pub geo_hash_code: [u8; GeoHash::CODE_LENGTH],
pub coordinates: (f64, f64),
}
#[inline]
fn arg(input: &ObjectInput, i: usize) -> &[u8] {
input.parse_state.get_arg_slice_by_ref(i).as_slice()
}
const RESP_ERR_ZSET_MEMBER: &[u8] = b"ERR could not decode requested zset member";
impl SortedSetObject {
pub(crate) fn geo_add(&mut self, input: &ObjectInput, output: &mut ObjectOutput) {
self.delete_expired_items();
let options = GeoAddOptions::from_bits_truncate(input.arg1 as u8);
let count = input.parse_state.count;
let mut curr_token_idx = 0;
let mut elements_added = 0_i64;
let mut elements_changed = 0_i64;
while curr_token_idx < count {
let (longitude, latitude) = if curr_token_idx + 1 < count {
try_get_geo_lon_lat(arg(input, curr_token_idx), arg(input, curr_token_idx + 1))
.unwrap_or((0.0, 0.0))
} else {
(0.0, 0.0)
};
curr_token_idx += 2;
if curr_token_idx >= count {
break;
}
let member = arg(input, curr_token_idx).to_vec();
curr_token_idx += 1;
let score = GeoHash::geo_to_long_value(latitude, longitude);
if score == -1 {
continue;
}
match self.sorted_set_dict.get(&member).copied() {
None => {
if !options.contains(GeoAddOptions::XX) {
self.sorted_set_dict.insert(member.clone(), score as f64);
self.sorted_set.insert(SortedSetEntry {
score: score as f64,
member: member.clone(),
});
elements_added += 1;
self.update_size(&member, true);
elements_changed += 1;
}
}
Some(score_stored) => {
if !options.contains(GeoAddOptions::NX) && score_stored != score as f64 {
self.sorted_set_dict.insert(member.clone(), score as f64);
self.sorted_set.remove(&SortedSetEntry {
score: score_stored,
member: member.clone(),
});
self.sorted_set.insert(SortedSetEntry {
score: score as f64,
member: member.clone(),
});
elements_changed += 1;
}
}
}
}
let result = if !options.contains(GeoAddOptions::CH) {
elements_added
} else {
elements_changed
};
output.write_int64(result);
}
pub(crate) fn geo_hash(
&mut self,
input: &ObjectInput,
output: &mut ObjectOutput,
resp_protocol_version: u8,
) {
output.write_array_length(input.parse_state.count);
for i in 0..input.parse_state.count {
let member = arg(input, i);
match self.sorted_set_dict.get(member).copied() {
Some(value52_int) => {
let geo_hash = GeoHash::get_geo_hash_code(value52_int as i64);
output.write_ascii_bulk_string(&geo_hash);
}
None => output.write_null(resp_protocol_version),
}
}
}
pub(crate) fn geo_distance(
&mut self,
input: &ObjectInput,
output: &mut ObjectOutput,
resp_protocol_version: u8,
) {
let member1 = arg(input, 0);
let member2 = arg(input, 1);
let mut units = GeoDistanceUnitType::M;
if input.parse_state.count > 2 {
units = try_get_geo_distance_unit(arg(input, 2)).unwrap_or(GeoDistanceUnitType::M);
}
match (
self.sorted_set_dict.get(member1).copied(),
self.sorted_set_dict.get(member2).copied(),
) {
(Some(score_member1), Some(score_member2)) => {
let first = GeoHash::get_coordinates_from_long(score_member1 as i64);
let second = GeoHash::get_coordinates_from_long(score_member2 as i64);
let distance = GeoHash::distance(first.0, first.1, second.0, second.1);
output.write_double_bulk_string(GeoHash::convert_meters_to_units(distance, units));
}
_ => output.write_null(resp_protocol_version),
}
}
pub(crate) fn geo_position(
&mut self,
input: &ObjectInput,
output: &mut ObjectOutput,
resp_protocol_version: u8,
) {
output.write_array_length(input.parse_state.count);
for i in 0..input.parse_state.count {
let member = arg(input, i);
match self.sorted_set_dict.get(member).copied() {
Some(score_member) => {
let (lat, lon) = GeoHash::get_coordinates_from_long(score_member as i64);
output.write_array_length(2);
output.write_double_numeric(lon, resp_protocol_version);
output.write_double_numeric(lat, resp_protocol_version);
}
None => output.write_null_array(resp_protocol_version),
}
}
}
pub fn geo_search(
&mut self,
opts: &mut GeoSearchOptions,
output: &mut ObjectOutput,
resp_protocol_version: u8,
read_only: bool,
) {
if opts.origin == GeoOriginType::FromMember {
let Some(center_point_score) = self.sorted_set_dict.get(&opts.from_member).copied() else {
output.write_error(RESP_ERR_ZSET_MEMBER);
return;
};
(opts.lat, opts.lon) = GeoHash::get_coordinates_from_long(center_point_score as i64);
}
let mut response_data: Vec<GeoSearchData> = Vec::with_capacity(
if opts.with_count_any
&& opts.count_value > 0
&& (opts.count_value as usize) < self.sorted_set.len()
{
opts.count_value as usize
} else {
self.sorted_set.len()
},
);
for point in &self.sorted_set {
let coor_in_item = GeoHash::get_coordinates_from_long(point.score as i64);
let distance = if opts.search_type == GeoSearchType::ByBox {
let Some(d) = GeoHash::get_distance_when_in_rectangle(
GeoHash::convert_value_to_meters(opts.box_width, opts.unit),
GeoHash::convert_value_to_meters(opts.radius, opts.unit),
opts.lat,
opts.lon,
coor_in_item.0,
coor_in_item.1,
) else {
continue;
};
d
} else {
let Some(d) = GeoHash::is_point_within_radius(
GeoHash::convert_value_to_meters(opts.radius, opts.unit),
opts.lat,
opts.lon,
coor_in_item.0,
coor_in_item.1,
) else {
continue;
};
d
};
response_data.push(GeoSearchData {
member: point.member.clone(),
distance,
geo_hash: point.score as i64,
geo_hash_code: GeoHash::get_geo_hash_code(point.score as i64),
coordinates: GeoHash::get_coordinates_from_long(point.score as i64),
});
if opts.with_count_any && response_data.len() == opts.count_value as usize {
break;
}
}
if response_data.is_empty() {
output.write_empty_array();
return;
}
let mut inner_array_length = 1_usize;
if opts.with_dist {
inner_array_length += 1;
}
if opts.with_hash {
inner_array_length += 1;
}
if opts.with_coord {
inner_array_length += 1;
}
match opts.sort {
GeoOrder::Descending => response_data.sort_by(|a, b| b.distance.total_cmp(&a.distance)),
GeoOrder::Ascending => response_data.sort_by(|a, b| a.distance.total_cmp(&b.distance)),
GeoOrder::None => {
if !opts.with_count_any && opts.count_value > 0 {
response_data.sort_by(|a, b| a.distance.total_cmp(&b.distance));
}
}
}
if opts.count_value > 0 && (opts.count_value as usize) < response_data.len() {
response_data.truncate(opts.count_value as usize);
output.write_array_length(opts.count_value as usize);
} else {
output.write_array_length(response_data.len());
}
for item in &response_data {
if inner_array_length > 1 {
output.write_array_length(inner_array_length);
}
output.write_bulk_string(&item.member);
if opts.with_dist {
output.write_double_bulk_string(GeoHash::convert_meters_to_units(item.distance, opts.unit));
}
if opts.with_hash {
if read_only {
output.write_int64(item.geo_hash);
} else {
output.write_array_item(item.geo_hash);
}
}
if opts.with_coord {
output.write_array_length(2);
output.write_double_numeric(item.coordinates.1, resp_protocol_version);
output.write_double_numeric(item.coordinates.0, resp_protocol_version);
}
}
}
}