use std::{ptr::null, str::from_utf8};
use crate::{arg_slice::ArgSlice, session_parse_state::SessionParseState};
pub const MAX_ARGUMENT_LENGTH_BYTES: usize = 512 * 1024 * 1024;
#[inline]
pub fn strict_i64(raw: &[u8]) -> Option<i64> {
let (digits, negative) = match raw {
[b'+', rest @ ..] => (rest, false),
[b'-', rest @ ..] => (rest, true),
rest => (rest, false),
};
if digits.is_empty() || digits.len() > 1 && digits[0] == b'0' {
return None;
}
let mut number: u64 = 0;
for &d in digits {
if !d.is_ascii_digit() {
return None;
}
number = number.checked_mul(10)?.checked_add(u64::from(d - b'0'))?;
}
if negative {
if number > i64::MAX as u64 + 1 {
return None;
}
if number == i64::MAX as u64 + 1 {
return Some(i64::MIN);
}
Some(-(number as i64))
} else if number <= i64::MAX as u64 {
Some(number as i64)
} else {
None
}
}
#[inline]
pub fn strict_i32(raw: &[u8]) -> Option<i32> {
i32::try_from(strict_i64(raw)?).ok()
}
#[inline]
fn is_inf_literal(raw: &[u8]) -> bool {
const LITERALS: [&[u8]; 6] = [
b"inf",
b"+inf",
b"-inf",
b"infinity",
b"+infinity",
b"-infinity",
];
LITERALS.iter().any(|l| raw.eq_ignore_ascii_case(l))
}
#[inline]
pub fn strict_f64(raw: &[u8], can_be_infinite: bool) -> Option<f64> {
if let Some(v) = from_utf8(raw).ok().and_then(|t| t.parse::<f64>().ok()) {
if !v.is_nan() && !(v.is_infinite() && is_inf_literal(raw)) {
return Some(v);
}
}
if can_be_infinite {
if raw.eq_ignore_ascii_case(b"INF") || raw.eq_ignore_ascii_case(b"+INF") {
return Some(f64::INFINITY);
}
if raw.eq_ignore_ascii_case(b"-INF") {
return Some(f64::NEG_INFINITY);
}
}
None
}
#[inline]
pub fn strict_f32(raw: &[u8], can_be_infinite: bool) -> Option<f32> {
if let Some(v) = from_utf8(raw).ok().and_then(|t| t.parse::<f32>().ok())
&& !v.is_nan()
&& !(v.is_infinite() && is_inf_literal(raw))
{
return Some(v);
}
if can_be_infinite {
if raw.eq_ignore_ascii_case(b"INF") || raw.eq_ignore_ascii_case(b"+INF") {
return Some(f32::INFINITY);
}
if raw.eq_ignore_ascii_case(b"-INF") {
return Some(f32::NEG_INFINITY);
}
}
None
}
#[inline]
pub fn initialize_with_argument(state: &mut SessionParseState, arg: ArgSlice) {
state.initialize_with_arg(arg);
}
#[inline]
pub fn initialize_with_arguments(state: &mut SessionParseState, args: &[ArgSlice]) {
state.initialize_with_args(args);
}
pub fn set_argument(state: &mut SessionParseState, i: usize, arg: ArgSlice) {
if i >= state.root_buffer.len() {
state.root_buffer.resize(i + 1, ArgSlice::new(null(), 0));
}
state.root_buffer[i] = arg;
if i >= state.count {
state.count = i + 1;
}
}
pub fn set_arguments(state: &mut SessionParseState, start: usize, args: &[ArgSlice]) {
debug_assert!(start + args.len() <= state.count);
for (j, &arg) in args.iter().enumerate() {
state.root_buffer[start + j] = arg;
}
}
#[inline]
pub fn get_int(state: &SessionParseState, i: usize) -> Option<i32> {
try_get_int(state, i)
}
#[inline]
pub fn try_get_int(state: &SessionParseState, i: usize) -> Option<i32> {
strict_i32(state.get_arg_slice_by_ref(i).as_slice())
}
#[inline]
pub fn get_long(state: &SessionParseState, i: usize) -> Option<i64> {
try_get_long(state, i)
}
#[inline]
pub fn try_get_long(state: &SessionParseState, i: usize) -> Option<i64> {
strict_i64(state.get_arg_slice_by_ref(i).as_slice())
}
#[inline]
pub fn get_double(state: &SessionParseState, i: usize, can_be_infinite: bool) -> Option<f64> {
try_get_double(state, i, can_be_infinite)
}
#[inline]
pub fn try_get_double(state: &SessionParseState, i: usize, can_be_infinite: bool) -> Option<f64> {
strict_f64(state.get_arg_slice_by_ref(i).as_slice(), can_be_infinite)
}
#[inline]
pub fn get_float(state: &SessionParseState, i: usize, can_be_infinite: bool) -> Option<f32> {
try_get_float(state, i, can_be_infinite)
}
#[inline]
pub fn try_get_float(state: &SessionParseState, i: usize, can_be_infinite: bool) -> Option<f32> {
strict_f32(state.get_arg_slice_by_ref(i).as_slice(), can_be_infinite)
}
#[inline]
pub fn get_string(state: &SessionParseState, i: usize) -> Option<&str> {
from_utf8(state.get_arg_slice_by_ref(i).as_slice()).ok()
}
#[inline]
pub fn get_bool(state: &SessionParseState, i: usize) -> Option<bool> {
try_get_bool(state, i)
}
#[inline]
pub fn try_get_bool(state: &SessionParseState, i: usize) -> Option<bool> {
match state.get_arg_slice_by_ref(i).as_slice() {
[b'1'] => Some(true),
[b'0'] => Some(false),
_ => None,
}
}
pub fn read(
state: &mut SessionParseState,
i: usize,
buffer: &[u8],
ptr: &mut usize,
end: usize,
) -> bool {
if *ptr + 3 > end || buffer[*ptr] != b'$' {
return false;
}
*ptr += 1;
let negative = matches!(buffer.get(*ptr), Some(b'-'));
if matches!(buffer.get(*ptr), Some(b'+') | Some(b'-')) {
*ptr += 1;
}
let digits_start = *ptr;
let mut length: i64 = 0;
while *ptr < end && buffer[*ptr].is_ascii_digit() {
length = length
.saturating_mul(10)
.saturating_add(i64::from(buffer[*ptr] - b'0'));
*ptr += 1;
}
if *ptr == digits_start || *ptr + 2 > end || &buffer[*ptr..*ptr + 2] != b"\r\n" {
return false;
}
*ptr += 2;
if negative || length < 0 || length > MAX_ARGUMENT_LENGTH_BYTES as i64 {
return false;
}
let length = length as usize;
if *ptr + length + 2 > end {
return false;
}
if &buffer[*ptr + length..*ptr + length + 2] != b"\r\n" {
return false;
}
if i >= state.root_buffer.len() {
state.root_buffer.resize(i + 1, ArgSlice::new(null(), 0));
}
state.root_buffer[i] = ArgSlice::new(buffer[*ptr..].as_ptr(), length);
if i >= state.count {
state.count = i + 1;
}
*ptr += length + 2;
true
}
#[cfg(test)]
mod tests {
use super::*;
fn state_of(args: &[&[u8]]) -> SessionParseState {
let slices: Vec<ArgSlice> = args
.iter()
.map(|a| ArgSlice::new(a.as_ptr(), a.len()))
.collect();
let mut state = SessionParseState::new();
state.initialize_with_args(&slices);
state
}
#[test]
fn typed_getters_roundtrip() {
let state = state_of(&[b"42", b"-7", b"3.5", b"1", b"abc", b"9223372036854775807"]);
assert_eq!(try_get_int(&state, 0), Some(42));
assert_eq!(try_get_int(&state, 1), Some(-7));
assert_eq!(try_get_int(&state, 5), None);
assert_eq!(try_get_long(&state, 5), Some(i64::MAX));
assert_eq!(try_get_long(&state, 4), None);
assert_eq!(try_get_double(&state, 2, false), Some(3.5));
assert_eq!(try_get_double(&state, 4, false), None);
assert_eq!(try_get_bool(&state, 3), Some(true));
assert_eq!(try_get_bool(&state, 0), None);
assert_eq!(get_string(&state, 1), Some("-7"));
}
#[test]
fn strict_int_rejects_leading_zeros_and_allows_sign() {
assert_eq!(strict_i64(b"007"), None);
assert_eq!(strict_i64(b"-007"), None);
assert_eq!(strict_i32(b"01"), None);
assert_eq!(strict_i64(b"0"), Some(0));
assert_eq!(strict_i64(b"-0"), Some(0));
assert_eq!(strict_i64(b"+0"), Some(0));
assert_eq!(strict_i64(b"+5"), Some(5));
assert_eq!(strict_i64(b"-5"), Some(-5));
assert_eq!(strict_i64(b"-9223372036854775808"), Some(i64::MIN));
assert_eq!(strict_i64(b"-9223372036854775809"), None);
assert_eq!(strict_i64(b"9223372036854775808"), None);
assert_eq!(strict_i64(b"5 "), None);
assert_eq!(strict_i64(b""), None);
assert_eq!(strict_i64(b"1x"), None);
assert_eq!(strict_i32(b"2147483647"), Some(i32::MAX));
assert_eq!(strict_i32(b"2147483648"), None);
assert_eq!(strict_i32(b"-2147483648"), Some(i32::MIN));
}
#[test]
fn infinite_gate_and_nan_rejection() {
let state = state_of(&[b"inf", b"-inf", b"nan", b"1e3", b"+inf", b"Infinity"]);
assert_eq!(try_get_double(&state, 0, true), Some(f64::INFINITY));
assert_eq!(try_get_double(&state, 0, false), None);
assert_eq!(try_get_double(&state, 1, true), Some(f64::NEG_INFINITY));
assert_eq!(try_get_double(&state, 4, true), Some(f64::INFINITY));
assert_eq!(try_get_double(&state, 2, true), None);
assert_eq!(try_get_double(&state, 5, true), None);
assert_eq!(try_get_float(&state, 3, false), Some(1000.0));
assert_eq!(strict_f64(b"1e999", false), Some(f64::INFINITY));
}
#[test]
fn read_parses_bulk_argument_and_advances() {
let buffer = b"$3\r\nSET\r\n$2\r\nv1\r\n";
let mut state = SessionParseState::new();
state.initialize(2);
let mut ptr = 0usize;
assert!(read(&mut state, 0, buffer, &mut ptr, buffer.len()));
assert_eq!(state.get_arg_slice_by_ref(0).as_slice(), b"SET");
assert_eq!(ptr, 9);
assert!(read(&mut state, 1, buffer, &mut ptr, buffer.len()));
assert_eq!(state.get_arg_slice_by_ref(1).as_slice(), b"v1");
assert_eq!(ptr, buffer.len());
let partial = b"$5\r\nab";
let mut ptr = 0usize;
assert!(!read(&mut state, 0, partial, &mut ptr, partial.len()));
let empty = b"$0\r\n\r\n";
let mut ptr = 0usize;
assert!(read(&mut state, 0, empty, &mut ptr, empty.len()));
assert_eq!(state.get_arg_slice_by_ref(0).as_slice(), b"");
let neg = b"$-1\r\n";
let mut ptr = 0usize;
assert!(!read(&mut state, 0, neg, &mut ptr, neg.len()));
let mut ptr = 0usize;
assert!(!read(&mut state, 0, b"$1", &mut ptr, 2));
}
#[test]
fn set_argument_grows_count() {
let mut state = SessionParseState::new();
state.initialize(1);
set_argument(&mut state, 2, ArgSlice::new(b"k".as_ptr(), 1));
assert_eq!(state.count, 3);
assert_eq!(state.get_arg_slice_by_ref(2).as_slice(), b"k");
set_arguments(&mut state, 0, &[ArgSlice::new(b"x".as_ptr(), 1)]);
assert_eq!(state.get_arg_slice_by_ref(0).as_slice(), b"x");
}
}