use std::io;
use kevy_resp::Reply;
use crate::{Connection, num_f64, num_u64, string, unexpected};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum IdxType {
I64,
F64,
Str,
}
impl IdxType {
fn tag(self) -> &'static [u8] {
match self {
Self::I64 => b"i64",
Self::F64 => b"f64",
Self::Str => b"str",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IdxRow {
pub key: Vec<u8>,
pub value: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IdxPage {
pub cursor: Option<Vec<u8>>,
pub rows: Vec<IdxRow>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IdxInfo {
pub name: Vec<u8>,
pub prefix: Vec<u8>,
pub kind: String,
pub state: String,
pub entries: u64,
pub bytes: u64,
}
impl Connection {
pub fn idx_create_range(
&mut self,
name: &[u8],
prefix: &[u8],
field: &[u8],
ty: IdxType,
) -> io::Result<()> {
let args: &[&[u8]] = &[
name, b"ON", b"PREFIX", prefix, b"FIELD", field, b"TYPE", ty.tag(), b"KIND",
b"range",
];
self.idx_create_raw(args)
}
pub fn idx_create_raw(&mut self, args: &[&[u8]]) -> io::Result<()> {
match self.remote("IDX.CREATE")?.request(&raw_argv(b"IDX.CREATE", args))? {
Reply::Simple(s) if s == b"OK" => Ok(()),
Reply::Error(e) => Err(io::Error::other(string(e))),
other => Err(unexpected(other)),
}
}
pub fn idx_drop(&mut self, name: &[u8]) -> io::Result<bool> {
match self.remote("IDX.DROP")?.request_borrowed(&[b"IDX.DROP", name])? {
Reply::Int(1) => Ok(true),
Reply::Int(0) => Ok(false),
Reply::Error(e) => Err(io::Error::other(string(e))),
other => Err(unexpected(other)),
}
}
pub fn idx_list(&mut self) -> io::Result<Vec<IdxInfo>> {
match self.remote("IDX.LIST")?.request_borrowed(&[b"IDX.LIST"])? {
Reply::Array(items) => items.into_iter().map(parse_info).collect(),
Reply::Error(e) => Err(io::Error::other(string(e))),
other => Err(unexpected(other)),
}
}
pub fn idx_query_range(
&mut self,
name: &[u8],
min: &[u8],
max: &[u8],
limit: usize,
cursor: Option<&[u8]>,
) -> io::Result<IdxPage> {
let lim = limit.to_string();
let mut args: Vec<&[u8]> = vec![name, b"RANGE", min, max, b"LIMIT", lim.as_bytes()];
if let Some(c) = cursor {
args.push(b"CURSOR");
args.push(c);
}
parse_page(self.idx_query_raw(&args)?)
}
pub fn idx_query_eq(&mut self, name: &[u8], value: &[u8], limit: usize) -> io::Result<IdxPage> {
let lim = limit.to_string();
parse_page(self.idx_query_raw(&[name, b"EQ", value, b"LIMIT", lim.as_bytes()])?)
}
pub fn idx_query_match(
&mut self,
name: &[u8],
text: &[u8],
limit: usize,
) -> io::Result<Vec<(Vec<u8>, f64)>> {
let lim = limit.to_string();
parse_ranked(self.idx_query_raw(&[name, b"MATCH", text, b"LIMIT", lim.as_bytes()])?)
}
pub fn idx_query_knn(
&mut self,
name: &[u8],
vector: &[f32],
k: usize,
) -> io::Result<Vec<(Vec<u8>, f64)>> {
let blob: Vec<u8> = vector.iter().flat_map(|f| f.to_le_bytes()).collect();
let lim = k.to_string();
parse_ranked(self.idx_query_raw(&[name, b"KNN", &blob, b"LIMIT", lim.as_bytes()])?)
}
pub fn idx_query_raw(&mut self, args: &[&[u8]]) -> io::Result<Reply> {
match self.remote("IDX.QUERY")?.request(&raw_argv(b"IDX.QUERY", args))? {
Reply::Error(e) => Err(io::Error::other(string(e))),
other => Ok(other),
}
}
}
fn raw_argv(verb: &[u8], args: &[&[u8]]) -> Vec<Vec<u8>> {
let mut argv = Vec::with_capacity(args.len() + 1);
argv.push(verb.to_vec());
argv.extend(args.iter().map(|a| a.to_vec()));
argv
}
fn parse_page(reply: Reply) -> io::Result<IdxPage> {
let Reply::Array(items) = reply else {
return Err(unexpected(reply));
};
if items.len() != 2 {
return Err(io::Error::other("IDX.QUERY page: expected [cursor, rows]"));
}
let mut it = items.into_iter();
let cursor = match it.next().unwrap() {
Reply::Bulk(c) if c == b"0" => None,
Reply::Bulk(c) => Some(c),
other => return Err(unexpected(other)),
};
let Reply::Array(flat) = it.next().unwrap() else {
return Err(io::Error::other("IDX.QUERY page: rows not an array"));
};
let mut rows = Vec::with_capacity(flat.len() / 2);
let mut flat = flat.into_iter();
while let Some(k) = flat.next() {
let (Reply::Bulk(key), Some(Reply::Bulk(value))) = (k, flat.next()) else {
return Err(io::Error::other("IDX.QUERY page: odd or non-bulk row pair"));
};
rows.push(IdxRow { key, value });
}
Ok(IdxPage { cursor, rows })
}
fn parse_ranked(reply: Reply) -> io::Result<Vec<(Vec<u8>, f64)>> {
let Reply::Array(items) = reply else {
return Err(unexpected(reply));
};
items
.into_iter()
.map(|row| {
let Reply::Array(cells) = row else {
return Err(unexpected(row));
};
let mut it = cells.into_iter();
match (it.next(), it.next()) {
(Some(Reply::Bulk(key)), Some(Reply::Bulk(score))) => {
Ok((key, num_f64(&score)?))
}
_ => Err(io::Error::other("IDX.QUERY ranked row: expected [key, score]")),
}
})
.collect()
}
fn parse_info(entry: Reply) -> io::Result<IdxInfo> {
let Reply::Array(cells) = entry else {
return Err(unexpected(entry));
};
let mut info = IdxInfo {
name: Vec::new(),
prefix: Vec::new(),
kind: String::new(),
state: String::new(),
entries: 0,
bytes: 0,
};
let mut it = cells.into_iter();
while let (Some(Reply::Bulk(label)), Some(Reply::Bulk(value))) = (it.next(), it.next()) {
match label.as_slice() {
b"name" => info.name = value,
b"prefix" => info.prefix = value,
b"kind" => info.kind = string(value),
b"state" => info.state = string(value),
b"entries" => info.entries = num_u64(&value)?,
b"bytes" => info.bytes = num_u64(&value)?,
_ => {} }
}
Ok(info)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn embedded_idx_is_unsupported() {
let mut c = Connection::open("mem://").unwrap();
let err = c.idx_list().unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Unsupported);
let err = c
.idx_create_range(b"i", b"user:", b"age", IdxType::I64)
.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Unsupported);
}
#[test]
fn page_parser_maps_cursor_and_rows() {
let reply = Reply::Array(vec![
Reply::Bulk(b"abc1".to_vec()),
Reply::Array(vec![
Reply::Bulk(b"user:1".to_vec()),
Reply::Bulk(b"21".to_vec()),
Reply::Bulk(b"user:2".to_vec()),
Reply::Bulk(b"22".to_vec()),
]),
]);
let page = parse_page(reply).unwrap();
assert_eq!(page.cursor, Some(b"abc1".to_vec()));
assert_eq!(page.rows.len(), 2);
assert_eq!(page.rows[0].key, b"user:1");
assert_eq!(page.rows[0].value, b"21");
let done = parse_page(Reply::Array(vec![
Reply::Bulk(b"0".to_vec()),
Reply::Array(vec![]),
]))
.unwrap();
assert_eq!(done.cursor, None);
assert!(done.rows.is_empty());
}
#[test]
fn ranked_parser_maps_key_score() {
let reply = Reply::Array(vec![Reply::Array(vec![
Reply::Bulk(b"doc:1".to_vec()),
Reply::Bulk(b"1.5".to_vec()),
])]);
assert_eq!(parse_ranked(reply).unwrap(), vec![(b"doc:1".to_vec(), 1.5)]);
}
}