Skip to main content

tla_syntax/
lexer.rs

1use crate::error::{Error, Result};
2use crate::token::{self, Kw, Op, Tok, Token};
3
4pub fn lex(src: &str) -> Result<Vec<Token>> {
5    Lexer::new(src).run()
6}
7
8struct Lexer {
9    chars: Vec<char>,
10    pos: usize,
11    line: u32,
12    col: u32,
13    out: Vec<Token>,
14}
15
16/// Where the module actually begins.
17///
18/// A `.tla` file is a module surrounded by prose: an explanation above the
19/// header, notes and shell transcripts below the terminator. That text is not
20/// TLA+ and must not be lexed as it — plenty of it contains `$`, `&` and `/`,
21/// which are not characters the language allows loose.
22fn module_start(src: &str) -> (usize, u32) {
23    let mut offset = 0;
24    if src.starts_with('\u{feff}') {
25        offset = '\u{feff}'.len_utf8();
26    }
27    for (index, text) in src[offset..].split_inclusive('\n').enumerate() {
28        let trimmed = text.trim_start();
29        if trimmed.starts_with("----")
30            && trimmed
31                .trim_start_matches('-')
32                .trim_start()
33                .starts_with("MODULE")
34        {
35            return (offset, u32::try_from(index).unwrap_or(u32::MAX) + 1);
36        }
37        offset += text.len();
38    }
39    (0, 1)
40}
41
42impl Lexer {
43    fn new(src: &str) -> Self {
44        let (offset, line) = module_start(src);
45        Self {
46            chars: src[offset..].chars().collect(),
47            pos: 0,
48            line,
49            col: 1,
50            out: Vec::new(),
51        }
52    }
53
54    fn peek(&self) -> Option<char> {
55        self.chars.get(self.pos).copied()
56    }
57
58    fn at(&self, offset: usize) -> Option<char> {
59        self.chars.get(self.pos + offset).copied()
60    }
61
62    fn bump(&mut self) -> Option<char> {
63        let c = self.chars.get(self.pos).copied()?;
64        self.pos += 1;
65        // `\r\n` is one line break, not two, so a carriage return only ends a
66        // line when no newline follows it.
67        let ends_line = c == '\n' || (c == '\r' && self.peek() != Some('\n'));
68        if ends_line {
69            self.line += 1;
70            self.col = 1;
71        } else {
72            self.col += 1;
73        }
74        Some(c)
75    }
76
77    fn advance(&mut self, n: usize) {
78        for _ in 0..n {
79            self.bump();
80        }
81    }
82
83    fn starts_with(&self, s: &str) -> bool {
84        s.chars().enumerate().all(|(i, c)| self.at(i) == Some(c))
85    }
86
87    fn run_of(&self, c: char) -> usize {
88        let mut n = 0;
89        while self.at(n) == Some(c) {
90            n += 1;
91        }
92        n
93    }
94
95    fn err(&self, msg: impl Into<String>) -> Error {
96        Error::lex(msg, self.line, self.col)
97    }
98
99    fn run(mut self) -> Result<Vec<Token>> {
100        // Modules nest, so the file ends at the terminator that closes the
101        // outermost one, not at the first one seen.
102        let mut depth = 0usize;
103        loop {
104            self.skip_trivia()?;
105            let (line, col) = (self.line, self.col);
106            let Some(c) = self.peek() else { break };
107            let tok = self.scan(c)?;
108            match tok {
109                Tok::Kw(Kw::Module) => depth += 1,
110                Tok::ModuleEnd => depth = depth.saturating_sub(1),
111                _ => {}
112            }
113            let closed = matches!(tok, Tok::ModuleEnd) && depth == 0;
114            self.out.push(Token { tok, line, col });
115            if closed {
116                break;
117            }
118        }
119        self.out.push(Token {
120            tok: Tok::Eof,
121            line: self.line,
122            col: self.col,
123        });
124        Ok(self.out)
125    }
126
127    fn skip_trivia(&mut self) -> Result<()> {
128        loop {
129            match self.peek() {
130                Some(c) if c.is_whitespace() => {
131                    self.bump();
132                }
133                Some('\\') if self.at(1) == Some('*') => {
134                    while !matches!(self.peek(), None | Some('\n')) {
135                        self.bump();
136                    }
137                }
138                Some('(') if self.at(1) == Some('*') => self.skip_block_comment()?,
139                _ => return Ok(()),
140            }
141        }
142    }
143
144    fn skip_block_comment(&mut self) -> Result<()> {
145        let (line, col) = (self.line, self.col);
146        let mut depth = 0usize;
147        loop {
148            if self.peek().is_none() {
149                return Err(Error::lex("unterminated (* comment", line, col));
150            }
151            if self.starts_with("(*") {
152                depth += 1;
153                self.advance(2);
154            } else if self.starts_with("*)") {
155                depth -= 1;
156                self.advance(2);
157                if depth == 0 {
158                    return Ok(());
159                }
160            } else {
161                self.bump();
162            }
163        }
164    }
165
166    fn scan(&mut self, c: char) -> Result<Tok> {
167        if c.is_ascii_digit() {
168            return self.scan_number();
169        }
170        // A lone `_` is the placeholder in `Op(_, _)`; followed by more, it is
171        // an ordinary identifier, and `]_vars` is disambiguated by the parser.
172        if c.is_ascii_alphabetic()
173            || (c == '_'
174                && self
175                    .at(1)
176                    .is_some_and(|n| n.is_ascii_alphanumeric() || n == '_'))
177        {
178            return Ok(self.scan_word());
179        }
180        match c {
181            '"' => self.scan_string(),
182            '=' => self.scan_equals(),
183            '-' => Ok(self.scan_dashes()),
184            '<' => Ok(self.scan_lt()),
185            '>' => Ok(self.scan_gt()),
186            '\\' => self.scan_backslash(),
187            '/' => self.scan_slash(),
188            '|' => {
189                if self.starts_with("|->") {
190                    self.advance(3);
191                    return Ok(Tok::MapsTo);
192                }
193                self.user_symbol()
194                    .map_or_else(|| Err(self.err("stray `|`")), Ok)
195            }
196            ':' => {
197                if self.starts_with(":>") {
198                    self.advance(2);
199                    return Ok(Tok::Op(Op::OneTo));
200                }
201                if let Some(tok) = self.user_symbol() {
202                    return Ok(tok);
203                }
204                if self.starts_with("::") {
205                    self.advance(2);
206                    return Ok(Tok::ColonColon);
207                }
208                self.bump();
209                Ok(Tok::Colon)
210            }
211            '@' => {
212                if self.starts_with("@@") {
213                    self.advance(2);
214                    Ok(Tok::Op(Op::AtAt))
215                } else {
216                    self.bump();
217                    Ok(Tok::At)
218                }
219            }
220            '~' => {
221                if self.starts_with("~>") {
222                    self.advance(2);
223                    Ok(Tok::Op(Op::LeadsTo))
224                } else {
225                    self.bump();
226                    Ok(Tok::Op(Op::Not))
227                }
228            }
229            '&' | '$' | '?' | '%' | '#' | '!' | '^' | '(' | '+' | '*' => self
230                .user_symbol()
231                .map_or_else(|| self.scan_punctuation(c), Ok),
232            '.' => {
233                if self.starts_with("...") {
234                    self.advance(3);
235                    return Ok(Tok::Op(Op::User("...")));
236                }
237                if self.starts_with("..") {
238                    self.advance(2);
239                    Ok(Tok::Op(Op::DotDot))
240                } else {
241                    self.bump();
242                    Ok(Tok::Dot)
243                }
244            }
245            '[' => {
246                if self.starts_with("[]") {
247                    self.advance(2);
248                    Ok(Tok::Op(Op::Always))
249                } else {
250                    self.bump();
251                    Ok(Tok::LBrack)
252                }
253            }
254            _ => self.scan_punctuation(c),
255        }
256    }
257
258    fn scan_punctuation(&mut self, c: char) -> Result<Tok> {
259        let at = (self.line, self.col);
260        self.bump();
261        Ok(match c {
262            '(' => Tok::LParen,
263            ')' => Tok::RParen,
264            ']' => Tok::RBrack,
265            '{' => Tok::LBrace,
266            '}' => Tok::RBrace,
267            ',' => Tok::Comma,
268            '!' => Tok::Bang,
269            '\'' => Tok::Prime,
270            '_' => Tok::Underscore,
271            '#' => Tok::Op(Op::Neq),
272            '+' => Tok::Op(Op::Plus),
273            '*' => Tok::Op(Op::Times),
274            '%' => Tok::Op(Op::Mod),
275            '^' => Tok::Op(Op::Pow),
276            _ => {
277                return Err(Error::lex(
278                    format!("unexpected character {c:?}"),
279                    at.0,
280                    at.1,
281                ));
282            }
283        })
284    }
285
286    /// The longest symbol from the language's user-definable operator table
287    /// that starts here. The named `\`-operators are matched elsewhere.
288    fn user_symbol(&mut self) -> Option<Tok> {
289        let mut best: Option<&'static str> = None;
290        for (symbol, _) in token::USER_OPERATORS {
291            if symbol.starts_with('\\') || !self.starts_with(symbol) {
292                continue;
293            }
294            if best.is_none_or(|found| symbol.len() > found.len()) {
295                best = Some(symbol);
296            }
297        }
298        let symbol = best?;
299        self.advance(symbol.chars().count());
300        Some(Tok::Op(Op::User(symbol)))
301    }
302
303    /// `\b1011`, `\o777`, `\h1F` — a number written in another base. The
304    /// digits must follow the letter directly, which is what keeps `\o` as
305    /// concatenation everywhere else.
306    fn scan_based_number(&mut self) -> Result<Option<Tok>> {
307        let (radix, digits) = match self.at(1) {
308            Some('b' | 'B') => (2u32, "01"),
309            Some('o' | 'O') => (8, "01234567"),
310            Some('h' | 'H') => (16, "0123456789abcdefABCDEF"),
311            _ => return Ok(None),
312        };
313        let mut end = self.pos + 2;
314        while self.chars.get(end).is_some_and(|c| digits.contains(*c)) {
315            end += 1;
316        }
317        if end == self.pos + 2 {
318            return Ok(None);
319        }
320        let text: String = self.chars[self.pos + 2..end].iter().collect();
321        let value = i64::from_str_radix(&text, radix)
322            .map_err(|_| self.err(format!("number out of range: {text}")))?;
323        self.advance(2 + text.chars().count());
324        Ok(Some(Tok::Num(value)))
325    }
326
327    fn scan_number(&mut self) -> Result<Tok> {
328        let start = self.pos;
329        while self.peek().is_some_and(|c| c.is_ascii_digit()) {
330            self.bump();
331        }
332        // `1aMessage` is a name, not a number followed by one: TLA+ only asks
333        // that an identifier contain a letter, not that it begin with one.
334        if self
335            .peek()
336            .is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
337        {
338            self.pos = start;
339            return Ok(self.scan_word());
340        }
341        // `123.456` is one number; `1..2` is a range between two.
342        if self.peek() == Some('.') && self.at(1).is_some_and(|c| c.is_ascii_digit()) {
343            self.bump();
344            while self.peek().is_some_and(|c| c.is_ascii_digit()) {
345                self.bump();
346            }
347            return Ok(Tok::Decimal(self.chars[start..self.pos].iter().collect()));
348        }
349        let text: String = self.chars[start..self.pos].iter().collect();
350        text.parse()
351            .map(Tok::Num)
352            .map_err(|_| self.err(format!("integer literal out of range: {text}")))
353    }
354
355    fn scan_word(&mut self) -> Tok {
356        let start = self.pos;
357        while self
358            .peek()
359            .is_some_and(|c| c.is_ascii_alphanumeric() || c == '_')
360        {
361            self.bump();
362        }
363        let word: String = self.chars[start..self.pos].iter().collect();
364
365        // `WF_vars` is one word to the lexer but an operator plus a subscript.
366        // Immediately after a `.` a word is a field name, whatever else it
367        // would otherwise spell: `bar.NEW` and `bar.SF_` name fields.
368        if matches!(self.out.last().map(|t| &t.tok), Some(Tok::Dot)) {
369            return Tok::Ident(word);
370        }
371        // `WF_vars` is one word; `WF_<<a, b>>` is the same operator with a
372        // subscript the lexer cannot see, so the name is left empty.
373        for (prefix, strong) in [("WF_", false), ("SF_", true)] {
374            if let Some(rest) = word.strip_prefix(prefix) {
375                return Tok::Fair {
376                    strong,
377                    subscript: rest.to_string(),
378                };
379            }
380        }
381        match Kw::lookup(&word) {
382            Some(kw) => Tok::Kw(kw),
383            None => Tok::Ident(word),
384        }
385    }
386
387    fn scan_string(&mut self) -> Result<Tok> {
388        let (line, col) = (self.line, self.col);
389        self.bump();
390        let mut s = String::new();
391        loop {
392            match self.bump() {
393                None | Some('\n') => return Err(Error::lex("unterminated string", line, col)),
394                Some('"') => return Ok(Tok::Str(s)),
395                Some('\\') => {
396                    let escaped = self
397                        .bump()
398                        .ok_or_else(|| Error::lex("unterminated string", line, col))?;
399                    s.push(match escaped {
400                        'n' => '\n',
401                        't' => '\t',
402                        other => other,
403                    });
404                }
405                Some(c) => s.push(c),
406            }
407        }
408    }
409
410    fn scan_equals(&mut self) -> Result<Tok> {
411        let run = self.run_of('=');
412        if run >= 4 {
413            self.advance(run);
414            return Ok(Tok::ModuleEnd);
415        }
416        if run == 2 {
417            self.advance(2);
418            return Ok(Tok::DefEq);
419        }
420        if run == 1 {
421            return Ok(match self.at(1) {
422                Some('>') => {
423                    self.advance(2);
424                    Tok::Op(Op::Implies)
425                }
426                Some('<') => {
427                    self.advance(2);
428                    Tok::Op(Op::Le)
429                }
430                Some('|') => {
431                    self.advance(2);
432                    Tok::Op(Op::User("=|"))
433                }
434                _ => {
435                    self.bump();
436                    Tok::Op(Op::Eq)
437                }
438            });
439        }
440        Err(self.err("`===` is neither a definition nor a module terminator"))
441    }
442
443    /// A run of four or more dashes is a separator, so the operators spelled
444    /// with dashes have to be recognised around it rather than before it.
445    fn scan_dashes(&mut self) -> Tok {
446        if self.starts_with("-+->") {
447            self.advance(4);
448            return Tok::Op(Op::User("-+->"));
449        }
450        let run = self.run_of('-');
451        if run >= 4 {
452            self.advance(run);
453            return Tok::Separator;
454        }
455        if run == 2 {
456            self.advance(2);
457            return Tok::Op(Op::User("--"));
458        }
459        for (text, tok) in [("->", Tok::Arrow), ("-|", Tok::Op(Op::User("-|")))] {
460            if self.starts_with(text) {
461                self.advance(2);
462                return tok;
463            }
464        }
465        self.bump();
466        Tok::Op(Op::Minus)
467    }
468
469    fn scan_lt(&mut self) -> Tok {
470        for (text, tok) in [
471            ("<<", Tok::LTup),
472            ("<=>", Tok::Op(Op::Equiv)),
473            ("<=", Tok::Op(Op::Le)),
474            ("<-", Tok::Gets),
475            ("<>", Tok::Op(Op::Eventually)),
476            ("<:", Tok::Op(Op::User("<:"))),
477        ] {
478            if self.starts_with(text) {
479                self.advance(text.len());
480                return tok;
481            }
482        }
483        self.bump();
484        Tok::Op(Op::Lt)
485    }
486
487    fn scan_gt(&mut self) -> Tok {
488        for (text, tok) in [(">>", Tok::RTup), (">=", Tok::Op(Op::Ge))] {
489            if self.starts_with(text) {
490                self.advance(text.len());
491                return tok;
492            }
493        }
494        self.bump();
495        Tok::Op(Op::Gt)
496    }
497
498    fn scan_slash(&mut self) -> Result<Tok> {
499        if self.starts_with("/\\") {
500            self.advance(2);
501            return Ok(Tok::Op(Op::And));
502        }
503        if self.starts_with("/=") {
504            self.advance(2);
505            return Ok(Tok::Op(Op::Neq));
506        }
507        self.user_symbol()
508            .map_or_else(|| Err(self.err("stray `/`")), Ok)
509    }
510
511    fn scan_backslash(&mut self) -> Result<Tok> {
512        if self.starts_with("\\/") {
513            self.advance(2);
514            return Ok(Tok::Op(Op::Or));
515        }
516        if let Some(based) = self.scan_based_number()? {
517            return Ok(based);
518        }
519        let start = self.pos + 1;
520        let mut end = start;
521        while self.chars.get(end).is_some_and(char::is_ascii_alphabetic) {
522            end += 1;
523        }
524        if end == start {
525            self.bump();
526            return Ok(Tok::Op(Op::SetMinus));
527        }
528        let name: String = self.chars[start..end].iter().collect();
529        let op = match name.as_str() {
530            "in" => Op::In,
531            "notin" => Op::NotIn,
532            "subseteq" => Op::Subseteq,
533            "supseteq" => Op::Supseteq,
534            "cup" | "union" => Op::Cup,
535            "cap" | "intersect" => Op::Cap,
536            "times" | "X" => Op::Cartesian,
537            "div" => Op::Div,
538            "o" | "circ" => Op::Concat,
539            "equiv" => Op::Equiv,
540            "lnot" | "neg" => Op::Not,
541            "land" => Op::And,
542            "lor" => Op::Or,
543            "leq" => Op::Le,
544            "geq" => Op::Ge,
545            "neq" => Op::Neq,
546            "A" | "forall" => Op::Forall,
547            "E" | "exists" => Op::Exists,
548            "AA" => Op::TemporalForall,
549            "EE" => Op::TemporalExists,
550            _ => match token::user_operator(&format!("\\{name}")) {
551                Some(op) => op,
552                None => return Err(self.err(format!("unknown operator `\\{name}`"))),
553            },
554        };
555        self.advance(1 + name.len());
556        Ok(Tok::Op(op))
557    }
558}