1use crate::error::MathError;
44use crate::node::Span;
45use crate::token::{Tok, TokKind, lex};
46use std::collections::BTreeMap;
47
48const EXPANSION_TOKEN_BUDGET: usize = 65_536;
50const EXPANSION_DEPTH_BUDGET: usize = 32;
52
53#[derive(Clone, Debug, Default, PartialEq, Eq)]
56pub struct MacroSet {
57 defs: BTreeMap<String, MacroDef>,
58}
59
60#[derive(Clone, Debug, PartialEq, Eq)]
61struct MacroDef {
62 params: u8,
64 body: String,
66}
67
68impl MacroSet {
69 #[must_use]
71 pub fn new() -> Self {
72 Self::default()
73 }
74
75 #[must_use]
81 pub fn pack(id: &str) -> Option<Self> {
82 match id {
83 "fmd-math/pack/default" | "default" => {
84 let mut defs = BTreeMap::new();
91 defs.insert(
92 "minus".to_owned(),
93 MacroDef {
94 params: 0,
95 body: "-".to_owned(),
96 },
97 );
98 Some(Self { defs })
99 }
100 "fmd-math/pack/basic" | "basic" | "fmd-math/pack/empty" | "empty" => Some(Self::new()),
101 _ => None,
102 }
103 }
104
105 pub fn define(&mut self, name: &str, params: u8, body: &str) -> Result<(), MathError> {
115 let malformed = |what: String| MathError::Malformed { what, at: 0 };
116 if name.is_empty() || !name.bytes().all(|b| b.is_ascii_alphabetic()) {
117 return Err(malformed(format!(
118 "macro name {name:?} must be one or more ASCII letters"
119 )));
120 }
121 if params > 9 {
122 return Err(malformed(format!(
123 "macro \\{name} declares {params} parameters; TeX allows at most 9"
124 )));
125 }
126 validate_body(name, params, body)?;
127 self.defs.insert(
128 name.to_owned(),
129 MacroDef {
130 params,
131 body: body.to_owned(),
132 },
133 );
134 Ok(())
135 }
136
137 pub fn names(&self) -> impl Iterator<Item = &str> {
139 self.defs.keys().map(String::as_str)
140 }
141
142 #[must_use]
144 pub fn len(&self) -> usize {
145 self.defs.len()
146 }
147
148 #[must_use]
150 pub fn is_empty(&self) -> bool {
151 self.defs.is_empty()
152 }
153
154 #[must_use]
161 pub fn canonical_bytes(&self) -> Vec<u8> {
162 let mut out = b"fmd-math-macroset-v1\x1e".to_vec();
163 for (name, def) in &self.defs {
164 out.extend_from_slice(name.as_bytes());
165 out.push(0x1f);
166 out.push(b'0' + def.params);
167 out.push(0x1f);
168 out.extend_from_slice(def.body.as_bytes());
169 out.push(0x1e);
170 }
171 out
172 }
173}
174
175fn validate_body(name: &str, params: u8, body: &str) -> Result<(), MathError> {
178 let malformed = |what: String| MathError::Malformed { what, at: 0 };
179 let mut depth = 0_i32;
180 let toks = lex(body);
181 let mut i = 0;
182 while i < toks.len() {
183 match toks[i].kind {
184 TokKind::BeginGroup => depth += 1,
185 TokKind::EndGroup => {
186 depth -= 1;
187 if depth < 0 {
188 return Err(malformed(format!(
189 "macro \\{name} body has an unmatched '}}'"
190 )));
191 }
192 }
193 TokKind::Char('#') => {
194 let param = toks.get(i + 1).and_then(|t| match t.kind {
195 TokKind::Char(c) => c.to_digit(10),
196 _ => None,
197 });
198 match param {
199 Some(d) if (1..=u32::from(params)).contains(&d) => i += 1,
200 Some(d) => {
201 return Err(malformed(format!(
202 "macro \\{name} body uses #{d} but declares {params} parameter(s)"
203 )));
204 }
205 None => {
206 return Err(malformed(format!(
207 "macro \\{name} body has a '#' not followed by a parameter digit"
208 )));
209 }
210 }
211 }
212 _ => {}
213 }
214 i += 1;
215 }
216 if depth != 0 {
217 return Err(malformed(format!(
218 "macro \\{name} body has {depth} unclosed '{{'"
219 )));
220 }
221 Ok(())
222}
223
224struct Live<'a> {
226 params: u8,
227 body: Vec<Tok<'a>>,
228}
229
230pub(crate) fn expand<'a>(
234 toks: Vec<Tok<'a>>,
235 set: &'a MacroSet,
236 src_len: usize,
237) -> Result<Vec<Tok<'a>>, MathError> {
238 let involved = !set.is_empty()
240 || toks
241 .iter()
242 .any(|t| matches!(t.kind, TokKind::ControlWord("newcommand" | "renewcommand")));
243 if !involved {
244 return Ok(toks);
245 }
246
247 let mut table: BTreeMap<&'a str, Live<'a>> = BTreeMap::new();
248 for (name, def) in &set.defs {
249 table.insert(
250 name.as_str(),
251 Live {
252 params: def.params,
253 body: lex(&def.body),
254 },
255 );
256 }
257
258 let mut cx = Expansion {
259 table,
260 budget: EXPANSION_TOKEN_BUDGET,
261 src_len,
262 };
263 let mut out = Vec::with_capacity(toks.len());
264 let mut i = 0;
265 while i < toks.len() {
266 let tok = &toks[i];
267 match tok.kind {
268 TokKind::ControlWord(cw @ ("newcommand" | "renewcommand")) => {
269 i = cx.definition(&toks, i, cw == "renewcommand")?;
270 }
271 TokKind::ControlWord(name) if cx.table.contains_key(name) => {
272 let mut active = Vec::new();
273 i = cx.call(&toks, i, name, &mut active, 0, &mut out)?;
274 }
275 _ => {
276 out.push(tok.clone());
277 i += 1;
278 }
279 }
280 }
281 Ok(out)
282}
283
284struct Expansion<'a> {
285 table: BTreeMap<&'a str, Live<'a>>,
286 budget: usize,
287 src_len: usize,
288}
289
290impl<'a> Expansion<'a> {
291 fn definition(&mut self, toks: &[Tok<'a>], i: usize, renew: bool) -> Result<usize, MathError> {
295 let cw_span = toks[i].span;
296 let which = if renew {
297 "\\renewcommand"
298 } else {
299 "\\newcommand"
300 };
301 let mut j = i + 1;
302 let skip_space = |j: &mut usize| {
303 while toks
304 .get(*j)
305 .is_some_and(|t| matches!(t.kind, TokKind::Space))
306 {
307 *j += 1;
308 }
309 };
310 skip_space(&mut j);
311 let braced = matches!(toks.get(j).map(|t| &t.kind), Some(TokKind::BeginGroup));
313 if braced {
314 j += 1;
315 skip_space(&mut j);
316 }
317 let Some(name_tok) = toks.get(j) else {
318 return Err(MathError::Malformed {
319 what: format!("{which} ends before its macro name"),
320 at: self.src_len,
321 });
322 };
323 let TokKind::ControlWord(name) = name_tok.kind else {
324 return Err(MathError::Malformed {
325 what: format!("{which} expects a \\name to define"),
326 at: name_tok.span.start,
327 });
328 };
329 j += 1;
330 if braced {
331 skip_space(&mut j);
332 let Some(Tok {
333 kind: TokKind::EndGroup,
334 ..
335 }) = toks.get(j)
336 else {
337 return Err(MathError::Malformed {
338 what: format!("{which}{{\\{name}}} has an unclosed name group"),
339 at: toks.get(j).map_or(self.src_len, |t| t.span.start),
340 });
341 };
342 j += 1;
343 }
344 skip_space(&mut j);
345 let mut params = 0_u8;
347 if matches!(toks.get(j).map(|t| &t.kind), Some(TokKind::Char('['))) {
348 let digit = toks.get(j + 1).and_then(|t| match t.kind {
349 TokKind::Char(c) => c.to_digit(10),
350 _ => None,
351 });
352 let close = matches!(toks.get(j + 2).map(|t| &t.kind), Some(TokKind::Char(']')));
353 match (digit, close) {
354 (Some(d @ 1..=9), true) => {
355 params = u8::try_from(d).unwrap_or(9);
356 j += 3;
357 }
358 _ => {
359 return Err(MathError::Malformed {
360 what: format!("{which}{{\\{name}}}: expected [1]..[9] parameter count"),
361 at: toks.get(j).map_or(self.src_len, |t| t.span.start),
362 });
363 }
364 }
365 skip_space(&mut j);
366 }
367 let Some(Tok {
369 kind: TokKind::BeginGroup,
370 ..
371 }) = toks.get(j)
372 else {
373 return Err(MathError::Malformed {
374 what: format!("{which}{{\\{name}}}: expected a {{body}} group"),
375 at: toks.get(j).map_or(self.src_len, |t| t.span.start),
376 });
377 };
378 let body_start = j + 1;
379 let mut depth = 1_i32;
380 let mut k = body_start;
381 while k < toks.len() {
382 match toks[k].kind {
383 TokKind::BeginGroup => depth += 1,
384 TokKind::EndGroup => {
385 depth -= 1;
386 if depth == 0 {
387 break;
388 }
389 }
390 _ => {}
391 }
392 k += 1;
393 }
394 if depth != 0 {
395 return Err(MathError::Malformed {
396 what: format!("{which}{{\\{name}}}: unclosed body group"),
397 at: self.src_len,
398 });
399 }
400 let exists = self.table.contains_key(name);
403 if !renew && exists {
404 return Err(MathError::Malformed {
405 what: format!(
406 "\\newcommand: \\{name} is already defined (use \\renewcommand to replace it)"
407 ),
408 at: cw_span.start,
409 });
410 }
411 if renew && !exists {
412 return Err(MathError::Malformed {
413 what: format!("\\renewcommand: \\{name} is not defined (use \\newcommand)"),
414 at: cw_span.start,
415 });
416 }
417 let body = toks[body_start..k].to_vec();
420 validate_body_tokens(name, params, &body, cw_span.start)?;
421 self.table
422 .insert(name_interned(toks, j, name), Live { params, body });
423 Ok(k + 1)
424 }
425
426 fn call(
429 &mut self,
430 toks: &[Tok<'a>],
431 i: usize,
432 name: &'a str,
433 active: &mut Vec<String>,
434 depth: usize,
435 out: &mut Vec<Tok<'a>>,
436 ) -> Result<usize, MathError> {
437 let call_start = toks[i].span;
438 if depth >= EXPANSION_DEPTH_BUDGET {
439 return Err(MathError::Malformed {
440 what: format!(
441 "macro expansion nests deeper than {EXPANSION_DEPTH_BUDGET} (at \\{name})"
442 ),
443 at: call_start.start,
444 });
445 }
446 if active.iter().any(|a| a == name) {
447 return Err(MathError::Malformed {
448 what: format!(
449 "recursive macro: \\{name} expands itself (macros are non-recursive substitutions)"
450 ),
451 at: call_start.start,
452 });
453 }
454 let params = self.table.get(name).map(|l| l.params).unwrap_or(0);
455 let mut j = i + 1;
458 let mut args: Vec<&[Tok<'a>]> = Vec::new();
459 let mut end_span = call_start;
460 for argn in 1..=params {
461 while toks
462 .get(j)
463 .is_some_and(|t| matches!(t.kind, TokKind::Space))
464 {
465 j += 1;
466 }
467 let Some(first) = toks.get(j) else {
468 return Err(MathError::Malformed {
469 what: format!("\\{name} needs {params} argument(s); input ends before #{argn}"),
470 at: self.src_len,
471 });
472 };
473 if matches!(first.kind, TokKind::BeginGroup) {
474 let start = j + 1;
475 let mut depth_b = 1_i32;
476 let mut k = start;
477 while k < toks.len() {
478 match toks[k].kind {
479 TokKind::BeginGroup => depth_b += 1,
480 TokKind::EndGroup => {
481 depth_b -= 1;
482 if depth_b == 0 {
483 break;
484 }
485 }
486 _ => {}
487 }
488 k += 1;
489 }
490 if depth_b != 0 {
491 return Err(MathError::Malformed {
492 what: format!("\\{name}: unclosed argument group for #{argn}"),
493 at: self.src_len,
494 });
495 }
496 args.push(&toks[start..k]);
497 end_span = toks[k].span;
498 j = k + 1;
499 } else {
500 args.push(core::slice::from_ref(first));
501 end_span = first.span;
502 j += 1;
503 }
504 }
505 let call_span = call_start.union(end_span);
506 self.splice(name, &args, call_span, active, depth, out)?;
507 Ok(j)
508 }
509
510 fn splice(
514 &mut self,
515 name: &'a str,
516 args: &[&[Tok<'a>]],
517 call_span: Span,
518 active: &mut Vec<String>,
519 depth: usize,
520 out: &mut Vec<Tok<'a>>,
521 ) -> Result<(), MathError> {
522 active.push(name.to_owned());
523 let body = self
527 .table
528 .get(name)
529 .map(|l| l.body.clone())
530 .unwrap_or_default();
531 let mut j = 0;
532 while j < body.len() {
533 let t = &body[j];
534 match t.kind {
535 TokKind::Char('#') => {
536 let d = body.get(j + 1).and_then(|n| match n.kind {
537 TokKind::Char(c) => c.to_digit(10),
538 _ => None,
539 });
540 let Some(d) = d else {
541 return Err(MathError::Malformed {
542 what: format!("macro \\{name} body has a stray '#'"),
543 at: call_span.start,
544 });
545 };
546 let arg = args.get(d as usize - 1).copied().unwrap_or(&[]);
547 let mut k = 0;
550 while k < arg.len() {
551 match arg[k].kind {
552 TokKind::ControlWord(n) if self.table.contains_key(n) => {
553 k = self.call(arg, k, n, active, depth + 1, out)?
554 }
555 _ => {
556 self.push(out, arg[k].clone(), call_span, false)?;
557 k += 1;
558 }
559 }
560 }
561 j += 2;
562 }
563 TokKind::ControlWord(n) if self.table.contains_key(n) => {
564 j = self.call(&body, j, n, active, depth + 1, out)?;
565 }
566 _ => {
567 self.push(out, t.clone(), call_span, true)?;
568 j += 1;
569 }
570 }
571 }
572 active.pop();
573 Ok(())
574 }
575
576 fn push(
579 &mut self,
580 out: &mut Vec<Tok<'a>>,
581 mut tok: Tok<'a>,
582 call_span: Span,
583 rewrite_span: bool,
584 ) -> Result<(), MathError> {
585 if self.budget == 0 {
586 return Err(MathError::Malformed {
587 what: format!("macro expansion produced more than {EXPANSION_TOKEN_BUDGET} tokens"),
588 at: call_span.start,
589 });
590 }
591 self.budget -= 1;
592 if rewrite_span {
593 tok.span = call_span;
594 }
595 out.push(tok);
596 Ok(())
597 }
598}
599
600fn name_interned<'a>(toks: &[Tok<'a>], upto: usize, name: &'a str) -> &'a str {
604 let _ = (toks, upto);
607 name
608}
609
610fn validate_body_tokens(
613 name: &str,
614 params: u8,
615 body: &[Tok<'_>],
616 at: usize,
617) -> Result<(), MathError> {
618 let mut j = 0;
619 while j < body.len() {
620 if let TokKind::Char('#') = body[j].kind {
621 let d = body.get(j + 1).and_then(|n| match n.kind {
622 TokKind::Char(c) => c.to_digit(10),
623 _ => None,
624 });
625 match d {
626 Some(d) if (1..=u32::from(params)).contains(&d) => j += 1,
627 Some(d) => {
628 return Err(MathError::Malformed {
629 what: format!(
630 "macro \\{name} body uses #{d} but declares {params} parameter(s)"
631 ),
632 at,
633 });
634 }
635 None => {
636 return Err(MathError::Malformed {
637 what: format!("macro \\{name} body has a '#' not followed by a digit"),
638 at,
639 });
640 }
641 }
642 }
643 j += 1;
644 }
645 Ok(())
646}
647
648#[cfg(test)]
649mod tests {
650 #![allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
651
652 use super::*;
653
654 fn expand_str<'a>(src: &'a str, set: &'a MacroSet) -> Result<String, MathError> {
655 let toks = expand(lex(src), set, src.len())?;
656 Ok(toks
657 .iter()
658 .map(|t| match &t.kind {
659 TokKind::ControlWord(w) => format!("\\{w} "),
660 TokKind::ControlSymbol(c) => format!("\\{c}"),
661 TokKind::BeginGroup => "{".into(),
662 TokKind::EndGroup => "}".into(),
663 TokKind::Sup => "^".into(),
664 TokKind::Sub => "_".into(),
665 TokKind::AlignTab => "&".into(),
666 TokKind::Tie => "~".into(),
667 TokKind::MathShift => "$".into(),
668 TokKind::Space => " ".into(),
669 TokKind::Char(c) => (*c).to_string(),
670 })
671 .collect())
672 }
673
674 #[test]
675 fn pack_macros_expand_with_call_site_spans() {
676 let set = MacroSet::pack("fmd-math/pack/default").unwrap();
677 let src = r"a\minus b";
678 let toks = expand(lex(src), &set, src.len()).unwrap();
679 let minus = toks
681 .iter()
682 .find(|t| matches!(t.kind, TokKind::Char('-')))
683 .expect("expanded minus");
684 assert_eq!((minus.span.start, minus.span.end), (1, 7));
685 }
686
687 #[test]
688 fn inline_definition_with_arguments() {
689 let set = MacroSet::new();
690 let out = expand_str(r"\newcommand{\half}[1]{\frac{#1}{2}}\half{x}", &set).unwrap();
691 assert_eq!(out, r"\frac {x}{2}");
692 }
693
694 #[test]
695 fn arguments_keep_their_own_spans_and_bodies_take_the_call() {
696 let set = MacroSet::new();
697 let src = r"\newcommand{\half}[1]{\frac{#1}{2}}\half{x}";
698 let toks = expand(lex(src), &set, src.len()).unwrap();
699 let call_start = src.find(r"\half{x}").unwrap();
700 let x = toks
702 .iter()
703 .find(|t| matches!(t.kind, TokKind::Char('x')))
704 .unwrap();
705 assert_eq!(&src[x.span.start..x.span.end], "x");
706 let frac = toks
708 .iter()
709 .find(|t| matches!(t.kind, TokKind::ControlWord("frac")))
710 .unwrap();
711 assert_eq!(frac.span.start, call_start);
712 assert_eq!(frac.span.end, src.len());
713 }
714
715 #[test]
716 fn macros_reference_other_macros() {
717 let mut set = MacroSet::new();
718 set.define("dd", 0, r"\mathrm{d}").unwrap();
719 set.define("dx", 0, r"\dd x").unwrap();
720 let out = expand_str(r"\dx", &set).unwrap();
721 assert_eq!(out, r"\mathrm {d}x");
722 }
723
724 #[test]
725 fn recursion_is_refused_with_the_macro_named() {
726 let mut set = MacroSet::new();
727 set.define("loop", 0, r"a\loop").unwrap();
728 let err = expand_str(r"\loop", &set).unwrap_err();
729 assert!(err.to_string().contains("recursive macro: \\loop"), "{err}");
730
731 let mut set = MacroSet::new();
733 set.define("ping", 0, r"\pong").unwrap();
734 set.define("pong", 0, r"\ping").unwrap();
735 let err = expand_str(r"\ping", &set).unwrap_err();
736 assert!(err.to_string().contains("recursive macro"), "{err}");
737 }
738
739 #[test]
740 fn expansion_bombs_hit_the_budget() {
741 let mut set = MacroSet::new();
744 set.define("a", 0, "xx").unwrap();
745 for (prev, name) in [
746 ("a", "b"),
747 ("b", "c"),
748 ("c", "d"),
749 ("d", "e"),
750 ("e", "f"),
751 ("f", "g"),
752 ("g", "h"),
753 ("h", "i"),
754 ("i", "j"),
755 ("j", "k"),
756 ("k", "l"),
757 ("l", "m"),
758 ("m", "n"),
759 ("n", "o"),
760 ("o", "p"),
761 ("p", "q"),
762 ("q", "r"),
763 ] {
764 let body = format!("\\{prev}\\{prev}");
765 set.define(name, 0, &body).unwrap();
766 }
767 let err = expand_str(r"\r", &set).unwrap_err();
768 let msg = err.to_string();
769 assert!(
770 msg.contains("more than") || msg.contains("nests deeper"),
771 "{msg}"
772 );
773 }
774
775 #[test]
776 fn shadowing_rules_are_latexs() {
777 let set = MacroSet::new();
778 let err = expand_str(r"\newcommand{\x}{a}\newcommand{\x}{b}", &set).unwrap_err();
780 assert!(err.to_string().contains("already defined"), "{err}");
781 let err = expand_str(r"\renewcommand{\y}{a}", &set).unwrap_err();
783 assert!(err.to_string().contains("not defined"), "{err}");
784 let out = expand_str(r"\newcommand{\x}{a}\renewcommand{\x}{b}\x", &set).unwrap();
786 assert_eq!(out, "b");
787 }
788
789 #[test]
790 fn definition_faults_are_precise() {
791 let set = MacroSet::new();
792 for (src, needle) in [
793 (r"\newcommand", "ends before its macro name"),
794 (r"\newcommand{x}{a}", "expects a \\name"),
795 (r"\newcommand{\x}[0]{a}", "expected [1]..[9]"),
796 (r"\newcommand{\x}[2]{#3}", "uses #3 but declares 2"),
797 (r"\newcommand{\x}", "expected a {body} group"),
798 (r"\newcommand{\x}{a", "unclosed body group"),
799 ] {
800 let err = expand_str(src, &set).unwrap_err();
801 assert!(err.to_string().contains(needle), "{src}: {err}");
802 }
803 }
804
805 #[test]
806 fn undelimited_single_token_arguments() {
807 let set = MacroSet::new();
808 let out = expand_str(r"\newcommand{\sq}[1]{#1^2}\sq x", &set).unwrap();
809 assert_eq!(out, "x^2");
810 }
811
812 #[test]
813 fn canonical_bytes_are_deterministic_and_content_sensitive() {
814 let mut a = MacroSet::new();
815 a.define("dd", 0, r"\mathrm{d}").unwrap();
816 a.define("half", 1, r"\frac{#1}{2}").unwrap();
817 let mut b = MacroSet::new();
818 b.define("half", 1, r"\frac{#1}{2}").unwrap();
820 b.define("dd", 0, r"\mathrm{d}").unwrap();
821 assert_eq!(a.canonical_bytes(), b.canonical_bytes());
822 let mut c = MacroSet::new();
824 c.define("dd", 0, r"\mathrm{D}").unwrap();
825 c.define("half", 1, r"\frac{#1}{2}").unwrap();
826 assert_ne!(a.canonical_bytes(), c.canonical_bytes());
827 }
828
829 #[test]
830 fn define_validation_is_precise() {
831 let mut set = MacroSet::new();
832 assert!(set.define("", 0, "x").is_err());
833 assert!(set.define("bad name", 0, "x").is_err());
834 assert!(set.define("x", 10, "y").is_err());
835 assert!(set.define("x", 1, "#2").is_err());
836 assert!(set.define("x", 0, "{unclosed").is_err());
837 assert!(set.define("x", 0, "}stray").is_err());
838 assert!(set.define("ok", 2, r"\frac{#1}{#2}").is_ok());
839 }
840
841 #[test]
842 fn packs_exist_by_content_id_and_name() {
843 for id in [
844 "fmd-math/pack/default",
845 "default",
846 "fmd-math/pack/basic",
847 "basic",
848 "fmd-math/pack/empty",
849 "empty",
850 ] {
851 assert!(MacroSet::pack(id).is_some(), "{id}");
852 }
853 assert!(MacroSet::pack("nonexistent").is_none());
854 assert_eq!(MacroSet::pack("default").unwrap().len(), 1);
855 assert!(MacroSet::pack("empty").unwrap().is_empty());
856 }
857}