| |
| |
| |
| |
|
|
| use anyhow::{bail, Context, Result}; |
| use rayon::prelude::*; |
| use serde::Deserialize; |
| use std::cmp::Reverse; |
| use std::collections::BinaryHeap; |
| use std::collections::HashMap; |
| use std::collections::HashSet; |
| use std::fs; |
| use std::path::Path; |
| use unicode_categories::UnicodeCategories; |
|
|
| const PAD_ID: u32 = 0; |
| const BASE_BYTES: u32 = 256; |
| |
| |
| |
| const MAX_SUPPORTED_VOCAB_ID: u32 = 1_048_575; |
| const CONTROL_TOKENS: [&str; 28] = [ |
| "<BOS>", |
| "<EOS>", |
| "<EOD>", |
| "<UNK_UNUSED>", |
| "<|endoftext|>", |
| "<|im_start|>", |
| "<|im_end|>", |
| "<|user|>", |
| "<|assistant|>", |
| "<|system|>", |
| "<|tool|>", |
| "<|fim_prefix|>", |
| "<|fim_middle|>", |
| "<|fim_suffix|>", |
| "<|reasoning|>", |
| "<|repo_context|>", |
| "<|file_context|>", |
| "<|cross_file|>", |
| "<|tool_call|>", |
| "<|scratchpad|>", |
| "<|plan|>", |
| "\n ", |
| "\n ", |
| "\n ", |
| "\n ", |
| "\n ", |
| "\n\t", |
| "\n\t\t", |
| ]; |
| const EXTRA_CONTROL_TOKEN: &str = "\n\t\t\t"; |
| const KEYWORDS: [&str; 48] = [ |
| "__peg_parser__", |
| "undefined", |
| "continue", |
| "function", |
| "nonlocal", |
| "extends", |
| "finally", |
| "assert", |
| "except", |
| "export", |
| "global", |
| "import", |
| "lambda", |
| "return", |
| "static", |
| "False", |
| "async", |
| "await", |
| "break", |
| "class", |
| "const", |
| "false", |
| "raise", |
| "while", |
| "yield", |
| "None", |
| "True", |
| "elif", |
| "else", |
| "enum", |
| "from", |
| "null", |
| "pass", |
| "true", |
| "with", |
| "and", |
| "def", |
| "del", |
| "for", |
| "let", |
| "not", |
| "try", |
| "var", |
| "as", |
| "if", |
| "in", |
| "is", |
| "or", |
| ]; |
| const OPERATORS: [&str; 36] = [ |
| ">>>=", "===", "!==", ">>=", "<<=", "&&=", "||=", "??=", "?.", "...", "**=", "//=", ":=", "->", |
| "=>", "==", "!=", "<=", ">=", "++", "--", ">>>", "<<", ">>", "**", "//", "+=", "-=", "*=", |
| "/=", "%=", "&=", "|=", "^=", "::", "@=", |
| ]; |
| const SHORT_OPERATORS: [&str; 3] = ["??", "&&", "||"]; |
| const CONTRACTIONS: [&str; 7] = ["'re", "'ve", "'ll", "'d", "'m", "'s", "'t"]; |
|
|
| #[derive(Debug, Deserialize)] |
| struct MergeFile { |
| version: u32, |
| rules: Vec<MergeRule>, |
| } |
|
|
| #[derive(Debug, Deserialize)] |
| struct MergeRule { |
| left: String, |
| right: String, |
| rank: u32, |
| } |
|
|
| |
| |
| |
| |
| pub struct StenToken { |
| bpe: StenBpe, |
| id_to_token: Vec<String>, |
| control_ids: HashMap<String, u32>, |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| struct StenBpe { |
| ranks: HashMap<(String, String), u32>, |
| piece_to_id: HashMap<String, u32>, |
| } |
|
|
| #[derive(Debug, Eq, Ord, PartialEq, PartialOrd)] |
| struct MergeCandidate { |
| rank: u32, |
| left: usize, |
| left_version: u64, |
| right: usize, |
| right_version: u64, |
| } |
|
|
| #[derive(Debug)] |
| struct BpeNode { |
| piece: String, |
| prev: Option<usize>, |
| next: Option<usize>, |
| alive: bool, |
| version: u64, |
| } |
|
|
| impl StenBpe { |
| fn new(ranks: HashMap<(String, String), u32>, piece_to_id: HashMap<String, u32>) -> Self { |
| Self { ranks, piece_to_id } |
| } |
|
|
| fn encode_internal(&self, internal: &str) -> Vec<u32> { |
| if internal.is_empty() { |
| return Vec::new(); |
| } |
| if let Some(&id) = self.piece_to_id.get(internal) { |
| return vec![id]; |
| } |
| let pieces: Vec<String> = internal.chars().map(|ch| ch.to_string()).collect(); |
| if pieces.len() == 1 { |
| return self.piece_to_ids(&pieces[0]); |
| } |
| let mut nodes: Vec<BpeNode> = pieces |
| .into_iter() |
| .enumerate() |
| .map(|(index, piece)| BpeNode { |
| piece, |
| prev: index.checked_sub(1), |
| next: None, |
| alive: true, |
| version: 0, |
| }) |
| .collect(); |
| for index in 0..nodes.len().saturating_sub(1) { |
| nodes[index].next = Some(index + 1); |
| } |
| let mut heap: BinaryHeap<Reverse<MergeCandidate>> = BinaryHeap::new(); |
| let mut blocked: HashSet<(String, String)> = HashSet::new(); |
| for index in 0..nodes.len().saturating_sub(1) { |
| self.push_candidate(&nodes, &mut heap, index); |
| } |
|
|
| while let Some(Reverse(candidate)) = heap.pop() { |
| let (left_alive, right_alive) = |
| (nodes[candidate.left].alive, nodes[candidate.right].alive); |
| if !left_alive || !right_alive { |
| continue; |
| } |
| if nodes[candidate.left].version != candidate.left_version |
| || nodes[candidate.right].version != candidate.right_version |
| || nodes[candidate.left].next != Some(candidate.right) |
| || nodes[candidate.right].prev != Some(candidate.left) |
| { |
| continue; |
| } |
| let pair = ( |
| nodes[candidate.left].piece.clone(), |
| nodes[candidate.right].piece.clone(), |
| ); |
| if blocked.contains(&pair) || self.ranks.get(&pair) != Some(&candidate.rank) { |
| continue; |
| } |
| let merged = format!("{}{}", pair.0, pair.1); |
| if surface_has_reserved_control(&merged) { |
| blocked.insert(pair); |
| continue; |
| } |
|
|
| let left_prev = nodes[candidate.left].prev; |
| let right_next = nodes[candidate.right].next; |
| if let Some(previous) = left_prev { |
| nodes[previous].next = Some(candidate.left); |
| nodes[previous].version += 1; |
| } |
| if let Some(next) = right_next { |
| nodes[next].prev = Some(candidate.left); |
| nodes[next].version += 1; |
| } |
| nodes[candidate.left].piece = merged; |
| nodes[candidate.left].next = right_next; |
| nodes[candidate.left].version += 1; |
| nodes[candidate.right].alive = false; |
| nodes[candidate.right].prev = None; |
| nodes[candidate.right].next = None; |
| nodes[candidate.right].version += 1; |
|
|
| |
| |
| |
| if let Some(previous) = left_prev { |
| self.push_candidate(&nodes, &mut heap, previous); |
| } |
| self.push_candidate(&nodes, &mut heap, candidate.left); |
| } |
|
|
| let mut start = 0; |
| while start < nodes.len() && nodes[start].prev.is_some() { |
| start += 1; |
| } |
| let mut ids = Vec::new(); |
| let mut current = (start < nodes.len()).then_some(start); |
| while let Some(index) = current { |
| ids.extend(self.piece_to_ids(&nodes[index].piece)); |
| current = nodes[index].next; |
| } |
| ids |
| } |
|
|
| fn push_candidate( |
| &self, |
| nodes: &[BpeNode], |
| heap: &mut BinaryHeap<Reverse<MergeCandidate>>, |
| left: usize, |
| ) { |
| if !nodes[left].alive { |
| return; |
| } |
| let Some(right) = nodes[left].next else { |
| return; |
| }; |
| if !nodes[right].alive { |
| return; |
| } |
| let pair = (nodes[left].piece.clone(), nodes[right].piece.clone()); |
| if let Some(&rank) = self.ranks.get(&pair) { |
| heap.push(Reverse(MergeCandidate { |
| rank, |
| left, |
| left_version: nodes[left].version, |
| right, |
| right_version: nodes[right].version, |
| })); |
| } |
| } |
|
|
| fn piece_to_ids(&self, piece: &str) -> Vec<u32> { |
| if let Some(&id) = self.piece_to_id.get(piece) { |
| return vec![id]; |
| } |
| let mut ids = Vec::with_capacity(piece.len()); |
| for ch in piece.chars() { |
| if let Some(byte) = internal_char_to_byte(ch) { |
| ids.push(byte as u32 + 1); |
| } else { |
| ids.extend(ch.to_string().bytes().map(|byte| byte as u32 + 1)); |
| } |
| } |
| ids |
| } |
| } |
|
|
| impl StenToken { |
| |
| |
| |
| pub fn from_bundle(bundle_dir: impl AsRef<Path>) -> Result<Self> { |
| let bundle = bundle_dir.as_ref(); |
| let vocab_path = bundle.join("vocab.json"); |
| let ranks_path = bundle.join("mergeable_ranks.json"); |
| let state_path = bundle.join("sten_tokenizer.pt"); |
| for required in [&vocab_path, &ranks_path, &state_path] { |
| if !required.is_file() { |
| bail!("StenToken bundle is missing {}", required.display()); |
| } |
| } |
| if fs::metadata(&state_path) |
| .with_context(|| format!("reading metadata for {}", state_path.display()))? |
| .len() |
| == 0 |
| { |
| bail!("sten_tokenizer.pt must not be empty"); |
| } |
|
|
| let vocab: HashMap<String, u32> = serde_json::from_slice( |
| &fs::read(&vocab_path).with_context(|| format!("reading {}", vocab_path.display()))?, |
| ) |
| .with_context(|| format!("parsing {}", vocab_path.display()))?; |
| if vocab.get("<PAD>") != Some(&PAD_ID) { |
| bail!("vocab.json must map <PAD> to ID 0"); |
| } |
| let max_id = vocab.values().copied().max().unwrap_or(PAD_ID); |
| if max_id > MAX_SUPPORTED_VOCAB_ID { |
| bail!( |
| "vocab.json contains ID {max_id}, above the supported maximum {MAX_SUPPORTED_VOCAB_ID}" |
| ); |
| } |
| let mut id_slots = vec![None::<String>; max_id as usize + 1]; |
| for (token, &id) in &vocab { |
| let destination = &mut id_slots[id as usize]; |
| if destination.is_some() { |
| bail!("vocab.json assigns ID {id} to multiple token surfaces"); |
| } |
| *destination = Some(token.clone()); |
| } |
| let id_to_token = id_slots |
| .into_iter() |
| .map(Option::unwrap_or_default) |
| .collect::<Vec<_>>(); |
|
|
| let ranks: MergeFile = serde_json::from_slice( |
| &fs::read(&ranks_path).with_context(|| format!("reading {}", ranks_path.display()))?, |
| ) |
| .with_context(|| format!("parsing {}", ranks_path.display()))?; |
| if ranks.version != 1 { |
| bail!("unsupported mergeable_ranks.json version {}", ranks.version); |
| } |
| let mut rules = ranks.rules; |
| |
| |
| |
| |
| rules.sort_by(|left, right| { |
| left.rank |
| .cmp(&right.rank) |
| .then_with(|| left.left.cmp(&right.left)) |
| .then_with(|| left.right.cmp(&right.right)) |
| }); |
| let control_ids = known_control_ids(); |
| let piece_to_id: HashMap<String, u32> = vocab |
| .iter() |
| .filter_map(|(piece, &id)| { |
| (id != PAD_ID && !control_ids.values().any(|control_id| *control_id == id)) |
| .then_some((piece.clone(), id)) |
| }) |
| .collect(); |
|
|
| let bpe = StenBpe::new( |
| rules |
| .into_iter() |
| .map(|rule| ((rule.left, rule.right), rule.rank)) |
| .collect(), |
| piece_to_id.clone(), |
| ); |
|
|
| for (token, id) in &control_ids { |
| if id_to_token.get(*id as usize).map(String::as_str) != Some(token.as_str()) { |
| bail!("vocab.json control token {token:?} is not at its canonical ID {id}"); |
| } |
| } |
| Ok(Self { |
| bpe, |
| id_to_token, |
| control_ids, |
| }) |
| } |
|
|
| |
| pub fn encode(&self, text: &str) -> Result<Vec<u32>> { |
| let capcoded = capcode_encode(text); |
| let chunks = pre_tokenize_capcoded(&capcoded); |
| let mut ids = Vec::with_capacity(text.len() / 2); |
| for chunk in chunks { |
| let internal = text_to_internal_bytespace(&chunk); |
| ids.extend(self.bpe.encode_internal(&internal)); |
| } |
| Ok(ids) |
| } |
|
|
| |
| pub fn encode_controls<'a>( |
| &self, |
| tokens: impl IntoIterator<Item = &'a str>, |
| ) -> Result<Vec<u32>> { |
| tokens |
| .into_iter() |
| .map(|token| { |
| self.control_ids |
| .get(token) |
| .copied() |
| .with_context(|| format!("unknown StenToken control token {token:?}")) |
| }) |
| .collect() |
| } |
|
|
| |
| |
| pub fn encode_batch(&self, texts: &[String]) -> Result<Vec<Vec<u32>>> { |
| if texts.len() < 32 { |
| return texts.iter().map(|text| self.encode(text)).collect(); |
| } |
| texts.par_iter().map(|text| self.encode(text)).collect() |
| } |
|
|
| |
| pub fn decode(&self, ids: &[u32]) -> String { |
| let mut result = String::new(); |
| let mut internal = String::new(); |
| let flush = |result: &mut String, internal: &mut String| { |
| if !internal.is_empty() { |
| result.push_str(&decode_internal_capcode(internal)); |
| internal.clear(); |
| } |
| }; |
| for &id in ids { |
| if let Some(control) = control_token_for_id(id) { |
| flush(&mut result, &mut internal); |
| result.push_str(control); |
| } else if id == PAD_ID { |
| |
| } else if (1..=BASE_BYTES).contains(&id) { |
| internal.push(byte_to_internal_char((id - 1) as u8)); |
| } else if let Some(piece) = self.id_to_token.get(id as usize) { |
| internal.push_str(piece); |
| } |
| } |
| flush(&mut result, &mut internal); |
| result |
| } |
|
|
| pub fn vocab_size(&self) -> usize { |
| self.id_to_token.len() |
| } |
| } |
|
|
| fn known_control_ids() -> HashMap<String, u32> { |
| CONTROL_TOKENS |
| .iter() |
| .copied() |
| .chain(std::iter::once(EXTRA_CONTROL_TOKEN)) |
| .enumerate() |
| .map(|(index, token)| (token.to_owned(), BASE_BYTES + 1 + index as u32)) |
| .collect() |
| } |
|
|
| fn control_token_for_id(id: u32) -> Option<&'static str> { |
| let index = id.checked_sub(BASE_BYTES + 1)? as usize; |
| CONTROL_TOKENS |
| .get(index) |
| .copied() |
| .or_else(|| (index == CONTROL_TOKENS.len()).then_some(EXTRA_CONTROL_TOKEN)) |
| } |
|
|
| fn byte_to_internal_char(byte: u8) -> char { |
| if (32..127).contains(&byte) { |
| byte as char |
| } else { |
| char::from_u32(0xE000 + byte as u32).expect("private-use byte mapping is valid") |
| } |
| } |
|
|
| fn internal_char_to_byte(ch: char) -> Option<u8> { |
| let code = ch as u32; |
| if (0xE000..=0xE0FF).contains(&code) { |
| Some((code - 0xE000) as u8) |
| } else if ch.is_ascii() { |
| Some(ch as u8) |
| } else { |
| None |
| } |
| } |
|
|
| fn text_to_internal_bytespace(text: &str) -> String { |
| text.bytes().map(byte_to_internal_char).collect() |
| } |
|
|
| fn decode_internal_capcode(internal: &str) -> String { |
| capcode_decode(&internal_to_text(internal)) |
| } |
|
|
| fn internal_to_text(internal: &str) -> String { |
| let mut bytes = Vec::with_capacity(internal.len()); |
| for ch in internal.chars() { |
| let code = ch as u32; |
| if (0xE000..=0xE0FF).contains(&code) { |
| bytes.push((code - 0xE000) as u8); |
| } else { |
| let mut encoded = [0; 4]; |
| bytes.extend_from_slice(ch.encode_utf8(&mut encoded).as_bytes()); |
| } |
| } |
| String::from_utf8_lossy(&bytes).into_owned() |
| } |
|
|
| fn surface_has_reserved_control(internal: &str) -> bool { |
| let visible = internal_to_text(internal); |
| visible.contains("<PAD>") |
| || CONTROL_TOKENS |
| .iter() |
| .any(|control| visible.contains(control)) |
| || visible.contains(EXTRA_CONTROL_TOKEN) |
| } |
|
|
| fn is_word_char(ch: char) -> bool { |
| ch.is_alphabetic() || ch.is_numeric() || ch == '_' |
| } |
|
|
| fn all_caps(word: &[char]) -> bool { |
| let mut saw_alpha = false; |
| for &ch in word { |
| if ch.is_alphabetic() { |
| saw_alpha = true; |
| if !ch.is_uppercase() { |
| return false; |
| } |
| } |
| } |
| saw_alpha |
| } |
|
|
| fn lowercase(chars: &[char]) -> String { |
| chars.iter().flat_map(|ch| ch.to_lowercase()).collect() |
| } |
|
|
| fn split_identifier(token: &str) -> Vec<String> { |
| let chars: Vec<char> = token.chars().collect(); |
| if chars.is_empty() || chars[0] == '\r' || chars[0] == '\n' { |
| return vec![token.to_owned()]; |
| } |
| let mut prefix = String::new(); |
| let mut start = 0; |
| let has_prefix = chars.len() > 1 |
| && ((matches!(chars[0], ' ' | '\t') && !chars[1].is_whitespace()) |
| || (!matches!(chars[0], ' ' | '\t' | '\r' | '\n') |
| && !chars[0].is_alphanumeric() |
| && chars[0] != '_')); |
| if has_prefix { |
| prefix.push(chars[0]); |
| start = 1; |
| } |
| let core = &chars[start..]; |
| if core.is_empty() |
| || !core |
| .iter() |
| .all(|ch| ch.is_ascii() && (ch.is_ascii_alphanumeric() || *ch == '_')) |
| { |
| return vec![token.to_owned()]; |
| } |
| let needs_split = core.iter().any(|ch| *ch == '_' || ch.is_ascii_digit()) |
| || core |
| .windows(2) |
| .any(|pair| pair[0].is_ascii_lowercase() && pair[1].is_ascii_uppercase()) |
| || (core.iter().any(|ch| ch.is_ascii_uppercase()) |
| && core.iter().any(|ch| ch.is_ascii_lowercase())); |
| if !needs_split { |
| return vec![token.to_owned()]; |
| } |
| let mut pieces = Vec::new(); |
| let mut i = 0; |
| while i < core.len() { |
| let start_piece = i; |
| let ch = core[i]; |
| if ch == '_' { |
| while i < core.len() && core[i] == '_' { |
| i += 1; |
| } |
| } else if ch.is_ascii_uppercase() { |
| while i < core.len() && core[i].is_ascii_uppercase() { |
| i += 1; |
| } |
| let run_length = i - start_piece; |
| if i == core.len() || (i < core.len() && core[i].is_ascii_digit()) { |
| |
| |
| } else if run_length > 1 && i < core.len() && core[i].is_ascii_lowercase() { |
| |
| |
| i -= 1; |
| } else if run_length == 1 && i < core.len() && core[i].is_ascii_lowercase() { |
| |
| |
| while i < core.len() && core[i].is_ascii_lowercase() { |
| i += 1; |
| } |
| } else { |
| |
| |
| |
| i = start_piece + 1; |
| continue; |
| } |
| } else if ch.is_ascii_lowercase() { |
| while i < core.len() && core[i].is_ascii_lowercase() { |
| i += 1; |
| } |
| } else if ch.is_ascii_digit() { |
| while i < core.len() && core[i].is_ascii_digit() { |
| i += 1; |
| } |
| } else { |
| return vec![token.to_owned()]; |
| } |
| pieces.push(core[start_piece..i].iter().collect()); |
| } |
| if pieces.len() <= 1 { |
| return vec![token.to_owned()]; |
| } |
| if prefix.is_empty() { |
| pieces |
| } else { |
| std::iter::once(prefix).chain(pieces).collect() |
| } |
| } |
|
|
| fn encode_capcode_word(word: &[char]) -> String { |
| let word: String = word.iter().collect(); |
| split_identifier(&word) |
| .into_iter() |
| .map(|piece| { |
| let chars: Vec<char> = piece.chars().collect(); |
| if chars.iter().any(|ch| ch.is_uppercase()) { |
| let alpha_count = chars.iter().filter(|ch| ch.is_alphabetic()).count(); |
| let marker = if all_caps(&chars) && alpha_count > 1 { |
| 'W' |
| } else { |
| 'C' |
| }; |
| format!("{marker}{}", lowercase(&chars)) |
| } else { |
| lowercase(&chars) |
| } |
| }) |
| .collect() |
| } |
|
|
| |
| pub fn capcode_encode(text: &str) -> String { |
| let chars: Vec<char> = text.chars().collect(); |
| let mut out = String::with_capacity(text.len() + text.len() / 8); |
| let mut i = 0; |
| while i < chars.len() { |
| if chars[i].is_whitespace() { |
| let start = i; |
| while i < chars.len() && chars[i].is_whitespace() { |
| i += 1; |
| } |
| out.extend(chars[start..i].iter()); |
| continue; |
| } |
| if is_word_char(chars[i]) { |
| let start = i; |
| while i < chars.len() && is_word_char(chars[i]) { |
| i += 1; |
| } |
| let word = &chars[start..i]; |
| if all_caps(word) { |
| let alpha_count = word.iter().filter(|ch| ch.is_alphabetic()).count(); |
| if alpha_count <= 1 { |
| out.push('C'); |
| out.push_str(&lowercase(word)); |
| continue; |
| } |
| let mut lookahead = i; |
| let mut words: Vec<(Vec<char>, Vec<char>)> = vec![(word.to_vec(), Vec::new())]; |
| let trailing: Vec<char>; |
| loop { |
| let ws_start = lookahead; |
| while lookahead < chars.len() && chars[lookahead].is_whitespace() { |
| lookahead += 1; |
| } |
| let separator = chars[ws_start..lookahead].to_vec(); |
| if lookahead >= chars.len() || !is_word_char(chars[lookahead]) { |
| trailing = separator; |
| break; |
| } |
| |
| |
| |
| |
| |
| let next_start = lookahead; |
| let mut next_end = lookahead; |
| while next_end < chars.len() && is_word_char(chars[next_end]) { |
| next_end += 1; |
| } |
| let next_word = &chars[next_start..next_end]; |
| if !all_caps(next_word) { |
| trailing = separator; |
| lookahead = next_start; |
| break; |
| } |
| words.push((next_word.to_vec(), separator)); |
| lookahead = next_end; |
| } |
| if words.len() > 1 { |
| out.push('B'); |
| for (index, (block_word, separator)) in words.iter().enumerate() { |
| if index > 0 { |
| out.extend(separator); |
| } |
| out.push_str(&lowercase(block_word)); |
| } |
| out.push('E'); |
| out.extend(trailing); |
| i = lookahead; |
| continue; |
| } |
| out.push('W'); |
| out.push_str(&lowercase(word)); |
| continue; |
| } |
| out.push_str(&encode_capcode_word(word)); |
| continue; |
| } |
| let start = i; |
| while i < chars.len() && !chars[i].is_whitespace() && !is_word_char(chars[i]) { |
| i += 1; |
| } |
| out.extend(chars[start..i].iter()); |
| } |
| out |
| } |
|
|
| |
| pub fn capcode_decode(text: &str) -> String { |
| let mut out = String::with_capacity(text.len()); |
| let mut capitalize_next = false; |
| let mut word_caps = false; |
| let mut block_caps = false; |
| for ch in text.chars() { |
| match ch { |
| 'B' => { |
| block_caps = true; |
| word_caps = false; |
| capitalize_next = false; |
| } |
| 'E' => { |
| block_caps = false; |
| word_caps = false; |
| capitalize_next = false; |
| } |
| 'W' => { |
| word_caps = true; |
| capitalize_next = false; |
| } |
| 'C' => { |
| capitalize_next = true; |
| word_caps = false; |
| } |
| _ => { |
| if capitalize_next && ch.is_alphabetic() { |
| out.extend(ch.to_uppercase()); |
| capitalize_next = false; |
| } else if word_caps && (ch.is_alphanumeric() || ch == '_') { |
| if ch.is_alphabetic() { |
| out.extend(ch.to_uppercase()); |
| } else { |
| out.push(ch); |
| } |
| } else { |
| if word_caps { |
| word_caps = false; |
| } |
| if block_caps && ch.is_alphabetic() { |
| out.extend(ch.to_uppercase()); |
| } else { |
| out.push(ch); |
| } |
| } |
| } |
| } |
| } |
| out |
| } |
|
|
| fn is_punctuation_or_symbol(ch: char) -> bool { |
| ch.is_punctuation() || ch.is_symbol() |
| } |
|
|
| fn is_scanner_word_char(ch: char) -> bool { |
| ch.is_alphabetic() || ch.is_numeric() || ch == '_' |
| } |
|
|
| fn starts_with_chars(chars: &[char], at: usize, needle: &str) -> bool { |
| needle |
| .chars() |
| .enumerate() |
| .all(|(offset, needle_ch)| chars.get(at + offset) == Some(&needle_ch)) |
| } |
|
|
| fn keyword_with_space(chars: &[char], at: usize) -> Option<usize> { |
| let mut candidates = KEYWORDS.to_vec(); |
| candidates.sort_unstable_by_key(|keyword| std::cmp::Reverse(keyword.len())); |
| for keyword in candidates { |
| let len = keyword.len(); |
| if at + len > chars.len() { |
| continue; |
| } |
| if !keyword |
| .chars() |
| .enumerate() |
| .all(|(offset, expected)| chars[at + offset].eq_ignore_ascii_case(&expected)) |
| { |
| continue; |
| } |
| if at + len < chars.len() && is_scanner_word_char(chars[at + len]) { |
| continue; |
| } |
| if at + len >= chars.len() || !matches!(chars[at + len], ' ' | '\t') { |
| continue; |
| } |
| let mut end = at + len; |
| while end < chars.len() && matches!(chars[end], ' ' | '\t') { |
| end += 1; |
| } |
| return Some(end); |
| } |
| None |
| } |
|
|
| |
| pub fn pre_tokenize_capcoded(text: &str) -> Vec<String> { |
| let chars: Vec<char> = text.chars().collect(); |
| let mut raw = Vec::new(); |
| let mut i = 0; |
| while i < chars.len() { |
| if chars[i] == '\'' { |
| if let Some(suffix) = CONTRACTIONS |
| .iter() |
| .find(|suffix| starts_with_chars(&chars, i, suffix)) |
| { |
| let end = i + suffix.chars().count(); |
| raw.push(chars[i..end].iter().collect()); |
| i = end; |
| continue; |
| } |
| } |
| if let Some(end) = keyword_with_space(&chars, i) { |
| raw.push(chars[i..end].iter().collect()); |
| i = end; |
| continue; |
| } |
| let ch = chars[i]; |
| if matches!(ch, '\r' | '\n') { |
| let start = i; |
| while i < chars.len() && matches!(chars[i], '\r' | '\n') { |
| if chars[i] == '\r' && chars.get(i + 1) == Some(&'\n') { |
| i += 2; |
| } else { |
| i += 1; |
| } |
| } |
| while i < chars.len() && matches!(chars[i], ' ' | '\t') { |
| i += 1; |
| } |
| raw.push(chars[start..i].iter().collect()); |
| continue; |
| } |
| if ch.is_whitespace() { |
| if ch == ' ' |
| && chars |
| .get(i + 1) |
| .is_some_and(|next| is_punctuation_or_symbol(*next)) |
| { |
| let start = i; |
| i += 1; |
| while i < chars.len() && is_punctuation_or_symbol(chars[i]) { |
| i += 1; |
| } |
| raw.push(chars[start..i].iter().collect()); |
| continue; |
| } |
| let start = i; |
| i += 1; |
| while i < chars.len() && chars[i].is_whitespace() && !matches!(chars[i], '\r' | '\n') { |
| if chars[i] == ' ' |
| && chars |
| .get(i + 1) |
| .is_some_and(|next| is_punctuation_or_symbol(*next)) |
| { |
| break; |
| } |
| i += 1; |
| } |
| raw.push(chars[start..i].iter().collect()); |
| continue; |
| } |
| let operator = OPERATORS |
| .iter() |
| .chain(SHORT_OPERATORS.iter()) |
| .find(|operator| starts_with_chars(&chars, i, operator)); |
| if let Some(operator) = operator { |
| let end = i + operator.chars().count(); |
| raw.push((*operator).to_owned()); |
| i = end; |
| continue; |
| } |
| if ch.is_numeric() { |
| let mut end = i + 1; |
| while end < chars.len() && chars[end].is_numeric() { |
| end += 1; |
| } |
| while i < end { |
| let next = (i + 3).min(end); |
| raw.push(chars[i..next].iter().collect()); |
| i = next; |
| } |
| continue; |
| } |
| if is_scanner_word_char(ch) { |
| let start = i; |
| i += 1; |
| while i < chars.len() && is_scanner_word_char(chars[i]) { |
| i += 1; |
| } |
| raw.push(chars[start..i].iter().collect()); |
| continue; |
| } |
| if is_punctuation_or_symbol(ch) { |
| let start = i; |
| i += 1; |
| while i < chars.len() && is_punctuation_or_symbol(chars[i]) { |
| i += 1; |
| } |
| raw.push(chars[start..i].iter().collect()); |
| continue; |
| } |
| raw.push(ch.to_string()); |
| i += 1; |
| } |
| raw.into_iter() |
| .flat_map(|token| split_identifier(&token)) |
| .filter(|token| !token.is_empty()) |
| .collect() |
| } |
|
|
| #[cfg(test)] |
| mod tests { |
| use super::*; |
|
|
| #[test] |
| fn capcode_common_cases_round_trip() { |
| for input in [ |
| "NASA API", |
| "getHTTPResponse", |
| "SK F\tx\u{fffd}4", |
| "from django.db import models", |
| "naïve café — 東京 😀", |
| ] { |
| assert_eq!(capcode_decode(&capcode_encode(input)), input, "{input:?}"); |
| } |
| } |
|
|
| #[test] |
| fn scanner_preserves_useful_code_pieces() { |
| assert_eq!( |
| pre_tokenize_capcoded(&capcode_encode("from django.db import models")), |
| vec!["from ", "django", ".", "db", " ", "import ", "models"] |
| ); |
| } |
|
|
| #[test] |
| fn scanner_preserves_legacy_operator_and_identifier_edges() { |
| assert_eq!( |
| pre_tokenize_capcoded("?.]5"), |
| vec!["?.", "]", "5"], |
| "?. must remain an atomic operator rather than merging with ]" |
| ); |
| assert_eq!( |
| pre_tokenize_capcoded("@=value"), |
| vec!["@=", "value"], |
| "@= is an atomic operator in the published scanner" |
| ); |
| assert_eq!( |
| split_identifier("X_a"), |
| vec!["_", "a"], |
| "match the original regex.findall behavior for unmatched X" |
| ); |
| assert_eq!( |
| split_identifier("getWhttpCresponse"), |
| vec!["get", "Whttp", "Cresponse"] |
| ); |
| } |
|
|
| #[test] |
| fn canonical_control_ids_are_structural_and_round_trip() { |
| let controls = known_control_ids(); |
| assert_eq!(controls.get("<BOS>"), Some(&(BASE_BYTES + 1))); |
| assert_eq!(controls.get("<|tool_call|>"), Some(&275)); |
| assert_eq!(controls.get(EXTRA_CONTROL_TOKEN), Some(&285)); |
| assert_eq!(control_token_for_id(275), Some("<|tool_call|>")); |
| assert_eq!(control_token_for_id(285), Some(EXTRA_CONTROL_TOKEN)); |
| assert_eq!(control_token_for_id(286), None); |
| } |
| } |
|
|