#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
use super::mask::{self, AsciiMasks};
use super::{decode_cp, is_ascii_ws, is_digit, is_letter, scan_numbers_max3};
use crate::pretokenize::unicode::{O200kCharClass, kimi_class_of, o200k_class_of};
#[inline(always)]
fn is_upper_ascii(b: u8) -> bool {
b.wrapping_sub(b'A') < 26
}
#[inline(always)]
fn is_tail_byte<const SLASH: bool>(b: u8) -> bool {
b == b'\r' || b == b'\n' || (SLASH && b == b'/')
}
#[inline(always)]
fn family_class_of<const HAN: bool>(cp: u32) -> (O200kCharClass, bool) {
if HAN {
let k = kimi_class_of(cp);
(k.base(), k.is_han())
} else {
(o200k_class_of(cp), false)
}
}
#[derive(Clone, Copy)]
enum CaseState {
U { last_cl_end: usize },
L,
}
#[inline(always)]
fn ascii_letter_state(b: u8) -> CaseState {
if is_upper_ascii(b) {
CaseState::U { last_cl_end: 0 }
} else {
CaseState::L
}
}
#[inline(always)]
fn letter_run_first<const HAN: bool>(bytes: &[u8], pos: usize) -> Option<(usize, CaseState)> {
let &b = bytes.get(pos)?;
if is_letter(b) {
return Some((pos + 1, ascii_letter_state(b)));
}
if b >= 0x80 {
let (cp, l) = unsafe { decode_cp(bytes, pos) };
match family_class_of::<HAN>(cp) {
(_, true) => {}
(O200kCharClass::Upper, _) => return Some((pos + l, CaseState::U { last_cl_end: 0 })),
(O200kCharClass::Lower, _) => return Some((pos + l, CaseState::L)),
(O200kCharClass::Caseless | O200kCharClass::Mark, _) => {
return Some((
pos + l,
CaseState::U {
last_cl_end: pos + l,
},
));
}
_ => {}
}
}
None
}
#[inline(always)]
fn scan_case_run<const HAN: bool>(bytes: &[u8], mut pos: usize, mut st: CaseState) -> usize {
let len = bytes.len();
loop {
while pos < len {
let b = unsafe { *bytes.get_unchecked(pos) };
if is_upper_ascii(b) {
if matches!(st, CaseState::L) {
return pos;
}
pos += 1;
} else if is_letter(b) {
st = CaseState::L;
pos += 1;
} else {
break;
}
}
if pos < len && unsafe { *bytes.get_unchecked(pos) } >= 0x80 {
let (cp, l) = unsafe { decode_cp(bytes, pos) };
match family_class_of::<HAN>(cp) {
(_, true) => break,
(O200kCharClass::Upper, _) => {
if matches!(st, CaseState::L) {
return pos;
}
pos += l;
}
(O200kCharClass::Lower, _) => {
st = CaseState::L;
pos += l;
}
(O200kCharClass::Caseless | O200kCharClass::Mark, _) => {
pos += l;
if let CaseState::U {
ref mut last_cl_end,
} = st
{
*last_cl_end = pos;
}
}
_ => break,
}
continue;
}
break;
}
match st {
CaseState::U { last_cl_end } if last_cl_end != 0 => last_cl_end,
_ => pos,
}
}
#[inline(always)]
fn try_suffix<const CONTRACTIONS: bool>(bytes: &[u8], end: usize) -> usize {
if !CONTRACTIONS || bytes.get(end) != Some(&b'\'') {
return end;
}
super::contraction_end(bytes, end).unwrap_or(end)
}
#[inline(always)]
fn scan_punct_from<const HAN: bool>(bytes: &[u8], pos: usize) -> usize {
let len = bytes.len();
let mut p = pos;
loop {
while p < len {
let b = unsafe { *bytes.get_unchecked(p) };
if b >= 0x80 {
break;
}
if is_letter(b) || is_digit(b) || is_ascii_ws(b) {
return p;
}
p += 1;
}
if p < len {
let (cp, l) = unsafe { decode_cp(bytes, p) };
if matches!(
family_class_of::<HAN>(cp).0,
O200kCharClass::Other | O200kCharClass::Mark
) {
p += l;
continue;
}
}
return p;
}
}
#[inline(always)]
fn scan_tail<const SLASH: bool>(bytes: &[u8], mut pos: usize) -> usize {
while pos < bytes.len() && is_tail_byte::<SLASH>(unsafe { *bytes.get_unchecked(pos) }) {
pos += 1;
}
pos
}
#[inline(always)]
fn scan_han_run(bytes: &[u8], mut pos: usize) -> usize {
while pos < bytes.len() {
let b = unsafe { *bytes.get_unchecked(pos) };
if b < 0x80 {
return pos;
}
let (cp, l) = unsafe { decode_cp(bytes, pos) };
if !kimi_class_of(cp).is_han() {
return pos;
}
pos += l;
}
pos
}
#[inline(always)]
fn ws_token_end(bytes: &[u8], start: usize) -> usize {
super::whitespace_token_end::<true>(bytes, start, |codepoint| {
o200k_class_of(codepoint) == O200kCharClass::Whitespace
})
}
#[inline(always)]
fn digit_token_end<const DIGITS3: bool>(bytes: &[u8], first_end: usize) -> usize {
if DIGITS3 {
scan_numbers_max3(bytes, first_end, 1)
} else {
first_end
}
}
#[inline(always)]
pub(crate) fn advance_pos<
const CONTRACTIONS: bool,
const DIGITS3: bool,
const SLASH: bool,
const HAN: bool,
>(
bytes: &[u8],
pos: usize,
) -> usize {
let b0 = unsafe { *bytes.get_unchecked(pos) };
if is_letter(b0) {
let e = scan_case_run::<HAN>(bytes, pos + 1, ascii_letter_state(b0));
return try_suffix::<CONTRACTIONS>(bytes, e);
}
if b0 == b' ' {
let Some(&b1) = bytes.get(pos + 1) else {
return pos + 1; };
if is_letter(b1) {
let e = scan_case_run::<HAN>(bytes, pos + 2, ascii_letter_state(b1));
return try_suffix::<CONTRACTIONS>(bytes, e);
}
if b1 < 0x80 {
if is_digit(b1) {
return pos + 1; }
if is_ascii_ws(b1) {
return ws_token_end(bytes, pos);
}
let p = scan_punct_from::<HAN>(bytes, pos + 2);
return scan_tail::<SLASH>(bytes, p);
}
let (cp, l) = unsafe { decode_cp(bytes, pos + 1) };
let p1 = pos + 1 + l;
return match family_class_of::<HAN>(cp) {
(O200kCharClass::Caseless, true) => pos + 1,
(O200kCharClass::Upper, _) => try_suffix::<CONTRACTIONS>(
bytes,
scan_case_run::<HAN>(bytes, p1, CaseState::U { last_cl_end: 0 }),
),
(O200kCharClass::Lower, _) => {
try_suffix::<CONTRACTIONS>(bytes, scan_case_run::<HAN>(bytes, p1, CaseState::L))
}
(O200kCharClass::Caseless | O200kCharClass::Mark, _) => try_suffix::<CONTRACTIONS>(
bytes,
scan_case_run::<HAN>(bytes, p1, CaseState::U { last_cl_end: p1 }),
),
(O200kCharClass::Whitespace, _) => ws_token_end(bytes, pos),
(O200kCharClass::Number, _) => pos + 1,
(O200kCharClass::Other, _) => {
scan_tail::<SLASH>(bytes, scan_punct_from::<HAN>(bytes, p1))
}
};
}
if b0 >= 0x80 {
let (cp, l) = unsafe { decode_cp(bytes, pos) };
let p0 = pos + l;
let (class, han) = family_class_of::<HAN>(cp);
if HAN && han {
return scan_han_run(bytes, p0);
}
return match class {
O200kCharClass::Upper => try_suffix::<CONTRACTIONS>(
bytes,
scan_case_run::<HAN>(bytes, p0, CaseState::U { last_cl_end: 0 }),
),
O200kCharClass::Lower => {
try_suffix::<CONTRACTIONS>(bytes, scan_case_run::<HAN>(bytes, p0, CaseState::L))
}
O200kCharClass::Caseless | O200kCharClass::Mark => try_suffix::<CONTRACTIONS>(
bytes,
scan_case_run::<HAN>(bytes, p0, CaseState::U { last_cl_end: p0 }),
),
O200kCharClass::Number => digit_token_end::<DIGITS3>(bytes, p0),
class => {
if let Some((e, st)) = letter_run_first::<HAN>(bytes, p0) {
return try_suffix::<CONTRACTIONS>(bytes, scan_case_run::<HAN>(bytes, e, st));
}
if class == O200kCharClass::Whitespace {
ws_token_end(bytes, pos)
} else {
scan_tail::<SLASH>(bytes, scan_punct_from::<HAN>(bytes, p0))
}
}
};
}
if is_digit(b0) {
return digit_token_end::<DIGITS3>(bytes, pos + 1);
}
if b0 == b'\r' || b0 == b'\n' {
return ws_token_end(bytes, pos);
}
if is_ascii_ws(b0) {
if let Some((e, st)) = letter_run_first::<HAN>(bytes, pos + 1) {
return try_suffix::<CONTRACTIONS>(bytes, scan_case_run::<HAN>(bytes, e, st));
}
return ws_token_end(bytes, pos);
}
if let Some((e, st)) = letter_run_first::<HAN>(bytes, pos + 1) {
return try_suffix::<CONTRACTIONS>(bytes, scan_case_run::<HAN>(bytes, e, st));
}
scan_tail::<SLASH>(bytes, scan_punct_from::<HAN>(bytes, pos + 1))
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[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
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[derive(Clone, Copy, Default)]
struct OUni {
l: u64,
u: u64,
cl: u64,
n: u64,
o: u64,
ws: u64,
w2: u64,
w3: u64,
lead2: u64,
lead3: u64,
lead4: u64,
cont: u64,
resid: u64,
mk: u64,
han: u64,
han_leads: u64,
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[derive(Clone, Copy, Default)]
struct OCarries {
pl: u64,
pu: u64,
pcl: u64,
ps: u64,
pwt: u64,
po: u64,
pws: u64,
pd: u64,
phan: u64,
c2_os: u64,
b2b_in: u64,
p_abs: bool,
force_bad_lead: bool,
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
fn prev_tail_absorbed<const SLASH: bool, const HAN: bool>(
bytes: &[u8],
scan: usize,
) -> Option<bool> {
debug_assert!(scan >= 1 && is_tail_byte::<SLASH>(bytes[scan - 1]));
let mut r = scan - 1;
let mut steps = 0;
while r > 0 && is_tail_byte::<SLASH>(bytes[r - 1]) {
r -= 1;
steps += 1;
if steps > 8 {
return None;
}
}
let run = &bytes[r..scan];
let mut trigger = usize::MAX;
let mut seen_slash = false;
for (j, &b) in run.iter().enumerate() {
if b == b'/' {
seen_slash = true;
continue;
}
if seen_slash {
trigger = j;
break;
}
if j == 0 {
if r == 0 {
continue;
}
let pb = bytes[r - 1];
let pred_punct = if pb < 0x80 {
if !is_letter(pb) && !is_digit(pb) && !is_ascii_ws(pb) {
Some(true)
} else {
Some(false)
}
} else {
let mut k = r - 1;
while k > 0 && bytes[k] & 0xC0 == 0x80 {
k -= 1;
}
let (cp, _) = unsafe { decode_cp(bytes, k) };
match family_class_of::<HAN>(cp) {
(O200kCharClass::Other, false) => Some(true),
(O200kCharClass::Mark, _) | (O200kCharClass::Other, true) => None,
_ => Some(false),
}
};
match pred_punct {
Some(true) => {
trigger = 0;
break;
}
Some(false) => {}
None => return None,
}
}
}
Some(scan - 1 - r >= trigger)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn c2_os_ascii<const SLASH: bool, const HAN: bool>(bytes: &[u8], idx: usize) -> Option<u64> {
let b2 = bytes[idx];
if SLASH && b2 == b'/' {
return prev_tail_absorbed::<SLASH, HAN>(bytes, idx + 1).map(|abs| u64::from(!abs));
}
Some(u64::from(
b2 == b' ' || (!is_letter(b2) && !is_digit(b2) && !is_ascii_ws(b2)),
))
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn ascii_carries<const SLASH: bool, const HAN: bool>(bytes: &[u8], scan: usize) -> OCarries {
let b = bytes[scan - 1];
debug_assert!(!is_tail_byte::<SLASH>(b));
let bit = |c: bool| u64::from(c);
let (l, d, w) = (is_letter(b), is_digit(b), is_ascii_ws(b));
let (c2_os, c2_unresolved) = if scan >= 2 {
match c2_os_ascii::<SLASH, HAN>(bytes, scan - 2) {
Some(v) => (v, false),
None => (0, true),
}
} else {
(0, false)
};
OCarries {
force_bad_lead: c2_unresolved,
pl: bit(l),
pu: bit(is_upper_ascii(b)),
pcl: 0,
ps: bit(b == b' '),
pwt: bit(w && b != b' '), po: bit(!l && !d && !w),
pws: bit(w),
pd: bit(d),
phan: 0,
c2_os,
b2b_in: 0,
p_abs: false,
}
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(never)]
fn tail_carries<const SLASH: bool, const HAN: bool>(bytes: &[u8], scan: usize) -> OCarries {
match prev_tail_absorbed::<SLASH, HAN>(bytes, scan) {
Some(true) => OCarries {
p_abs: true,
..OCarries::default()
},
Some(false) => {
let b = bytes[scan - 1];
let bit = |c: bool| u64::from(c);
if scan >= 2 && bytes[scan - 2] >= 0x80 {
return OCarries {
force_bad_lead: true,
..OCarries::default()
};
}
let c2_os = if scan >= 2 {
let b2 = bytes[scan - 2];
bit(b2 == b' ' || (!is_letter(b2) && !is_digit(b2) && !is_ascii_ws(b2)))
} else {
0
};
OCarries {
po: bit(b == b'/'),
pws: bit(b != b'/'),
c2_os,
..OCarries::default()
}
}
None => OCarries {
force_bad_lead: true,
..OCarries::default()
},
}
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
unsafe fn classify_uni_o200k<const HAN: bool>(bytes: &[u8], scan: usize, mut m: u64) -> OUni {
use super::decode_cp_inbounds;
let mut u = OUni::default();
while m != 0 {
let i = m.trailing_zeros() as usize;
m &= m - 1;
let b = bytes[scan + i];
if b < 0xC2 {
u.resid |= 1 << i; continue;
}
let l = if b < 0xE0 {
2
} else if b < 0xF0 {
3
} else {
4
};
let chm = ((1u64 << l) - 1) << i; let lead = 1u64 << i;
let (cp, _) = unsafe { decode_cp_inbounds(bytes, scan + i) };
match family_class_of::<HAN>(cp) {
(O200kCharClass::Caseless, true) => {
u.han |= chm;
u.han_leads |= lead;
}
(O200kCharClass::Upper, _) => {
u.l |= chm;
u.u |= chm;
}
(O200kCharClass::Lower, _) => u.l |= chm,
(O200kCharClass::Caseless, _) => {
u.l |= chm;
u.cl |= chm;
}
(O200kCharClass::Mark, _) | (O200kCharClass::Other, true) => {
u.o |= chm;
u.mk |= chm;
}
(O200kCharClass::Number, _) => {
u.n |= chm;
u.resid |= chm; }
(O200kCharClass::Other, _) => u.o |= chm,
(O200kCharClass::Whitespace, _) => {
u.ws |= chm;
if i + l > 64 || l == 4 {
u.resid |= chm; } else if l == 2 {
u.w2 |= lead;
} else {
u.w3 |= lead;
}
}
}
match l {
2 => u.lead2 |= lead,
3 => u.lead3 |= lead,
_ => u.lead4 |= lead,
}
u.cont |= chm & !lead;
m &= !chm;
}
u
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[derive(Clone, Copy, Default)]
struct OAsciiExtra {
up: u64,
sl: u64,
}
#[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 o200k_extended_masks<
const CONTRACTIONS: bool,
const DIGITS3: bool,
const SLASH: bool,
const HAN: bool,
>(
bytes: &[u8],
scan: usize,
am: AsciiMasks,
ax: OAsciiExtra,
) -> (u64, u64) {
use super::decode_cp_inbounds;
#[inline(always)]
unsafe fn char_through_o200k<const HAN: bool>(
bytes: &[u8],
pos: usize,
) -> (O200kCharClass, bool, usize, usize) {
let b = bytes[pos - 1];
if b < 0x80 {
let cls = if is_upper_ascii(b) {
O200kCharClass::Upper
} else if is_letter(b) {
O200kCharClass::Lower
} else if is_digit(b) {
O200kCharClass::Number
} else if is_ascii_ws(b) {
O200kCharClass::Whitespace
} else {
O200kCharClass::Other
};
return (cls, false, pos - 1, pos);
}
let mut j = pos - 1;
while j > 0 && bytes[j] & 0xC0 == 0x80 {
j -= 1;
}
let (cp, l) = unsafe { decode_cp_inbounds(bytes, j) };
let (cls, han) = family_class_of::<HAN>(cp);
(cls, han, j, j + l)
}
let mut cl = OUni::default();
let cr = if scan == 0 {
OCarries::default()
} else if bytes[scan - 1] < 0x80 && is_tail_byte::<SLASH>(bytes[scan - 1]) {
tail_carries::<SLASH, HAN>(bytes, scan)
} else if bytes[scan - 1] < 0x80 && (scan < 2 || bytes[scan - 2] < 0x80) {
ascii_carries::<SLASH, HAN>(bytes, scan)
} else {
let (c1, h1, j1, e1) = unsafe { char_through_o200k::<HAN>(bytes, scan) };
let chm = if e1 > scan {
(1u64 << (e1 - scan)) - 1
} else {
0
};
cl.cont = chm;
let (c2v, c2_defer) = if j1 == 0 {
(0, false)
} else if SLASH && bytes[j1 - 1] == b'/' {
match prev_tail_absorbed::<SLASH, HAN>(bytes, j1) {
Some(abs) => (u64::from(!abs), false),
None => (0, true),
}
} else {
let (c2c, h2, _, _) = unsafe { char_through_o200k::<HAN>(bytes, j1) };
(
u64::from(
bytes[j1 - 1] == b' '
|| (matches!(c2c, O200kCharClass::Other | O200kCharClass::Mark) && !h2),
),
c2c == O200kCharClass::Mark || (h2 && c2c == O200kCharClass::Other),
)
};
let mut c = OCarries::default();
if e1 > scan {
c.b2b_in = c2v << (e1 - scan);
} else {
c.c2_os = c2v;
}
if c2_defer {
c.force_bad_lead = true;
}
c.pd = u64::from(c1 == O200kCharClass::Number);
match (c1, h1) {
(O200kCharClass::Caseless, true) => {
cl.han |= chm;
c.phan = 1;
}
(O200kCharClass::Upper, _) => {
cl.l |= chm;
cl.u |= chm;
c.pl = 1;
c.pu = 1;
}
(O200kCharClass::Lower, _) => {
cl.l |= chm;
c.pl = 1;
}
(O200kCharClass::Caseless, _) => {
cl.l |= chm;
cl.cl |= chm;
c.pl = 1;
c.pcl = 1;
}
(O200kCharClass::Mark, _) | (O200kCharClass::Other, true) => {
cl.o |= chm;
cl.mk |= chm | 1; c.po = 1;
}
(O200kCharClass::Number, h) => {
cl.n |= chm;
cl.resid |= chm;
if h {
cl.resid |= 1;
}
}
(O200kCharClass::Other, _) => {
cl.o |= chm;
c.po = 1;
}
(O200kCharClass::Whitespace, _) => {
cl.ws |= chm;
if e1 > scan {
cl.resid |= chm;
}
let pb = bytes[scan - 1];
c.ps = u64::from(pb == b' ');
let nl = pb == b'\r' || pb == b'\n';
c.pwt = u64::from(pb < 0x80 && pb != b' ' && !nl || pb >= 0x80);
c.pws = 1;
}
}
c
};
let mut uni = if am.hi != 0 {
unsafe { classify_uni_o200k::<HAN>(bytes, scan, am.hi & !cl.cont) }
} else {
OUni::default()
};
uni.l |= cl.l;
uni.u |= cl.u;
uni.cl |= cl.cl;
uni.n |= cl.n;
uni.o |= cl.o;
uni.ws |= cl.ws;
uni.cont |= cl.cont;
uni.resid |= cl.resid;
uni.mk |= cl.mk;
uni.han |= cl.han;
o200k_algebra::<CONTRACTIONS, DIGITS3, SLASH, HAN>(bytes, scan, am, ax, cr, uni)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn o200k_algebra<
const CONTRACTIONS: bool,
const DIGITS3: bool,
const SLASH: bool,
const HAN: bool,
>(
bytes: &[u8],
scan: usize,
am: AsciiMasks,
ax: OAsciiExtra,
cr: OCarries,
uni: OUni,
) -> (u64, u64) {
let OCarries {
pl,
pu,
pcl,
ps,
pwt,
po,
pws,
pd,
phan,
c2_os,
b2b_in,
p_abs,
force_bad_lead,
} = cr;
let contm = uni.cont;
let lb = am.l | uni.l;
let ub = ax.up | uni.u;
let clb = uni.cl;
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 tcls = if SLASH { am.n | ax.sl } else { am.n };
let abs_seed = (am.n & ((ob << 1) | po)) | (u64::from(p_abs) & tcls);
let abs_t = if abs_seed == 0 {
0
} else {
smear_up(abs_seed, tcls)
};
let ob_eff = ob & !abs_t;
let len1 = !(contm | uni.lead2 | uni.lead3 | uni.lead4);
let c_test = ((ob_eff | 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_u = (ub << 1) | pu;
let p_cl = (clb << 1) | pcl;
let p_s = (sb << 1) | ps;
let p_wt = (wtb << 1) | pwt;
let p_o = (ob_eff << 1) | po;
let absorb = p_o & !b2back;
let p_sl = p_l & !p_u & !p_cl;
let b_su = ub & !contm & p_sl;
let b_letters = (lb & !contm & !p_l & !p_s & !p_wt & !absorb) | b_su;
let b_digits = if DIGITS3 && am.d & (am.d >> 1) != 0 {
mask::digit_run_splits3(am.d)
} else {
am.d
};
let b_punct = ob_eff & !contm & !p_o & !p_s;
let resid = uni.resid;
let mut bad = resid | resid << 1 | resid >> 1;
let mk = uni.mk;
if mk != 0 {
bad |= mk | mk << 1 | mk << 2 | mk << 3 | mk << 4 | mk >> 1;
}
bad |= ub & !contm & ((clb << 1) | pcl);
if force_bad_lead {
bad |= smear_up(tcls & 1, tcls) << 1 | 0b11;
}
let ws_eff = ws_all & !abs_t;
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();
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_t;
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 b_han = if HAN {
uni.han_leads & !((uni.han << 1) | phan)
} else {
0
};
let mut boundary = b_letters | b_digits | b_punct | b_ws | b_han;
if CONTRACTIONS {
let mut cand = am.ap & boundary & p_l & !bad;
let mut last_forced = usize::MAX;
while cand != 0 {
let i = cand.trailing_zeros() as usize;
cand &= cand - 1;
if i <= 2 {
bad |= 0b111u64 << i;
continue;
}
if i >= 61 {
bad |= u64::MAX << i;
break;
}
if i == last_forced {
continue;
}
let p = scan + i;
let prev_suffix_possible = (bytes[p - 2] == b'\''
&& matches!(bytes[p - 1] | 0x20, b's' | b'd' | b'm' | b't'))
|| (p >= 3
&& bytes[p - 3] == b'\''
&& (matches!(
(bytes[p - 2] | 0x20, bytes[p - 1] | 0x20),
(b'l', b'l') | (b'v', b'e') | (b'r', b'e')
) || (bytes[p - 2] == 0xC5 && bytes[p - 1] == 0xBF)));
if prev_suffix_possible {
bad |= 0b111u64 << i;
continue;
}
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);
boundary &= !(((1u64 << (k - 1)) - 1) << (i + 1));
boundary |= 1u64 << (i + k);
last_forced = i + k;
}
}
}
(boundary & !bad, bad)
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[inline(always)]
fn ascii_batch_carries<const SLASH: bool, const HAN: bool>(bytes: &[u8], scan: usize) -> OCarries {
if scan == 0 {
OCarries::default()
} else if is_tail_byte::<SLASH>(bytes[scan - 1]) {
tail_carries::<SLASH, HAN>(bytes, scan)
} else {
ascii_carries::<SLASH, HAN>(bytes, scan)
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
pub(crate) fn batch_masks<
const CONTRACTIONS: bool,
const DIGITS3: bool,
const SLASH: bool,
const HAN: bool,
>(
bytes: &[u8],
scan: usize,
) -> (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 uv = [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];
let mut slv = [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));
uv[i] = vcleq_u8(vsubq_u8(v, 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'\''));
slv[i] = vceqq_u8(v, vdupq_n_u8(b'/'));
}
let l64 = mask::movemask64(lv[0], lv[1], lv[2], lv[3]);
let u64_ = mask::movemask64(uv[0], uv[1], uv[2], uv[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 ap64 = if CONTRACTIONS {
let ap_any = vorrq_u8(vorrq_u8(apv[0], apv[1]), vorrq_u8(apv[2], apv[3]));
if vmaxvq_u8(ap_any) != 0 {
mask::movemask64(apv[0], apv[1], apv[2], apv[3])
} else {
0
}
} else {
0
};
let sl64 = if SLASH {
let sl_any = vorrq_u8(vorrq_u8(slv[0], slv[1]), vorrq_u8(slv[2], slv[3]));
if vmaxvq_u8(sl_any) != 0 {
mask::movemask64(slv[0], slv[1], slv[2], slv[3])
} else {
0
}
} else {
0
};
let am = AsciiMasks {
l: l64,
d: d64,
s: s64,
wt: wsa & !s64 & !n64,
n: n64,
hi: 0,
ap: ap64,
};
let ax = OAsciiExtra { up: u64_, sl: sl64 };
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 o200k_extended_masks::<CONTRACTIONS, DIGITS3, SLASH, HAN>(bytes, scan, am, ax);
}
let cr = ascii_batch_carries::<SLASH, HAN>(bytes, scan);
o200k_algebra::<CONTRACTIONS, DIGITS3, SLASH, HAN>(bytes, scan, am, ax, cr, OUni::default())
}
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
pub(crate) unsafe fn batch_masks_x86<
const AVX512: bool,
const CONTRACTIONS: bool,
const DIGITS3: bool,
const SLASH: bool,
const HAN: bool,
>(
bytes: &[u8],
scan: usize,
) -> (u64, u64) {
let len = bytes.len();
if scan + 70 > len {
return (0, u64::MAX);
}
let (am, ax) = if AVX512 {
unsafe {
(
mask::ascii_masks_avx512(bytes, scan),
ascii_extra_avx512(bytes, scan),
)
}
} else {
unsafe {
(
mask::ascii_masks_avx2(bytes, scan),
ascii_extra_avx2(bytes, scan),
)
}
};
if am.hi != 0
|| (scan >= 1 && bytes[scan - 1] >= 0x80)
|| (scan >= 2 && bytes[scan - 2] >= 0x80)
{
return unsafe {
o200k_extended_masks::<CONTRACTIONS, DIGITS3, SLASH, HAN>(bytes, scan, am, ax)
};
}
let cr = ascii_batch_carries::<SLASH, HAN>(bytes, scan);
o200k_algebra::<CONTRACTIONS, DIGITS3, SLASH, HAN>(bytes, scan, am, ax, cr, OUni::default())
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vl,bmi1,bmi2,lzcnt,popcnt")]
#[inline]
fn ascii_extra_avx512(bytes: &[u8], scan: usize) -> OAsciiExtra {
use std::arch::x86_64::*;
unsafe {
let v = _mm512_loadu_si512(bytes.as_ptr().add(scan) as *const _);
let up = _mm512_cmple_epu8_mask(
_mm512_sub_epi8(v, _mm512_set1_epi8(b'A' as i8)),
_mm512_set1_epi8(25),
);
let sl = _mm512_cmpeq_epi8_mask(v, _mm512_set1_epi8(b'/' as i8));
OAsciiExtra { up, sl }
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,bmi1,bmi2,lzcnt,popcnt")]
#[inline(never)]
fn ascii_extra_avx2(bytes: &[u8], scan: usize) -> OAsciiExtra {
use std::arch::x86_64::*;
unsafe {
let le =
|v: __m256i, lim: __m256i| -> __m256i { _mm256_cmpeq_epi8(_mm256_min_epu8(v, lim), v) };
let mm = |m0: __m256i, m1: __m256i| -> u64 {
(_mm256_movemask_epi8(m0) as u32 as u64)
| ((_mm256_movemask_epi8(m1) as u32 as u64) << 32)
};
let p = bytes.as_ptr().add(scan);
let v0 = _mm256_loadu_si256(p as *const _);
let v1 = _mm256_loadu_si256(p.add(32) as *const _);
let ca = _mm256_set1_epi8(b'A' as i8);
let c25 = _mm256_set1_epi8(25);
let up = mm(
le(_mm256_sub_epi8(v0, ca), c25),
le(_mm256_sub_epi8(v1, ca), c25),
);
let slc = _mm256_set1_epi8(b'/' as i8);
let sl = mm(_mm256_cmpeq_epi8(v0, slc), _mm256_cmpeq_epi8(v1, slc));
OAsciiExtra { up, sl }
}
}