Skip to main content

mzp_peg_macro/
lib.rs

1use std::ops::RangeInclusive;
2
3use kw::{ANY, EOI};
4use proc_macro::TokenStream;
5use proc_macro2::Span;
6use quote::quote;
7use syn::{
8    Ident, LitChar, LitStr, Token, parenthesized,
9    parse::{Parse, ParseStream},
10    parse_macro_input,
11    punctuated::Punctuated,
12};
13
14struct Grammar {
15    rules: Punctuated<Rule, Token![;]>,
16}
17
18struct Rule {
19    name: Ident,
20    definition: Term,
21}
22
23#[derive(Debug)]
24enum Term {
25    AnyChar,
26    Capture(String, Box<Term>),
27    Choice(Vec<Term>),
28    EOI,
29    Literal(String, bool),
30    NegLookahead(Box<Term>),
31    Optional(Box<Term>),
32    Plus(Box<Term>),
33    PosLookahead(Box<Term>),
34    Range(RangeInclusive<char>, bool),
35    Rule(Ident),
36    Sequence(Vec<Term>),
37    Star(Box<Term>),
38}
39
40mod kw {
41    syn::custom_keyword!(ANY);
42    syn::custom_keyword!(EOI);
43    syn::custom_keyword!(icase);
44}
45
46impl Parse for Grammar {
47    fn parse(input: ParseStream) -> syn::Result<Self> {
48        Ok(Grammar {
49            rules: Punctuated::parse_terminated(input)?,
50        })
51    }
52}
53
54impl Parse for Rule {
55    fn parse(input: ParseStream) -> syn::Result<Self> {
56        let mut icase = false;
57        if input.parse::<Token![@]>().is_ok() {
58            let look = input.lookahead1();
59            if look.peek(kw::icase) {
60                input.parse::<kw::icase>()?;
61                icase = true;
62            } else {
63                return Err(look.error());
64            }
65        }
66
67        let name = input.parse()?;
68        input.parse::<Token![=]>()?;
69        let mut definition: Term = input.parse()?;
70        if icase {
71            definition.set_icase();
72        }
73        Ok(Self { name, definition })
74    }
75}
76
77impl Parse for Term {
78    fn parse(input: ParseStream) -> syn::Result<Self> {
79        fn parse_range(input: ParseStream) -> syn::Result<Term> {
80            let start_lit = input.parse::<LitChar>()?;
81            input.parse::<Token![..]>()?;
82            let end_lit = input.parse::<LitChar>()?;
83            let range = start_lit.value()..=end_lit.value();
84            let icase = end_lit.suffix() == "i";
85            Ok(Term::Range(range, icase))
86        }
87
88        fn parse_atom(input: ParseStream) -> syn::Result<Term> {
89            let look = input.lookahead1();
90            if look.peek(Ident) {
91                if input.parse::<ANY>().is_ok() {
92                    Ok(Term::AnyChar)
93                } else if input.parse::<EOI>().is_ok() {
94                    Ok(Term::EOI)
95                } else {
96                    input.parse().map(Term::Rule)
97                }
98            } else if look.peek(LitStr) {
99                let lit = input.parse::<LitStr>()?;
100                let icase = lit.suffix() == "i";
101                Ok(Term::Literal(lit.value(), icase))
102            } else if look.peek(LitChar) {
103                parse_range(input)
104            } else if look.peek(syn::token::Paren) {
105                let content;
106                parenthesized!(content in input);
107                parse_choice(&content)
108            } else {
109                Err(look.error())
110            }
111        }
112
113        fn parse_repeat(input: ParseStream) -> syn::Result<Term> {
114            let mut result = parse_atom(input)?;
115            loop {
116                if input.parse::<Token![?]>().is_ok() {
117                    result = Term::Optional(Box::new(result));
118                } else if input.parse::<Token![+]>().is_ok() {
119                    result = Term::Plus(Box::new(result));
120                } else if input.parse::<Token![*]>().is_ok() {
121                    result = Term::Star(Box::new(result));
122                } else {
123                    break;
124                }
125            }
126            Ok(result)
127        }
128
129        fn parse_prefix(input: ParseStream) -> syn::Result<Term> {
130            if input.parse::<Token![!]>().is_ok() {
131                parse_repeat(input).map(|x| Term::NegLookahead(x.into()))
132            } else if input.parse::<Token![&]>().is_ok() {
133                parse_repeat(input).map(|x| Term::PosLookahead(x.into()))
134            } else if input.parse::<Token![#]>().is_ok() {
135                let tag: Ident = input.parse()?;
136                input.parse::<Token![:]>()?;
137                let expr = parse_repeat(input)?;
138                Ok(Term::Capture(tag.to_string(), expr.into()))
139            } else {
140                parse_repeat(input)
141            }
142        }
143
144        fn parse_sequence(input: ParseStream) -> syn::Result<Term> {
145            let mut terms = vec![parse_prefix(input)?];
146            while !input.is_empty() && !input.peek(Token![/]) && !input.peek(Token![;]) {
147                terms.push(parse_prefix(input)?);
148            }
149            if terms.len() == 1 {
150                Ok(terms.pop().unwrap())
151            } else {
152                Ok(Term::Sequence(terms))
153            }
154        }
155
156        fn parse_choice(input: ParseStream) -> syn::Result<Term> {
157            let mut choices = vec![parse_sequence(input)?];
158            while input.peek(Token![/]) {
159                input.parse::<Token![/]>()?;
160                choices.push(parse_sequence(input)?);
161            }
162            if choices.len() == 1 {
163                Ok(choices.pop().unwrap())
164            } else {
165                Ok(Term::Choice(choices))
166            }
167        }
168
169        parse_choice(input)
170    }
171}
172
173impl Term {
174    fn generate_code(&self) -> proc_macro2::TokenStream {
175        match self {
176            Term::AnyChar => quote! {
177                p.any()
178            },
179            Term::Capture(name, pat) => {
180                let tag = Ident::new(&name, Span::call_site());
181                let code = pat.generate_code();
182                quote! {
183                    {
184                        let save = p.begin_capture(Tag::#tag);
185                        if !#code {
186                            p.restore(save);
187                            false
188                        } else {
189                            p.commit_capture(save);
190                            true
191                        }
192                    }
193                }
194            }
195            Term::EOI => quote! {
196                p.eoi()
197            },
198            Term::Rule(ident) => quote! {
199                #ident(p)
200            },
201            Term::Literal(lit_str, icase) => {
202                let method = if *icase {
203                    quote! { literal_i }
204                } else {
205                    quote! { literal }
206                };
207                quote! {
208                    p.#method(#lit_str)
209                }
210            }
211            Term::Sequence(terms) => {
212                let expr = terms
213                    .iter()
214                    .map(|t| t.generate_code())
215                    .reduce(|x, y| quote! { #x && #y })
216                    .unwrap();
217                quote! {
218                    {
219                        let save = p.save();
220                        if #expr {
221                            true
222                        } else {
223                            p.restore(save);
224                            false
225                        }
226                    }
227                }
228            }
229            Term::Choice(terms) => {
230                let code = terms
231                    .iter()
232                    .map(|t| t.generate_code())
233                    .reduce(|x, y| quote! { #x || #y })
234                    .unwrap();
235                quote! {
236                    ( #code )
237                }
238            }
239            Term::Optional(term) => {
240                let expr = term.generate_code();
241                quote! {
242                    ( #expr || true )
243                }
244            }
245            Term::Star(term) => {
246                let expr = term.generate_code();
247                quote! {
248                    { while #expr {}; true }
249                }
250            }
251            Term::Plus(term) => {
252                let expr = term.generate_code();
253                quote! {
254                    {
255                        let mut closure = || #expr;
256                        if closure() {
257                            while closure() {}
258                            true
259                        } else {
260                            false
261                        }
262                    }
263                }
264            }
265            Term::Range(range, icase) => {
266                let (lo, hi) = (range.start(), range.end());
267                let method = if *icase {
268                    quote! { range_i }
269                } else {
270                    quote! { range }
271                };
272                quote! {
273                    p.#method(#lo..=#hi)
274                }
275            }
276            Term::NegLookahead(term) => {
277                let code = term.generate_code();
278                quote! {
279                    {
280                        let save = p.save();
281                        if #code {
282                            p.restore(save);
283                            false
284                        } else {
285                            true
286                        }
287                    }
288                }
289            }
290            Term::PosLookahead(term) => {
291                let code = term.generate_code();
292                quote! {
293                    {
294                        let save = p.save();
295                        if #code {
296                            p.restore(save);
297                            true
298                        } else {
299                            false
300                        }
301                    }
302                }
303            }
304        }
305    }
306
307    fn set_icase(&mut self) {
308        match self {
309            Term::Literal(_, icase) | Term::Range(_, icase) => {
310                *icase = true;
311            }
312            Term::Choice(terms) | Term::Sequence(terms) => {
313                terms.iter_mut().for_each(|x| x.set_icase());
314            }
315            Term::Capture(_, term)
316            | Term::NegLookahead(term)
317            | Term::Optional(term)
318            | Term::Plus(term)
319            | Term::PosLookahead(term)
320            | Term::Star(term) => {
321                term.set_icase();
322            }
323            Term::AnyChar | Term::EOI | Term::Rule(_) => {}
324        }
325    }
326
327    fn get_capture_names(&self) -> Vec<&str> {
328        let mut result = vec![];
329        match self {
330            Term::AnyChar | Term::EOI | Term::Literal(_, _) | Term::Range(_, _) | Term::Rule(_) => {
331            }
332            Term::Capture(name, term) => {
333                result.push(name.as_str());
334                result.extend(term.get_capture_names());
335            }
336            Term::Choice(terms) | Term::Sequence(terms) => {
337                terms
338                    .iter()
339                    .for_each(|x| result.extend(x.get_capture_names()));
340            }
341            Term::NegLookahead(term)
342            | Term::Optional(term)
343            | Term::Plus(term)
344            | Term::PosLookahead(term)
345            | Term::Star(term) => {
346                result.extend(term.get_capture_names());
347            }
348        }
349        result
350    }
351}
352
353#[proc_macro]
354pub fn grammar(ts: TokenStream) -> TokenStream {
355    let input = parse_macro_input!(ts as Grammar);
356
357    let mut capture_names: Vec<_> = input
358        .rules
359        .iter()
360        .flat_map(|r| r.definition.get_capture_names())
361        .collect();
362    capture_names.sort();
363    capture_names.dedup();
364    let tag_idents: Vec<Ident> = capture_names
365        .iter()
366        .map(|x| Ident::new(x, proc_macro2::Span::call_site()))
367        .collect();
368    let enum_tag = quote! {
369        #[derive(Copy, Clone, Debug, Eq, PartialEq)]
370        pub enum Tag {
371            #(#tag_idents),*
372        }
373    };
374
375    let fns: Vec<_> = input
376        .rules
377        .iter()
378        .map(|r| {
379            let fn_name = &r.name;
380            let generated = r.definition.generate_code();
381            quote! {
382                pub fn #fn_name(p: &mut crate::peg::ParseState<Tag>) -> bool {
383                    use crate::peg::backend::LowLevel;
384                    #generated
385                }
386            }
387        })
388        .collect();
389    quote! {
390        #enum_tag
391        #(#fns)*
392    }
393    .into()
394}