1use 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
236pub 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}