use super::is_ascii_ws;
use super::mask::{self, AsciiMasks};
use crate::pretokenize::unicode::CharClass;
#[inline(always)]
fn smear_up(seed: u64, within: u64) -> u64 {
let mut a = seed;
let mut m = within;
let mut sh = 1u32;
while sh < 64 {
a |= (a << sh) & m;
m &= m << sh;
sh <<= 1;
}
a
}
#[derive(Clone, Copy, Default)]
struct Carries {
pl: u64,
ps: u64,
pwt: u64,
po: u64,
pws: u64,
pd: u64,
c2_os: u64,
b2b_in: u64,
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn ascii_carries(bytes: &[u8], scan: usize) -> Carries {
let b = bytes[scan - 1];
let bit = |c: bool| u64::from(c);
let (l, d, w) = (super::is_letter(b), super::is_digit(b), is_ascii_ws(b));
let n = b == b'\r' || b == b'\n';
let c2_os = if scan >= 2 {
let b2 = bytes[scan - 2];
bit(b2 == b' ' || (!super::is_letter(b2) && !super::is_digit(b2) && !is_ascii_ws(b2)))
} else {
0
};
Carries {
pl: bit(l),
ps: bit(b == b' '),
pwt: bit(w && !n && b != b' '),
po: bit(!l && !d && !w),
pws: bit(w),
pd: bit(d),
c2_os,
b2b_in: 0,
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
pub(crate) fn batch_masks(
bytes: &[u8],
scan: usize,
digits3: bool,
class: impl Fn(u32) -> CharClass + Copy,
) -> (u64, u64) {
use std::arch::aarch64::*;
let len = bytes.len();
if scan + 70 > len {
return (0, u64::MAX);
}
unsafe {
let p = bytes.as_ptr().add(scan);
let zero = vdupq_n_u8(0);
let mut lv = [zero; 4];
let mut dv = [zero; 4];
let mut sv = [zero; 4];
let mut wsv = [zero; 4];
let mut nv = [zero; 4];
let mut hiv = [zero; 4];
let mut apv = [zero; 4];
for i in 0..4 {
let v = vld1q_u8(p.add(16 * i));
let lowered = vorrq_u8(v, vdupq_n_u8(0x20));
lv[i] = vcleq_u8(vsubq_u8(lowered, vdupq_n_u8(b'a')), vdupq_n_u8(25));
dv[i] = vcleq_u8(vsubq_u8(v, vdupq_n_u8(b'0')), vdupq_n_u8(9));
sv[i] = vceqq_u8(v, vdupq_n_u8(b' '));
wsv[i] = vorrq_u8(sv[i], vcleq_u8(vsubq_u8(v, vdupq_n_u8(9)), vdupq_n_u8(4)));
nv[i] = vorrq_u8(
vceqq_u8(v, vdupq_n_u8(b'\r')),
vceqq_u8(v, vdupq_n_u8(b'\n')),
);
hiv[i] = vcltzq_s8(vreinterpretq_s8_u8(v));
apv[i] = vceqq_u8(v, vdupq_n_u8(b'\''));
}
let l64 = mask::movemask64(lv[0], lv[1], lv[2], lv[3]);
let d64 = mask::movemask64(dv[0], dv[1], dv[2], dv[3]);
let s64 = mask::movemask64(sv[0], sv[1], sv[2], sv[3]);
let wsa = mask::movemask64(wsv[0], wsv[1], wsv[2], wsv[3]);
let n64 = mask::movemask64(nv[0], nv[1], nv[2], nv[3]);
let ap_any = vorrq_u8(vorrq_u8(apv[0], apv[1]), vorrq_u8(apv[2], apv[3]));
let ap64 = if vmaxvq_u8(ap_any) != 0 {
mask::movemask64(apv[0], apv[1], apv[2], apv[3])
} else {
0
};
let am = mask::AsciiMasks {
l: l64,
d: d64,
s: s64,
wt: wsa & !s64 & !n64,
n: n64,
hi: 0,
ap: ap64,
};
let hi_any = vorrq_u8(vorrq_u8(hiv[0], hiv[1]), vorrq_u8(hiv[2], hiv[3]));
if vmaxvq_u8(hi_any) != 0
|| (scan >= 1 && bytes[scan - 1] >= 0x80)
|| (scan >= 2 && bytes[scan - 2] >= 0x80)
{
let mut am = am;
am.hi = mask::movemask64(hiv[0], hiv[1], hiv[2], hiv[3]);
return family_extended_masks(bytes, scan, digits3, class, am);
}
let cr = if scan == 0 {
Carries::default()
} else {
ascii_carries(bytes, scan)
};
family_algebra(bytes, scan, digits3, am, cr, mask::UniClasses::default())
}
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
pub(crate) unsafe fn batch_masks_x86<const AVX512: bool>(
bytes: &[u8],
scan: usize,
digits3: bool,
class: impl Fn(u32) -> CharClass + Copy,
) -> (u64, u64) {
let len = bytes.len();
if scan + 70 > len {
return (0, u64::MAX);
}
let am = if AVX512 {
unsafe { mask::ascii_masks_avx512(bytes, scan) }
} else {
unsafe { mask::ascii_masks_avx2(bytes, scan) }
};
if am.hi != 0
|| (scan >= 1 && bytes[scan - 1] >= 0x80)
|| (scan >= 2 && bytes[scan - 2] >= 0x80)
{
return unsafe { family_extended_masks(bytes, scan, digits3, class, am) };
}
let cr = if scan == 0 {
Carries::default()
} else {
ascii_carries(bytes, scan)
};
family_algebra(bytes, scan, digits3, am, cr, mask::UniClasses::default())
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[cfg_attr(
target_arch = "x86_64",
target_feature(enable = "bmi1,bmi2,lzcnt,popcnt")
)]
#[inline(never)]
fn family_extended_masks(
bytes: &[u8],
scan: usize,
digits3: bool,
class: impl Fn(u32) -> CharClass + Copy,
am: mask::AsciiMasks,
) -> (u64, u64) {
let mut cl = mask::UniClasses::default();
let cr = if scan == 0 {
Carries::default()
} else if bytes[scan - 1] < 0x80 && (scan < 2 || bytes[scan - 2] < 0x80) {
ascii_carries(bytes, scan)
} else {
let (c1, j1, e1) = unsafe { mask::char_through(bytes, scan, class) };
let pb = bytes[scan - 1];
let chm = if e1 > scan {
(1u64 << (e1 - scan)) - 1
} else {
0
};
cl.cont = chm;
let c2v = if j1 == 0 {
0
} else {
let c2c = unsafe { mask::char_through(bytes, j1, class) }.0;
u64::from(bytes[j1 - 1] == b' ' || c2c == CharClass::Other)
};
let mut c = Carries::default();
if e1 > scan {
c.b2b_in = c2v << (e1 - scan);
} else {
c.c2_os = c2v;
}
c.pd = u64::from(c1 == CharClass::Number);
match c1 {
CharClass::Letter => {
cl.l = chm;
c.pl = 1;
}
CharClass::Number => {
cl.n = chm;
cl.resid |= chm;
}
CharClass::Other => {
cl.o = chm;
c.po = 1;
}
CharClass::Whitespace => {
cl.ws = chm;
if e1 > scan {
cl.resid = chm;
}
c.ps = u64::from(pb == b' ');
let nl = pb == b'\r' || pb == b'\n';
c.pwt = u64::from(pb != b' ' && !nl);
c.pws = 1;
}
}
c
};
let mut uni = if am.hi != 0 {
unsafe { mask::classify_uni_chars::<false, true>(bytes, scan, am.hi & !cl.cont, class) }
} else {
mask::UniClasses::default()
};
uni.l |= cl.l;
uni.n |= cl.n;
uni.o |= cl.o;
uni.ws |= cl.ws;
uni.cont |= cl.cont;
uni.resid |= cl.resid;
family_algebra(bytes, scan, digits3, am, cr, uni)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn family_algebra(
bytes: &[u8],
scan: usize,
digits3: bool,
am: mask::AsciiMasks,
cr: Carries,
uni: mask::UniClasses,
) -> (u64, u64) {
let Carries {
pl,
ps,
pwt,
po,
pws,
pd,
c2_os,
b2b_in,
} = cr;
let contm = uni.cont;
let resid = uni.resid;
let lb = am.l | uni.l;
let sb = am.s; let wtb = am.wt | uni.ws;
let ob = !(am.l | am.d | am.s | am.wt | am.n | am.hi) | uni.o;
let ws_all = sb | wtb | am.n;
let len1 = !(contm | uni.lead2 | uni.lead3 | uni.lead4);
let c_test = ((ob | sb) << 1) | po | ps; let b2back = ((c_test & len1) << 1)
| ((c_test & uni.lead2) << 2)
| ((c_test & uni.lead3) << 3)
| ((c_test & uni.lead4) << 4)
| c2_os | b2b_in; let p_l = (lb << 1) | pl;
let p_s = (sb << 1) | ps;
let p_wt = (wtb << 1) | pwt;
let p_o = (ob << 1) | po;
let absorb = p_o & !b2back;
let b_letters = lb & !contm & !p_l & !p_s & !p_wt & !absorb;
let b_digits = if digits3 && am.d & (am.d >> 1) != 0 {
mask::digit_run_splits3(am.d)
} else {
am.d
};
let b_punct = ob & !contm & !p_o & !p_s;
let abs_seed = am.n & ((ob << 1) | po);
let abs_n = if abs_seed == 0 {
0
} else {
smear_up(abs_seed, am.n)
};
let ws_eff = ws_all & !abs_n;
let mut bad = resid | resid << 1 | resid >> 1;
let nb64 = bytes[scan + 64]; let nn64 = if nb64 < 0x80 {
!is_ascii_ws(nb64)
} else {
bad >> 63 == 0 && unsafe { mask::nn_at_full(bytes, scan + 64) }
};
let nn64m = u64::from(nn64).wrapping_neg();
if abs_n >> 63 != 0 && !nn64 {
bad |= 1u64 << 63;
}
let nonws = !ws_eff;
if ws_eff >> 63 != 0 && !nn64 {
if nonws == 0 {
return (0, u64::MAX); }
let h = 63 - nonws.leading_zeros(); bad |= u64::MAX << (h + 1);
}
if digits3 {
let seed = (am.d & (bad << 1)) | (am.d & pd);
if seed != 0 {
bad |= smear_up(seed, am.d);
}
}
let ws_leads1 = (am.s | am.wt | am.n) & ws_eff;
let ws_leads = (ws_leads1 | uni.w2 | uni.w3) & !abs_n;
let p_ws = (ws_eff << 1) | pws; let edge_last = (ws_leads1 & (1 << 63)) | (uni.w2 & (1 << 62)) | (uni.w3 & (1 << 61));
let split_ok = (ws_leads1 & (nonws >> 1))
| (uni.w2 & (nonws >> 2))
| (uni.w3 & (nonws >> 3))
| (edge_last & nn64m);
let mut b_ws = ws_leads & (!p_ws | split_ok);
let mut runs_n = am.n & ws_eff & !bad;
while runs_n != 0 {
let f = runs_n.trailing_zeros();
let below_gap = nonws & ((1u64 << f) - 1);
let a = if below_gap == 0 {
0
} else {
64 - below_gap.leading_zeros()
};
let e = (nonws & (u64::MAX << f)).trailing_zeros();
let run_mask = (u64::MAX << a) & !u64::MAX.unbounded_shl(e);
b_ws &= !run_mask;
b_ws |= 1u64 << a;
let q = 63 - (am.n & run_mask).leading_zeros(); if (q + 1) < e {
b_ws |= 1u64 << (q + 1);
let tail = run_mask & (u64::MAX << (q + 1));
let tail_leads = ws_leads & tail;
b_ws |= 1u64 << (63 - tail_leads.leading_zeros());
}
runs_n &= !run_mask;
}
let mut boundary = b_letters | b_digits | b_punct | b_ws;
let mut cand = am.ap & boundary & !bad;
while cand != 0 {
let i = cand.trailing_zeros() as usize;
cand &= cand - 1;
if i >= 61 {
bad |= u64::MAX << i;
break;
}
let b1 = bytes[scan + i + 1];
if b1 >= 0x80 {
bad |= 0b111u64 << i;
continue;
}
let k = match b1 | 0x20 {
b's' | b'd' | b'm' | b't' => 2,
b'l' if bytes[scan + i + 2] | 0x20 == b'l' => 3,
b'v' if bytes[scan + i + 2] | 0x20 == b'e' => 3,
b'r' if bytes[scan + i + 2] | 0x20 == b'e' => 3,
_ => 0,
};
if k != 0 {
boundary &= !(1u64 << (i + 1));
boundary |= 1u64 << (i + k);
}
}
(boundary & !bad, bad)
}