Skip to main content

latex_rust/parser/
parse.rs

1//! LaTeX math string → [`MathNode`].
2//!
3//! Binding of `_` / `^` / primes is Pratt-style (tight postfix on the nucleus).
4//! The math list itself is a TeX-style row of atoms, not an arithmetic tree.
5
6use super::ast::{
7    AccentKind, AtomKind, ColSpec, DelimSize, Delimiter, EnvRow, EqNumber, IntegralKind, MathNode,
8    MatrixStyle, PhantomKind, SpaceKind, TextStyle,
9};
10use super::preproc::preprocess;
11use super::token::{tokenize, Token};
12use crate::color::{parse_color_spec, Color, ColorTable};
13use crate::dim::Dim;
14use crate::error::{Error, ParseError};
15use crate::symbols::{lookup, SymbolKind as CatalogKind};
16
17/// Parse a LaTeX math string into a [`MathNode`] using a fresh color table.
18///
19/// Accepts raw math, `$...$`, `$$...$$`, `\(...\)`, or `\[...\]`. Binding of
20/// `_` / `^` / primes is Pratt-style (tight postfix on the nucleus).
21///
22/// # Arguments
23///
24/// * `input` — math source, with or without delimiter fences.
25///
26/// # Returns
27///
28/// The typed AST, or a [`ParseError`] naming the problem.
29///
30/// # Errors
31///
32/// * [`ParseError::Unknown`] — command is not in the catalog.
33/// * [`ParseError::Unsupported`] — known construct this crate will not invent.
34/// * [`ParseError::Malformed`] — syntactically invalid input.
35/// * [`ParseError::UnmatchedDelimiter`] — `\left` without `\right` (or vice versa).
36/// * [`ParseError::TrailingBackslash`] — input ended with a stray `\`.
37///
38/// # Examples
39///
40/// ```
41/// use latex_rust::parse;
42///
43/// let ast = parse(r"\frac{1}{2}").unwrap();
44/// assert_eq!(ast.gold(), r#"(frac (atom Ord "1") (atom Ord "2"))"#);
45/// ```
46pub fn parse(input: &str) -> Result<MathNode, ParseError> {
47    parse_with_colors(input).map(|(n, _)| n)
48}
49
50/// Parse a math string, returning the AST and the color table after `\definecolor`.
51///
52/// # Arguments
53///
54/// * `input` — math source, with or without delimiter fences.
55///
56/// # Returns
57///
58/// The AST and the color table including any `\definecolor` names from `input`.
59///
60/// # Errors
61///
62/// Same as [`parse`].
63///
64/// # Examples
65///
66/// ```
67/// use latex_rust::parse_with_colors;
68///
69/// let (ast, table) = parse_with_colors(r"\definecolor{ok}{named}{red}x").unwrap();
70/// assert!(table.get("ok").is_ok());
71/// let _ = ast;
72/// ```
73pub fn parse_with_colors(input: &str) -> Result<(MathNode, ColorTable), ParseError> {
74    let sanitized = preprocess(input);
75    let tokens = tokenize(&sanitized)?;
76    let tokens = strip_fences(&tokens)?;
77    let mut p = Parser {
78        tokens,
79        pos: 0,
80        colors: ColorTable::new(),
81    };
82    let node = p.parse_list(Stop::eof())?;
83    p.skip_ws();
84    if p.pos < p.tokens.len() {
85        return Err(ParseError::Malformed(format!(
86            "unexpected leftover token {}",
87            p.tokens[p.pos]
88        )));
89    }
90    Ok((node, p.colors))
91}
92
93#[derive(Clone, Copy)]
94struct Stop {
95    end_group: bool,
96    amp: bool,
97    cr: bool,
98    right: bool,
99    end_env: bool,
100    rbracket: bool,
101}
102
103impl Stop {
104    fn eof() -> Self {
105        Self {
106            end_group: false,
107            amp: false,
108            cr: false,
109            right: false,
110            end_env: false,
111            rbracket: false,
112        }
113    }
114
115    fn group() -> Self {
116        Self {
117            end_group: true,
118            ..Self::eof()
119        }
120    }
121
122    fn cell() -> Self {
123        Self {
124            amp: true,
125            cr: true,
126            end_env: true,
127            ..Self::eof()
128        }
129    }
130
131    fn delim() -> Self {
132        Self {
133            right: true,
134            ..Self::eof()
135        }
136    }
137
138    fn index() -> Self {
139        Self {
140            rbracket: true,
141            ..Self::eof()
142        }
143    }
144
145    fn substack_line() -> Self {
146        Self {
147            end_group: true,
148            cr: true,
149            ..Self::eof()
150        }
151    }
152}
153
154struct Parser {
155    tokens: Vec<Token>,
156    pos: usize,
157    colors: ColorTable,
158}
159
160impl Parser {
161    fn skip_ws(&mut self) {
162        while matches!(self.tokens.get(self.pos), Some(Token::Space)) {
163            self.pos += 1;
164        }
165    }
166
167    fn peek(&self) -> Option<&Token> {
168        self.tokens.get(self.pos)
169    }
170
171    fn peek_ws(&mut self) -> Option<&Token> {
172        self.skip_ws();
173        self.peek()
174    }
175
176    fn bump(&mut self) -> Option<Token> {
177        self.skip_ws();
178        self.bump_raw()
179    }
180
181    fn bump_raw(&mut self) -> Option<Token> {
182        let t = self.tokens.get(self.pos).cloned()?;
183        self.pos += 1;
184        Some(t)
185    }
186
187    fn is_stop(&self, tok: &Token, stop: Stop) -> bool {
188        match tok {
189            Token::EndGroup if stop.end_group => true,
190            Token::AlignmentTab if stop.amp => true,
191            Token::Command(s) if s == "\\" && stop.cr => true,
192            Token::Command(s) if s == "cr" && stop.cr => true,
193            Token::Command(s) if s == "right" && stop.right => true,
194            Token::Command(s) if s == "end" && stop.end_env => true,
195            Token::Char(']') if stop.rbracket => true,
196            _ => false,
197        }
198    }
199
200    fn parse_list(&mut self, stop: Stop) -> Result<MathNode, ParseError> {
201        let mut items = Vec::new();
202        loop {
203            self.skip_ws();
204            let Some(tok) = self.peek().cloned() else {
205                break;
206            };
207            if self.is_stop(&tok, stop) {
208                break;
209            }
210            match &tok {
211                Token::MathShift | Token::DisplayShift => {
212                    return Err(ParseError::Malformed("unexpected math shift".into()));
213                }
214                Token::Command(n) if n == "color" => {
215                    self.bump();
216                    let c = self.parse_color_from_cmd()?;
217                    let rest = self.parse_list(stop)?;
218                    items.push(MathNode::Color(c, Box::new(rest)));
219                    break;
220                }
221                Token::Command(n) if n == "definecolor" => {
222                    self.bump();
223                    self.parse_definecolor()?;
224                    continue;
225                }
226                _ => {}
227            }
228            items.push(self.parse_atom()?);
229        }
230        Ok(wrap_row(items))
231    }
232
233    fn parse_atom(&mut self) -> Result<MathNode, ParseError> {
234        let mut nucleus = self.parse_nucleus()?;
235        let mut limits = None;
236        loop {
237            match self.peek_ws() {
238                Some(Token::Command(n)) if n == "limits" => {
239                    self.bump();
240                    limits = Some(true);
241                }
242                Some(Token::Command(n)) if n == "nolimits" => {
243                    self.bump();
244                    limits = Some(false);
245                }
246                _ => break,
247            }
248        }
249        if let Some(flag) = limits {
250            if let MathNode::Operator(name, _) = nucleus {
251                nucleus = MathNode::Operator(name, flag);
252            }
253        }
254        self.bind_scripts(nucleus)
255    }
256
257    fn parse_nucleus(&mut self) -> Result<MathNode, ParseError> {
258        self.skip_ws();
259        match self.peek().cloned() {
260            None => Err(ParseError::Malformed("unexpected end of input".into())),
261            Some(Token::Superscript | Token::Subscript | Token::Char('\'')) => {
262                Ok(MathNode::Row(Vec::new()))
263            }
264            Some(Token::BeginGroup) => self.parse_group(),
265            Some(Token::Char(c)) => {
266                self.bump();
267                Ok(MathNode::Atom(c, atom_kind(c)))
268            }
269            Some(Token::Command(name)) => {
270                self.bump();
271                self.parse_command(&name)
272            }
273            Some(other) => Err(ParseError::Malformed(format!("unexpected token {other}"))),
274        }
275    }
276
277    fn parse_group(&mut self) -> Result<MathNode, ParseError> {
278        match self.bump() {
279            Some(Token::BeginGroup) => {}
280            _ => return Err(ParseError::Malformed("expected '{'".into())),
281        }
282        let inner = self.parse_list(Stop::group())?;
283        match self.bump() {
284            Some(Token::EndGroup) => Ok(inner),
285            _ => Err(ParseError::Malformed("unmatched '{'".into())),
286        }
287    }
288
289    fn parse_arg(&mut self) -> Result<MathNode, ParseError> {
290        self.skip_ws();
291        match self.peek() {
292            Some(Token::BeginGroup) => self.parse_group(),
293            Some(_) => self.parse_nucleus(),
294            None => Err(ParseError::Malformed("missing argument".into())),
295        }
296    }
297
298    fn parse_script(&mut self) -> Result<MathNode, ParseError> {
299        self.parse_arg()
300    }
301
302    fn bind_scripts(&mut self, nucleus: MathNode) -> Result<MathNode, ParseError> {
303        let mut sub: Option<MathNode> = None;
304        let mut sup: Option<MathNode> = None;
305        let mut sup_from_prime = false;
306        loop {
307            match self.peek_ws() {
308                Some(Token::Subscript) => {
309                    self.bump();
310                    if sub.is_some() {
311                        return Err(ParseError::Malformed("double subscript".into()));
312                    }
313                    sub = Some(self.parse_script()?);
314                }
315                Some(Token::Superscript) => {
316                    self.bump();
317                    let s = self.parse_script()?;
318                    if let Some(prev) = sup.take() {
319                        if !sup_from_prime {
320                            return Err(ParseError::Malformed("double superscript".into()));
321                        }
322                        sup = Some(wrap_row(vec![prev, s]));
323                        sup_from_prime = false;
324                    } else {
325                        sup = Some(s);
326                    }
327                }
328                Some(Token::Char('\'')) => {
329                    self.bump();
330                    let prime = MathNode::Atom('′', AtomKind::Ord);
331                    sup = Some(match sup.take() {
332                        None => prime,
333                        Some(prev) => wrap_row(vec![prev, prime]),
334                    });
335                    sup_from_prime = true;
336                }
337                _ => break,
338            }
339        }
340        Ok(apply_scripts(nucleus, sub, sup))
341    }
342
343    fn parse_command(&mut self, name: &str) -> Result<MathNode, ParseError> {
344        match name {
345            "frac" | "dfrac" | "tfrac" | "cfrac" => {
346                let n = self.parse_arg()?;
347                let d = self.parse_arg()?;
348                Ok(MathNode::Fraction(Box::new(n), Box::new(d)))
349            }
350            "binom" | "dbinom" | "tbinom" => {
351                let n = self.parse_arg()?;
352                let k = self.parse_arg()?;
353                Ok(MathNode::Delimited(
354                    Delimiter::Char('('),
355                    Box::new(MathNode::Fraction(Box::new(n), Box::new(k))),
356                    Delimiter::Char(')'),
357                ))
358            }
359            "genfrac" => self.parse_genfrac(),
360            "sqrt" => {
361                let index = if matches!(self.peek_ws(), Some(Token::Char('['))) {
362                    self.bump();
363                    let idx = self.parse_list(Stop::index())?;
364                    match self.bump() {
365                        Some(Token::Char(']')) => {}
366                        _ => {
367                            return Err(ParseError::Malformed(
368                                "expected ']' after \\sqrt index".into(),
369                            ))
370                        }
371                    }
372                    Some(Box::new(idx))
373                } else {
374                    None
375                };
376                let rad = self.parse_arg()?;
377                Ok(MathNode::Radical(index, Box::new(rad)))
378            }
379            "left" => self.parse_delimited(),
380            "right" => Err(ParseError::UnmatchedDelimiter),
381            "begin" => self.parse_begin(),
382            "end" => Err(ParseError::Malformed("unexpected \\end".into())),
383            "over" => Err(ParseError::Malformed("\\over outside a group".into())),
384            "choose" => Err(ParseError::Malformed("\\choose outside a group".into())),
385            "hat" => self.accent(AccentKind::Hat),
386            "check" => self.accent(AccentKind::Check),
387            "breve" => self.accent(AccentKind::Breve),
388            "acute" => self.accent(AccentKind::Acute),
389            "grave" => self.accent(AccentKind::Grave),
390            "tilde" => self.accent(AccentKind::Tilde),
391            "bar" => self.accent(AccentKind::Bar),
392            "vec" => self.accent(AccentKind::Vec),
393            "dot" => self.accent(AccentKind::Dot),
394            "ddot" => self.accent(AccentKind::Ddot),
395            "dddot" => self.accent(AccentKind::Dddot),
396            "ddddot" => self.accent(AccentKind::Ddddot),
397            "widehat" => self.accent(AccentKind::WideHat),
398            "widetilde" => self.accent(AccentKind::WideTilde),
399            "overline" => self.accent(AccentKind::Overline),
400            "underline" => self.accent(AccentKind::Underline),
401            "overbrace" => self.accent(AccentKind::Overbrace),
402            "underbrace" => self.accent(AccentKind::Underbrace),
403            "overleftarrow" => self.accent(AccentKind::Overleftarrow),
404            "overrightarrow" => self.accent(AccentKind::Overrightarrow),
405            "overleftrightarrow" => self.accent(AccentKind::Overleftrightarrow),
406            "underleftarrow" => self.accent(AccentKind::Underleftarrow),
407            "underrightarrow" => self.accent(AccentKind::Underrightarrow),
408            "underleftrightarrow" => self.accent(AccentKind::Underleftrightarrow),
409            "cancel" => self.accent(AccentKind::Cancel),
410            "bcancel" => self.accent(AccentKind::BCancel),
411            "xcancel" => self.accent(AccentKind::XCancel),
412            "boxed" | "fbox" => self.accent(AccentKind::Boxed),
413            "mathring" => self.accent(AccentKind::Ring),
414            "cancelto" => {
415                let value = self.parse_arg()?;
416                let expr = self.parse_arg()?;
417                if is_empty_node(&expr) {
418                    return Err(ParseError::Malformed("empty accent base".into()));
419                }
420                Ok(MathNode::CancelTo(Box::new(value), Box::new(expr)))
421            }
422            "not" => {
423                let body = self.parse_nucleus()?;
424                Ok(MathNode::Accent(Box::new(body), AccentKind::Not))
425            }
426            "overset" => {
427                let over = self.parse_arg()?;
428                let base = self.parse_arg()?;
429                Ok(MathNode::OverUnder(
430                    Box::new(base),
431                    Some(Box::new(over)),
432                    None,
433                ))
434            }
435            "underset" => {
436                let under = self.parse_arg()?;
437                let base = self.parse_arg()?;
438                Ok(MathNode::OverUnder(
439                    Box::new(base),
440                    None,
441                    Some(Box::new(under)),
442                ))
443            }
444            "stackrel" => {
445                let over = self.parse_arg()?;
446                let base = self.parse_arg()?;
447                Ok(MathNode::OverUnder(
448                    Box::new(base),
449                    Some(Box::new(over)),
450                    None,
451                ))
452            }
453            "mathrm" | "textrm" => self.font(TextStyle::Rm),
454            "mathbf" | "textbf" => self.font(TextStyle::Bf),
455            "mathit" | "textit" => self.font(TextStyle::It),
456            "mathsf" | "textsf" => self.font(TextStyle::Sf),
457            "mathtt" | "texttt" => self.font(TextStyle::Tt),
458            "mathbb" => self.font(TextStyle::Bb),
459            "mathcal" => self.font(TextStyle::Cal),
460            "mathfrak" => self.font(TextStyle::Frak),
461            "mathscr" => self.font(TextStyle::Scr),
462            "boldsymbol" => self.font(TextStyle::Boldsymbol),
463            "pmb" => self.font(TextStyle::Pmb),
464            "xrightarrow" => self.parse_xarrow("longrightarrow"),
465            "xleftarrow" => self.parse_xarrow("longleftarrow"),
466            "text" | "mbox" => self.parse_text(TextStyle::Text),
467            "operatorname" => {
468                let name = self.collect_group_text()?;
469                Ok(MathNode::Operator(name, false))
470            }
471            "," => Ok(MathNode::Space(SpaceKind::Thin)),
472            ":" | ">" => Ok(MathNode::Space(SpaceKind::Medium)),
473            ";" => Ok(MathNode::Space(SpaceKind::Thick)),
474            "!" => Ok(MathNode::Space(SpaceKind::NegThin)),
475            "quad" => Ok(MathNode::Space(SpaceKind::Quad)),
476            "qquad" => Ok(MathNode::Space(SpaceKind::Qquad)),
477            " " => Ok(MathNode::Space(SpaceKind::ControlSpace)),
478            "hspace" => {
479                let spec = self.collect_group_text()?;
480                let d = parse_tex_dim(&spec)?;
481                Ok(MathNode::Space(SpaceKind::Hspace(d)))
482            }
483            "phantom" => {
484                let b = self.parse_arg()?;
485                Ok(MathNode::Phantom(PhantomKind::Full, Box::new(b)))
486            }
487            "vphantom" => {
488                let b = self.parse_arg()?;
489                Ok(MathNode::Phantom(PhantomKind::Vertical, Box::new(b)))
490            }
491            "hphantom" => {
492                let b = self.parse_arg()?;
493                Ok(MathNode::Phantom(PhantomKind::Horizontal, Box::new(b)))
494            }
495            "strut" => Ok(MathNode::Strut(Dim::ratio(7, 10), Dim::ratio(3, 10))),
496            "rule" => {
497                let _w = parse_tex_dim(&self.collect_group_text()?)?;
498                let h = parse_tex_dim(&self.collect_group_text()?)?;
499                Ok(MathNode::Strut(h, Dim::zero()))
500            }
501            "textcolor" => {
502                let c = self.parse_color_from_cmd()?;
503                let body = self.parse_arg()?;
504                Ok(MathNode::TextColor(c, Box::new(body)))
505            }
506            "colorbox" => {
507                let c = self.parse_color_from_cmd()?;
508                let body = self.parse_arg()?;
509                Ok(MathNode::ColorBox(c, Box::new(body)))
510            }
511            "fcolorbox" => {
512                let border = self.parse_color_from_cmd()?;
513                let fill = self.parse_color_from_cmd()?;
514                let body = self.parse_arg()?;
515                Ok(MathNode::FColorBox(border, fill, Box::new(body)))
516            }
517            "sum" => Ok(MathNode::Sum(None, None)),
518            "prod" => Ok(MathNode::Product(None, None)),
519            "int" => Ok(MathNode::Integral(IntegralKind::Int, None, None)),
520            "iint" => Ok(MathNode::Integral(IntegralKind::Iint, None, None)),
521            "iiint" => Ok(MathNode::Integral(IntegralKind::Iiint, None, None)),
522            "oint" => Ok(MathNode::Integral(IntegralKind::Oint, None, None)),
523            "oiint" => Ok(MathNode::Integral(IntegralKind::Oiint, None, None)),
524            "lim" => Ok(MathNode::Limit(None)),
525            "sin" | "cos" | "tan" | "cot" | "sec" | "csc" | "arcsin" | "arccos" | "arctan"
526            | "sinh" | "cosh" | "tanh" | "coth" | "log" | "ln" | "lg" | "exp" | "limsup"
527            | "liminf" | "sup" | "inf" | "max" | "min" | "det" | "dim" | "ker" | "deg" | "gcd"
528            | "lcm" | "Pr" | "arg" => Ok(MathNode::Operator(name.to_string(), false)),
529            "coprod" | "bigcup" | "bigcap" | "bigsqcup" | "bigvee" | "bigwedge" | "bigoplus"
530            | "bigotimes" | "biguplus" => Ok(MathNode::Operator(name.to_string(), true)),
531            "big" | "Big" | "bigg" | "Bigg" | "bigl" | "bigr" | "Bigl" | "Bigr" | "biggl"
532            | "biggr" | "Biggl" | "Biggr" | "bigm" | "Bigm" | "biggm" | "Biggm" => {
533                self.parse_sized_delim(name)
534            }
535            "tag" => {
536                let star = matches!(self.peek_ws(), Some(Token::Char('*')));
537                if star {
538                    self.bump();
539                }
540                let body = self.parse_arg()?;
541                Ok(MathNode::Tag {
542                    star,
543                    body: Box::new(body),
544                })
545            }
546            "label" => {
547                let key = self.collect_group_text()?;
548                Ok(MathNode::Label(key))
549            }
550            "ref" => {
551                let key = self.collect_group_text()?;
552                Ok(MathNode::Ref(key))
553            }
554            "nonumber" | "notag" => Ok(MathNode::NoNumber),
555            "hline" => Ok(MathNode::Hline),
556            "intertext" => {
557                let s = self.collect_group_text()?;
558                Ok(MathNode::Intertext(Box::new(MathNode::Text(
559                    s,
560                    TextStyle::Text,
561                ))))
562            }
563            "substack" => self.parse_substack(),
564            "displaystyle" | "textstyle" | "scriptstyle" | "scriptscriptstyle" | "limits"
565            | "nolimits" => self.parse_nucleus(),
566            "{" | "}" => {
567                let c = name.chars().next().unwrap_or('{');
568                Ok(MathNode::Atom(
569                    c,
570                    if name == "{" {
571                        AtomKind::Open
572                    } else {
573                        AtomKind::Close
574                    },
575                ))
576            }
577            "|" => Ok(MathNode::Symbol("Vert".into())),
578            "backslash" => Ok(MathNode::Symbol("backslash".into())),
579            _ => {
580                if name.starts_with("math")
581                    && name.len() > 4
582                    && name.chars().all(|c| c.is_ascii_alphabetic())
583                {
584                    return Err(ParseError::Unsupported(format!("font style {name}")));
585                }
586                if name.starts_with("wide") {
587                    return Err(ParseError::Unsupported(format!("accent {name}")));
588                }
589                self.parse_symbol_or_unknown(name)
590            }
591        }
592    }
593
594    fn accent(&mut self, kind: AccentKind) -> Result<MathNode, ParseError> {
595        let body = self.parse_arg()?;
596        if is_empty_node(&body) {
597            return Err(ParseError::Malformed("empty accent base".into()));
598        }
599        Ok(MathNode::Accent(Box::new(body), kind))
600    }
601
602    fn parse_xarrow(&mut self, arrow: &str) -> Result<MathNode, ParseError> {
603        let under = if matches!(self.peek_ws(), Some(Token::Char('['))) {
604            self.bump();
605            let u = self.parse_list(Stop::index())?;
606            match self.bump() {
607                Some(Token::Char(']')) => Some(Box::new(u)),
608                _ => {
609                    return Err(ParseError::Malformed(
610                        "expected ']' after x-arrow optional argument".into(),
611                    ))
612                }
613            }
614        } else {
615            None
616        };
617        let over = self.parse_arg()?;
618        Ok(MathNode::OverUnder(
619            Box::new(MathNode::Symbol(arrow.to_string())),
620            Some(Box::new(over)),
621            under,
622        ))
623    }
624
625    fn font(&mut self, style: TextStyle) -> Result<MathNode, ParseError> {
626        let inner = self.parse_arg()?;
627        Ok(collapse_text(apply_text_style(inner, style)))
628    }
629
630    fn parse_text(&mut self, style: TextStyle) -> Result<MathNode, ParseError> {
631        let s = self.collect_group_text()?;
632        Ok(MathNode::Text(s, style))
633    }
634
635    fn parse_delimited(&mut self) -> Result<MathNode, ParseError> {
636        let open = self.parse_delimiter()?;
637        let body = self.parse_list(Stop::delim())?;
638        match self.bump() {
639            Some(Token::Command(n)) if n == "right" => {}
640            _ => return Err(ParseError::UnmatchedDelimiter),
641        }
642        let close = self.parse_delimiter()?;
643        Ok(MathNode::Delimited(open, Box::new(body), close))
644    }
645
646    fn parse_sized_delim(&mut self, name: &str) -> Result<MathNode, ParseError> {
647        let size = DelimSize::from_command(name)
648            .ok_or_else(|| ParseError::Malformed(format!("unknown delimiter size \\{name}")))?;
649        let d = self.parse_delimiter()?;
650        let class = DelimSize::class_from_command(name).unwrap_or_else(|| match &d {
651            Delimiter::Char(c) => atom_kind(*c),
652            Delimiter::Named(n) if n == "{" => AtomKind::Open,
653            Delimiter::Named(n) if n == "}" => AtomKind::Close,
654            _ => AtomKind::Open,
655        });
656        Ok(MathNode::SizedDelim(d, size, class))
657    }
658
659    fn parse_delimiter(&mut self) -> Result<Delimiter, ParseError> {
660        self.skip_ws();
661        match self.bump() {
662            Some(Token::Char('.')) => Ok(Delimiter::Empty),
663            Some(Token::Char(c)) if matches!(c, '(' | ')' | '[' | ']' | '|' | '/' | '<' | '>') => {
664                Ok(Delimiter::Char(c))
665            }
666            Some(Token::Command(n)) => match n.as_str() {
667                "." => Ok(Delimiter::Empty),
668                "{" | "}" | "|" => Ok(Delimiter::Named(n)),
669                "langle" | "rangle" | "lfloor" | "rfloor" | "lceil" | "rceil" | "lvert"
670                | "rvert" | "lVert" | "rVert" | "vert" | "Vert" | "uparrow" | "downarrow"
671                | "Uparrow" | "Downarrow" | "updownarrow" | "Updownarrow" | "backslash"
672                | "lgroup" | "rgroup" | "lmoustache" | "rmoustache" => Ok(Delimiter::Named(n)),
673                other => Err(ParseError::Malformed(format!(
674                    "unknown delimiter \\{other}"
675                ))),
676            },
677            Some(other) => Err(ParseError::Malformed(format!(
678                "expected delimiter, found {other}"
679            ))),
680            None => Err(ParseError::Malformed("expected delimiter".into())),
681        }
682    }
683
684    fn parse_begin(&mut self) -> Result<MathNode, ParseError> {
685        let name = self.collect_group_text()?;
686        let colspec = if name == "array" {
687            let preamble = self.collect_group_text()?;
688            parse_colspec(&preamble)?
689        } else {
690            Vec::new()
691        };
692        let style = match name.as_str() {
693            "matrix" => MatrixStyle::Matrix,
694            "pmatrix" => MatrixStyle::Pmatrix,
695            "bmatrix" => MatrixStyle::Bmatrix,
696            "vmatrix" => MatrixStyle::Vmatrix,
697            "Vmatrix" => MatrixStyle::VVmatrix,
698            "Bmatrix" => MatrixStyle::BBmatrix,
699            "cases" => MatrixStyle::Cases,
700            "array" => MatrixStyle::Array,
701            "aligned" => MatrixStyle::Aligned,
702            "align" => MatrixStyle::Align,
703            "gather" => MatrixStyle::Gather,
704            "multline" => MatrixStyle::Multline,
705            "equation" => MatrixStyle::Equation,
706            "split" => MatrixStyle::Split,
707            other => {
708                return Err(ParseError::Unsupported(format!("environment {other}")));
709            }
710        };
711        let rows = self.parse_rows()?;
712        self.expect_end(&name)?;
713        Ok(MathNode::Matrix(style, colspec, rows))
714    }
715
716    fn parse_substack(&mut self) -> Result<MathNode, ParseError> {
717        self.skip_ws();
718        match self.bump() {
719            Some(Token::BeginGroup) => {}
720            _ => {
721                return Err(ParseError::Malformed(
722                    "expected '{' after \\substack".into(),
723                ))
724            }
725        }
726        let mut lines = Vec::new();
727        loop {
728            self.skip_ws();
729            if matches!(self.peek(), Some(Token::EndGroup)) {
730                self.bump();
731                break;
732            }
733            let line = self.parse_list(Stop::substack_line())?;
734            lines.push(line);
735            self.skip_ws();
736            match self.peek() {
737                Some(Token::Command(n)) if n == "\\" || n == "cr" => {
738                    self.bump();
739                }
740                Some(Token::EndGroup) => {
741                    self.bump();
742                    break;
743                }
744                None => return Err(ParseError::Malformed("unmatched '{' in \\substack".into())),
745                Some(other) => {
746                    return Err(ParseError::Malformed(format!(
747                        "unexpected token {other} in \\substack"
748                    )))
749                }
750            }
751        }
752        if lines.is_empty() {
753            return Err(ParseError::Malformed("empty \\substack".into()));
754        }
755        Ok(MathNode::Substack(lines))
756    }
757
758    fn parse_rows(&mut self) -> Result<Vec<EnvRow>, ParseError> {
759        self.skip_ws();
760        if matches!(self.peek(), Some(Token::Command(n)) if n == "end") {
761            return Ok(Vec::new());
762        }
763        let mut rows = Vec::new();
764        loop {
765            self.skip_ws();
766            if matches!(self.peek(), Some(Token::Command(n)) if n == "end") {
767                return Ok(rows);
768            }
769            if matches!(self.peek(), Some(Token::Command(n)) if n == "hline") {
770                self.bump();
771                rows.push(EnvRow::Hline);
772                continue;
773            }
774            if matches!(self.peek(), Some(Token::Command(n)) if n == "intertext") {
775                self.bump();
776                let s = self.collect_group_text()?;
777                rows.push(EnvRow::Intertext(Box::new(MathNode::Text(
778                    s,
779                    TextStyle::Text,
780                ))));
781                continue;
782            }
783            let mut cells = Vec::new();
784            let mut number = EqNumber::Default;
785            let mut labels = Vec::new();
786            loop {
787                let cell = self.parse_list(Stop::cell())?;
788                let cell = peel_row_meta(cell, &mut number, &mut labels);
789                cells.push(cell);
790                self.skip_ws();
791                match self.peek() {
792                    Some(Token::AlignmentTab) => {
793                        self.bump();
794                    }
795                    Some(Token::Command(n)) if n == "\\" || n == "cr" => {
796                        self.bump();
797                        rows.push(finish_env_row(cells, number, labels));
798                        self.skip_ws();
799                        if matches!(self.peek(), Some(Token::Command(e)) if e == "end") {
800                            return Ok(rows);
801                        }
802                        break;
803                    }
804                    Some(Token::Command(n)) if n == "end" => {
805                        rows.push(finish_env_row(cells, number, labels));
806                        return Ok(rows);
807                    }
808                    None => {
809                        return Err(ParseError::Malformed(
810                            "unmatched \\begin (missing \\end)".into(),
811                        ));
812                    }
813                    Some(other) => {
814                        return Err(ParseError::Malformed(format!(
815                            "unexpected token {other} in environment body"
816                        )));
817                    }
818                }
819            }
820        }
821    }
822
823    fn expect_end(&mut self, name: &str) -> Result<(), ParseError> {
824        match self.bump() {
825            Some(Token::Command(n)) if n == "end" => {}
826            _ => return Err(ParseError::Malformed(format!("expected \\end{{{name}}}"))),
827        }
828        let got = self.collect_group_text()?;
829        if got != name {
830            return Err(ParseError::Malformed(format!(
831                "\\begin{{{name}}} closed by \\end{{{got}}}"
832            )));
833        }
834        Ok(())
835    }
836
837    fn parse_genfrac(&mut self) -> Result<MathNode, ParseError> {
838        let ldel = self.collect_group_text()?;
839        let rdel = self.collect_group_text()?;
840        let _thickness = self.collect_group_text()?;
841        let _style = self.collect_group_text()?;
842        let num = self.parse_arg()?;
843        let den = self.parse_arg()?;
844        let frac = MathNode::Fraction(Box::new(num), Box::new(den));
845        if ldel.is_empty() && rdel.is_empty() {
846            return Ok(frac);
847        }
848        Ok(MathNode::Delimited(
849            delim_from_text(&ldel)?,
850            Box::new(frac),
851            delim_from_text(&rdel)?,
852        ))
853    }
854
855    fn parse_color_from_cmd(&mut self) -> Result<Color, ParseError> {
856        self.skip_ws();
857        let model = if matches!(self.peek(), Some(Token::Char('['))) {
858            self.bump();
859            let m = self.collect_until_char(']')?;
860            match self.bump() {
861                Some(Token::Char(']')) => {}
862                _ => {
863                    return Err(ParseError::Malformed(
864                        "expected ']' after color model".into(),
865                    ))
866                }
867            }
868            m
869        } else {
870            "named".into()
871        };
872        let spec = self.collect_group_text()?;
873        parse_color_spec(&model, &spec, Some(&self.colors)).map_err(color_err)
874    }
875
876    fn parse_definecolor(&mut self) -> Result<(), ParseError> {
877        let name = self.collect_group_text()?;
878        let model = self.collect_group_text()?;
879        let spec = self.collect_group_text()?;
880        self.colors
881            .define(&name, &model, &spec)
882            .map_err(color_err)?;
883        Ok(())
884    }
885
886    fn collect_group_text(&mut self) -> Result<String, ParseError> {
887        self.skip_ws();
888        match self.bump_raw() {
889            Some(Token::Space) => {
890                self.pos -= 1;
891                self.skip_ws();
892                return self.collect_group_text();
893            }
894            Some(Token::BeginGroup) => {}
895            _ => return Err(ParseError::Malformed("expected '{'".into())),
896        }
897        let mut s = String::new();
898        let mut depth = 1;
899        while depth > 0 {
900            match self.bump_raw() {
901                None => return Err(ParseError::Malformed("unmatched '{'".into())),
902                Some(Token::BeginGroup) => {
903                    depth += 1;
904                    s.push('{');
905                }
906                Some(Token::EndGroup) => {
907                    depth -= 1;
908                    if depth > 0 {
909                        s.push('}');
910                    }
911                }
912                Some(Token::Space) => s.push(' '),
913                Some(Token::Char(c)) => s.push(c),
914                Some(Token::Command(n)) => {
915                    if n.len() == 1 {
916                        s.push(n.chars().next().unwrap_or('\\'));
917                    } else {
918                        s.push('\\');
919                        s.push_str(&n);
920                    }
921                }
922                Some(other) => {
923                    return Err(ParseError::Malformed(format!(
924                        "unexpected token {other} in group text"
925                    )))
926                }
927            }
928        }
929        Ok(s)
930    }
931
932    fn collect_until_char(&mut self, end: char) -> Result<String, ParseError> {
933        let mut s = String::new();
934        loop {
935            match self.peek() {
936                None => return Err(ParseError::Malformed(format!("expected '{end}'"))),
937                Some(Token::Char(c)) if *c == end => break,
938                Some(Token::Char(c)) => {
939                    s.push(*c);
940                    self.bump_raw();
941                }
942                Some(Token::Space) => {
943                    s.push(' ');
944                    self.bump_raw();
945                }
946                Some(other) => {
947                    return Err(ParseError::Malformed(format!(
948                        "unexpected token {other} in optional argument"
949                    )))
950                }
951            }
952        }
953        Ok(s)
954    }
955
956    fn parse_symbol_or_unknown(&self, name: &str) -> Result<MathNode, ParseError> {
957        let canon = alias(name);
958        if let Some(e) = lookup(canon).or_else(|| lookup(name)) {
959            match e.kind {
960                CatalogKind::Container | CatalogKind::Modifier => {
961                    return Err(ParseError::Unsupported(format!("\\{name}")));
962                }
963                CatalogKind::Symbol | CatalogKind::Operator => {
964                    return Ok(MathNode::Symbol(canon.to_string()));
965                }
966            }
967        }
968        if is_extra_symbol(canon) {
969            return Ok(MathNode::Symbol(canon.to_string()));
970        }
971        Err(ParseError::Unknown(format!("\\{name}")))
972    }
973}
974
975fn strip_fences(tokens: &[Token]) -> Result<Vec<Token>, ParseError> {
976    let t = trim_spaces(tokens);
977    if t.len() >= 2 {
978        let inner = match (&t[0], &t[t.len() - 1]) {
979            (Token::MathShift, Token::MathShift) => Some(&t[1..t.len() - 1]),
980            (Token::DisplayShift, Token::DisplayShift) => Some(&t[1..t.len() - 1]),
981            (Token::Command(a), Token::Command(b)) if a == "[" && b == "]" => {
982                Some(&t[1..t.len() - 1])
983            }
984            (Token::Command(a), Token::Command(b)) if a == "(" && b == ")" => {
985                Some(&t[1..t.len() - 1])
986            }
987            _ => None,
988        };
989        if let Some(inner) = inner {
990            return Ok(trim_spaces(inner).to_vec());
991        }
992        if matches!(t[0], Token::MathShift | Token::DisplayShift)
993            || matches!(&t[0], Token::Command(s) if s == "[" || s == "(")
994        {
995            return Err(ParseError::Malformed("unmatched math delimiter".into()));
996        }
997    }
998    Ok(t.to_vec())
999}
1000
1001fn trim_spaces(tokens: &[Token]) -> &[Token] {
1002    let mut a = 0;
1003    let mut b = tokens.len();
1004    while a < b && tokens[a] == Token::Space {
1005        a += 1;
1006    }
1007    while b > a && tokens[b - 1] == Token::Space {
1008        b -= 1;
1009    }
1010    &tokens[a..b]
1011}
1012
1013fn wrap_row(mut items: Vec<MathNode>) -> MathNode {
1014    if items.len() == 1 {
1015        items.remove(0)
1016    } else {
1017        MathNode::Row(items)
1018    }
1019}
1020
1021fn parse_colspec(s: &str) -> Result<Vec<ColSpec>, ParseError> {
1022    let mut out = Vec::new();
1023    for c in s.chars() {
1024        match c {
1025            'l' => out.push(ColSpec::Left),
1026            'c' => out.push(ColSpec::Center),
1027            'r' => out.push(ColSpec::Right),
1028            '|' => out.push(ColSpec::VRule),
1029            ' ' | '\t' => {}
1030            '@' | '!' | '>' | '<' | 'p' | 'm' | 'b' | '*' => {
1031                return Err(ParseError::Unsupported(format!("array preamble `{c}`")))
1032            }
1033            other => return Err(ParseError::Malformed(format!("array preamble `{other}`"))),
1034        }
1035    }
1036    if out.is_empty() {
1037        return Err(ParseError::Malformed("empty array preamble".into()));
1038    }
1039    Ok(out)
1040}
1041
1042fn peel_row_meta(node: MathNode, number: &mut EqNumber, labels: &mut Vec<String>) -> MathNode {
1043    match node {
1044        MathNode::NoNumber => {
1045            *number = EqNumber::Suppress;
1046            MathNode::Row(Vec::new())
1047        }
1048        MathNode::Tag { star, body } => {
1049            *number = EqNumber::Tag { star, body };
1050            MathNode::Row(Vec::new())
1051        }
1052        MathNode::Label(k) => {
1053            labels.push(k);
1054            MathNode::Row(Vec::new())
1055        }
1056        MathNode::Hline => MathNode::Hline,
1057        MathNode::Intertext(n) => MathNode::Intertext(n),
1058        MathNode::Row(items) => {
1059            let mut kept = Vec::new();
1060            for it in items {
1061                let p = peel_row_meta(it, number, labels);
1062                if !is_empty_node(&p) {
1063                    kept.push(p);
1064                }
1065            }
1066            wrap_row(kept)
1067        }
1068        other => other,
1069    }
1070}
1071
1072fn finish_env_row(cells: Vec<MathNode>, number: EqNumber, labels: Vec<String>) -> EnvRow {
1073    if cells.len() == 1 && matches!(cells[0], MathNode::Hline) {
1074        return EnvRow::Hline;
1075    }
1076    if cells.len() == 1 {
1077        if let MathNode::Intertext(n) = &cells[0] {
1078            return EnvRow::Intertext(n.clone());
1079        }
1080    }
1081    EnvRow::Cells {
1082        cells,
1083        number,
1084        labels,
1085    }
1086}
1087
1088fn is_empty_node(n: &MathNode) -> bool {
1089    match n {
1090        MathNode::Row(v) => v.is_empty() || v.iter().all(is_empty_node),
1091        MathNode::Space(_) | MathNode::NoNumber | MathNode::Label(_) => true,
1092        _ => false,
1093    }
1094}
1095
1096fn apply_scripts(nucleus: MathNode, sub: Option<MathNode>, sup: Option<MathNode>) -> MathNode {
1097    match nucleus {
1098        MathNode::Sum(None, None) => MathNode::Sum(sub.map(Box::new), sup.map(Box::new)),
1099        MathNode::Product(None, None) => MathNode::Product(sub.map(Box::new), sup.map(Box::new)),
1100        MathNode::Integral(k, None, None) => {
1101            MathNode::Integral(k, sub.map(Box::new), sup.map(Box::new))
1102        }
1103        MathNode::Limit(None) => {
1104            let lim = MathNode::Limit(sub.map(Box::new));
1105            match sup {
1106                Some(s) => MathNode::Superscript(Box::new(lim), Box::new(s)),
1107                None => lim,
1108            }
1109        }
1110        MathNode::Accent(b, k @ (AccentKind::Overbrace | AccentKind::Underbrace)) => {
1111            match (sub, sup) {
1112                (None, None) => MathNode::Accent(b, k),
1113                (s, e) => MathNode::OverUnder(
1114                    Box::new(MathNode::Accent(b, k)),
1115                    e.map(Box::new),
1116                    s.map(Box::new),
1117                ),
1118            }
1119        }
1120        other => match (sub, sup) {
1121            (None, None) => other,
1122            (Some(s), None) => MathNode::Subscript(Box::new(other), Box::new(s)),
1123            (None, Some(e)) => MathNode::Superscript(Box::new(other), Box::new(e)),
1124            (Some(s), Some(e)) => MathNode::SubSup(Box::new(other), Box::new(s), Box::new(e)),
1125        },
1126    }
1127}
1128
1129fn apply_text_style(node: MathNode, style: TextStyle) -> MathNode {
1130    match node {
1131        MathNode::Atom(c, _) if crate::style_map::is_stylable(c) => {
1132            MathNode::Text(c.to_string(), style)
1133        }
1134        MathNode::Text(s, _) => MathNode::Text(s, style),
1135        MathNode::Symbol(name) => {
1136            if let Some(ch) = crate::symbols::glyph_char(&name) {
1137                if crate::style_map::is_stylable(ch) {
1138                    MathNode::Text(ch.to_string(), style)
1139                } else {
1140                    MathNode::Symbol(name)
1141                }
1142            } else {
1143                MathNode::Symbol(name)
1144            }
1145        }
1146        MathNode::Row(v) => collapse_text(MathNode::Row(
1147            v.into_iter().map(|n| apply_text_style(n, style)).collect(),
1148        )),
1149        MathNode::Substack(v) => {
1150            MathNode::Substack(v.into_iter().map(|n| apply_text_style(n, style)).collect())
1151        }
1152        other => other,
1153    }
1154}
1155
1156fn collapse_text(node: MathNode) -> MathNode {
1157    let MathNode::Row(v) = node else {
1158        return node;
1159    };
1160    let mut out: Vec<MathNode> = Vec::new();
1161    for n in v {
1162        match (out.last_mut(), &n) {
1163            (Some(MathNode::Text(a, sa)), MathNode::Text(b, sb)) if sa == sb => {
1164                a.push_str(b);
1165            }
1166            _ => out.push(n),
1167        }
1168    }
1169    wrap_row(out)
1170}
1171
1172fn atom_kind(c: char) -> AtomKind {
1173    match c {
1174        '+' | '-' | '*' | '±' | '∓' | '·' | '×' | '÷' => AtomKind::Bin,
1175        '=' | '<' | '>' | '≠' | '≤' | '≥' | '≈' | '≡' => AtomKind::Rel,
1176        '(' | '[' | '{' => AtomKind::Open,
1177        ')' | ']' | '}' => AtomKind::Close,
1178        ',' | ';' | '!' | '?' | ':' => AtomKind::Punct,
1179        _ => AtomKind::Ord,
1180    }
1181}
1182
1183fn delim_from_text(s: &str) -> Result<Delimiter, ParseError> {
1184    let s = s.trim();
1185    if s.is_empty() || s == "." {
1186        return Ok(Delimiter::Empty);
1187    }
1188    if s.chars().count() == 1 {
1189        let c = s.chars().next().unwrap();
1190        return Ok(Delimiter::Char(c));
1191    }
1192    let name = s.strip_prefix('\\').unwrap_or(s);
1193    Ok(Delimiter::Named(name.to_string()))
1194}
1195
1196fn parse_tex_dim(s: &str) -> Result<Dim, ParseError> {
1197    let s = s.trim();
1198    if s.is_empty() {
1199        return Err(ParseError::Malformed("empty dimension".into()));
1200    }
1201    let mut i = 0;
1202    let b = s.as_bytes();
1203    if i < b.len() && (b[i] == b'+' || b[i] == b'-') {
1204        i += 1;
1205    }
1206    while i < b.len() && (b[i].is_ascii_digit() || b[i] == b'.') {
1207        i += 1;
1208    }
1209    if i == 0 || (i == 1 && (b[0] == b'+' || b[0] == b'-')) {
1210        return Err(ParseError::Malformed(format!("invalid dimension `{s}`")));
1211    }
1212    let num = Dim::parse(&s[..i]);
1213    let unit = s[i..].trim();
1214    match unit {
1215        "" | "em" => Ok(num),
1216        "mu" => Ok(Dim::from_mu(&num)),
1217        "pt" | "bp" => Ok(num / Dim::from_i64(10)),
1218        other => Err(ParseError::Unsupported(format!("dimension unit {other}"))),
1219    }
1220}
1221
1222fn color_err(e: Error) -> ParseError {
1223    match e {
1224        Error::Unsupported { what } => ParseError::Unsupported(what),
1225        Error::Parse(p) => p,
1226        Error::Malformed { what } => ParseError::Malformed(what),
1227        Error::InvalidOption { what } => ParseError::Malformed(what),
1228        other => ParseError::Malformed(other.to_string()),
1229    }
1230}
1231
1232fn alias(name: &str) -> &str {
1233    match name {
1234        "le" => "leq",
1235        "ge" => "geq",
1236        "ne" => "neq",
1237        "dots" => "ldots",
1238        "lnot" => "neg",
1239        "dag" => "dagger",
1240        "ddag" => "ddagger",
1241        "owns" => "ni",
1242        _ => name,
1243    }
1244}
1245
1246fn is_extra_symbol(name: &str) -> bool {
1247    matches!(
1248        name,
1249        "Gamma"
1250            | "Delta"
1251            | "Theta"
1252            | "Lambda"
1253            | "Xi"
1254            | "Pi"
1255            | "Sigma"
1256            | "Upsilon"
1257            | "Phi"
1258            | "Psi"
1259            | "Omega"
1260            | "varepsilon"
1261            | "vartheta"
1262            | "varpi"
1263            | "varrho"
1264            | "varsigma"
1265            | "varphi"
1266            | "ldots"
1267            | "cdots"
1268            | "vdots"
1269            | "ddots"
1270            | "colon"
1271            | "mid"
1272            | "lvert"
1273            | "rvert"
1274            | "lVert"
1275            | "rVert"
1276            | "vert"
1277            | "Vert"
1278            | "implies"
1279            | "iff"
1280            | "to"
1281            | "gets"
1282            | "neq"
1283            | "leq"
1284            | "geq"
1285    )
1286}
1287
1288#[cfg(test)]
1289mod tests {
1290    use super::*;
1291
1292    #[test]
1293    fn frac_gold() {
1294        let n = parse(r"\frac{1}{2}").unwrap();
1295        assert_eq!(n.gold(), r#"(frac (atom Ord "1") (atom Ord "2"))"#);
1296    }
1297}