use regex::Regex;
pub trait PreTokenizer: Send + Sync {
fn next_match(&self, text: &str, pos: usize) -> Option<(usize, usize)>;
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum FastPath {
None,
Cl100k,
O200k,
Qwen2,
Deepseek,
Tekken,
MiniMax,
Kimi,
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum WhitespaceRules {
Generic,
NewlineFirst,
NewlineFirstSplitOnNumCjk,
}
#[inline]
fn is_deepseek_cjk(c: u32) -> bool {
matches!(c, 0x4E00..=0x9FA5 | 0x3040..=0x30FF)
}
#[inline]
fn is_deepseek_split_boundary(c: char) -> bool {
c.is_numeric() || matches!(c, '一'..='龥' | '\u{3040}'..='\u{309F}' | '\u{30A0}'..='\u{30FF}')
}
pub struct RegexPreTokenizer {
regex: Regex,
fast: FastPath,
ws: WhitespaceRules,
}
impl RegexPreTokenizer {
pub(crate) fn new(pattern: &str, fast: FastPath, ws: WhitespaceRules) -> Self {
Self {
regex: Regex::new(pattern).expect("invalid regex pattern"),
fast,
ws,
}
}
}
impl PreTokenizer for RegexPreTokenizer {
#[inline]
fn next_match(&self, text: &str, pos: usize) -> Option<(usize, usize)> {
let bytes = text.as_bytes();
let fast = match self.fast {
FastPath::Cl100k => cl100k_ascii_next::<3>(bytes, pos),
FastPath::Qwen2 => cl100k_ascii_next::<1>(bytes, pos),
FastPath::O200k => o200k_like_ascii_next::<true, 3, true, false>(bytes, pos),
FastPath::Tekken => o200k_like_ascii_next::<false, 1, true, false>(bytes, pos),
FastPath::MiniMax => o200k_like_ascii_next::<true, 3, true, false>(bytes, pos),
FastPath::Kimi => o200k_like_ascii_next::<true, 3, false, true>(bytes, pos),
FastPath::Deepseek => deepseek_ascii_next(bytes, pos),
FastPath::None => None,
};
if let Some(r) = fast {
return Some(r);
}
let mat = self.regex.find_at(text, pos)?;
let start = mat.start();
let end = adjust_whitespace_end(bytes, start, mat.end(), self.ws);
Some((start, end))
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum CjkClass {
Han,
Caseless,
Upper,
Lower,
Punct,
Num,
Ws,
Other,
}
#[inline]
fn cjk_class(c: u32) -> CjkClass {
use CjkClass::*;
match c {
0x4E00..=0x9FFF | 0x3400..=0x4DBF => Han,
0x3041..=0x3096 | 0x309D..=0x309F => Caseless, 0x30A1..=0x30FA | 0x30FC..=0x30FF => Caseless, 0xAC00..=0xD7A3 => Caseless, 0xFF66..=0xFF9F => Caseless, 0xFF21..=0xFF3A => Upper, 0xFF41..=0xFF5A => Lower, 0x3001..=0x3004 | 0x3008..=0x3020 | 0x3030 | 0x3036 | 0x303D => Punct,
0x30A0 | 0x30FB => Punct, 0xFF01..=0xFF0F | 0xFF1A..=0xFF20 | 0xFF3B..=0xFF40 | 0xFF5B..=0xFF65 => Punct,
0x2014 | 0x2018..=0x201D | 0x2025..=0x2026 => Punct, 0xFF10..=0xFF19 | 0x3007 | 0x3021..=0x3029 | 0x3038..=0x303A => Num,
0x3000 => Ws,
_ => Other,
}
}
#[inline]
fn decode_char(b: &[u8], i: usize) -> (u32, usize) {
let c0 = b[i];
if c0 < 0x80 {
(c0 as u32, 1)
} else if c0 < 0xE0 {
((((c0 & 0x1F) as u32) << 6) | (b[i + 1] & 0x3F) as u32, 2)
} else if c0 < 0xF0 {
(
(((c0 & 0x0F) as u32) << 12)
| (((b[i + 1] & 0x3F) as u32) << 6)
| (b[i + 2] & 0x3F) as u32,
3,
)
} else {
(
(((c0 & 0x07) as u32) << 18)
| (((b[i + 1] & 0x3F) as u32) << 12)
| (((b[i + 2] & 0x3F) as u32) << 6)
| (b[i + 3] & 0x3F) as u32,
4,
)
}
}
#[cold]
#[inline(never)]
fn scan_letter_run_mixed(b: &[u8], mut j: usize) -> Option<usize> {
let n = b.len();
while j < n {
let c = b[j];
if c < 0x80 {
if c.is_ascii_alphabetic() {
j += 1;
continue;
}
return Some(j);
}
let (ch, len) = decode_char(b, j);
match cjk_class(ch) {
CjkClass::Han | CjkClass::Caseless | CjkClass::Upper | CjkClass::Lower => j += len,
CjkClass::Punct | CjkClass::Num | CjkClass::Ws => return Some(j),
CjkClass::Other => return None,
}
}
Some(j)
}
#[cold]
#[inline(never)]
fn scan_punct_run_mixed(b: &[u8], mut j: usize) -> Option<usize> {
let n = b.len();
while j < n {
let c = b[j];
if c < 0x80 {
if !is_ascii_ws(c) && !c.is_ascii_alphanumeric() {
j += 1;
continue;
}
return Some(j);
}
let (ch, len) = decode_char(b, j);
match cjk_class(ch) {
CjkClass::Punct => j += len,
CjkClass::Han
| CjkClass::Caseless
| CjkClass::Upper
| CjkClass::Lower
| CjkClass::Num
| CjkClass::Ws => return Some(j),
CjkClass::Other => return None,
}
}
Some(j)
}
#[inline]
fn take_line_tail<const SLASH: bool>(b: &[u8], mut k: usize) -> usize {
while k < b.len() && (b[k] == b'\r' || b[k] == b'\n' || (SLASH && b[k] == b'/')) {
k += 1;
}
k
}
#[inline(always)]
fn ascii_num_punct<const MAX_DIGITS: usize, const SLASH_TAIL: bool>(
b: &[u8],
i: usize,
) -> Option<(usize, usize)> {
let n = b.len();
let c0 = b[i];
if c0.is_ascii_digit() {
let mut j = i;
let mut k = 0;
while j < n && k < MAX_DIGITS && b[j] < 0x80 && b[j].is_ascii_digit() {
j += 1;
k += 1;
}
if k < MAX_DIGITS && j < n && b[j] >= 0x80 {
return None;
}
return Some((i, j));
}
let mut j = i;
if c0 == b' ' {
match b.get(i + 1) {
Some(&c1)
if c1 < 0x80
&& !is_ascii_ws(c1)
&& !c1.is_ascii_alphabetic()
&& !c1.is_ascii_digit() =>
{
j = i + 1;
}
Some(&c1) if c1 >= 0x80 => {
return space_cjk_punct::<SLASH_TAIL>(b, i);
}
_ => return None,
}
}
let cj = b[j];
if cj < 0x80 && !is_ascii_ws(cj) && !cj.is_ascii_alphabetic() && !cj.is_ascii_digit() {
let mut k = j;
while k < n
&& b[k] < 0x80
&& !is_ascii_ws(b[k])
&& !b[k].is_ascii_alphabetic()
&& !b[k].is_ascii_digit()
{
k += 1;
}
if k < n && b[k] >= 0x80 {
let k = scan_punct_run_mixed(b, k)?;
return Some((i, take_line_tail::<SLASH_TAIL>(b, k)));
}
k = take_line_tail::<SLASH_TAIL>(b, k);
return Some((i, k));
}
None
}
#[cold]
#[inline(never)]
fn space_cjk_punct<const SLASH_TAIL: bool>(b: &[u8], i: usize) -> Option<(usize, usize)> {
if cjk_class(decode_char(b, i + 1).0) != CjkClass::Punct {
return None;
}
let e = scan_punct_run_mixed(b, i + 1)?;
Some((i, take_line_tail::<SLASH_TAIL>(b, e)))
}
#[cold]
#[inline(never)]
fn cl100k_cjk_next(b: &[u8], i: usize) -> Option<(usize, usize)> {
let n = b.len();
let (ch, len) = decode_char(b, i);
match cjk_class(ch) {
CjkClass::Han | CjkClass::Caseless | CjkClass::Upper | CjkClass::Lower => {
scan_letter_run_mixed(b, i + len).map(|e| (i, e))
}
CjkClass::Punct | CjkClass::Ws => {
let j = i + len;
let next_is_letter = j < n && {
let c1 = b[j];
if c1 < 0x80 {
c1.is_ascii_alphabetic()
} else {
matches!(
cjk_class(decode_char(b, j).0),
CjkClass::Han | CjkClass::Caseless | CjkClass::Upper | CjkClass::Lower
)
}
};
if next_is_letter {
return scan_letter_run_mixed(b, j).map(|e| (i, e));
}
if cjk_class(ch) == CjkClass::Ws {
return None; }
let e = scan_punct_run_mixed(b, i + len)?;
Some((i, take_line_tail::<false>(b, e)))
}
CjkClass::Num | CjkClass::Other => None,
}
}
#[cold]
#[inline(never)]
fn cl100k_cjk_after_lead<const MAX_DIGITS: usize>(b: &[u8], i: usize) -> Option<(usize, usize)> {
if matches!(
cjk_class(decode_char(b, i + 1).0),
CjkClass::Han | CjkClass::Caseless | CjkClass::Upper | CjkClass::Lower
) {
return scan_letter_run_mixed(b, i + 1).map(|e| (i, e));
}
ascii_num_punct::<MAX_DIGITS, false>(b, i)
}
#[inline(always)]
fn cl100k_ascii_next<const MAX_DIGITS: usize>(b: &[u8], i: usize) -> Option<(usize, usize)> {
let n = b.len();
if i >= n {
return None;
}
let c0 = b[i];
if c0 >= 0x80 {
return cl100k_cjk_next(b, i);
}
if c0 == b'\''
&& let Some(len) = match_contraction(b, i)
{
return Some((i, i + len));
}
if c0 != b'\r'
&& c0 != b'\n'
&& !c0.is_ascii_alphabetic()
&& !c0.is_ascii_digit()
&& let Some(&c1) = b.get(i + 1)
{
if c1 < 0x80 && c1.is_ascii_alphabetic() {
let mut j = i + 2;
while j < n && b[j] < 0x80 && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return scan_letter_run_mixed(b, j).map(|e| (i, e));
}
return Some((i, j));
}
if c1 >= 0x80 {
return cl100k_cjk_after_lead::<MAX_DIGITS>(b, i);
}
}
if c0.is_ascii_alphabetic() {
let mut j = i + 1;
while j < n && b[j] < 0x80 && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return scan_letter_run_mixed(b, j).map(|e| (i, e));
}
return Some((i, j));
}
ascii_num_punct::<MAX_DIGITS, false>(b, i)
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum LetterKind {
Upper,
Lower,
Both,
End,
Defer,
}
#[inline(always)]
fn o200k_letter_kind<const HAN_APART: bool>(b: &[u8], j: usize) -> (LetterKind, usize) {
let c = b[j];
if c < 0x80 {
if c.is_ascii_uppercase() {
return (LetterKind::Upper, 1);
}
if c.is_ascii_lowercase() {
return (LetterKind::Lower, 1);
}
return (LetterKind::End, 1);
}
o200k_letter_kind_cjk::<HAN_APART>(b, j)
}
#[cold]
#[inline(never)]
fn o200k_letter_kind_cjk<const HAN_APART: bool>(b: &[u8], j: usize) -> (LetterKind, usize) {
let (ch, len) = decode_char(b, j);
let kind = match cjk_class(ch) {
CjkClass::Han => {
if HAN_APART {
LetterKind::End
} else {
LetterKind::Both
}
}
CjkClass::Caseless => LetterKind::Both,
CjkClass::Upper => LetterKind::Upper,
CjkClass::Lower => LetterKind::Lower,
CjkClass::Punct | CjkClass::Num | CjkClass::Ws => LetterKind::End,
CjkClass::Other => LetterKind::Defer,
};
(kind, len)
}
enum CjkStart {
Piece(usize),
Defer,
Punct,
Letters(usize),
}
#[cold]
#[inline(never)]
fn o200k_cjk_start<const HAN_APART: bool>(b: &[u8], i: usize) -> CjkStart {
let n = b.len();
let (ch, len) = decode_char(b, i);
let cls = cjk_class(ch);
if HAN_APART && cls == CjkClass::Han {
let mut j = i + len;
while j < n {
if b[j] < 0x80 {
break;
}
let (c2, l2) = decode_char(b, j);
match cjk_class(c2) {
CjkClass::Han => j += l2,
CjkClass::Other => return CjkStart::Defer,
_ => break,
}
}
return CjkStart::Piece(j);
}
match cls {
CjkClass::Han | CjkClass::Caseless | CjkClass::Upper | CjkClass::Lower => {
CjkStart::Letters(i)
}
CjkClass::Punct | CjkClass::Ws => {
let j = i + len;
let follows_word = j < n && {
let (k1, _) = o200k_letter_kind::<HAN_APART>(b, j);
matches!(k1, LetterKind::Upper | LetterKind::Lower | LetterKind::Both)
};
if follows_word {
CjkStart::Letters(j)
} else if cls == CjkClass::Ws {
CjkStart::Defer
} else {
CjkStart::Punct
}
}
CjkClass::Num | CjkClass::Other => CjkStart::Defer,
}
}
#[inline(always)]
fn o200k_like_ascii_next<
const CONTRACTIONS: bool,
const MAX_DIGITS: usize,
const SLASH_TAIL: bool,
const HAN_APART: bool,
>(
b: &[u8],
i: usize,
) -> Option<(usize, usize)> {
let n = b.len();
if i >= n {
return None;
}
let c0 = b[i];
if c0 >= 0x80 {
return o200k_cjk_next::<CONTRACTIONS, MAX_DIGITS, SLASH_TAIL, HAN_APART>(b, i);
}
let p = if c0.is_ascii_alphabetic() {
i
} else if c0 != b'\r' && c0 != b'\n' && !c0.is_ascii_digit() {
match b.get(i + 1) {
Some(&c1) if c1 < 0x80 && c1.is_ascii_alphabetic() => i + 1,
Some(&c1) if c1 >= 0x80 => {
return o200k_cjk_after_lead::<CONTRACTIONS, MAX_DIGITS, SLASH_TAIL, HAN_APART>(
b, i,
);
}
_ => return ascii_num_punct::<MAX_DIGITS, SLASH_TAIL>(b, i),
}
} else {
return ascii_num_punct::<MAX_DIGITS, SLASH_TAIL>(b, i);
};
let mut q = p;
while q < n && b[q] < 0x80 && b[q].is_ascii_uppercase() {
q += 1;
}
if q < n && b[q] >= 0x80 {
return o200k_word_mixed::<CONTRACTIONS, HAN_APART>(b, i, p);
}
let letters_end = if q > p {
if q < n && b[q].is_ascii_lowercase() {
let mut r = q;
while r < n && b[r] < 0x80 && b[r].is_ascii_lowercase() {
r += 1;
}
if r < n && b[r] >= 0x80 {
return o200k_word_mixed::<CONTRACTIONS, HAN_APART>(b, i, p);
}
r
} else {
q
}
} else {
let mut r = p;
while r < n && b[r] < 0x80 && b[r].is_ascii_lowercase() {
r += 1;
}
if r < n && b[r] >= 0x80 {
return o200k_word_mixed::<CONTRACTIONS, HAN_APART>(b, i, p);
}
r
};
let mut end = letters_end;
if CONTRACTIONS
&& end < n
&& b[end] == b'\''
&& let Some(len) = match_contraction(b, end)
{
end += len;
}
Some((i, end))
}
#[cold]
#[inline(never)]
fn o200k_cjk_next<
const CONTRACTIONS: bool,
const MAX_DIGITS: usize,
const SLASH_TAIL: bool,
const HAN_APART: bool,
>(
b: &[u8],
i: usize,
) -> Option<(usize, usize)> {
match o200k_cjk_start::<HAN_APART>(b, i) {
CjkStart::Piece(e) => Some((i, e)),
CjkStart::Defer => None,
CjkStart::Punct => ascii_num_punct::<MAX_DIGITS, SLASH_TAIL>(b, i),
CjkStart::Letters(p) => o200k_word_mixed::<CONTRACTIONS, HAN_APART>(b, i, p),
}
}
#[cold]
#[inline(never)]
fn o200k_cjk_after_lead<
const CONTRACTIONS: bool,
const MAX_DIGITS: usize,
const SLASH_TAIL: bool,
const HAN_APART: bool,
>(
b: &[u8],
i: usize,
) -> Option<(usize, usize)> {
if matches!(
o200k_letter_kind::<HAN_APART>(b, i + 1).0,
LetterKind::Upper | LetterKind::Lower | LetterKind::Both
) {
return o200k_word_mixed::<CONTRACTIONS, HAN_APART>(b, i, i + 1);
}
ascii_num_punct::<MAX_DIGITS, SLASH_TAIL>(b, i)
}
#[cold]
#[inline(never)]
fn o200k_word_mixed<const CONTRACTIONS: bool, const HAN_APART: bool>(
b: &[u8],
i: usize,
p: usize,
) -> Option<(usize, usize)> {
let n = b.len();
let mut q = p;
let mut last_both_end: Option<usize> = None;
while q < n {
let (kind, len) = o200k_letter_kind::<HAN_APART>(b, q);
match kind {
LetterKind::Upper => q += len,
LetterKind::Both => {
q += len;
last_both_end = Some(q);
}
LetterKind::Lower | LetterKind::End => break,
LetterKind::Defer => return None,
}
}
let next_kind = if q < n {
o200k_letter_kind::<HAN_APART>(b, q).0
} else {
LetterKind::End
};
let letters_end = if q > p {
match next_kind {
LetterKind::Lower => scan_lower_run::<HAN_APART>(b, q)?,
LetterKind::End => last_both_end.unwrap_or(q),
_ => return None,
}
} else {
scan_lower_run::<HAN_APART>(b, p)?
};
let mut end = letters_end;
if CONTRACTIONS
&& end < n
&& b[end] == b'\''
&& let Some(len) = match_contraction(b, end)
{
end += len;
}
Some((i, end))
}
#[inline]
fn scan_lower_run<const HAN_APART: bool>(b: &[u8], mut r: usize) -> Option<usize> {
let n = b.len();
while r < n {
let (kind, len) = o200k_letter_kind::<HAN_APART>(b, r);
match kind {
LetterKind::Lower | LetterKind::Both => r += len,
LetterKind::Upper | LetterKind::End => return Some(r),
LetterKind::Defer => return None,
}
}
Some(r)
}
#[inline]
fn deepseek_ascii_next(b: &[u8], i: usize) -> Option<(usize, usize)> {
let n = b.len();
if i >= n {
return None;
}
let c0 = b[i];
if c0 >= 0x80 {
let (ch, len) = decode_char(b, i);
if is_deepseek_cjk(ch) {
let mut j = i + len;
while j < n && b[j] >= 0x80 {
let (c2, l2) = decode_char(b, j);
if is_deepseek_cjk(c2) {
j += l2;
} else {
break;
}
}
return Some((i, j));
}
return None; }
if c0.is_ascii_digit() {
let mut j = i;
let mut k = 0;
while j < n && k < 3 && b[j].is_ascii_digit() {
j += 1;
k += 1;
}
if k < 3 && j < n && b[j] >= 0x80 {
return None; }
return Some((i, j));
}
if c0.is_ascii_punctuation() {
if let Some(&c1) = b.get(i + 1)
&& c1 < 0x80
&& c1.is_ascii_alphabetic()
{
let mut j = i + 1;
while j < n && b[j].is_ascii_alphabetic() {
j += 1;
}
return Some((i, j));
}
let mut k = i;
while k < n && b[k] < 0x80 && b[k].is_ascii_punctuation() {
k += 1;
}
if k < n && b[k] >= 0x80 {
return None; }
k = take_line_tail::<false>(b, k);
return Some((i, k));
}
if c0.is_ascii_alphabetic() {
let mut j = i;
while j < n && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return None; }
return Some((i, j));
}
if c0 == b' ' {
match b.get(i + 1) {
Some(&c1) if c1 >= 0x80 => return None, Some(&c1) if c1.is_ascii_alphabetic() => {
let mut j = i + 1;
while j < n && b[j].is_ascii_alphabetic() {
j += 1;
}
if j < n && b[j] >= 0x80 {
return None;
}
return Some((i, j));
}
Some(&c1) if c1.is_ascii_punctuation() => {
let mut k = i + 1;
while k < n && b[k] < 0x80 && b[k].is_ascii_punctuation() {
k += 1;
}
if k < n && b[k] >= 0x80 {
return None;
}
k = take_line_tail::<false>(b, k);
return Some((i, k));
}
_ => return None,
}
}
None
}
#[inline]
fn match_contraction(b: &[u8], i: usize) -> Option<usize> {
let c1 = b.get(i + 1).copied()?.to_ascii_lowercase();
match c1 {
b's' | b't' | b'm' | b'd' => Some(2),
b'r' if b.get(i + 2).map(|c| c.to_ascii_lowercase()) == Some(b'e') => Some(3),
b'v' if b.get(i + 2).map(|c| c.to_ascii_lowercase()) == Some(b'e') => Some(3),
b'l' if b.get(i + 2).map(|c| c.to_ascii_lowercase()) == Some(b'l') => Some(3),
_ => None,
}
}
#[inline]
fn adjust_whitespace_end(bytes: &[u8], start: usize, end: usize, ws: WhitespaceRules) -> usize {
if end - start <= 1 || end >= bytes.len() {
return end;
}
if ws != WhitespaceRules::Generic && matches!(bytes[end - 1], b'\r' | b'\n') {
return end;
}
let first = bytes[start];
if first > 0x20 && first < 0x7F {
return end;
}
if ws == WhitespaceRules::NewlineFirstSplitOnNumCjk
&& let Some(next) = bytes[end..].iter().next()
&& (next.is_ascii_digit() || *next >= 0x80)
&& let Some(c) = std::str::from_utf8(&bytes[end..])
.ok()
.and_then(|s| s.chars().next())
&& is_deepseek_split_boundary(c)
{
return end;
}
let piece = &bytes[start..end];
if piece.iter().all(|&b| is_ascii_ws(b)) {
let next = bytes[end];
if is_ascii_ws(next) {
return end;
}
return end - 1;
}
let matched = std::str::from_utf8(&bytes[start..end]).unwrap();
if !matched.chars().all(|c| c.is_whitespace()) {
return end;
}
let tail = std::str::from_utf8(&bytes[end..]).unwrap();
let next_char = match tail.chars().next() {
Some(c) => c,
None => return end,
};
if next_char.is_whitespace() {
return end;
}
let last_len = matched.chars().next_back().unwrap().len_utf8();
if end - last_len <= start {
return end;
}
end - last_len
}
#[inline(always)]
const fn is_ascii_ws(b: u8) -> bool {
matches!(b, b' ' | b'\t' | b'\n' | b'\r' | 0x0B | 0x0C)
}
#[cfg(test)]
mod tests {
use super::*;
fn collect_matches(pt: &dyn PreTokenizer, text: &str) -> Vec<(usize, usize)> {
let mut result = vec![];
let mut pos = 0;
while let Some((start, end)) = pt.next_match(text, pos) {
result.push((start, end));
pos = end;
}
result
}
use crate::encoding::{
CL100K_PATTERN, DEEPSEEK_V3_PATTERN, KIMI_PATTERN, MISTRAL_V3_PATTERN, O200K_PATTERN,
P50K_PATTERN, QWEN2_PATTERN,
};
#[derive(Clone, Copy)]
struct Spec {
pattern: &'static str,
fast: FastPath,
ws: WhitespaceRules,
}
const CL100K: Spec = Spec {
pattern: CL100K_PATTERN,
fast: FastPath::Cl100k,
ws: WhitespaceRules::NewlineFirst,
};
const O200K: Spec = Spec {
pattern: O200K_PATTERN,
fast: FastPath::O200k,
ws: WhitespaceRules::NewlineFirst,
};
const QWEN2: Spec = Spec {
pattern: QWEN2_PATTERN,
fast: FastPath::Qwen2,
ws: WhitespaceRules::NewlineFirst,
};
const DEEPSEEK: Spec = Spec {
pattern: DEEPSEEK_V3_PATTERN,
fast: FastPath::Deepseek,
ws: WhitespaceRules::NewlineFirst,
};
const MISTRAL: Spec = Spec {
pattern: MISTRAL_V3_PATTERN,
fast: FastPath::Tekken,
ws: WhitespaceRules::NewlineFirst,
};
const KIMI: Spec = Spec {
pattern: KIMI_PATTERN,
fast: FastPath::Kimi,
ws: WhitespaceRules::NewlineFirst,
};
const P50K: Spec = Spec {
pattern: P50K_PATTERN,
fast: FastPath::None,
ws: WhitespaceRules::Generic,
};
impl Spec {
fn tokenizer(self) -> RegexPreTokenizer {
RegexPreTokenizer::new(self.pattern, self.fast, self.ws)
}
}
fn reference_matches(spec: Spec, text: &str) -> Vec<(usize, usize)> {
let regex = Regex::new(spec.pattern).unwrap();
let bytes = text.as_bytes();
let mut result = vec![];
let mut pos = 0;
while pos < text.len() {
let mat = match regex.find_at(text, pos) {
Some(m) => m,
None => break,
};
let start = mat.start();
let end = adjust_whitespace_end(bytes, start, mat.end(), spec.ws);
result.push((start, end));
pos = end;
}
result
}
fn assert_fast_matches_reference(spec: Spec, text: &str) {
let pt = spec.tokenizer();
assert_eq!(
reference_matches(spec, text),
collect_matches(&pt, text),
"fast/regex mismatch for {text:?}"
);
}
#[test]
fn test_cl100k_english() {
assert_fast_matches_reference(CL100K, "Hello, world!");
}
#[test]
fn test_cl100k_cjk() {
assert_fast_matches_reference(CL100K, "你好世界");
}
#[test]
fn test_cl100k_contractions() {
assert_fast_matches_reference(CL100K, "I'm don't they're we've she'll it'd");
}
#[test]
fn test_o200k_english() {
assert_fast_matches_reference(O200K, "Hello, world! CamelCase mixedScript123");
}
#[test]
fn test_p50k_english() {
assert_fast_matches_reference(P50K, "Hello world, I'm testing!");
}
#[test]
fn test_empty_input() {
let pt = CL100K.tokenizer();
assert_eq!(collect_matches(&pt, ""), vec![]);
}
#[test]
fn test_only_whitespace() {
assert_fast_matches_reference(CL100K, " \n \t ");
}
#[test]
fn test_emoji() {
assert_fast_matches_reference(CL100K, "🎉🚀💡");
}
#[test]
fn test_mixed_script() {
assert_fast_matches_reference(CL100K, "Hello 你好 World 🌍");
}
use WhitespaceRules::{Generic, NewlineFirst};
#[test]
fn test_adjust_whitespace_single_byte() {
assert_eq!(adjust_whitespace_end(b"a b", 0, 1, Generic), 1);
}
#[test]
fn test_adjust_whitespace_at_end_of_input() {
assert_eq!(adjust_whitespace_end(b" ", 0, 2, Generic), 2);
}
#[test]
fn test_adjust_whitespace_non_ws_piece() {
assert_eq!(adjust_whitespace_end(b"hello world", 0, 5, Generic), 5);
}
#[test]
fn test_adjust_whitespace_trim_before_nonws() {
let bytes = b" x";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 1);
}
#[test]
fn test_adjust_whitespace_no_trim_before_ws() {
let bytes = b" ";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 2);
}
#[test]
fn test_adjust_whitespace_unicode_slow_path() {
let input = "\u{3000}\u{3000}x";
let bytes = input.as_bytes();
assert_eq!(adjust_whitespace_end(bytes, 0, 6, Generic), 3);
}
#[test]
fn test_adjust_whitespace_unicode_followed_by_unicode_ws() {
let input = "\u{3000}\u{3000}\u{3000}";
let bytes = input.as_bytes();
assert_eq!(adjust_whitespace_end(bytes, 0, 6, Generic), 6);
}
#[test]
fn test_adjust_whitespace_single_multibyte_ws_before_nonws() {
let input = "\u{3000}x";
let bytes = input.as_bytes();
assert_eq!(adjust_whitespace_end(bytes, 0, 3, Generic), 3);
}
#[test]
fn test_adjust_whitespace_newline_branch_keeps_double_newline() {
let bytes = b"\n\nx";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, NewlineFirst), 2);
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 1);
}
#[test]
fn test_adjust_whitespace_newline_branch_keeps_crlf() {
let bytes = b"\r\n@";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, NewlineFirst), 2);
assert_eq!(adjust_whitespace_end(bytes, 0, 2, Generic), 1);
}
#[test]
fn test_adjust_whitespace_newline_branch_still_trims_spaces() {
let bytes = b" x";
assert_eq!(adjust_whitespace_end(bytes, 0, 2, NewlineFirst), 1);
}
#[test]
fn test_adjust_whitespace_newline_branch_trims_trailing_spaces_after_newline() {
let bytes = b"\n x";
assert_eq!(adjust_whitespace_end(bytes, 1, 3, NewlineFirst), 2);
}
#[test]
fn test_all_patterns_match_reference() {
let texts = vec![
"Hello, world!",
"你好世界",
"fn main() { }",
" hello ",
"line1\nline2\n",
"café résumé",
"100% of $1,000",
"a@b.com",
" \t\n ",
"",
"a",
"hello world! 你好 🚀 test 123",
"word\n\nnext",
"\r\n@rem",
"a\n\n\nb",
"a \n\n b",
];
for spec in [CL100K, O200K, QWEN2, DEEPSEEK, MISTRAL, KIMI, P50K] {
for text in &texts {
assert_fast_matches_reference(spec, text);
}
}
}
#[test]
fn cjk_class_matches_regex_tables() {
let letter = Regex::new(r"^\p{L}$").unwrap();
let han = Regex::new(r"^\p{Han}$").unwrap();
let caseless = Regex::new(r"^[\p{Lo}\p{Lm}]$").unwrap();
let upper = Regex::new(r"^\p{Lu}$").unwrap();
let lower = Regex::new(r"^\p{Ll}$").unwrap();
let num = Regex::new(r"^\p{N}$").unwrap();
let ws = Regex::new(r"^\s$").unwrap();
let punct = Regex::new(r"^[^\s\p{L}\p{N}]$").unwrap();
let mut buf = [0u8; 4];
for cp in 0x80..=0x10FFFF_u32 {
let Some(c) = char::from_u32(cp) else {
continue;
};
let s: &str = c.encode_utf8(&mut buf);
match cjk_class(cp) {
CjkClass::Han => {
assert!(
han.is_match(s) && caseless.is_match(s),
"U+{cp:04X} claimed Han"
);
}
CjkClass::Caseless => {
assert!(
caseless.is_match(s) && !han.is_match(s),
"U+{cp:04X} claimed caseless letter"
);
}
CjkClass::Upper => assert!(upper.is_match(s), "U+{cp:04X} claimed Lu"),
CjkClass::Lower => assert!(lower.is_match(s), "U+{cp:04X} claimed Ll"),
CjkClass::Num => assert!(num.is_match(s), "U+{cp:04X} claimed N"),
CjkClass::Ws => assert!(ws.is_match(s), "U+{cp:04X} claimed whitespace"),
CjkClass::Punct => {
assert!(
punct.is_match(s),
"U+{cp:04X} claimed [^\\s\\p{{L}}\\p{{N}}]"
);
}
CjkClass::Other => {}
}
let _ = is_deepseek_cjk(cp);
if matches!(
cjk_class(cp),
CjkClass::Han | CjkClass::Caseless | CjkClass::Upper | CjkClass::Lower
) {
assert!(
letter.is_match(s),
"U+{cp:04X} claimed letter but is not \\p{{L}}"
);
}
}
}
#[test]
fn test_cjk_pieces_match_reference() {
let texts = [
"世界",
"你好,世界!",
"、你好",
" 世界",
"ハロー・ワールド",
"世A",
"世AB",
"A世",
"AB世A",
"abc世界",
"世界abc",
"カタカナー",
"パーティー",
"が", "か\u{3099}", "안녕하세요 세계",
"アイウエオ゙",
"ABCabc",
"世's",
"世界。。。",
"……你好……",
"「引用」",
"(括号)",
"第123号",
"3.14",
"一二三四五六七八九十",
"〇一二", "々仕事", "\u{20000}好", "深圳市-广州市",
"FULLwidth",
"ガギグ",
"日本語テスト123テスト",
"。\n、",
"税込1,000円",
"「こんにちは」と言った",
];
for spec in [CL100K, O200K, QWEN2, DEEPSEEK, MISTRAL, KIMI, P50K] {
for text in &texts {
assert_fast_matches_reference(spec, text);
}
}
}
proptest::proptest! {
#![proptest_config(proptest::prelude::ProptestConfig::with_cases(20000))]
#[test]
fn prop_cl100k_fast_matches_regex(text in ".*") {
let pt = CL100K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(CL100K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_cl100k_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = CL100K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(CL100K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_cl100k_fast_matches_regex_newlines(text in "[\r\n \tabc.!]*") {
let pt = CL100K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(CL100K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex(text in ".*") {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex_newlines(text in "[\r\n \tabc.!]*") {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_qwen2_fast_matches_regex(text in ".*") {
let pt = QWEN2.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(QWEN2, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_qwen2_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = QWEN2.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(QWEN2, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_deepseek_fast_matches_regex(text in ".*") {
let pt = DEEPSEEK.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(DEEPSEEK, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_deepseek_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = DEEPSEEK.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(DEEPSEEK, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_mistral_fast_matches_regex(text in ".*") {
let pt = MISTRAL.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(MISTRAL, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_mistral_fast_matches_regex_ascii(text in "[ -~ \t\r\n]*") {
let pt = MISTRAL.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(MISTRAL, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex_slashes(text in "[/\r\n .!abcAB0]*") {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_mistral_fast_matches_regex_slashes(text in "[/\r\n .!abcAB0]*") {
let pt = MISTRAL.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(MISTRAL, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_kimi_fast_matches_regex(text in ".*") {
let pt = KIMI.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(KIMI, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_kimi_fast_matches_regex_slashes(text in "[/\r\n .!abcAB0]*") {
let pt = KIMI.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(KIMI, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_p50k_fast_matches_regex(text in "[ -~ \t\r\n]*") {
let pt = P50K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(P50K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_cl100k_fast_matches_regex_cjk(
text in "[世界你好日本語謎アイウエオぁあんーゟABab한글、。!?()「」・… a-cA-C0-9'\r\n々〇\u{3099}\u{20000}é]*"
) {
let pt = CL100K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(CL100K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_o200k_fast_matches_regex_cjk(
text in "[世界你好日本語謎アイウエオぁあんーゟABab한글、。!?()「」・… a-cA-C0-9'\r\n々〇\u{3099}\u{20000}é]*"
) {
let pt = O200K.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(O200K, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_qwen2_fast_matches_regex_cjk(
text in "[世界你好ぁーア12、。! a-cA-C0-9'\r\n々〇é]*"
) {
let pt = QWEN2.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(QWEN2, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_deepseek_fast_matches_regex_cjk(
text in "[世界你好龥龦ぁゟ゠アヿー、。1 a-c0-9\r\n々é]*"
) {
let pt = DEEPSEEK.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(DEEPSEEK, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_mistral_fast_matches_regex_cjk(
text in "[世界アAa、。! a-cA-C0-9/\r\né]*"
) {
let pt = MISTRAL.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(MISTRAL, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
#[test]
fn prop_kimi_fast_matches_regex_cjk(
text in "[世界你好龥アイぁーAa한、。! a-cA-C0-9'\r\n々\u{20000}é]*"
) {
let pt = KIMI.tokenizer();
let fast = collect_matches(&pt, &text);
let reference = reference_matches(KIMI, &text);
proptest::prop_assert_eq!(fast, reference, "fast/regex mismatch for {:?}", text);
}
}
}