use kevy_resp::{Argv, ArgvView};
use kevy_rt::{BlockHint, BlockKind, Route, Store, XGroupCtx};
pub(crate) fn block_hint_for_verb<A: ArgvView + ?Sized>(upper: &[u8], args: &A) -> BlockHint {
match upper {
b"BLPOP" => blpop_hint(BlockKind::Blpop, args),
b"BRPOP" => blpop_hint(BlockKind::Brpop, args),
b"BZPOPMIN" => blpop_hint(BlockKind::Bzpopmin, args),
b"BRPOPLPUSH" => brpoplpush_hint(args),
b"XREAD" => xread_block_hint(args),
b"XREADGROUP" => xreadgroup_block_hint(args),
_ => BlockHint::None,
}
}
fn brpoplpush_hint<A: ArgvView + ?Sized>(args: &A) -> BlockHint {
if args.len() != 4 {
return BlockHint::None;
}
let Ok(s) = std::str::from_utf8(&args[3]) else {
return BlockHint::None;
};
let Ok(secs) = s.parse::<f64>() else {
return BlockHint::None;
};
if !secs.is_finite() || secs < 0.0 {
return BlockHint::None;
}
BlockHint::Block {
kind: BlockKind::Brpoplpush,
keys: vec![args[1].to_vec()],
timeout_ms: (secs * 1000.0) as u64,
}
}
fn blpop_hint<A: ArgvView + ?Sized>(kind: BlockKind, args: &A) -> BlockHint {
if args.len() < 3 {
return BlockHint::None;
}
let timeout_idx = args.len() - 1;
let Ok(timeout_str) = std::str::from_utf8(&args[timeout_idx]) else {
return BlockHint::None;
};
let Ok(secs) = timeout_str.parse::<f64>() else {
return BlockHint::None;
};
if !secs.is_finite() || secs < 0.0 {
return BlockHint::None;
}
let timeout_ms = (secs * 1000.0) as u64;
let keys = (1..timeout_idx).map(|i| args[i].to_vec()).collect();
BlockHint::Block { kind, keys, timeout_ms }
}
fn xread_block_hint<A: ArgvView + ?Sized>(args: &A) -> BlockHint {
let mut block_ms: Option<u64> = None;
let mut i = 1usize;
while i < args.len() {
let upper = args[i].to_ascii_uppercase();
match upper.as_slice() {
b"COUNT" => i = i.saturating_add(2),
b"BLOCK" => {
let Some(ms_arg) = args.get(i + 1) else {
return BlockHint::None;
};
let Ok(s) = std::str::from_utf8(ms_arg) else {
return BlockHint::None;
};
let Ok(ms) = s.parse::<u64>() else {
return BlockHint::None;
};
block_ms = Some(ms);
i = i.saturating_add(2);
}
b"STREAMS" => {
let Some(bm) = block_ms else {
return BlockHint::None;
};
let Some(keys) = streams_keys(args, i + 1) else {
return BlockHint::None;
};
return BlockHint::Block { kind: BlockKind::XReadBlock, keys, timeout_ms: bm };
}
_ => return BlockHint::None,
}
}
BlockHint::None
}
fn streams_keys<A: ArgvView + ?Sized>(args: &A, start: usize) -> Option<Vec<Vec<u8>>> {
let rest = args.len().checked_sub(start)?;
if rest == 0 || !rest.is_multiple_of(2) {
return None;
}
let n = rest / 2;
Some((start..start + n).map(|i| args[i].to_vec()).collect())
}
fn xreadgroup_block_hint<A: ArgvView + ?Sized>(args: &A) -> BlockHint {
if args.len() < 4 || !args[1].eq_ignore_ascii_case(b"GROUP") {
return BlockHint::None;
}
let mut block_ms: Option<u64> = None;
let mut i = 4usize;
while i < args.len() {
let upper = args[i].to_ascii_uppercase();
match upper.as_slice() {
b"COUNT" => i = i.saturating_add(2),
b"BLOCK" => {
let Some(ms_arg) = args.get(i + 1) else {
return BlockHint::None;
};
let Ok(s) = std::str::from_utf8(ms_arg) else {
return BlockHint::None;
};
let Ok(ms) = s.parse::<u64>() else {
return BlockHint::None;
};
block_ms = Some(ms);
i = i.saturating_add(2);
}
b"NOACK" => i = i.saturating_add(1),
b"STREAMS" => {
let Some(bm) = block_ms else {
return BlockHint::None;
};
let Some(keys) = streams_keys(args, i + 1) else {
return BlockHint::None;
};
return BlockHint::Block { kind: BlockKind::XReadGroupBlock, keys, timeout_ms: bm };
}
_ => return BlockHint::None,
}
}
BlockHint::None
}
pub(crate) fn xread_route<A: ArgvView + ?Sized>(args: &A) -> Route {
let mut count: Option<usize> = None;
let mut i = 1usize;
while i < args.len() {
let upper = args[i].to_ascii_uppercase();
match upper.as_slice() {
b"BLOCK" => return Route::Local,
b"COUNT" => {
match args
.get(i + 1)
.and_then(|b| std::str::from_utf8(b).ok())
.and_then(|s| s.parse::<usize>().ok())
{
Some(c) => count = Some(c),
None => return Route::Single(1),
}
i = i.saturating_add(2);
}
b"STREAMS" => return xread_streams_route(args, i + 1, count, None),
_ => return Route::Single(1),
}
}
Route::Local
}
fn xread_streams_route<A: ArgvView + ?Sized>(
args: &A,
start: usize,
count: Option<usize>,
group: Option<XGroupCtx>,
) -> Route {
let Some(rest) = args.len().checked_sub(start) else {
return Route::Single(1);
};
if rest == 0 || !rest.is_multiple_of(2) {
return Route::Single(1);
}
let n = rest / 2;
if n == 1 {
return Route::Single(start);
}
let streams =
(0..n).map(|j| (args[start + j].to_vec(), args[start + n + j].to_vec())).collect();
Route::XReadGather { streams, count, group }
}
pub(crate) fn xreadgroup_route<A: ArgvView + ?Sized>(args: &A) -> Route {
if args.len() < 4 || !args[1].eq_ignore_ascii_case(b"GROUP") {
return if args.len() >= 2 { Route::Single(1) } else { Route::Local };
}
let mut count: Option<usize> = None;
let mut noack = false;
let mut i = 4usize;
while i < args.len() {
let upper = args[i].to_ascii_uppercase();
match upper.as_slice() {
b"BLOCK" => return Route::Local,
b"STREAMS" => {
if i + 1 >= args.len() {
return Route::Local; }
let group =
XGroupCtx { group: args[2].to_vec(), consumer: args[3].to_vec(), noack };
return xread_streams_route(args, i + 1, count, Some(group));
}
b"COUNT" => {
match args
.get(i + 1)
.and_then(|b| std::str::from_utf8(b).ok())
.and_then(|s| s.parse::<usize>().ok())
{
Some(c) => count = Some(c),
None => return Route::Single(1),
}
i = i.saturating_add(2);
}
b"NOACK" => {
noack = true;
i = i.saturating_add(1);
}
_ => return Route::Single(1),
}
}
Route::Local
}
pub(crate) fn wake_idx_for_verb(upper: &[u8]) -> Option<u8> {
matches!(upper, b"LPUSH" | b"RPUSH" | b"XADD" | b"ZADD" | b"ZINCRBY").then_some(1)
}
pub(crate) fn xread_resolve_argv<A: ArgvView + ?Sized>(store: &mut Store, args: &A) -> Argv {
let Some(streams_at) = find_xread_streams_token(args) else {
return args.to_argv();
};
let after = args.len() - (streams_at + 1);
if after == 0 || !after.is_multiple_of(2) {
return args.to_argv();
}
let n = after / 2;
let keys_start = streams_at + 1;
let ids_start = keys_start + n;
let mut out = Argv::default();
for j in 0..args.len() {
let arg = args.get(j).expect("in range");
let pos = j.wrapping_sub(ids_start);
if pos < n && arg == b"$" {
let key = args.get(keys_start + pos).expect("in range");
let resolved = store
.xread_dollar_last_id(key)
.map_or_else(|_| arg.to_vec(), kevy_store::StreamId::encode);
out.push(&resolved);
} else {
out.push(arg);
}
}
out
}
fn find_xread_streams_token<A: ArgvView + ?Sized>(args: &A) -> Option<usize> {
let mut i = 1usize;
while i < args.len() {
let upper = args[i].to_ascii_uppercase();
match upper.as_slice() {
b"STREAMS" => return Some(i),
b"COUNT" | b"BLOCK" => i = i.saturating_add(2),
_ => return None,
}
}
None
}