use yo_common::{Code, Error, Result};
use yo_kv::{Db, Keyspace};
use super::args::{self, Args};
use crate::reply::Out;
const DUPLICATE: &str = "duplicate field name in fieldset";
const NO_FIELDSET: &str = "no such fieldset";
const BAD_COUNT: &str = "value count does not match fieldset field count";
struct Fieldset {
name: Vec<u8>,
fields: Vec<(Vec<u8>, usize)>,
}
#[derive(Default)]
pub(super) struct Fieldsets(Vec<Fieldset>);
impl Fieldsets {
fn prepare(&mut self, args: Args<'_>) -> Result<()> {
let mut fields: Vec<(Vec<u8>, usize)> = yo_alloc::allow(|| {
(3..args.len())
.map(|i| (args.get(i).to_vec(), i - 3))
.collect()
});
fields.sort_by(|(a, _), (b, _)| (a.len(), a.as_slice()).cmp(&(b.len(), b.as_slice())));
if fields.windows(2).any(|w| w[0].0 == w[1].0) {
return Err(Error::new(Code::Invalid, DUPLICATE));
}
let name = args.get(2);
yo_alloc::allow(|| match self.0.iter_mut().find(|f| f.name == name) {
Some(old) => old.fields = fields,
None => self.0.push(Fieldset {
name: name.to_vec(),
fields,
}),
});
Ok(())
}
fn get(&self, name: &[u8]) -> Option<&Fieldset> {
self.0.iter().find(|f| f.name == name)
}
fn discard(&mut self, name: &[u8]) -> bool {
let Some(at) = self.0.iter().position(|f| f.name == name) else {
return false;
};
self.0.swap_remove(at);
true
}
fn discard_all(&mut self) -> usize {
let had = self.0.len();
self.0.clear();
had
}
pub(super) fn clear(&mut self) {
self.0.clear();
}
}
pub(super) fn execute(db: &Db, sets: &mut Fieldsets, args: Args<'_>, out: &mut Out) -> Result<()> {
let sub = args.get(1);
if args::is(sub, b"prepare") {
if args.len() < 4 {
return Err(args::wrong_arity_sub("himport", "prepare"));
}
sets.prepare(args)?;
out.ok();
} else if args::is(sub, b"set") {
if args.len() < 5 {
return Err(args::wrong_arity_sub("himport", "set"));
}
set(&mut db.hold(args.get(2)), sets, args, out)?;
} else if args::is(sub, b"discard") {
if args.len() != 3 {
return Err(args::wrong_arity_sub("himport", "discard"));
}
out.int(i64::from(sets.discard(args.get(2))));
} else if args::is(sub, b"discardall") {
if args.len() != 2 {
return Err(args::wrong_arity_sub("himport", "discardall"));
}
out.int(i64::try_from(sets.discard_all()).unwrap_or(i64::MAX));
} else {
return Err(args::unknown_subcommand(sub, "HIMPORT"));
}
Ok(())
}
fn set(db: &mut Keyspace, sets: &Fieldsets, args: Args<'_>, out: &mut Out) -> Result<()> {
let key = args.get(2);
db.hlen(key)?;
let Some(fs) = sets.get(args.get(3)) else {
return Err(Error::new(Code::Invalid, NO_FIELDSET));
};
if args.len() - 4 != fs.fields.len() {
return Err(Error::new(Code::Invalid, BAD_COUNT));
}
db.hreplace(
key,
fs.fields
.iter()
.map(|(field, at)| (field.as_slice(), args.get(at + 4))),
)?;
out.ok();
Ok(())
}