Skip to main content

lean_ctx/core/
structural_tokenizer.rs

1//! Structural tokenizer treating idiomatic multi-token spans as single motifs.
2
3use std::collections::HashSet;
4use std::sync::OnceLock;
5
6use regex::Regex;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum TokenKind {
10    Keyword,
11    Identifier,
12    Operator,
13    Literal,
14    Pattern,
15    Noise,
16}
17
18#[derive(Debug, Clone, PartialEq)]
19pub struct StructuralToken {
20    pub kind: TokenKind,
21    pub text: String,
22    pub weight: f64,
23}
24
25const W_PATTERN: f64 = 3.0;
26const W_KEYWORD: f64 = 2.0;
27const W_LITERAL: f64 = 1.5;
28const W_IDENTIFIER: f64 = 1.0;
29const W_OPERATOR: f64 = 0.8;
30const W_NOISE: f64 = 0.15;
31
32fn for_in_rust_re() -> &'static Regex {
33    static CELL: OnceLock<Regex> = OnceLock::new();
34    CELL.get_or_init(|| Regex::new(r"for\s+[a-zA-Z_][a-zA-Z0-9_]*\s+in\s+").expect("for-in regex"))
35}
36
37fn keywords_for(lang: &str) -> &'static HashSet<&'static str> {
38    static RUST: OnceLock<HashSet<&str>> = OnceLock::new();
39    static GO: OnceLock<HashSet<&str>> = OnceLock::new();
40    static GENERIC: OnceLock<HashSet<&str>> = OnceLock::new();
41
42    match lang {
43        "rust" | "rs" => RUST.get_or_init(|| {
44            HashSet::from([
45                "pub", "fn", "let", "mut", "struct", "enum", "impl", "trait", "use", "mod",
46                "crate", "super", "self", "where", "type", "const", "static", "async", "await",
47                "match", "if", "else", "for", "while", "loop", "break", "continue", "return",
48                "unsafe", "move", "ref", "dyn", "extern", "in", "as",
49            ])
50        }),
51        "go" => GO.get_or_init(|| {
52            HashSet::from([
53                "func",
54                "package",
55                "import",
56                "var",
57                "const",
58                "type",
59                "struct",
60                "interface",
61                "map",
62                "chan",
63                "defer",
64                "go",
65                "select",
66                "switch",
67                "case",
68                "default",
69                "if",
70                "else",
71                "for",
72                "range",
73                "return",
74                "break",
75                "continue",
76                "fallthrough",
77                "nil",
78                "make",
79                "new",
80                "len",
81                "cap",
82            ])
83        }),
84        _ => GENERIC.get_or_init(|| {
85            HashSet::from([
86                "if", "else", "for", "while", "return", "fn", "func", "let", "var", "const", "pub",
87                "import", "class", "def",
88            ])
89        }),
90    }
91}
92
93fn try_pattern(rest: &str, lang: &str) -> Option<(usize, String)> {
94    let ascii_patterns: &[(&str, &[&str])] = &[
95        ("if err != nil", &["go"]),
96        ("pub async fn", &["rust", "rs"]),
97        ("async fn", &["rust", "rs"]),
98        ("pub fn", &["rust", "rs"]),
99        ("fn main()", &["rust", "rs", "generic", ""]),
100        ("match ", &["rust", "rs"]),
101    ];
102
103    for (pat, langs) in ascii_patterns {
104        if !langs.iter().any(|&l| l == lang || l.is_empty()) {
105            continue;
106        }
107        if rest.starts_with(pat) {
108            return Some((pat.len(), (*pat).to_string()));
109        }
110    }
111
112    if (lang == "rust" || lang == "rs")
113        && let Some(m) = for_in_rust_re().find(rest)
114        && m.start() == 0
115    {
116        return Some((m.end(), m.as_str().to_string()));
117    }
118
119    None
120}
121
122fn skip_line_comment(bytes: &[u8], mut i: usize) -> usize {
123    while i < bytes.len() && bytes[i] != b'\n' {
124        i += 1;
125    }
126    i
127}
128
129fn skip_block_comment(bytes: &[u8], mut i: usize) -> Option<usize> {
130    if i + 1 >= bytes.len() || bytes[i] != b'/' || bytes[i + 1] != b'*' {
131        return None;
132    }
133    i += 2;
134    while i + 1 < bytes.len() {
135        if bytes[i] == b'*' && bytes[i + 1] == b'/' {
136            return Some(i + 2);
137        }
138        i += 1;
139    }
140    Some(bytes.len())
141}
142
143fn scan_string(bytes: &[u8], quote: u8, mut i: usize) -> usize {
144    i += 1;
145    while i < bytes.len() {
146        let b = bytes[i];
147        if b == b'\\' && i + 1 < bytes.len() {
148            i += 2;
149            continue;
150        }
151        if b == quote {
152            return i + 1;
153        }
154        i += 1;
155    }
156    bytes.len()
157}
158
159fn scan_raw_string(bytes: &[u8], i: usize) -> usize {
160    if i + 1 >= bytes.len() || bytes[i] != b'r' {
161        return i;
162    }
163    let mut j = i + 1;
164    let mut hashes = 0usize;
165    while j < bytes.len() && bytes[j] == b'#' {
166        hashes += 1;
167        j += 1;
168    }
169    if j >= bytes.len() || bytes[j] != b'"' {
170        return i;
171    }
172    j += 1;
173    while j < bytes.len() {
174        if bytes[j] == b'"' {
175            let mut k = j + 1;
176            let mut ok = true;
177            for _ in 0..hashes {
178                if k >= bytes.len() || bytes[k] != b'#' {
179                    ok = false;
180                    break;
181                }
182                k += 1;
183            }
184            if ok && hashes == 0 {
185                return k;
186            }
187            if ok {
188                return k;
189            }
190        }
191        j += 1;
192    }
193    bytes.len()
194}
195
196fn scan_number(bytes: &[u8], mut i: usize) -> usize {
197    let start = i;
198    if bytes.get(i) == Some(&b'0') && bytes.get(i + 1).is_some_and(|b| *b == b'x' || *b == b'X') {
199        i += 2;
200        while i < bytes.len() && bytes[i].is_ascii_hexdigit() {
201            i += 1;
202        }
203        return i.max(start + 1);
204    }
205    while i < bytes.len() && (bytes[i].is_ascii_digit() || bytes[i] == b'_' || bytes[i] == b'.') {
206        i += 1;
207    }
208    if bytes.get(i) == Some(&b'e') || bytes.get(i) == Some(&b'E') {
209        i += 1;
210        if bytes.get(i) == Some(&b'+') || bytes.get(i) == Some(&b'-') {
211            i += 1;
212        }
213        while i < bytes.len() && bytes[i].is_ascii_digit() {
214            i += 1;
215        }
216    }
217    i.max(start + 1)
218}
219
220fn scan_identifier(bytes: &[u8], mut i: usize) -> usize {
221    let start = i;
222    while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') {
223        i += 1;
224    }
225    i.max(start + 1)
226}
227
228fn push_op(out: &mut Vec<StructuralToken>, text: &str) {
229    out.push(StructuralToken {
230        kind: TokenKind::Operator,
231        text: text.to_string(),
232        weight: W_OPERATOR,
233    });
234}
235
236/// Tokenize source into weighted structural tokens (motifs, keywords, literals, …).
237pub fn structural_tokenize(code: &str, lang: &str) -> Vec<StructuralToken> {
238    let lang_lower = lang.to_lowercase();
239    let lang_k = match lang_lower.as_str() {
240        "rust" | "rs" => "rust",
241        "go" | "golang" => "go",
242        _ => "generic",
243    };
244
245    let kw = keywords_for(lang_k);
246    let bytes = code.as_bytes();
247    let mut i = 0usize;
248    let mut out = Vec::new();
249
250    while i < bytes.len() {
251        if bytes[i].is_ascii_whitespace() {
252            let start = i;
253            while i < bytes.len() && bytes[i].is_ascii_whitespace() {
254                i += 1;
255            }
256            if start != i {
257                out.push(StructuralToken {
258                    kind: TokenKind::Noise,
259                    text: code[start..i].to_string(),
260                    weight: W_NOISE,
261                });
262            }
263            continue;
264        }
265
266        let rest = &code[i..];
267        if let Some((len, text)) = try_pattern(rest, lang_k) {
268            out.push(StructuralToken {
269                kind: TokenKind::Pattern,
270                text,
271                weight: W_PATTERN,
272            });
273            i += len;
274            continue;
275        }
276
277        if bytes[i] == b'/' && bytes.get(i + 1) == Some(&b'/') {
278            let start = i;
279            i = skip_line_comment(bytes, i);
280            out.push(StructuralToken {
281                kind: TokenKind::Noise,
282                text: code[start..i].to_string(),
283                weight: W_NOISE,
284            });
285            continue;
286        }
287
288        if let Some(next) = skip_block_comment(bytes, i) {
289            let start = i;
290            i = next;
291            out.push(StructuralToken {
292                kind: TokenKind::Noise,
293                text: code[start..i].to_string(),
294                weight: W_NOISE,
295            });
296            continue;
297        }
298
299        if lang_k == "rust"
300            && bytes[i] == b'r'
301            && (bytes.get(i + 1) == Some(&b'#') || bytes.get(i + 1) == Some(&b'"'))
302        {
303            let start = i;
304            i = scan_raw_string(bytes, i);
305            out.push(StructuralToken {
306                kind: TokenKind::Literal,
307                text: code[start..i].to_string(),
308                weight: W_LITERAL,
309            });
310            continue;
311        }
312
313        if bytes[i] == b'"' || bytes[i] == b'\'' {
314            let quote = bytes[i];
315            let start = i;
316            i = scan_string(bytes, quote, i);
317            out.push(StructuralToken {
318                kind: TokenKind::Literal,
319                text: code[start..i].to_string(),
320                weight: W_LITERAL,
321            });
322            continue;
323        }
324
325        if bytes[i].is_ascii_digit() {
326            let start = i;
327            i = scan_number(bytes, i);
328            out.push(StructuralToken {
329                kind: TokenKind::Literal,
330                text: code[start..i].to_string(),
331                weight: W_LITERAL,
332            });
333            continue;
334        }
335
336        if bytes[i].is_ascii_alphabetic() || bytes[i] == b'_' {
337            let start = i;
338            i = scan_identifier(bytes, i);
339            let word = &code[start..i];
340            let kind = if kw.contains(word) {
341                TokenKind::Keyword
342            } else {
343                TokenKind::Identifier
344            };
345            let weight = if kind == TokenKind::Keyword {
346                W_KEYWORD
347            } else {
348                W_IDENTIFIER
349            };
350            out.push(StructuralToken {
351                kind,
352                text: word.to_string(),
353                weight,
354            });
355            continue;
356        }
357
358        let two = i + 1 < bytes.len();
359        if two {
360            let pair = [bytes[i], bytes[i + 1]];
361            let s = std::str::from_utf8(&pair).unwrap_or("??");
362            match pair {
363                [b'!' | b'=' | b'<' | b'>' | b'+' | b'-', b'=']
364                | [b'-' | b'=', b'>']
365                | [b':', b':']
366                | [b'&', b'&']
367                | [b'|', b'|'] => {
368                    push_op(&mut out, s);
369                    i += 2;
370                    continue;
371                }
372                _ => {}
373            }
374        }
375
376        let ch = bytes[i] as char;
377        push_op(&mut out, &ch.to_string());
378        i += 1;
379    }
380
381    out
382}
383
384#[cfg(test)]
385mod tests {
386    use super::*;
387
388    #[test]
389    fn rust_pub_fn_pattern() {
390        let toks = structural_tokenize("pub fn foo() {}", "rust");
391        assert_eq!(toks[0].kind, TokenKind::Pattern);
392        assert_eq!(toks[0].text, "pub fn");
393        assert_eq!(toks[0].weight, W_PATTERN);
394    }
395
396    #[test]
397    fn rust_async_fn_pattern() {
398        let toks = structural_tokenize("pub async fn bar() {}", "rust");
399        assert!(
400            toks.iter()
401                .any(|t| t.kind == TokenKind::Pattern && t.text.starts_with("pub async fn")),
402            "{toks:?}"
403        );
404    }
405
406    #[test]
407    fn rust_match_pattern_prefix() {
408        let toks = structural_tokenize("match x {", "rust");
409        assert_eq!(toks[0].kind, TokenKind::Pattern);
410        assert_eq!(toks[0].text, "match ");
411    }
412
413    #[test]
414    fn rust_for_in_loop_pattern() {
415        let src = "for item in items.iter() {";
416        let toks = structural_tokenize(src, "rust");
417        assert!(
418            toks.iter()
419                .any(|t| t.kind == TokenKind::Pattern && t.text.starts_with("for "))
420        );
421    }
422
423    #[test]
424    fn go_err_nil_pattern() {
425        let toks = structural_tokenize("if err != nil { return err }", "go");
426        assert!(
427            toks.iter()
428                .any(|t| t.kind == TokenKind::Pattern && t.text.contains("err"))
429        );
430        let pat = toks
431            .iter()
432            .find(|t| t.kind == TokenKind::Pattern)
433            .expect("pattern");
434        assert_eq!(pat.text, "if err != nil");
435        assert_eq!(pat.weight, W_PATTERN);
436    }
437
438    #[test]
439    fn weights_pattern_above_identifier() {
440        let toks = structural_tokenize("pub fn main() {}", "rust");
441        let p = toks.iter().find(|t| t.kind == TokenKind::Pattern).unwrap();
442        let id = toks
443            .iter()
444            .find(|t| t.kind == TokenKind::Identifier && t.text == "main")
445            .unwrap();
446        assert!(p.weight > id.weight);
447        assert!(p.weight > W_KEYWORD);
448    }
449
450    #[test]
451    fn comment_is_noise() {
452        let toks = structural_tokenize("// hello\nlet x = 1;", "rust");
453        assert!(
454            toks.iter()
455                .any(|t| t.kind == TokenKind::Noise && t.text.starts_with("//"))
456        );
457        assert!(
458            toks.iter()
459                .any(|t| t.kind == TokenKind::Keyword && t.text == "let")
460        );
461    }
462
463    #[test]
464    fn string_literal_kind() {
465        let toks = structural_tokenize(r#"let s = "ab";"#, "rust");
466        let lit = toks
467            .iter()
468            .find(|t| t.kind == TokenKind::Literal && t.text.starts_with('"'));
469        assert!(lit.is_some());
470        assert_eq!(lit.unwrap().weight, W_LITERAL);
471    }
472}