use kevy_resp::{Argv, ArgvView};
use kevy_rt::{BlockHint, BlockKind, Route, Store};
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"XREAD" => xread_block_hint(args),
b"XREADGROUP" => xreadgroup_block_hint(args),
_ => BlockHint::None,
}
}
fn blpop_hint<A: ArgvView + ?Sized>(kind: BlockKind, args: &A) -> BlockHint {
if args.len() != 3 {
return BlockHint::None;
}
let Ok(timeout_str) = std::str::from_utf8(&args[2]) 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;
BlockHint::Block {
kind,
key: args[1].to_vec(),
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(key) = args.get(i + 1) else {
return BlockHint::None;
};
return BlockHint::Block {
kind: BlockKind::XReadBlock,
key: key.to_vec(),
timeout_ms: bm,
};
}
_ => return BlockHint::None,
}
}
BlockHint::None
}
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(key) = args.get(i + 1) else {
return BlockHint::None;
};
return BlockHint::Block {
kind: BlockKind::XReadGroupBlock,
key: key.to_vec(),
timeout_ms: bm,
};
}
_ => return BlockHint::None,
}
}
BlockHint::None
}
pub(crate) fn xread_route<A: ArgvView + ?Sized>(args: &A) -> Route {
let mut i = 1usize;
while i < args.len() {
let upper = args[i].to_ascii_uppercase();
match upper.as_slice() {
b"STREAMS" => {
let key_idx = i + 1;
return if key_idx < args.len() {
Route::Single(key_idx)
} else {
Route::Local
};
}
b"COUNT" | b"BLOCK" => i = i.saturating_add(2),
_ => return Route::Single(1),
}
}
Route::Local
}
pub(crate) fn xreadgroup_route<A: ArgvView + ?Sized>(args: &A) -> Route {
if args.len() < 4 || !args[1].eq_ignore_ascii_case(b"GROUP") {
return Route::Single(1);
}
let mut i = 4usize;
while i < args.len() {
let upper = args[i].to_ascii_uppercase();
match upper.as_slice() {
b"STREAMS" => {
let key_idx = i + 1;
return if key_idx < args.len() {
Route::Single(key_idx)
} else {
Route::Local
};
}
b"COUNT" | b"BLOCK" => i = i.saturating_add(2),
b"NOACK" => 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").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(|id| id.encode())
.unwrap_or_else(|_| arg.to_vec());
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
}