use crate::db::Db;
use crate::keyspace::wrong_type;
use crate::value::Kind;
use crate::zsets::Window;
use std::cmp::Ordering;
use yo_common::{Code, Error, Result, num};
const NOT_A_DOUBLE: &str = "One or more scores can't be converted into double";
#[derive(Debug, Clone, Copy, Default)]
pub struct Sort<'a> {
pub by: Option<&'a [u8]>,
pub get: &'a [&'a [u8]],
pub limit: Option<(i64, i64)>,
pub desc: bool,
pub alpha: bool,
pub store: bool,
}
struct Weighted {
elem: Vec<u8>,
score: f64,
text: Option<Vec<u8>>,
}
impl Db {
pub fn sort(&self, key: &[u8], opts: &Sort<'_>) -> Result<Vec<Option<Vec<u8>>>> {
let mut opts = *opts;
opts.store = false;
self.sorted(key, &opts)
}
pub fn sort_store(&self, key: &[u8], dest: &[u8], opts: &Sort<'_>) -> Result<usize> {
let mut opts = *opts;
opts.store = true;
let rows = self.sorted(key, &opts)?;
if rows.is_empty() {
self.hold(dest).del(dest);
return Ok(0);
}
self.hold(dest).del(dest);
let owned: Vec<Vec<u8>> = rows.into_iter().map(Option::unwrap_or_default).collect();
self.hold(dest).push(
dest,
crate::lists::End::Right,
owned.iter().map(Vec::as_slice),
)
}
fn sorted(&self, key: &[u8], opts: &Sort<'_>) -> Result<Vec<Option<Vec<u8>>>> {
let kind = self.hold(key).kind_of(key);
let elems = self.elements(key, kind)?;
let mut alpha = opts.alpha;
let mut by = opts.by;
let mut dontsort = by.is_some_and(|p| !p.contains(&b'*'));
if dontsort && kind == Some(Kind::Set) && opts.store {
dontsort = false;
alpha = true;
by = None;
}
let ordered = if dontsort {
let mut e = elems;
if opts.desc {
e.reverse();
}
e
} else {
self.order(elems, by, alpha, opts.desc)?
};
let window = limit(ordered.len(), opts.limit);
self.emit(&ordered[window], opts.get)
}
fn elements(&self, key: &[u8], kind: Option<Kind>) -> Result<Vec<Vec<u8>>> {
let mut out = Vec::new();
let mut db = self.hold(key);
match kind {
None => {}
Some(Kind::List) => {
for e in db.lrange(key, 0, -1)? {
let mut v = Vec::new();
e.write_to(&mut v);
out.push(v);
}
}
Some(Kind::Set) => {
if let Some(members) = db.smembers(key)? {
for m in members {
let mut v = Vec::new();
m.write_to(&mut v);
out.push(v);
}
}
}
Some(Kind::Zset) => {
let n = db.zcard(key)?;
let w = Window {
from: 0,
count: n,
rev: false,
};
out.reserve(n);
db.zwalk(key, w, |m, _| {
let mut v = Vec::new();
m.write_to(&mut v);
out.push(v);
})?;
}
Some(_) => return Err(wrong_type()),
}
Ok(out)
}
fn order(
&self,
elems: Vec<Vec<u8>>,
by: Option<&[u8]>,
alpha: bool,
desc: bool,
) -> Result<Vec<Vec<u8>>> {
let mut weighed = Vec::with_capacity(elems.len());
for elem in elems {
let looked = match by {
Some(pattern) => self.by_pattern(pattern, &elem),
None => None,
};
let (score, text) = if alpha {
(0.0, if by.is_some() { looked } else { None })
} else {
let raw = match by {
Some(_) => match looked {
Some(v) => v,
None => {
weighed.push(Weighted {
elem,
score: 0.0,
text: None,
});
continue;
}
},
None => elem.clone(),
};
let n =
num::parse_f64(&raw).ok_or_else(|| Error::new(Code::Invalid, NOT_A_DOUBLE))?;
if n.is_nan() {
return Err(Error::new(Code::Invalid, NOT_A_DOUBLE));
}
(n, None)
};
weighed.push(Weighted { elem, score, text });
}
weighed.sort_by(|a, b| {
let cmp = if alpha {
match (&a.text, &b.text) {
(None, None) => a.elem.cmp(&b.elem),
(None, Some(_)) => Ordering::Less,
(Some(_), None) => Ordering::Greater,
(Some(x), Some(y)) => x.cmp(y).then_with(|| a.elem.cmp(&b.elem)),
}
} else {
a.score
.partial_cmp(&b.score)
.unwrap_or(Ordering::Equal)
.then_with(|| a.elem.cmp(&b.elem))
};
if desc { cmp.reverse() } else { cmp }
});
Ok(weighed.into_iter().map(|w| w.elem).collect())
}
fn emit(&self, elems: &[Vec<u8>], get: &[&[u8]]) -> Result<Vec<Option<Vec<u8>>>> {
if get.is_empty() {
return Ok(elems.iter().map(|e| Some(e.clone())).collect());
}
let mut out = Vec::with_capacity(elems.len() * get.len());
for elem in elems {
for pattern in get {
if *pattern == b"#" {
out.push(Some(elem.clone()));
} else {
out.push(self.by_pattern(pattern, elem));
}
}
}
Ok(out)
}
fn by_pattern(&self, pattern: &[u8], elem: &[u8]) -> Option<Vec<u8>> {
let star = pattern.iter().position(|&c| c == b'*')?;
let arrow = pattern[star + 1..]
.windows(2)
.position(|w| w == b"->")
.map(|i| star + 1 + i)
.filter(|&i| i + 2 < pattern.len());
let (key_part, field) = match arrow {
Some(i) => (&pattern[..i], Some(&pattern[i + 2..])),
None => (pattern, None),
};
let mut key = Vec::with_capacity(key_part.len() + elem.len());
key.extend_from_slice(&key_part[..star]);
key.extend_from_slice(elem);
key.extend_from_slice(&key_part[star + 1..]);
let mut stripe = self.hold(&key);
match field {
Some(f) => stripe
.hget(&key, f, |t| {
t.map(|t| {
let mut v = Vec::new();
t.write_to(&mut v);
v
})
})
.unwrap_or(None),
None => stripe.get(&key).ok().flatten().map(|s| s.to_vec()),
}
}
}
fn limit(len: usize, limit: Option<(i64, i64)>) -> std::ops::Range<usize> {
let Some((offset, count)) = limit else {
return 0..len;
};
let start = usize::try_from(offset).unwrap_or(0).min(len);
let end = if count < 0 {
len
} else {
start
.saturating_add(usize::try_from(count).unwrap_or(0))
.min(len)
};
start..end
}
#[cfg(test)]
mod tests {
use super::*;
use crate::lists::End;
use crate::strings::SetOptions;
use crate::zsets::ZAdd;
fn flat(rows: Vec<Option<Vec<u8>>>) -> Vec<String> {
rows.into_iter()
.map(|r| match r {
Some(v) => String::from_utf8_lossy(&v).into_owned(),
None => "nil".to_string(),
})
.collect()
}
fn text(e: crate::listpack::Entry<'_>) -> String {
let mut v = Vec::new();
e.write_to(&mut v);
String::from_utf8_lossy(&v).into_owned()
}
fn list(db: &mut Db, key: &[u8], items: &[&str]) {
db.at(key)
.push(key, End::Right, items.iter().map(|s| s.as_bytes()))
.expect("a fresh list takes elements");
}
#[test]
fn numbers_sort_as_numbers_and_not_as_text() {
let mut db = Db::new();
list(&mut db, b"l", &["10", "9", "100", "1"]);
let opts = Sort::default();
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["1", "9", "10", "100"]);
let alpha = Sort {
alpha: true,
..Sort::default()
};
assert_eq!(
flat(db.sort(b"l", &alpha).unwrap()),
["1", "10", "100", "9"]
);
}
#[test]
fn an_element_that_is_not_a_number_fails_the_whole_command() {
let mut db = Db::new();
list(&mut db, b"l", &["1", "two", "3"]);
let err = db.sort(b"l", &Sort::default()).unwrap_err();
assert_eq!(err.message(), NOT_A_DOUBLE);
let alpha = Sort {
alpha: true,
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &alpha).unwrap()), ["1", "3", "two"]);
}
#[test]
fn desc_reverses_and_limit_takes_a_window_of_what_is_left() {
let mut db = Db::new();
list(&mut db, b"l", &["3", "1", "5", "2", "4"]);
let opts = Sort {
desc: true,
limit: Some((1, 2)),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["4", "3"]);
let rest = Sort {
limit: Some((3, -1)),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &rest).unwrap()), ["4", "5"]);
let past = Sort {
limit: Some((99, 5)),
..Sort::default()
};
assert!(db.sort(b"l", &past).unwrap().is_empty());
let before = Sort {
limit: Some((-4, 2)),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &before).unwrap()), ["1", "2"]);
}
#[test]
fn by_reads_a_key_for_every_element() {
let mut db = Db::new();
list(&mut db, b"l", &["a", "b", "c"]);
db.at(b"w_a").set(b"w_a", b"3", SetOptions::PLAIN).unwrap();
db.at(b"w_b").set(b"w_b", b"1", SetOptions::PLAIN).unwrap();
db.at(b"w_c").set(b"w_c", b"2", SetOptions::PLAIN).unwrap();
let opts = Sort {
by: Some(b"w_*"),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["b", "c", "a"]);
}
#[test]
fn a_by_lookup_that_missed_weighs_nothing_and_the_element_breaks_the_tie() {
let mut db = Db::new();
list(&mut db, b"l", &["c", "a", "b"]);
db.at(b"w_b").set(b"w_b", b"5", SetOptions::PLAIN).unwrap();
let opts = Sort {
by: Some(b"w_*"),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["a", "c", "b"]);
}
#[test]
fn under_alpha_a_missed_by_sorts_before_every_hit() {
let mut db = Db::new();
list(&mut db, b"l", &["c", "a", "b"]);
db.at(b"w_b")
.set(b"w_b", b"zzz", SetOptions::PLAIN)
.unwrap();
db.at(b"w_c")
.set(b"w_c", b"aaa", SetOptions::PLAIN)
.unwrap();
let opts = Sort {
by: Some(b"w_*"),
alpha: true,
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["a", "c", "b"]);
}
#[test]
fn a_pattern_can_reach_into_a_hash() {
let mut db = Db::new();
list(&mut db, b"l", &["a", "b"]);
db.at(b"h_a")
.hset(b"h_a", [(&b"w"[..], &b"2"[..])].into_iter())
.unwrap();
db.at(b"h_b")
.hset(b"h_b", [(&b"w"[..], &b"1"[..])].into_iter())
.unwrap();
let opts = Sort {
by: Some(b"h_*->w"),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["b", "a"]);
db.at(b"h_a->")
.set(b"h_a->", b"9", SetOptions::PLAIN)
.unwrap();
let trailing = Sort {
by: Some(b"h_*->"),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &trailing).unwrap()), ["b", "a"]);
}
#[test]
fn get_answers_other_keys_and_a_hash_of_them() {
let mut db = Db::new();
list(&mut db, b"l", &["2", "1"]);
db.at(b"d_1")
.set(b"d_1", b"one", SetOptions::PLAIN)
.unwrap();
db.at(b"d_2")
.set(b"d_2", b"two", SetOptions::PLAIN)
.unwrap();
let get: [&[u8]; 2] = [b"#", b"d_*"];
let opts = Sort {
get: &get,
..Sort::default()
};
assert_eq!(
flat(db.sort(b"l", &opts).unwrap()),
["1", "one", "2", "two"]
);
db.at(b"d_2").del(b"d_2");
assert_eq!(
flat(db.sort(b"l", &opts).unwrap()),
["1", "one", "2", "nil"]
);
}
#[test]
fn a_lookup_at_the_wrong_type_is_a_miss_and_not_an_error() {
let mut db = Db::new();
list(&mut db, b"l", &["a"]);
list(&mut db, b"d_a", &["x"]);
let get: [&[u8]; 1] = [b"d_*"];
let opts = Sort {
get: &get,
alpha: true,
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["nil"]);
}
#[test]
fn by_without_a_star_leaves_the_order_alone() {
let mut db = Db::new();
list(&mut db, b"l", &["3", "1", "2"]);
let opts = Sort {
by: Some(b"nosort"),
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &opts).unwrap()), ["3", "1", "2"]);
let back = Sort {
by: Some(b"nosort"),
desc: true,
..Sort::default()
};
assert_eq!(flat(db.sort(b"l", &back).unwrap()), ["2", "1", "3"]);
}
#[test]
fn a_set_stored_without_a_sort_is_sorted_anyway() {
let mut db = Db::new();
for m in ["c", "a", "b"] {
db.at(b"s").sadd(b"s", [m.as_bytes()].into_iter()).unwrap();
}
let opts = Sort {
by: Some(b"nosort"),
..Sort::default()
};
assert_eq!(db.sort_store(b"s", b"out", &opts).unwrap(), 3);
let got: Vec<String> = db
.at(b"out")
.lrange(b"out", 0, -1)
.unwrap()
.map(text)
.collect();
assert_eq!(got, ["a", "b", "c"]);
}
#[test]
fn a_sorted_set_comes_out_in_score_order_when_nothing_says_otherwise() {
let mut db = Db::new();
db.at(b"z")
.zadd(
b"z",
[(3.0, &b"c"[..]), (1.0, &b"a"[..]), (2.0, &b"b"[..])].into_iter(),
ZAdd::default(),
)
.unwrap();
let opts = Sort {
by: Some(b"nosort"),
..Sort::default()
};
assert_eq!(flat(db.sort(b"z", &opts).unwrap()), ["a", "b", "c"]);
}
#[test]
fn storing_an_empty_result_removes_the_destination() {
let mut db = Db::new();
list(&mut db, b"out", &["stale"]);
assert_eq!(
db.sort_store(b"missing", b"out", &Sort::default()).unwrap(),
0
);
assert!(!db.at(b"out").exists(b"out"));
}
#[test]
fn a_stored_get_that_missed_is_an_empty_string() {
let mut db = Db::new();
list(&mut db, b"l", &["1"]);
let get: [&[u8]; 1] = [b"d_*"];
let opts = Sort {
get: &get,
..Sort::default()
};
assert_eq!(db.sort_store(b"l", b"out", &opts).unwrap(), 1);
assert_eq!(db.at(b"out").llen(b"out").unwrap(), 1);
}
#[test]
fn a_missing_key_is_empty_and_a_wrong_type_is_an_error() {
let mut db = Db::new();
assert!(db.sort(b"nosuchkey", &Sort::default()).unwrap().is_empty());
db.at(b"str").set(b"str", b"x", SetOptions::PLAIN).unwrap();
assert_eq!(
db.sort(b"str", &Sort::default()).unwrap_err().code(),
Code::WrongType
);
}
#[test]
fn sorting_into_the_key_being_sorted_works() {
let mut db = Db::new();
list(&mut db, b"l", &["3", "1", "2"]);
assert_eq!(db.sort_store(b"l", b"l", &Sort::default()).unwrap(), 3);
let got: Vec<String> = db.at(b"l").lrange(b"l", 0, -1).unwrap().map(text).collect();
assert_eq!(got, ["1", "2", "3"]);
}
}