pub fn batch_decode_varints(data: &[u8], count: usize) -> Vec<(u64, usize)> {
let mut results = Vec::with_capacity(count);
let mut offset = 0;
for _ in 0..count {
if offset >= data.len() {
break;
}
match crous_core::varint::decode_varint(data, offset) {
Ok((val, consumed)) => {
results.push((val, consumed));
offset += consumed;
}
Err(_) => break,
}
}
results
}
pub fn batch_decode_total_consumed(data: &[u8], count: usize) -> usize {
let mut offset = 0;
for _ in 0..count {
if offset >= data.len() {
break;
}
match crous_core::varint::decode_varint(data, offset) {
Ok((_val, consumed)) => offset += consumed,
Err(_) => break,
}
}
offset
}
#[cfg(all(feature = "simd-varint", target_arch = "aarch64"))]
mod simd_varint_neon {
use std::arch::aarch64::*;
#[inline]
pub(crate) unsafe fn varint_len_neon(data: &[u8], offset: usize) -> Option<usize> {
let remaining = data.len() - offset;
if remaining == 0 {
return None;
}
if remaining >= 16 {
let ptr = data.as_ptr().add(offset);
let chunk = unsafe { vld1q_u8(ptr) };
let high_bits = unsafe { vshrq_n_u8::<7>(chunk) }; let zero_vec = unsafe { vdupq_n_u8(0) };
let is_terminator = unsafe { vceqq_u8(high_bits, zero_vec) };
let max_val = unsafe { vmaxvq_u8(is_terminator) };
if max_val != 0 {
let mut mask = [0u8; 16];
unsafe { vst1q_u8(mask.as_mut_ptr(), is_terminator) };
for (j, &m) in mask.iter().enumerate() {
if m != 0 {
let len = j + 1;
if len <= 10 {
return Some(len);
} else {
return None; }
}
}
}
None
} else {
scalar_varint_len(data, offset)
}
}
fn scalar_varint_len(data: &[u8], offset: usize) -> Option<usize> {
for i in 0..10.min(data.len() - offset) {
if data[offset + i] & 0x80 == 0 {
return Some(i + 1);
}
}
None
}
}
pub fn batch_decode_varints_simd(data: &[u8], count: usize) -> Vec<(u64, usize)> {
#[cfg(all(feature = "simd-varint", target_arch = "aarch64"))]
{
let mut results = Vec::with_capacity(count);
let mut offset = 0;
for _ in 0..count {
if offset >= data.len() {
break;
}
let vlen = unsafe { simd_varint_neon::varint_len_neon(data, offset) };
match vlen {
Some(len) => {
match crous_core::varint::decode_varint(data, offset) {
Ok((val, consumed)) => {
debug_assert_eq!(consumed, len);
results.push((val, consumed));
offset += consumed;
}
Err(_) => break,
}
}
None => {
match crous_core::varint::decode_varint(data, offset) {
Ok((val, consumed)) => {
results.push((val, consumed));
offset += consumed;
}
Err(_) => break,
}
}
}
}
results
}
#[cfg(not(all(feature = "simd-varint", target_arch = "aarch64")))]
{
batch_decode_varints(data, count)
}
}
#[cfg(target_arch = "aarch64")]
mod neon {
use std::arch::aarch64::*;
#[inline]
pub(crate) unsafe fn find_byte_neon(data: &[u8], needle: u8) -> Option<usize> {
let len = data.len();
let ptr = data.as_ptr();
let needle_vec = unsafe { vdupq_n_u8(needle) };
let mut i = 0;
while i + 16 <= len {
let chunk = unsafe { vld1q_u8(ptr.add(i)) };
let cmp = unsafe { vceqq_u8(chunk, needle_vec) };
let max = unsafe { vmaxvq_u8(cmp) };
if max != 0 {
let mut mask_bytes = [0u8; 16];
unsafe { vst1q_u8(mask_bytes.as_mut_ptr(), cmp) };
for (j, &m) in mask_bytes.iter().enumerate() {
if m != 0 {
return Some(i + j);
}
}
}
i += 16;
}
while i < len {
if unsafe { *ptr.add(i) } == needle {
return Some(i);
}
i += 1;
}
None
}
#[inline]
pub(crate) unsafe fn count_byte_neon(data: &[u8], needle: u8) -> usize {
let len = data.len();
let ptr = data.as_ptr();
let needle_vec = unsafe { vdupq_n_u8(needle) };
let mut total: usize = 0;
let mut i = 0;
while i + 16 <= len {
let chunk = unsafe { vld1q_u8(ptr.add(i)) };
let cmp = unsafe { vceqq_u8(chunk, needle_vec) };
let sum = unsafe { vaddlvq_u8(cmp) } as usize;
total += sum / 255;
i += 16;
}
while i < len {
if unsafe { *ptr.add(i) } == needle {
total += 1;
}
i += 1;
}
total
}
#[inline]
pub(crate) unsafe fn find_non_ascii_neon(data: &[u8]) -> Option<usize> {
let len = data.len();
let ptr = data.as_ptr();
let threshold = unsafe { vdupq_n_u8(0x80) };
let mut i = 0;
while i + 16 <= len {
let chunk = unsafe { vld1q_u8(ptr.add(i)) };
let high_bits = unsafe { vcgeq_u8(chunk, threshold) };
let max = unsafe { vmaxvq_u8(high_bits) };
if max != 0 {
let mut mask_bytes = [0u8; 16];
unsafe { vst1q_u8(mask_bytes.as_mut_ptr(), high_bits) };
for (j, &m) in mask_bytes.iter().enumerate() {
if m != 0 {
return Some(i + j);
}
}
}
i += 16;
}
while i < len {
if unsafe { *ptr.add(i) } >= 0x80 {
return Some(i);
}
i += 1;
}
None
}
}
#[inline]
pub fn find_byte(data: &[u8], needle: u8) -> Option<usize> {
#[cfg(target_arch = "aarch64")]
{
unsafe { neon::find_byte_neon(data, needle) }
}
#[cfg(not(target_arch = "aarch64"))]
{
data.iter().position(|&b| b == needle)
}
}
#[inline]
pub fn count_byte(data: &[u8], needle: u8) -> usize {
#[cfg(target_arch = "aarch64")]
{
unsafe { neon::count_byte_neon(data, needle) }
}
#[cfg(not(target_arch = "aarch64"))]
{
data.iter().filter(|&&b| b == needle).count()
}
}
#[inline]
pub fn find_non_ascii(data: &[u8]) -> Option<usize> {
#[cfg(target_arch = "aarch64")]
{
unsafe { neon::find_non_ascii_neon(data) }
}
#[cfg(not(target_arch = "aarch64"))]
{
data.iter().position(|&b| b >= 0x80)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn batch_decode_basic() {
let mut data = Vec::new();
for v in [0u64, 1, 127, 128, 300] {
crous_core::varint::encode_varint_vec(v, &mut data);
}
let results = batch_decode_varints(&data, 5);
assert_eq!(results.len(), 5);
assert_eq!(results[0].0, 0);
assert_eq!(results[1].0, 1);
assert_eq!(results[2].0, 127);
assert_eq!(results[3].0, 128);
assert_eq!(results[4].0, 300);
}
#[test]
fn batch_decode_simd_matches_scalar() {
let mut data = Vec::new();
let values = [0u64, 1, 42, 127, 128, 255, 300, 16384, u64::MAX];
for v in &values {
crous_core::varint::encode_varint_vec(*v, &mut data);
}
let scalar = batch_decode_varints(&data, values.len());
let simd = batch_decode_varints_simd(&data, values.len());
assert_eq!(scalar.len(), simd.len());
for (s, d) in scalar.iter().zip(simd.iter()) {
assert_eq!(s.0, d.0, "value mismatch");
assert_eq!(s.1, d.1, "consumed mismatch");
}
}
#[test]
fn find_byte_basic() {
assert_eq!(find_byte(b"hello", b'l'), Some(2));
assert_eq!(find_byte(b"hello", b'z'), None);
}
#[test]
fn find_byte_long() {
let data: Vec<u8> = (0..256).map(|i| i as u8).collect();
assert_eq!(find_byte(&data, 0), Some(0));
assert_eq!(find_byte(&data, 42), Some(42));
assert_eq!(find_byte(&data, 255), Some(255));
let zeros = vec![0u8; 100];
assert_eq!(find_byte(&zeros, 1), None);
}
#[test]
fn count_byte_basic() {
assert_eq!(count_byte(b"hello", b'l'), 2);
assert_eq!(count_byte(b"hello", b'z'), 0);
assert_eq!(count_byte(b"hello", b'o'), 1);
}
#[test]
fn count_byte_long() {
let data = vec![0xABu8; 200];
assert_eq!(count_byte(&data, 0xAB), 200);
assert_eq!(count_byte(&data, 0x00), 0);
}
#[test]
fn find_non_ascii_basic() {
assert_eq!(find_non_ascii(b"hello"), None);
assert_eq!(find_non_ascii(b"hello\x80"), Some(5));
assert_eq!(find_non_ascii(b"\xff"), Some(0));
}
#[test]
fn find_non_ascii_long() {
let mut data = vec![b'a'; 100];
assert_eq!(find_non_ascii(&data), None);
data[50] = 0x80;
assert_eq!(find_non_ascii(&data), Some(50));
}
}