Skip to main content

rlean_search/
lexer.rs

1//! Lexer for a useful fragment of Lean 4 type / declaration surface syntax.
2
3use std::fmt;
4
5#[derive(Debug, Clone, PartialEq, Eq)]
6pub enum Token {
7    Ident(String),
8    Nat(String),
9    /// String or character literal, including quotes.
10    Literal(String),
11    /// `_`
12    Underscore,
13    /// `?name`
14    NamedHole(String),
15    /// Keywords / symbols
16    Theorem,
17    Lemma,
18    Axiom,
19    Forall,
20    Exists,
21    Fun,
22    /// `:`
23    Colon,
24    /// `:=`
25    Assign,
26    /// `,`
27    Comma,
28    /// `.`
29    Dot,
30    /// `→` or `->`
31    Arrow,
32    /// `=>` or `↦`
33    MapsTo,
34    /// `|` (pattern / match; also conclusion search prefix before `-`)
35    Pipe,
36    /// `|-` conclusion-search marker (produced by lexer when seeing `|` `-`)
37    Turnstile,
38    LParen,
39    RParen,
40    LBrace,
41    RBrace,
42    LBracket,
43    RBracket,
44    /// `⦃`
45    LStrict,
46    /// `⦄`
47    RStrict,
48    /// Binary / unary operators and other symbol tokens
49    Op(String),
50    /// End of input
51    Eof,
52}
53
54impl Token {
55    pub fn is_ident_like(&self) -> bool {
56        matches!(
57            self,
58            Token::Ident(_) | Token::Underscore | Token::NamedHole(_) | Token::Nat(_)
59        )
60    }
61}
62
63impl fmt::Display for Token {
64    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
65        match self {
66            Token::Ident(s) => write!(f, "{s}"),
67            Token::Nat(s) => write!(f, "{s}"),
68            Token::Literal(s) => write!(f, "{s}"),
69            Token::Underscore => write!(f, "_"),
70            Token::NamedHole(s) => write!(f, "?{s}"),
71            Token::Theorem => write!(f, "theorem"),
72            Token::Lemma => write!(f, "lemma"),
73            Token::Axiom => write!(f, "axiom"),
74            Token::Forall => write!(f, "∀"),
75            Token::Exists => write!(f, "∃"),
76            Token::Fun => write!(f, "fun"),
77            Token::Colon => write!(f, ":"),
78            Token::Assign => write!(f, ":="),
79            Token::Comma => write!(f, ","),
80            Token::Dot => write!(f, "."),
81            Token::Arrow => write!(f, "→"),
82            Token::MapsTo => write!(f, "=>"),
83            Token::Pipe => write!(f, "|"),
84            Token::Turnstile => write!(f, "|-"),
85            Token::LParen => write!(f, "("),
86            Token::RParen => write!(f, ")"),
87            Token::LBrace => write!(f, "{{"),
88            Token::RBrace => write!(f, "}}"),
89            Token::LBracket => write!(f, "["),
90            Token::RBracket => write!(f, "]"),
91            Token::LStrict => write!(f, "⦃"),
92            Token::RStrict => write!(f, "⦄"),
93            Token::Op(s) => write!(f, "{s}"),
94            Token::Eof => write!(f, "<eof>"),
95        }
96    }
97}
98
99#[derive(Debug, Clone)]
100pub struct Lexer<'a> {
101    src: &'a str,
102    bytes: &'a [u8],
103    pos: usize,
104}
105
106impl<'a> Lexer<'a> {
107    pub fn new(src: &'a str) -> Self {
108        Self {
109            src,
110            bytes: src.as_bytes(),
111            pos: 0,
112        }
113    }
114
115    pub fn remaining(&self) -> &'a str {
116        &self.src[self.pos..]
117    }
118
119    pub fn position(&self) -> usize {
120        self.pos
121    }
122
123    fn peek_char(&self) -> Option<char> {
124        self.src[self.pos..].chars().next()
125    }
126
127    fn bump(&mut self) -> Option<char> {
128        let ch = self.peek_char()?;
129        self.pos += ch.len_utf8();
130        Some(ch)
131    }
132
133    fn starts_with(&self, s: &str) -> bool {
134        self.src[self.pos..].starts_with(s)
135    }
136
137    fn skip_ws_and_comments(&mut self) {
138        loop {
139            while let Some(c) = self.peek_char() {
140                if c.is_whitespace() {
141                    self.bump();
142                } else {
143                    break;
144                }
145            }
146            if self.starts_with("--") {
147                while let Some(c) = self.bump() {
148                    if c == '\n' {
149                        break;
150                    }
151                }
152                continue;
153            }
154            if self.starts_with("/-") {
155                self.pos += 2;
156                let mut depth = 1;
157                while depth > 0 && self.pos < self.bytes.len() {
158                    if self.starts_with("/-") {
159                        self.pos += 2;
160                        depth += 1;
161                    } else if self.starts_with("-/") {
162                        self.pos += 2;
163                        depth -= 1;
164                    } else {
165                        let ch = self.peek_char().unwrap();
166                        self.pos += ch.len_utf8();
167                    }
168                }
169                continue;
170            }
171            break;
172        }
173    }
174
175    pub fn next_token(&mut self) -> Token {
176        self.skip_ws_and_comments();
177        if self.pos >= self.bytes.len() {
178            return Token::Eof;
179        }
180
181        // Multi-char unicode / ascii operators first
182        let multi = [
183            ("|-", Token::Turnstile),
184            (":=", Token::Assign),
185            ("->", Token::Arrow),
186            ("→", Token::Arrow),
187            ("=>", Token::MapsTo),
188            ("↦", Token::MapsTo),
189            ("∀", Token::Forall),
190            ("∃", Token::Exists),
191            ("λ", Token::Fun),
192            ("⦃", Token::LStrict),
193            ("⦄", Token::RStrict),
194            ("≠", Token::Op("≠".into())),
195            ("≤", Token::Op("≤".into())),
196            ("≥", Token::Op("≥".into())),
197            ("<", Token::Op("<".into())),
198            (">", Token::Op(">".into())),
199            ("↔", Token::Op("↔".into())),
200            ("∧", Token::Op("∧".into())),
201            ("∨", Token::Op("∨".into())),
202            ("¬", Token::Op("¬".into())),
203            ("∘", Token::Op("∘".into())),
204            ("⁻¹", Token::Op("⁻¹".into())),
205            ("∈", Token::Op("∈".into())),
206            ("∉", Token::Op("∉".into())),
207            ("⊆", Token::Op("⊆".into())),
208            ("⊂", Token::Op("⊂".into())),
209            ("∪", Token::Op("∪".into())),
210            ("∩", Token::Op("∩".into())),
211            ("∑", Token::Op("∑".into())),
212            ("∏", Token::Op("∏".into())),
213            ("∫", Token::Op("∫".into())),
214            ("∥", Token::Op("∥".into())),
215            ("≈", Token::Op("≈".into())),
216            ("≃", Token::Op("≃".into())),
217            ("≅", Token::Op("≅".into())),
218            ("≡", Token::Op("≡".into())),
219            ("⋅", Token::Op("⋅".into())),
220            ("•", Token::Op("•".into())),
221            ("⋆", Token::Op("⋆".into())),
222            ("·", Token::Ident("·".into())),
223            ("∣", Token::Op("∣".into())),
224            ("ℕ", Token::Ident("ℕ".into())),
225            ("ℤ", Token::Ident("ℤ".into())),
226            ("ℚ", Token::Ident("ℚ".into())),
227            ("ℝ", Token::Ident("ℝ".into())),
228            ("ℂ", Token::Ident("ℂ".into())),
229            ("▸", Token::Op("▸".into())),
230            ("|>.", Token::Op("|>.".into())),
231            ("|>", Token::Op("|>".into())),
232            ("<|", Token::Op("<|".into())),
233            (">>", Token::Op(">>".into())),
234            ("<<", Token::Op("<<".into())),
235            ("++", Token::Op("++".into())),
236            ("::", Token::Op("::".into())),
237            ("..", Token::Op("..".into())),
238            ("$", Token::Op("$".into())),
239            ("@", Token::Op("@".into())),
240            ("^", Token::Op("^".into())),
241            ("+", Token::Op("+".into())),
242            ("-", Token::Op("-".into())),
243            ("*", Token::Op("*".into())),
244            ("/", Token::Op("/".into())),
245            ("%", Token::Op("%".into())),
246            ("=", Token::Op("=".into())),
247            ("|", Token::Pipe),
248        ];
249        for (s, tok) in multi {
250            if self.starts_with(s) {
251                // Don't treat single `-` as op when part of `|-` already handled.
252                self.pos += s.len();
253                return tok.clone();
254            }
255        }
256
257        let ch = self.peek_char().unwrap();
258
259        match ch {
260            '(' => {
261                self.bump();
262                Token::LParen
263            }
264            ')' => {
265                self.bump();
266                Token::RParen
267            }
268            '{' => {
269                self.bump();
270                Token::LBrace
271            }
272            '}' => {
273                self.bump();
274                Token::RBrace
275            }
276            '[' => {
277                self.bump();
278                Token::LBracket
279            }
280            ']' => {
281                self.bump();
282                Token::RBracket
283            }
284            ':' => {
285                self.bump();
286                Token::Colon
287            }
288            ',' => {
289                self.bump();
290                Token::Comma
291            }
292            '.' => {
293                self.bump();
294                // field / number projection handled by parser; keep Dot
295                Token::Dot
296            }
297            '_' => {
298                self.bump();
299                Token::Underscore
300            }
301            '?' => {
302                self.bump();
303                let name = self.lex_ident_tail();
304                if name.is_empty() {
305                    Token::Op("?".into())
306                } else {
307                    Token::NamedHole(name)
308                }
309            }
310            '"' => self.lex_string(),
311            '\'' => self.lex_char_or_prime(),
312            '«' => self.lex_escaped_ident(),
313            c if c.is_ascii_digit() => self.lex_number(),
314            c if is_ident_start(c) => {
315                let id = self.lex_ident();
316                match id.as_str() {
317                    "theorem" => Token::Theorem,
318                    "lemma" => Token::Lemma,
319                    "axiom" => Token::Axiom,
320                    "forall" => Token::Forall,
321                    "exists" => Token::Exists,
322                    "fun" | "λ" => Token::Fun,
323                    "Prop" | "Type" | "Sort" => Token::Ident(id),
324                    _ => Token::Ident(id),
325                }
326            }
327            _ => {
328                // Unknown symbol as operator token
329                let start = self.pos;
330                self.bump();
331                // glue common multi-byte leftover
332                Token::Op(self.src[start..self.pos].to_string())
333            }
334        }
335    }
336
337    fn lex_ident_tail(&mut self) -> String {
338        let start = self.pos;
339        while let Some(c) = self.peek_char() {
340            if is_ident_continue(c) {
341                self.bump();
342            } else {
343                break;
344            }
345        }
346        self.src[start..self.pos].to_string()
347    }
348
349    fn lex_ident(&mut self) -> String {
350        let start = self.pos;
351        if let Some(c) = self.peek_char() {
352            if is_ident_start(c) {
353                self.bump();
354            }
355        }
356        while let Some(c) = self.peek_char() {
357            if is_ident_continue(c) {
358                self.bump();
359            } else {
360                break;
361            }
362        }
363        // trailing `?` sometimes used; keep primes as part of name
364        self.src[start..self.pos].to_string()
365    }
366
367    fn lex_escaped_ident(&mut self) -> Token {
368        // «name with spaces»
369        self.bump(); // «
370        let start = self.pos;
371        while let Some(c) = self.peek_char() {
372            if c == '»' {
373                let name = self.src[start..self.pos].to_string();
374                self.bump();
375                return Token::Ident(name);
376            }
377            self.bump();
378        }
379        Token::Ident(self.src[start..self.pos].to_string())
380    }
381
382    fn lex_number(&mut self) -> Token {
383        let start = self.pos;
384        while let Some(c) = self.peek_char() {
385            if c.is_ascii_digit() {
386                self.bump();
387            } else {
388                break;
389            }
390        }
391        // scientific / decimal rare in types; keep integer token
392        Token::Nat(self.src[start..self.pos].to_string())
393    }
394
395    fn lex_string(&mut self) -> Token {
396        let start = self.pos;
397        self.bump(); // "
398        while let Some(c) = self.peek_char() {
399            if c == '\\' {
400                self.bump();
401                self.bump();
402                continue;
403            }
404            if c == '"' {
405                self.bump();
406                break;
407            }
408            self.bump();
409        }
410        Token::Literal(self.src[start..self.pos].to_string())
411    }
412
413    fn lex_char_or_prime(&mut self) -> Token {
414        // Could be 'a' char literal or trailing prime on previous token — as standalone, treat as op/prime
415        let start = self.pos;
416        self.bump();
417        if let Some(c) = self.peek_char() {
418            if c != '\'' && c != '\\' {
419                // likely prime operator suffix used alone
420                return Token::Op("'".into());
421            }
422            if c == '\\' {
423                self.bump();
424                self.bump();
425            } else {
426                self.bump();
427            }
428            if self.peek_char() == Some('\'') {
429                self.bump();
430            }
431            return Token::Literal(self.src[start..self.pos].to_string());
432        }
433        Token::Op("'".into())
434    }
435
436    /// Tokenize entire input into a vector (excluding trailing Eof unless empty).
437    pub fn tokenize(src: &str) -> Vec<Token> {
438        let mut lx = Lexer::new(src);
439        let mut toks = Vec::new();
440        loop {
441            let t = lx.next_token();
442            if t == Token::Eof {
443                toks.push(t);
444                break;
445            }
446            toks.push(t);
447        }
448        toks
449    }
450}
451
452fn is_ident_start(c: char) -> bool {
453    c.is_alphabetic() || c == '_' || c == 'ℕ' || c == 'ℤ' || c == 'ℚ' || c == 'ℝ' || c == 'ℂ'
454        || c == 'α' || c == 'β' || c == 'γ' || c == 'δ' || c == 'ε' || c == 'ζ' || c == 'η'
455        || c == 'θ' || c == 'ι' || c == 'κ' || c == 'λ' || c == 'μ' || c == 'ν' || c == 'ξ'
456        || c == 'π' || c == 'ρ' || c == 'σ' || c == 'τ' || c == 'υ' || c == 'φ' || c == 'χ'
457        || c == 'ψ' || c == 'ω' || c == 'Γ' || c == 'Δ' || c == 'Θ' || c == 'Λ' || c == 'Ξ'
458        || c == 'Π' || c == 'Σ' || c == 'Φ' || c == 'Ψ' || c == 'Ω' || c == '𝒜' || c == 'ℳ'
459        || ('\u{0370}'..='\u{03FF}').contains(&c) // Greek
460        || ('\u{1D400}'..='\u{1D7FF}').contains(&c) // math alphanumerics
461}
462
463fn is_ident_continue(c: char) -> bool {
464    is_ident_start(c)
465        || c.is_ascii_digit()
466        || c == '\''
467        || c == '?'
468        || c == '!'
469        || c == '₀'
470        || c == '₁'
471        || c == '₂'
472        || c == '₃'
473        || c == '₄'
474        || c == '₅'
475        || c == '₆'
476        || c == '₇'
477        || c == '₈'
478        || c == '₉'
479        || c == 'ₙ'
480        || c == 'ₘ'
481        || c == 'ᵢ'
482        || c == 'ⱼ'
483        || c == 'ₖ'
484        || ('\u{2080}'..='\u{209F}').contains(&c) // subscripts
485        || ('\u{2070}'..='\u{209F}').contains(&c)
486}
487
488#[cfg(test)]
489mod tests {
490    use super::*;
491
492    #[test]
493    fn lex_simple_eq() {
494        let toks = Lexer::tokenize("n + m = m + n");
495        assert!(toks.iter().any(|t| matches!(t, Token::Op(s) if s == "+")));
496        assert!(toks.iter().any(|t| matches!(t, Token::Op(s) if s == "=")));
497    }
498
499    #[test]
500    fn lex_holes() {
501        let toks = Lexer::tokenize("_ + ?a = 0");
502        assert_eq!(toks[0], Token::Underscore);
503        assert!(matches!(&toks[2], Token::NamedHole(s) if s == "a"));
504    }
505
506    #[test]
507    fn lex_turnstile() {
508        let toks = Lexer::tokenize("|- tsum _ = _");
509        assert_eq!(toks[0], Token::Turnstile);
510    }
511}