Skip to main content

driven/binary/
simd_tokenizer.rs

1//! SIMD-Accelerated Tokenizer
2//!
3//! Uses SIMD instructions for fast pattern matching in rule parsing.
4
5use crate::Result;
6
7/// Token types
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9#[repr(u8)]
10pub enum TokenType {
11    /// Heading (## ...)
12    Heading = 0,
13    /// Bullet point (- ...)
14    Bullet = 1,
15    /// Code block
16    CodeBlock = 2,
17    /// Plain text
18    Text = 3,
19    /// Whitespace/newline
20    Whitespace = 4,
21    /// Comment
22    Comment = 5,
23    /// Key-value pair
24    KeyValue = 6,
25    /// End of input
26    End = 255,
27}
28
29/// A token with position information
30#[derive(Debug, Clone, Copy)]
31pub struct Token {
32    /// Token type
33    pub ty: TokenType,
34    /// Start offset in input
35    pub start: u32,
36    /// Length in bytes
37    pub len: u32,
38}
39
40impl Token {
41    /// Create a new token
42    pub fn new(ty: TokenType, start: u32, len: u32) -> Self {
43        Self { ty, start, len }
44    }
45
46    /// Get end offset
47    pub fn end(&self) -> u32 {
48        self.start + self.len
49    }
50
51    /// Extract text from input
52    pub fn text<'a>(&self, input: &'a [u8]) -> Option<&'a str> {
53        let start = self.start as usize;
54        let end = self.end() as usize;
55        if end > input.len() {
56            return None;
57        }
58        std::str::from_utf8(&input[start..end]).ok()
59    }
60}
61
62/// SIMD-accelerated tokenizer
63#[derive(Debug)]
64pub struct SimdTokenizer<'a> {
65    /// Input data
66    input: &'a [u8],
67    /// Current position
68    pos: usize,
69    /// Line number
70    line: u32,
71    /// Column number
72    col: u32,
73}
74
75impl<'a> SimdTokenizer<'a> {
76    /// Create a new tokenizer
77    pub fn new(input: &'a [u8]) -> Self {
78        Self {
79            input,
80            pos: 0,
81            line: 1,
82            col: 1,
83        }
84    }
85
86    /// Get next token
87    pub fn next_token(&mut self) -> Result<Token> {
88        self.skip_whitespace();
89
90        if self.pos >= self.input.len() {
91            return Ok(Token::new(TokenType::End, self.pos as u32, 0));
92        }
93
94        let start = self.pos as u32;
95
96        // Check for heading
97        if self.starts_with(b"#") {
98            return Ok(self.tokenize_heading(start));
99        }
100
101        // Check for bullet
102        if self.starts_with(b"- ") || self.starts_with(b"* ") || self.starts_with(b"\xE2\x80\xA2 ")
103        {
104            return Ok(self.tokenize_bullet(start));
105        }
106
107        // Check for code block
108        if self.starts_with(b"```") {
109            return Ok(self.tokenize_code_block(start));
110        }
111
112        // Check for key-value
113        if let Some(token) = self.try_tokenize_key_value(start) {
114            return Ok(token);
115        }
116
117        // Default: text until newline
118        Ok(self.tokenize_text(start))
119    }
120
121    /// Tokenize all input
122    pub fn tokenize_all(&mut self) -> Result<Vec<Token>> {
123        let mut tokens = Vec::new();
124
125        loop {
126            let token = self.next_token()?;
127            let is_end = token.ty == TokenType::End;
128            tokens.push(token);
129            if is_end {
130                break;
131            }
132        }
133
134        Ok(tokens)
135    }
136
137    /// Find all occurrences of a pattern (SIMD-accelerated when available)
138    pub fn find_all(&self, pattern: &[u8]) -> Vec<usize> {
139        // Use memchr for SIMD-accelerated single-byte search
140        if pattern.len() == 1 {
141            return memchr::memchr_iter(pattern[0], self.input).collect();
142        }
143
144        // Multi-byte pattern search
145        let mut positions = Vec::new();
146        let mut pos = 0;
147
148        while pos + pattern.len() <= self.input.len() {
149            if &self.input[pos..pos + pattern.len()] == pattern {
150                positions.push(pos);
151                pos += pattern.len();
152            } else {
153                pos += 1;
154            }
155        }
156
157        positions
158    }
159
160    /// Count lines efficiently
161    pub fn count_lines(&self) -> usize {
162        memchr::memchr_iter(b'\n', self.input).count() + 1
163    }
164
165    // Private helpers
166
167    fn starts_with(&self, pattern: &[u8]) -> bool {
168        self.input[self.pos..].starts_with(pattern)
169    }
170
171    fn skip_whitespace(&mut self) {
172        while self.pos < self.input.len() {
173            match self.input[self.pos] {
174                b' ' | b'\t' => {
175                    self.pos += 1;
176                    self.col += 1;
177                }
178                b'\n' => {
179                    self.pos += 1;
180                    self.line += 1;
181                    self.col = 1;
182                }
183                b'\r' => {
184                    self.pos += 1;
185                    // Handle CRLF
186                    if self.pos < self.input.len() && self.input[self.pos] == b'\n' {
187                        self.pos += 1;
188                    }
189                    self.line += 1;
190                    self.col = 1;
191                }
192                _ => break,
193            }
194        }
195    }
196
197    fn advance_to_newline(&mut self) {
198        while self.pos < self.input.len() && self.input[self.pos] != b'\n' {
199            self.pos += 1;
200            self.col += 1;
201        }
202    }
203
204    fn tokenize_heading(&mut self, start: u32) -> Token {
205        self.advance_to_newline();
206        Token::new(TokenType::Heading, start, (self.pos as u32) - start)
207    }
208
209    fn tokenize_bullet(&mut self, start: u32) -> Token {
210        self.advance_to_newline();
211        Token::new(TokenType::Bullet, start, (self.pos as u32) - start)
212    }
213
214    fn tokenize_code_block(&mut self, start: u32) -> Token {
215        // Skip opening ```
216        self.pos += 3;
217
218        // Skip to end of opening line
219        self.advance_to_newline();
220        if self.pos < self.input.len() {
221            self.pos += 1; // Skip newline
222        }
223
224        // Find closing ```
225        while self.pos + 3 <= self.input.len() {
226            if &self.input[self.pos..self.pos + 3] == b"```" {
227                self.pos += 3;
228                self.advance_to_newline();
229                break;
230            }
231            self.pos += 1;
232        }
233
234        Token::new(TokenType::CodeBlock, start, (self.pos as u32) - start)
235    }
236
237    fn try_tokenize_key_value(&mut self, start: u32) -> Option<Token> {
238        // Look for : in the current line
239        let line_end = self.input[self.pos..]
240            .iter()
241            .position(|&b| b == b'\n')
242            .map(|p| self.pos + p)
243            .unwrap_or(self.input.len());
244
245        let line = &self.input[self.pos..line_end];
246
247        if line.contains(&b':') {
248            self.pos = line_end;
249            return Some(Token::new(
250                TokenType::KeyValue,
251                start,
252                (self.pos as u32) - start,
253            ));
254        }
255
256        None
257    }
258
259    fn tokenize_text(&mut self, start: u32) -> Token {
260        self.advance_to_newline();
261        Token::new(TokenType::Text, start, (self.pos as u32) - start)
262    }
263}
264
265#[cfg(test)]
266mod tests {
267    use super::*;
268
269    #[test]
270    fn test_tokenize_heading() {
271        let input = b"## Test Heading\nsome text";
272        let mut tokenizer = SimdTokenizer::new(input);
273
274        let token = tokenizer.next_token().unwrap();
275        assert_eq!(token.ty, TokenType::Heading);
276        assert_eq!(token.text(input), Some("## Test Heading"));
277    }
278
279    #[test]
280    fn test_tokenize_bullet() {
281        let input = b"- Item one\n- Item two";
282        let mut tokenizer = SimdTokenizer::new(input);
283
284        let token = tokenizer.next_token().unwrap();
285        assert_eq!(token.ty, TokenType::Bullet);
286        assert_eq!(token.text(input), Some("- Item one"));
287    }
288
289    #[test]
290    fn test_find_all() {
291        let input = b"hello world hello rust hello";
292        let tokenizer = SimdTokenizer::new(input);
293
294        let positions = tokenizer.find_all(b"hello");
295        assert_eq!(positions.len(), 3);
296    }
297
298    #[test]
299    fn test_count_lines() {
300        let input = b"line 1\nline 2\nline 3";
301        let tokenizer = SimdTokenizer::new(input);
302
303        assert_eq!(tokenizer.count_lines(), 3);
304    }
305}