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("bool"), |_| PrimitiveType::Bool),
115        map(tag("string"), |_| PrimitiveType::String),
116        box_ty,
117        list_ty,
118        option_ty,
119        result_ty,
120    ))(i)
121}
122
123pub fn type_params_list(i: &str) -> IResult<&str, Vec<String>> {
124    let (i, _) = tag("<")(i)?;
125    let (i, params) = separated_list0(
126        preceded(multispace0, char(',')),
127        preceded(multispace0, identifier),
128    )(i)?;
129    let (i, _) = multispace0(i)?;
130    let (i, _) = opt(char(','))(i)?;
131    let (i, _) = multispace0(i)?;
132    let (i, _) = tag(">")(i)?;
133
134    let params = params.into_iter().map(|p| p.into()).collect();
135
136    Ok((i, params))
137}
138
139pub fn type_args_list(i: &str) -> IResult<&str, Vec<Type>> {
140    let (i, _) = tag("<")(i)?;
141    let (i, args) =
142        separated_list0(preceded(multispace0, char(',')), preceded(multispace0, ty))(i)?;
143    let (i, _) = multispace0(i)?;
144    let (i, _) = opt(char(','))(i)?;
145    let (i, _) = multispace0(i)?;
146    let (i, _) = tag(">")(i)?;
147
148    let args = args.into_iter().map(|p| p.into()).collect();
149
150    Ok((i, args))
151}
152
153pub fn struct_def(i: &str) -> IResult<&str, TypeDef> {
154    let (i, _) = tag("struct")(i)?;
155    let (i, _) = multispace0(i)?;
156    let (i, name) = identifier(i)?;
157    let (i, _) = multispace0(i)?;
158    let (i, maybe_params) = opt(type_params_list)(i)?;
159    let (i, _) = multispace0(i)?;
160    let (i, fields) = named_fields(i)?;
161
162    let params = maybe_params.unwrap_or(Vec::new());
163
164    let type_def = TypeDef {
165        name: name.into(),
166        params,
167        body: TypeBody::Struct(Struct { fields }),
168    };
169
170    Ok((i, type_def))
171}
172
173fn enum_def<'a>(i: &'a str) -> IResult<&'a str, TypeDef> {
174    let (i, _) = tag("enum")(i)?;
175    let (i, _) = multispace0(i)?;
176    let (i, name) = identifier(i)?;
177    let (i, _) = multispace0(i)?;
178    let (i, maybe_params) = opt(type_params_list)(i)?;
179    let (i, _) = multispace0(i)?;
180    let (i, variants) = enum_variants(i)?;
181
182    let params = maybe_params.unwrap_or(Vec::new());
183
184    let type_def = TypeDef {
185        name: name.into(),
186        params,
187        body: TypeBody::Enum(Enum { variants }),
188    };
189
190    Ok((i, type_def))
191}
192
193fn enum_variants<'a>(i: &'a str) -> IResult<&'a str, Vec<(String, EnumVariant)>> {
194    context(
195        "map",
196        preceded(
197            char('{'),
198            cut(terminated(
199                opt_trailing_comma(map(
200                    separated_list0(
201                        preceded(multispace0, char(',')),
202                        preceded(multispace0, enum_variant),
203                    ),
204                    |tuple_vec| {
205                        tuple_vec
206                            .into_iter()
207                            .map(|(name, variant)| (name.into(), variant))
208                            .collect()
209                    },
210                )),
211                preceded(multispace0, char('}')),
212            )),
213        ),
214    )(i)
215}
216
217fn enum_variant<'a>(i: &'a str) -> IResult<&'a str, (&'a str, EnumVariant)> {
218    alt((
219        separated_pair(
220            identifier,
221            multispace0,
222            map(named_fields, |fields| EnumVariant::NamedFields { fields }),
223        ),
224        map(identifier, |x| (x, EnumVariant::Empty)),
225    ))(i)
226}
227
228pub fn defined_ty<'a>(i: &'a str) -> IResult<&'a str, Type> {
229    let (i, ident) = qualified_identifier(i)?;
230    let (i, _) = multispace0(i)?;
231    let (i, maybe_args) = opt(type_args_list)(i)?;
232
233    let args = maybe_args.unwrap_or(Vec::new());
234    let defined_type = Type::Defined { ident, args };
235
236    Ok((i, defined_type))
237}
238
239pub fn ty<'a>(i: &'a str) -> IResult<&'a str, Type> {
240    alt((map(builtin_ty, |x| Type::Primitive(x)), defined_ty))(i)
241}
242
243pub fn type_def<'a>(i: &'a str) -> IResult<&'a str, TypeDef> {
244    alt((struct_def, enum_def))(i)
245}
246
247fn named_field<'a>(i: &'a str) -> IResult<&'a str, (&'a str, Type)> {
248    separated_pair(
249        identifier,
250        cut(preceded(multispace0, char(':'))),
251        preceded(multispace0, ty),
252    )(i)
253}
254
255fn named_fields<'a>(i: &'a str) -> IResult<&'a str, Vec<NamedField>> {
256    context(
257        "map",
258        preceded(
259            char('{'),
260            cut(terminated(
261                opt_trailing_comma(map(
262                    separated_list0(
263                        preceded(multispace0, char(',')),
264                        preceded(multispace0, named_field),
265                    ),
266                    |tuple_vec| {
267                        tuple_vec
268                            .into_iter()
269                            .map(|(k, v)| NamedField {
270                                name: k.to_owned(),
271                                ty: v,
272                            })
273                            .collect()
274                    },
275                )),
276                preceded(multispace0, char('}')),
277            )),
278        ),
279    )(i)
280}
281
282pub fn root<'a>(i: &'a str) -> IResult<&'a str, Vec<TypeDef>> {
283    separated_list0(multispace0, type_def)(i)
284}
285
286pub fn parse_file(path: impl AsRef<std::path::Path>) -> Result<Vec<TypeDef>, String> {
287    use std::io::Read;
288
289    // Open input file
290    let Ok(mut file) = std::fs::File::open(path.as_ref()) else {
291        return Err(format!("Failed to open file '{}'", path.as_ref().display()));
292    };
293
294    // Load file to string
295    let mut file_str = String::new();
296    if let Err(e) = file.read_to_string(&mut file_str) {
297        return Err(format!(
298            "Failed to read file '{}': {}",
299            path.as_ref().display(),
300            e
301        ));
302    }
303
304    // Remove comments
305    let mut schema_str = file_str.lines()
306        .map(|line| {
307            if let Some(index) = line.find("//") {
308                // Return the slice from the start of the line up to the comment marker
309                &line[..index]
310            } else {
311                // If no comment is found, return the entire line
312                line
313            }
314        })
315        .collect::<Vec<&str>>()
316        .join("\n");
317    schema_str += "\n";
318
319    let (_, type_defs) = root(&schema_str).map_err(|e| format!("mproto schema parse error: {e}"))?;
320
321    Ok(type_defs)
322}
323
324#[cfg(test)]
325mod tests {
326    use super::*;
327
328    #[test]
329    fn test_builtin_u8() {
330        let data = "u8";
331        let (_, parsed) = builtin_ty(data).unwrap();
332
333        assert_eq!(parsed, PrimitiveType::U8);
334    }
335
336    #[test]
337    fn test_list_u8() {
338        let data = "[u8]";
339        let (_, parsed) = list_ty(data).unwrap();
340
341        assert_eq!(
342            parsed,
343            PrimitiveType::List(Box::new(Type::Primitive(PrimitiveType::U8))),
344        );
345    }
346
347    #[test]
348    fn test_box_u8() {
349        let data = "box<u8>";
350        let (_, parsed) = box_ty(data).unwrap();
351
352        assert_eq!(
353            parsed,
354            PrimitiveType::Box(Box::new(Type::Primitive(PrimitiveType::U8))),
355        );
356    }
357
358    #[test]
359    fn test_option_u8() {
360        let data = "option<u8>";
361        let (_, parsed) = option_ty(data).unwrap();
362
363        assert_eq!(
364            parsed,
365            PrimitiveType::Option(Box::new(Type::Primitive(PrimitiveType::U8))),
366        );
367    }
368
369    #[test]
370    fn test_result() {
371        let data = "result<void, string>";
372        let (_, parsed) = result_ty(data).unwrap();
373
374        assert_eq!(
375            parsed,
376            PrimitiveType::Result(
377                Box::new(Type::Primitive(PrimitiveType::Void)),
378                Box::new(Type::Primitive(PrimitiveType::String)),
379            ),
380        );
381    }
382
383    #[test]
384    fn test_struct_named_fields() {
385        use PrimitiveType::*;
386
387        let data = "struct Foo { bar : u32, baz : i8 }";
388        let (_, parsed) = struct_def(data).unwrap();
389
390        assert_eq!(
391            parsed,
392            TypeDef {
393                name: "Foo".into(),
394                params: vec![],
395                body: TypeBody::Struct(Struct {
396                    fields: vec![
397                        NamedField {
398                            name: "bar".into(),
399                            ty: Type::Primitive(U32)
400                        },
401                        NamedField {
402                            name: "baz".into(),
403                            ty: Type::Primitive(I8)
404                        },
405                    ]
406                }),
407            }
408        );
409    }
410
411    #[test]
412    fn test_struct_type_param() {
413        use PrimitiveType::*;
414
415        let data = "struct Foo <T, F>  { bar : u32, baz : i8 }";
416        let (_, parsed) = struct_def(data).unwrap();
417
418        assert_eq!(
419            parsed,
420            TypeDef {
421                name: "Foo".into(),
422                params: vec!["T".into(), "F".into()],
423                body: TypeBody::Struct(Struct {
424                    fields: vec![
425                        NamedField {
426                            name: "bar".into(),
427                            ty: Type::Primitive(U32)
428                        },
429                        NamedField {
430                            name: "baz".into(),
431                            ty: Type::Primitive(I8)
432                        },
433                    ]
434                }),
435            }
436        );
437    }
438
439    #[test]
440    fn test_defined_type() {
441        use PrimitiveType::*;
442
443        let data = "struct Foo { bar : bar_proto.Bar, baz : i8 }";
444        let (_, parsed) = struct_def(data).unwrap();
445
446        assert_eq!(
447            parsed,
448            TypeDef {
449                name: "Foo".into(),
450                params: vec![],
451                body: TypeBody::Struct(Struct {
452                    fields: vec![
453                        NamedField {
454                            name: "bar".into(),
455                            ty: Type::Defined {
456                                ident: QualifiedIdentifier {
457                                    module: Some("bar_proto".into()),
458                                    name: "Bar".into(),
459                                },
460                                args: vec![],
461                            },
462                        },
463                        NamedField {
464                            name: "baz".into(),
465                            ty: Type::Primitive(I8)
466                        },
467                    ]
468                }),
469            }
470        );
471    }
472
473    #[test]
474    fn test_enum_named_fields() {
475        use PrimitiveType::*;
476
477        let data = "enum Foo { Bar { x: u32, y: u8 }, Baz { bip: i8 } }";
478        let (_, parsed) = enum_def(data).unwrap();
479
480        assert_eq!(
481            parsed,
482            TypeDef {
483                name: "Foo".into(),
484                params: vec![],
485                body: TypeBody::Enum(Enum {
486                    variants: vec![
487                        (
488                            "Bar".into(),
489                            EnumVariant::NamedFields {
490                                fields: vec![
491                                    NamedField {
492                                        name: "x".into(),
493                                        ty: Type::Primitive(U32)
494                                    },
495                                    NamedField {
496                                        name: "y".into(),
497                                        ty: Type::Primitive(U8)
498                                    },
499                                ],
500                            }
501                        ),
502                        (
503                            "Baz".into(),
504                            EnumVariant::NamedFields {
505                                fields: vec![NamedField {
506                                    name: "bip".into(),
507                                    ty: Type::Primitive(I8)
508                                },],
509                            }
510                        ),
511                    ]
512                }),
513            }
514        );
515    }
516}