Skip to main content

mproto_codegen/
parse.rs

1use nom::{
2    branch::alt,
3    bytes::complete::{tag, take_while1},
4    character::complete::{char, multispace0},
5    combinator::{cut, map, opt},
6    error::{context, ParseError},
7    multi::separated_list0,
8    sequence::{preceded, separated_pair, terminated},
9    IResult,
10};
11
12use crate::ast::{
13    Enum, EnumVariant, NamedField, PrimitiveType, QualifiedIdentifier, Struct, Type, TypeBody,
14    TypeDef,
15};
16
17fn identifier(i: &str) -> IResult<&str, &str> {
18    take_while1(|c: char| c.is_alphanumeric() || c == '_')(i)
19}
20
21fn qualified_identifier(i: &str) -> IResult<&str, QualifiedIdentifier> {
22    let (i, maybe_module): (_, Option<&str>) = opt(|i| {
23        let (i, module) = identifier(i)?;
24        let (i, _) = char('.')(i)?;
25        Ok((i, module))
26    })(i)?;
27
28    let (i, name) = identifier(i)?;
29
30    Ok((
31        i,
32        QualifiedIdentifier {
33            name: name.to_string(),
34            module: maybe_module.map(|x| x.to_string()),
35        },
36    ))
37}
38
39fn opt_trailing_comma<I, O, E: ParseError<I>>(
40    f: impl FnMut(I) -> IResult<I, O, E>,
41) -> impl FnMut(I) -> IResult<I, O, E>
42where
43    I: nom::Slice<std::ops::RangeFrom<usize>> + nom::InputIter + Clone,
44    <I as nom::InputIter>::Item: nom::AsChar,
45{
46    terminated(f, opt(char(',')))
47}
48
49fn box_ty(i: &str) -> IResult<&str, PrimitiveType> {
50    let (i, _) = tag("box")(i)?;
51    let (i, _) = multispace0(i)?;
52    let (i, _) = tag("<")(i)?;
53    let (i, _) = multispace0(i)?;
54    let (i, inner_ty) = ty(i)?;
55    let (i, _) = multispace0(i)?;
56    let (i, _) = tag(">")(i)?;
57
58    Ok((i, PrimitiveType::Box(inner_ty.into())))
59}
60
61fn list_ty(i: &str) -> IResult<&str, PrimitiveType> {
62    let (i, _) = char('[')(i)?;
63    let (i, _) = multispace0(i)?;
64    let (i, ty) = ty(i)?;
65    let (i, _) = multispace0(i)?;
66    let (i, _) = char(']')(i)?;
67
68    Ok((i, PrimitiveType::List(ty.into())))
69}
70
71fn option_ty(i: &str) -> IResult<&str, PrimitiveType> {
72    let (i, _) = tag("option")(i)?;
73    let (i, _) = multispace0(i)?;
74    let (i, _) = tag("<")(i)?;
75    let (i, _) = multispace0(i)?;
76    let (i, ty) = ty(i)?;
77    let (i, _) = multispace0(i)?;
78    let (i, _) = tag(">")(i)?;
79
80    Ok((i, PrimitiveType::Option(ty.into())))
81}
82
83fn result_ty(i: &str) -> IResult<&str, PrimitiveType> {
84    let (i, _) = tag("result")(i)?;
85    let (i, _) = multispace0(i)?;
86    let (i, _) = tag("<")(i)?;
87    let (i, ok_ty) = ty(i)?;
88    let (i, _) = multispace0(i)?;
89    let (i, _) = char(',')(i)?;
90    let (i, _) = multispace0(i)?;
91    let (i, err_ty) = ty(i)?;
92    let (i, _) = multispace0(i)?;
93    let (i, _) = opt(char(','))(i)?;
94    let (i, _) = multispace0(i)?;
95    let (i, _) = tag(">")(i)?;
96
97    Ok((i, PrimitiveType::Result(ok_ty.into(), err_ty.into())))
98}
99
100fn builtin_ty(i: &str) -> IResult<&str, PrimitiveType> {
101    alt((
102        map(tag("void"), |_| PrimitiveType::Void),
103        map(tag("u8"), |_| PrimitiveType::U8),
104        map(tag("u16"), |_| PrimitiveType::U16),
105        map(tag("u32"), |_| PrimitiveType::U32),
106        map(tag("u64"), |_| PrimitiveType::U64),
107        map(tag("u128"), |_| PrimitiveType::U128),
108        map(tag("i8"), |_| PrimitiveType::I8),
109        map(tag("i16"), |_| PrimitiveType::I16),
110        map(tag("i32"), |_| PrimitiveType::I32),
111        map(tag("i64"), |_| PrimitiveType::I64),
112        map(tag("i128"), |_| PrimitiveType::I128),
113        map(tag("f32"), |_| PrimitiveType::F32),
114        map(tag("f64"), |_| PrimitiveType::F64),
115        map(tag("bool"), |_| PrimitiveType::Bool),
116        map(tag("string"), |_| PrimitiveType::String),
117        box_ty,
118        list_ty,
119        option_ty,
120        result_ty,
121    ))(i)
122}
123
124pub fn type_params_list(i: &str) -> IResult<&str, Vec<String>> {
125    let (i, _) = tag("<")(i)?;
126    let (i, params) = separated_list0(
127        preceded(multispace0, char(',')),
128        preceded(multispace0, identifier),
129    )(i)?;
130    let (i, _) = multispace0(i)?;
131    let (i, _) = opt(char(','))(i)?;
132    let (i, _) = multispace0(i)?;
133    let (i, _) = tag(">")(i)?;
134
135    let params = params.into_iter().map(|p| p.into()).collect();
136
137    Ok((i, params))
138}
139
140pub fn type_args_list(i: &str) -> IResult<&str, Vec<Type>> {
141    let (i, _) = tag("<")(i)?;
142    let (i, args) =
143        separated_list0(preceded(multispace0, char(',')), preceded(multispace0, ty))(i)?;
144    let (i, _) = multispace0(i)?;
145    let (i, _) = opt(char(','))(i)?;
146    let (i, _) = multispace0(i)?;
147    let (i, _) = tag(">")(i)?;
148
149    let args = args.into_iter().map(|p| p.into()).collect();
150
151    Ok((i, args))
152}
153
154pub fn struct_def(i: &str) -> IResult<&str, TypeDef> {
155    let (i, _) = tag("struct")(i)?;
156    let (i, _) = multispace0(i)?;
157    let (i, name) = identifier(i)?;
158    let (i, _) = multispace0(i)?;
159    let (i, maybe_params) = opt(type_params_list)(i)?;
160    let (i, _) = multispace0(i)?;
161    let (i, fields) = named_fields(i)?;
162
163    let params = maybe_params.unwrap_or(Vec::new());
164
165    let type_def = TypeDef {
166        name: name.into(),
167        params,
168        body: TypeBody::Struct(Struct { fields }),
169    };
170
171    Ok((i, type_def))
172}
173
174fn enum_def<'a>(i: &'a str) -> IResult<&'a str, TypeDef> {
175    let (i, _) = tag("enum")(i)?;
176    let (i, _) = multispace0(i)?;
177    let (i, name) = identifier(i)?;
178    let (i, _) = multispace0(i)?;
179    let (i, maybe_params) = opt(type_params_list)(i)?;
180    let (i, _) = multispace0(i)?;
181    let (i, variants) = enum_variants(i)?;
182
183    let params = maybe_params.unwrap_or(Vec::new());
184
185    let type_def = TypeDef {
186        name: name.into(),
187        params,
188        body: TypeBody::Enum(Enum { variants }),
189    };
190
191    Ok((i, type_def))
192}
193
194fn enum_variants<'a>(i: &'a str) -> IResult<&'a str, Vec<(String, EnumVariant)>> {
195    context(
196        "map",
197        preceded(
198            char('{'),
199            cut(terminated(
200                opt_trailing_comma(map(
201                    separated_list0(
202                        preceded(multispace0, char(',')),
203                        preceded(multispace0, enum_variant),
204                    ),
205                    |tuple_vec| {
206                        tuple_vec
207                            .into_iter()
208                            .map(|(name, variant)| (name.into(), variant))
209                            .collect()
210                    },
211                )),
212                preceded(multispace0, char('}')),
213            )),
214        ),
215    )(i)
216}
217
218fn enum_variant<'a>(i: &'a str) -> IResult<&'a str, (&'a str, EnumVariant)> {
219    alt((
220        separated_pair(
221            identifier,
222            multispace0,
223            map(named_fields, |fields| EnumVariant::NamedFields { fields }),
224        ),
225        map(identifier, |x| (x, EnumVariant::Empty)),
226    ))(i)
227}
228
229pub fn defined_ty<'a>(i: &'a str) -> IResult<&'a str, Type> {
230    let (i, ident) = qualified_identifier(i)?;
231    let (i, _) = multispace0(i)?;
232    let (i, maybe_args) = opt(type_args_list)(i)?;
233
234    let args = maybe_args.unwrap_or(Vec::new());
235    let defined_type = Type::Defined { ident, args };
236
237    Ok((i, defined_type))
238}
239
240pub fn ty<'a>(i: &'a str) -> IResult<&'a str, Type> {
241    alt((map(builtin_ty, |x| Type::Primitive(x)), defined_ty))(i)
242}
243
244pub fn type_def<'a>(i: &'a str) -> IResult<&'a str, TypeDef> {
245    alt((struct_def, enum_def))(i)
246}
247
248fn named_field<'a>(i: &'a str) -> IResult<&'a str, (&'a str, Type)> {
249    separated_pair(
250        identifier,
251        cut(preceded(multispace0, char(':'))),
252        preceded(multispace0, ty),
253    )(i)
254}
255
256fn named_fields<'a>(i: &'a str) -> IResult<&'a str, Vec<NamedField>> {
257    context(
258        "map",
259        preceded(
260            char('{'),
261            cut(terminated(
262                opt_trailing_comma(map(
263                    separated_list0(
264                        preceded(multispace0, char(',')),
265                        preceded(multispace0, named_field),
266                    ),
267                    |tuple_vec| {
268                        tuple_vec
269                            .into_iter()
270                            .map(|(k, v)| NamedField {
271                                name: k.to_owned(),
272                                ty: v,
273                            })
274                            .collect()
275                    },
276                )),
277                preceded(multispace0, char('}')),
278            )),
279        ),
280    )(i)
281}
282
283pub fn strip_comments(s: &str) -> String {
284    s.lines().map(|line| {
285        if let Some(index) = line.find("//") {
286            &line[..index]
287        } else {
288            line
289        }
290    })
291    .collect::<Vec<&str>>()
292    .join("\n")
293}
294
295pub fn root<'a>(i: &'a str) -> IResult<&'a str, Vec<TypeDef>> {
296    separated_list0(multispace0, type_def)(i)
297}
298
299pub fn parse_file(path: impl AsRef<std::path::Path>) -> Result<Vec<TypeDef>, String> {
300    use std::io::Read;
301
302    let Ok(mut file) = std::fs::File::open(path.as_ref()) else {
303        return Err(format!("Failed to open file '{}'", path.as_ref().display()));
304    };
305
306    let mut file_str = String::new();
307    if let Err(e) = file.read_to_string(&mut file_str) {
308        return Err(format!(
309            "Failed to read file '{}': {}",
310            path.as_ref().display(),
311            e
312        ));
313    }
314
315    let uncommented_schema = strip_comments(&file_str) + "\n";
316    let (_, type_defs) = root(&uncommented_schema).map_err(|e| format!("mproto schema parse error: {e}"))?;
317
318    Ok(type_defs)
319}
320
321#[cfg(test)]
322mod tests {
323    use super::*;
324
325    #[test]
326    fn test_strip_comments() {
327        assert_eq!(
328            strip_comments("foo // bar\n// bip boop\nbazz"),
329            "foo \n\nbazz".to_string(),
330        );
331    }
332
333    #[test]
334    fn test_builtin_u8() {
335        let data = "u8";
336        let (_, parsed) = builtin_ty(data).unwrap();
337
338        assert_eq!(parsed, PrimitiveType::U8);
339    }
340
341    #[test]
342    fn test_list_u8() {
343        let data = "[u8]";
344        let (_, parsed) = list_ty(data).unwrap();
345
346        assert_eq!(
347            parsed,
348            PrimitiveType::List(Box::new(Type::Primitive(PrimitiveType::U8))),
349        );
350    }
351
352    #[test]
353    fn test_box_u8() {
354        let data = "box<u8>";
355        let (_, parsed) = box_ty(data).unwrap();
356
357        assert_eq!(
358            parsed,
359            PrimitiveType::Box(Box::new(Type::Primitive(PrimitiveType::U8))),
360        );
361    }
362
363    #[test]
364    fn test_option_u8() {
365        let data = "option<u8>";
366        let (_, parsed) = option_ty(data).unwrap();
367
368        assert_eq!(
369            parsed,
370            PrimitiveType::Option(Box::new(Type::Primitive(PrimitiveType::U8))),
371        );
372    }
373
374    #[test]
375    fn test_result() {
376        let data = "result<void, string>";
377        let (_, parsed) = result_ty(data).unwrap();
378
379        assert_eq!(
380            parsed,
381            PrimitiveType::Result(
382                Box::new(Type::Primitive(PrimitiveType::Void)),
383                Box::new(Type::Primitive(PrimitiveType::String)),
384            ),
385        );
386    }
387
388    #[test]
389    fn test_struct_named_fields() {
390        use PrimitiveType::*;
391
392        let data = "struct Foo { bar : u32, baz : i8 }";
393        let (_, parsed) = struct_def(data).unwrap();
394
395        assert_eq!(
396            parsed,
397            TypeDef {
398                name: "Foo".into(),
399                params: vec![],
400                body: TypeBody::Struct(Struct {
401                    fields: vec![
402                        NamedField {
403                            name: "bar".into(),
404                            ty: Type::Primitive(U32)
405                        },
406                        NamedField {
407                            name: "baz".into(),
408                            ty: Type::Primitive(I8)
409                        },
410                    ]
411                }),
412            }
413        );
414    }
415
416    #[test]
417    fn test_struct_type_param() {
418        use PrimitiveType::*;
419
420        let data = "struct Foo <T, F>  { bar : u32, baz : i8 }";
421        let (_, parsed) = struct_def(data).unwrap();
422
423        assert_eq!(
424            parsed,
425            TypeDef {
426                name: "Foo".into(),
427                params: vec!["T".into(), "F".into()],
428                body: TypeBody::Struct(Struct {
429                    fields: vec![
430                        NamedField {
431                            name: "bar".into(),
432                            ty: Type::Primitive(U32)
433                        },
434                        NamedField {
435                            name: "baz".into(),
436                            ty: Type::Primitive(I8)
437                        },
438                    ]
439                }),
440            }
441        );
442    }
443
444    #[test]
445    fn test_defined_type() {
446        use PrimitiveType::*;
447
448        let data = "struct Foo { bar : bar_proto.Bar, baz : i8 }";
449        let (_, parsed) = struct_def(data).unwrap();
450
451        assert_eq!(
452            parsed,
453            TypeDef {
454                name: "Foo".into(),
455                params: vec![],
456                body: TypeBody::Struct(Struct {
457                    fields: vec![
458                        NamedField {
459                            name: "bar".into(),
460                            ty: Type::Defined {
461                                ident: QualifiedIdentifier {
462                                    module: Some("bar_proto".into()),
463                                    name: "Bar".into(),
464                                },
465                                args: vec![],
466                            },
467                        },
468                        NamedField {
469                            name: "baz".into(),
470                            ty: Type::Primitive(I8)
471                        },
472                    ]
473                }),
474            }
475        );
476    }
477
478    #[test]
479    fn test_enum_named_fields() {
480        use PrimitiveType::*;
481
482        let data = "enum Foo { Bar { x: u32, y: u8 }, Baz { bip: i8 } }";
483        let (_, parsed) = enum_def(data).unwrap();
484
485        assert_eq!(
486            parsed,
487            TypeDef {
488                name: "Foo".into(),
489                params: vec![],
490                body: TypeBody::Enum(Enum {
491                    variants: vec![
492                        (
493                            "Bar".into(),
494                            EnumVariant::NamedFields {
495                                fields: vec![
496                                    NamedField {
497                                        name: "x".into(),
498                                        ty: Type::Primitive(U32)
499                                    },
500                                    NamedField {
501                                        name: "y".into(),
502                                        ty: Type::Primitive(U8)
503                                    },
504                                ],
505                            }
506                        ),
507                        (
508                            "Baz".into(),
509                            EnumVariant::NamedFields {
510                                fields: vec![NamedField {
511                                    name: "bip".into(),
512                                    ty: Type::Primitive(I8)
513                                },],
514                            }
515                        ),
516                    ]
517                }),
518            }
519        );
520    }
521}