use super::lookup::narrow;
use crate::core::dictionary::DictionaryView;
use crate::core::types::{MAX_TOKEN_SIZE, Token, TokenRange};
pub fn tokenize<V: DictionaryView>(text: &[u8], dict: V) -> Vec<Token> {
let mut tokens = Vec::with_capacity(text.len());
let mut pos = 0;
while pos < text.len() {
let max_len = (text.len() - pos).min(MAX_TOKEN_SIZE);
let mut best: Token = 0;
let mut range = TokenRange {
begin: 0,
last: (dict.num_tokens() - 1) as Token,
};
for k in 0..max_len {
let next = narrow(dict, range, k, text[pos + k]);
if next.is_empty() {
break;
}
if dict.token_len(next.begin) == k + 1 {
best = next.begin;
}
range = next;
}
tokens.push(best);
pos += dict.token_len(best);
}
tokens
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::dictionary::{CompactDictionary, Dictionary, pad_raw};
use crate::test_corpus::{binary_strings, make_raw, random_ascii_strings, user_strings};
use crate::{DECODE_PADDING, DEFAULT_CONFIG, MAX_TOKEN_SIZE, Parser, decode_into, decoded_len};
fn decode_to_vec<V: DictionaryView>(codes: &[Token], dict: V) -> Vec<u8> {
let n = decoded_len(codes, dict);
let mut out = Vec::with_capacity(n + DECODE_PADDING);
let w = unsafe { decode_into(codes, dict, out.spare_capacity_mut()) };
unsafe { out.set_len(w) };
out
}
fn build(extra: &[&[u8]]) -> (CompactDictionary, Vec<Vec<u8>>) {
let mut toks: Vec<Vec<u8>> = (0u16..256).map(|b| vec![b as u8]).collect();
for t in extra {
toks.push(t.to_vec());
}
toks.sort();
toks.dedup();
let mut bytes = Vec::new();
let mut offsets = vec![0u32];
for t in &toks {
bytes.extend_from_slice(t);
offsets.push(bytes.len() as u32);
}
pad_raw(&mut bytes, &offsets);
let dict = CompactDictionary::from_raw(bytes, offsets);
(dict, toks)
}
fn reference(text: &[u8], toks: &[Vec<u8>]) -> Vec<Token> {
let mut out = Vec::new();
let mut pos = 0;
while pos < text.len() {
let (mut best_id, mut best_len) = (0usize, 0usize);
for (id, t) in toks.iter().enumerate() {
if t.len() <= MAX_TOKEN_SIZE
&& t.len() <= text.len() - pos
&& t.len() > best_len
&& text[pos..pos + t.len()] == t[..]
{
best_id = id;
best_len = t.len();
}
}
out.push(best_id as Token);
pos += best_len.max(1); }
out
}
fn check(text: &[u8], extra: &[&[u8]]) {
let (dict, toks) = build(extra);
let want = reference(text, &toks);
let compact = dict.as_view();
assert_eq!(tokenize(text, compact), want, "compact tokens");
let wide = dict.to_wide();
assert_eq!(tokenize(text, wide.as_view()), want, "wide tokens");
assert_eq!(decode_to_vec(&want, compact).as_slice(), text, "round-trip");
}
#[test]
fn empty_text_yields_no_tokens() {
let (dict, _) = build(&[]);
assert!(tokenize(b"", dict.as_view()).is_empty());
}
#[test]
fn single_byte_only_dictionary() {
check(b"hello, world!", &[]);
}
#[test]
fn longest_match_prefers_multibyte_token() {
let (dict, toks) = build(&[b"ab"]);
let ab = toks.iter().position(|t| t == b"ab").unwrap() as Token;
assert_eq!(tokenize(b"abab", dict.as_view()), vec![ab, ab]);
}
#[test]
fn greedy_over_overlapping_tokens() {
let extra: &[&[u8]] = &[b"ab", b"abc", b"abcd", b"bc", b"ca"];
check(b"ababcabcabcd", extra);
}
#[test]
fn matches_reference_on_mixed_corpus() {
let extra: &[&[u8]] = &[
b"th", b"the", b"he", b"er", b"ing", b"in", b"on", b"onpair", b"pair",
];
check(b"the onpair thing is interesting in the morning", extra);
}
#[test]
fn handles_max_length_token() {
let long = [b'z'; MAX_TOKEN_SIZE];
let extra: &[&[u8]] = &[&long, b"zz"];
let text = vec![b'z'; MAX_TOKEN_SIZE + 2]; check(&text, extra);
}
fn to_bytes(strings: Vec<String>) -> Vec<Vec<u8>> {
strings.into_iter().map(String::into_bytes).collect()
}
fn check_against_encoder(train: &[Vec<u8>], inputs: &[Vec<u8>]) {
let traw = make_raw(train);
let parser = Parser::train(&traw.data, &traw.offsets, DEFAULT_CONFIG).unwrap();
let dict = parser.dict.as_view();
let wide = parser.dict.to_wide();
let iraw = make_raw(inputs);
let col = parser.parse(&iraw.data, &iraw.offsets).unwrap();
let view = col.view();
for (i, s) in inputs.iter().enumerate() {
let toks = tokenize(s, dict);
assert_eq!(
toks.as_slice(),
view.row_codes(i),
"row {i}: tokenize vs encoder"
);
assert_eq!(
tokenize(s, wide.as_view()),
toks,
"row {i}: wide vs compact"
);
assert_eq!(
decode_to_vec(&toks, dict).as_slice(),
s.as_slice(),
"row {i}: round-trip"
);
assert!(toks.len() <= s.len(), "row {i}: more tokens than bytes");
}
}
#[test]
fn matches_encoder_on_repetitive_corpus() {
let corpus = to_bytes(user_strings(50));
check_against_encoder(&corpus, &corpus);
}
#[test]
fn matches_encoder_on_binary_strings() {
let corpus = binary_strings(50, 30, 13);
check_against_encoder(&corpus, &corpus);
}
#[test]
fn matches_encoder_on_unseen_strings() {
let train = random_ascii_strings(100, 50, 77);
let unseen = random_ascii_strings(20, 40, 999);
check_against_encoder(&train, &unseen);
}
#[test]
fn every_byte_value_tokenizes_to_its_base_token() {
let corpus = to_bytes(user_strings(20));
let raw = make_raw(&corpus);
let parser = Parser::train(&raw.data, &raw.offsets, DEFAULT_CONFIG).unwrap();
let dict = parser.dict.as_view();
for b in 0u16..256 {
let s = [b as u8];
let toks = tokenize(&s, dict);
assert_eq!(toks.len(), 1, "byte {b}: expected exactly one token");
assert_eq!(dict.token(toks[0]), &s, "byte {b}: reconstruction mismatch");
}
}
}