use std::{
array::from_fn,
ops::{RangeFrom, RangeInclusive, RangeToInclusive},
};
use crate::cmd_strings::abort_with_wrong_number_of_arguments;
pub trait ArgCountMatcher {
fn matches(&self, count: usize) -> bool;
}
impl ArgCountMatcher for usize {
#[inline]
fn matches(&self, count: usize) -> bool {
*self == count
}
}
impl ArgCountMatcher for RangeInclusive<usize> {
#[inline]
fn matches(&self, count: usize) -> bool {
self.contains(&count)
}
}
impl ArgCountMatcher for RangeFrom<usize> {
#[inline]
fn matches(&self, count: usize) -> bool {
self.contains(&count)
}
}
impl ArgCountMatcher for RangeToInclusive<usize> {
#[inline]
fn matches(&self, count: usize) -> bool {
self.contains(&count)
}
}
#[macro_export]
macro_rules! check_arg_count {
($cond:expr, $output:expr, $cmd:expr) => {
$crate::check_arg_count!($cond, $output, $cmd, return Ok(true));
};
($cond:expr, $output:expr, $cmd:expr, return $ret:expr) => {
if !($cond) {
$crate::cmd_strings::abort_with_wrong_number_of_arguments($output, $cmd);
return $ret;
}
};
($args:expr, $matcher:expr, $output:expr, $cmd:expr) => {
$crate::check_arg_count!($args, $matcher, $output, $cmd, return Ok(true));
};
($args:expr, $matcher:expr, $output:expr, $cmd:expr, return $ret:expr) => {
if !$crate::check_args::ArgCountMatcher::matches(&($matcher), $args.len()) {
$crate::cmd_strings::abort_with_wrong_number_of_arguments($output, $cmd);
return $ret;
}
};
}
#[inline]
pub fn unpack_args<'a, const N: usize>(
args: &[&'a [u8]],
output: &mut Vec<u8>,
cmd: &str,
) -> Option<[&'a [u8]; N]> {
if args.len() != N {
abort_with_wrong_number_of_arguments(output, cmd);
return None;
}
Some(from_fn(|i| unsafe { *args.get_unchecked(i) }))
}
pub type ArgsWithRest<'a, const N: usize> = ([&'a [u8]; N], &'a [&'a [u8]]);
#[inline]
pub fn unpack_args_rest<'a, const N: usize>(
args: &'a [&'a [u8]],
output: &mut Vec<u8>,
cmd: &str,
) -> Option<ArgsWithRest<'a, N>> {
if args.len() < N {
abort_with_wrong_number_of_arguments(output, cmd);
return None;
}
let (prefix, rest) = args.split_at(N);
Some((from_fn(|i| unsafe { *prefix.get_unchecked(i) }), rest))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_check_arg_count_macro() {
let mut out = Vec::new();
fn run_exact(args: &[&[u8]], out: &mut Vec<u8>) -> Result<bool, ()> {
check_arg_count!(args, 2, out, "CMD");
Ok(false)
}
let a: &[&[u8]] = &[b"k", b"v"];
assert_eq!(run_exact(a, &mut out), Ok(false));
let b: &[&[u8]] = &[b"k"];
assert_eq!(run_exact(b, &mut out), Ok(true));
assert!(out.starts_with(b"-ERR wrong number of arguments"));
out.clear();
fn run_min(args: &[&[u8]], out: &mut Vec<u8>) -> Result<bool, ()> {
check_arg_count!(args, 2.., out, "CMD");
Ok(false)
}
assert_eq!(run_min(b, &mut out), Ok(true));
assert_eq!(run_min(a, &mut out), Ok(false));
out.clear();
fn run_range(args: &[&[u8]], out: &mut Vec<u8>) -> Result<bool, ()> {
check_arg_count!(args, 1..=2, out, "AUTH");
Ok(false)
}
assert_eq!(run_range(a, &mut out), Ok(false));
let c: &[&[u8]] = &[b"1", b"2", b"3"];
assert_eq!(run_range(c, &mut out), Ok(true));
out.clear();
fn run_custom_ret(args: &[&[u8]], out: &mut Vec<u8>) -> bool {
check_arg_count!(args, 1.., out, "CMD", return false);
true
}
let empty: &[&[u8]] = &[];
assert!(!run_custom_ret(empty, &mut out));
assert!(run_custom_ret(a, &mut out));
}
#[test]
fn test_unpack_args() {
let mut out = Vec::new();
fn run_unpack(args: &[&[u8]], out: &mut Vec<u8>) -> Result<(&'static str, usize), ()> {
let Some([k, v]) = unpack_args(args, out, "SET") else {
return Err(());
};
assert_eq!(k, b"key");
assert_eq!(v, b"val");
Ok(("ok", 2))
}
let valid: &[&[u8]] = &[b"key", b"val"];
assert_eq!(run_unpack(valid, &mut out), Ok(("ok", 2)));
let invalid: &[&[u8]] = &[b"key"];
assert_eq!(run_unpack(invalid, &mut out), Err(()));
assert!(out.starts_with(b"-ERR wrong number of arguments"));
out.clear();
fn run_unpack_rest(args: &[&[u8]], out: &mut Vec<u8>) -> Result<(&'static str, usize), ()> {
let Some(([k], rest)) = unpack_args_rest(args, out, "MGET") else {
return Err(());
};
assert_eq!(k, b"key");
Ok(("ok", rest.len()))
}
let multi: &[&[u8]] = &[b"key", b"v1", b"v2"];
assert_eq!(run_unpack_rest(multi, &mut out), Ok(("ok", 2)));
assert_eq!(run_unpack_rest(invalid, &mut out), Ok(("ok", 0)));
assert_eq!(run_unpack_rest(&[], &mut out), Err(()));
}
}