1use std::collections::HashMap;
12use std::sync::RwLock;
13
14#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
16pub enum TokenizerError {
17 #[error("Invalid token ID: {0}")]
18 InvalidToken(i32),
19 #[error("Encoding error: {0}")]
20 EncodingError(String),
21 #[error("Decoding error: {0}")]
22 DecodingError(String),
23}
24
25pub type TokenizerResult<T> = std::result::Result<T, TokenizerError>;
26
27pub trait Tokenizer: Send + Sync {
29 fn encode(&self, text: &str) -> TokenizerResult<Vec<i32>>;
31
32 fn decode(&self, tokens: &[i32]) -> TokenizerResult<String>;
34
35 fn decode_token(&self, token: i32, state: &mut DecodingState) -> TokenizerResult<String>;
38
39 fn vocab_size(&self) -> usize;
41}
42
43#[derive(Debug, Clone, Default)]
45pub struct DecodingState {
46 buffer: String,
47 pending_utf8: Vec<u8>,
48 emitted_any: bool,
49}
50
51impl DecodingState {
52 pub fn new() -> Self {
53 Self::default()
54 }
55
56 pub fn buffer(&self) -> &str {
57 &self.buffer
58 }
59
60 pub fn clear(&mut self) {
61 self.buffer.clear();
62 self.pending_utf8.clear();
63 self.emitted_any = false;
64 }
65}
66
67pub struct WhitespaceTokenizer {
74 state: RwLock<VocabState>,
75}
76
77#[derive(Debug, Default)]
78struct VocabState {
79 vocab: HashMap<i32, String>,
80 reverse_vocab: HashMap<String, i32>,
81 next_id: i32,
82}
83
84impl WhitespaceTokenizer {
85 pub fn new() -> Self {
86 Self {
87 state: RwLock::new(VocabState::default()),
88 }
89 }
90
91 fn decode_id(&self, token: i32) -> TokenizerResult<String> {
92 let state = self
93 .state
94 .read()
95 .map_err(|_| TokenizerError::DecodingError("tokenizer lock poisoned".to_string()))?;
96
97 state
98 .vocab
99 .get(&token)
100 .cloned()
101 .ok_or(TokenizerError::InvalidToken(token))
102 }
103}
104
105impl Default for WhitespaceTokenizer {
106 fn default() -> Self {
107 Self::new()
108 }
109}
110
111impl Tokenizer for WhitespaceTokenizer {
112 fn encode(&self, text: &str) -> TokenizerResult<Vec<i32>> {
113 let mut state = self
114 .state
115 .write()
116 .map_err(|_| TokenizerError::EncodingError("tokenizer lock poisoned".to_string()))?;
117
118 let mut ids = Vec::new();
119 for word in text.split_whitespace() {
120 let id = if let Some(id) = state.reverse_vocab.get(word) {
121 *id
122 } else {
123 let id = state.next_id;
124 state.next_id += 1;
125 state.reverse_vocab.insert(word.to_string(), id);
126 state.vocab.insert(id, word.to_string());
127 id
128 };
129 ids.push(id);
130 }
131
132 Ok(ids)
133 }
134
135 fn decode(&self, tokens: &[i32]) -> TokenizerResult<String> {
136 let mut words = Vec::with_capacity(tokens.len());
137 for &id in tokens {
138 words.push(self.decode_id(id)?);
139 }
140 Ok(words.join(" "))
141 }
142
143 fn decode_token(&self, token: i32, state: &mut DecodingState) -> TokenizerResult<String> {
144 let word = self.decode_id(token)?;
145 let emitted = if state.emitted_any {
146 format!(" {}", word)
147 } else {
148 word
149 };
150 state.buffer.push_str(&emitted);
151 state.emitted_any = true;
152 Ok(emitted)
153 }
154
155 fn vocab_size(&self) -> usize {
156 self.state.read().map(|s| s.vocab.len()).unwrap_or(0)
157 }
158}
159
160#[cfg(test)]
161mod tests {
162 use super::*;
163
164 #[test]
165 fn encode_whitespace_simple() {
166 let tok = WhitespaceTokenizer::new();
167 let ids = tok.encode("hello world").unwrap();
168 assert_eq!(ids.len(), 2);
169 }
170
171 #[test]
172 fn encode_empty_string() {
173 let tok = WhitespaceTokenizer::new();
174 let ids = tok.encode("").unwrap();
175 assert!(ids.is_empty());
176 }
177
178 #[test]
179 fn encode_single_word() {
180 let tok = WhitespaceTokenizer::new();
181 let ids = tok.encode("hello").unwrap();
182 assert_eq!(ids.len(), 1);
183 }
184
185 #[test]
186 fn encode_multiple_spaces() {
187 let tok = WhitespaceTokenizer::new();
188 let ids = tok.encode("hello world").unwrap();
189 assert_eq!(ids.len(), 2);
190 }
191
192 #[test]
193 fn decode_roundtrip() {
194 let tok = WhitespaceTokenizer::new();
195 let original = "hello world test";
196 let encoded = tok.encode(original).unwrap();
197 let decoded = tok.decode(&encoded).unwrap();
198 assert_eq!(decoded, original);
199 }
200
201 #[test]
202 fn decode_empty_tokens() {
203 let tok = WhitespaceTokenizer::new();
204 let decoded = tok.decode(&[]).unwrap();
205 assert_eq!(decoded, "");
206 }
207
208 #[test]
209 fn streaming_decode_state() {
210 let tok: &dyn Tokenizer = &WhitespaceTokenizer::new();
211 let encoded = tok.encode("hello world").unwrap();
212 let mut state = DecodingState::new();
213
214 assert_eq!(state.buffer(), "");
215 assert_eq!(tok.decode_token(encoded[0], &mut state).unwrap(), "hello");
216 assert_eq!(state.buffer(), "hello");
217 assert_eq!(tok.decode_token(encoded[1], &mut state).unwrap(), " world");
218 assert_eq!(state.buffer(), "hello world");
219
220 state.clear();
221 assert_eq!(state.buffer(), "");
222 }
223
224 #[test]
225 fn decode_invalid_token_errors() {
226 let tok = WhitespaceTokenizer::new();
227 tok.encode("hello").unwrap();
228 assert_eq!(
229 tok.decode(&[999]).unwrap_err(),
230 TokenizerError::InvalidToken(999)
231 );
232 }
233
234 #[test]
235 fn vocab_size_reflects_built_vocab() {
236 let tok = WhitespaceTokenizer::new();
237 assert_eq!(tok.vocab_size(), 0);
238 tok.encode("hello world hello").unwrap();
239 assert_eq!(tok.vocab_size(), 2);
240 }
241}