use core::{
mem::size_of,
simd::{
Simd,
cmp::{SimdPartialEq, SimdPartialOrd},
},
};
use crate::unit::Unit;
#[cfg(target_arch = "aarch64")]
pub const TRANS_WIDE: usize = 64;
#[cfg(target_arch = "x86_64")]
pub(crate) const TRANS_WIDE: usize = 64;
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
pub(crate) const TRANS_WIDE: usize = 16;
pub const CMP_WIDE: usize = 16;
pub const NARROW: usize = 8;
#[inline(always)]
pub fn ascii_run<U: Unit, const N: usize>(
data: &[U],
limit: usize,
mut f: impl FnMut(Simd<U, N>, usize),
) -> usize {
debug_assert!(N <= limit && limit <= data.len());
let mut i = 0;
loop {
let v = Simd::<U, N>::from_slice(&data[i..]);
let bad = U::non_ascii_index(v);
f(v, i);
if let Some(p) = bad {
return i + p;
}
i += N;
if i + N > limit {
if i == limit {
return i;
}
i = limit - N;
}
}
}
#[inline]
pub fn raw_eq<U: Unit>(left: &[U], right: &[U]) -> bool {
if left.len() != right.len() {
return false;
}
let left = unsafe { core::slice::from_raw_parts(left.as_ptr().cast::<u8>(), size_of_val(left)) };
let right =
unsafe { core::slice::from_raw_parts(right.as_ptr().cast::<u8>(), size_of_val(right)) };
#[cfg(target_arch = "x86_64")]
{
left == right
}
#[cfg(not(target_arch = "x86_64"))]
{
let byte_len = left.len();
let mut offset = 0;
while byte_len - offset >= 256 {
let difference = ((Simd::<u8, 64>::from_slice(&left[offset..])
^ Simd::from_slice(&right[offset..]))
| (Simd::<u8, 64>::from_slice(&left[offset + 64..])
^ Simd::from_slice(&right[offset + 64..])))
| ((Simd::<u8, 64>::from_slice(&left[offset + 128..])
^ Simd::from_slice(&right[offset + 128..]))
| (Simd::<u8, 64>::from_slice(&left[offset + 192..])
^ Simd::from_slice(&right[offset + 192..])));
if difference != Simd::splat(0) {
return false;
}
offset += 256;
}
while byte_len - offset >= 128 {
let difference = (Simd::<u8, 64>::from_slice(&left[offset..])
^ Simd::from_slice(&right[offset..]))
| (Simd::<u8, 64>::from_slice(&left[offset + 64..])
^ Simd::from_slice(&right[offset + 64..]));
if difference != Simd::splat(0) {
return false;
}
offset += 128;
}
while byte_len - offset >= 16 {
let difference =
Simd::<u8, 16>::from_slice(&left[offset..]) ^ Simd::from_slice(&right[offset..]);
if difference != Simd::splat(0) {
return false;
}
offset += 16;
}
if offset == byte_len {
return true;
}
if byte_len >= 16 {
let difference = Simd::<u8, 16>::from_slice(&left[byte_len - 16..])
^ Simd::from_slice(&right[byte_len - 16..]);
return difference == Simd::splat(0);
}
left[offset..] == right[offset..]
}
}
#[inline(always)]
pub fn raw_eq_ignore_ascii_case<U: Unit, const N: usize>(a: &[U], b: &[U]) -> bool {
if a.len() != b.len() {
return false;
}
if a.len() < N {
return a.iter().zip(b).all(|(&x, &y)| {
let x = x.to_u32();
let y = y.to_u32();
let fx = if x.wrapping_sub('A' as u32) <= 25 {
x ^ 0x20
} else {
x
};
let fy = if y.wrapping_sub('A' as u32) <= 25 {
y ^ 0x20
} else {
y
};
fx == fy
});
}
let mut i = 0;
loop {
let va = U::fold_case(Simd::<U, N>::from_slice(&a[i..]), 'A' as u32);
let vb = U::fold_case(Simd::<U, N>::from_slice(&b[i..]), 'A' as u32);
if va != vb {
return false;
}
i += N;
if i + N > a.len() {
if i == a.len() {
return true;
}
i = a.len() - N;
}
}
}
pub enum CmpStep {
NonAscii,
Equal,
Diff(i64),
}
#[inline(always)]
pub fn ascii_eq_u8_u16<const N: usize>(
a: &[u8],
b: &[u16],
limit: usize,
caseless: bool,
advanced: &mut usize,
) -> CmpStep {
debug_assert!(N <= limit && limit <= a.len() && limit <= b.len());
let mut i = 0;
let step = loop {
let mut va = Simd::<u8, N>::from_slice(&a[i..]);
let mut vb = Simd::<u16, N>::from_slice(&b[i..]);
if u8::non_ascii_index(va).is_some() || u16::non_ascii_index(vb).is_some() {
break CmpStep::NonAscii;
}
if caseless {
va = u8::fold_case(va, 'A' as u32);
vb = u16::fold_case(vb, 'A' as u32);
}
if u8::cast::<u16, N>(va).simd_ne(vb).any() {
break CmpStep::Diff(1);
}
i += N;
if i + N > limit {
if i == limit {
break CmpStep::Equal;
}
i = limit - N;
}
};
*advanced = i;
step
}
#[inline(always)]
pub fn ascii_cmp<A: Unit, B: Unit, const N: usize>(
a: &[A],
b: &[B],
limit: usize,
caseless: bool,
for_eq: bool,
advanced: &mut usize,
) -> CmpStep {
debug_assert!(N <= limit && limit <= a.len() && limit <= b.len());
let mut i = 0;
let step = loop {
let va = A::widen(Simd::<A, N>::from_slice(&a[i..]));
let vb = B::widen(Simd::<B, N>::from_slice(&b[i..]));
if (va | vb).simd_gt(Simd::splat(0x7f)).any() {
break CmpStep::NonAscii;
}
let (fa, fb) = if caseless {
(u32::fold_case(va, 'A' as u32), u32::fold_case(vb, 'A' as u32))
} else {
(va, vb)
};
let ne = fa.simd_ne(fb);
if ne.any() {
if for_eq {
break CmpStep::Diff(1);
}
let p = ne.first_set().unwrap();
break CmpStep::Diff(fa[p] as i64 - fb[p] as i64);
}
i += N;
if i + N > limit {
if i == limit {
break CmpStep::Equal;
}
i = limit - N;
}
};
*advanced = i;
step
}
#[inline]
pub fn plain_prefix<U: Unit>(data: &[U]) -> usize {
if data
.first()
.is_none_or(|u| !(0x20..=0x7e).contains(&u.to_u32()))
{
return 0;
}
if size_of::<U>() == 1 {
if data.len() >= CMP_WIDE {
let v = Simd::<U, CMP_WIDE>::from_slice(data);
if let Some(p) = U::non_plain_index(v) {
return p;
}
}
plain_prefix_n::<U, TRANS_WIDE>(data)
} else {
plain_prefix_n::<U, 16>(data)
}
}
#[inline(always)]
fn plain_prefix_n<U: Unit, const N: usize>(data: &[U]) -> usize {
let mut i = 0;
while i + N <= data.len() {
let v = Simd::<U, N>::from_slice(&data[i..i + N]);
if let Some(p) = U::non_plain_index(v) {
return i + p;
}
i += N;
}
while i < data.len() && (0x20..=0x7e).contains(&data[i].to_u32()) {
i += 1;
}
i
}