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 root<'a>(i: &'a str) -> IResult<&'a str, Vec<TypeDef>> {
284    separated_list0(multispace0, type_def)(i)
285}
286
287pub fn parse_file(path: impl AsRef<std::path::Path>) -> Result<Vec<TypeDef>, String> {
288    use std::io::Read;
289
290    // Open input file
291    let Ok(mut file) = std::fs::File::open(path.as_ref()) else {
292        return Err(format!("Failed to open file '{}'", path.as_ref().display()));
293    };
294
295    // Load file to string
296    let mut file_str = String::new();
297    if let Err(e) = file.read_to_string(&mut file_str) {
298        return Err(format!(
299            "Failed to read file '{}': {}",
300            path.as_ref().display(),
301            e
302        ));
303    }
304
305    // Remove comments
306    let mut schema_str = file_str.lines()
307        .map(|line| {
308            if let Some(index) = line.find("//") {
309                // Return the slice from the start of the line up to the comment marker
310                &line[..index]
311            } else {
312                // If no comment is found, return the entire line
313                line
314            }
315        })
316        .collect::<Vec<&str>>()
317        .join("\n");
318    schema_str += "\n";
319
320    let (_, type_defs) = root(&schema_str).map_err(|e| format!("mproto schema parse error: {e}"))?;
321
322    Ok(type_defs)
323}
324
325#[cfg(test)]
326mod tests {
327    use super::*;
328
329    #[test]
330    fn test_builtin_u8() {
331        let data = "u8";
332        let (_, parsed) = builtin_ty(data).unwrap();
333
334        assert_eq!(parsed, PrimitiveType::U8);
335    }
336
337    #[test]
338    fn test_list_u8() {
339        let data = "[u8]";
340        let (_, parsed) = list_ty(data).unwrap();
341
342        assert_eq!(
343            parsed,
344            PrimitiveType::List(Box::new(Type::Primitive(PrimitiveType::U8))),
345        );
346    }
347
348    #[test]
349    fn test_box_u8() {
350        let data = "box<u8>";
351        let (_, parsed) = box_ty(data).unwrap();
352
353        assert_eq!(
354            parsed,
355            PrimitiveType::Box(Box::new(Type::Primitive(PrimitiveType::U8))),
356        );
357    }
358
359    #[test]
360    fn test_option_u8() {
361        let data = "option<u8>";
362        let (_, parsed) = option_ty(data).unwrap();
363
364        assert_eq!(
365            parsed,
366            PrimitiveType::Option(Box::new(Type::Primitive(PrimitiveType::U8))),
367        );
368    }
369
370    #[test]
371    fn test_result() {
372        let data = "result<void, string>";
373        let (_, parsed) = result_ty(data).unwrap();
374
375        assert_eq!(
376            parsed,
377            PrimitiveType::Result(
378                Box::new(Type::Primitive(PrimitiveType::Void)),
379                Box::new(Type::Primitive(PrimitiveType::String)),
380            ),
381        );
382    }
383
384    #[test]
385    fn test_struct_named_fields() {
386        use PrimitiveType::*;
387
388        let data = "struct Foo { bar : u32, baz : i8 }";
389        let (_, parsed) = struct_def(data).unwrap();
390
391        assert_eq!(
392            parsed,
393            TypeDef {
394                name: "Foo".into(),
395                params: vec![],
396                body: TypeBody::Struct(Struct {
397                    fields: vec![
398                        NamedField {
399                            name: "bar".into(),
400                            ty: Type::Primitive(U32)
401                        },
402                        NamedField {
403                            name: "baz".into(),
404                            ty: Type::Primitive(I8)
405                        },
406                    ]
407                }),
408            }
409        );
410    }
411
412    #[test]
413    fn test_struct_type_param() {
414        use PrimitiveType::*;
415
416        let data = "struct Foo <T, F>  { bar : u32, baz : i8 }";
417        let (_, parsed) = struct_def(data).unwrap();
418
419        assert_eq!(
420            parsed,
421            TypeDef {
422                name: "Foo".into(),
423                params: vec!["T".into(), "F".into()],
424                body: TypeBody::Struct(Struct {
425                    fields: vec![
426                        NamedField {
427                            name: "bar".into(),
428                            ty: Type::Primitive(U32)
429                        },
430                        NamedField {
431                            name: "baz".into(),
432                            ty: Type::Primitive(I8)
433                        },
434                    ]
435                }),
436            }
437        );
438    }
439
440    #[test]
441    fn test_defined_type() {
442        use PrimitiveType::*;
443
444        let data = "struct Foo { bar : bar_proto.Bar, baz : i8 }";
445        let (_, parsed) = struct_def(data).unwrap();
446
447        assert_eq!(
448            parsed,
449            TypeDef {
450                name: "Foo".into(),
451                params: vec![],
452                body: TypeBody::Struct(Struct {
453                    fields: vec![
454                        NamedField {
455                            name: "bar".into(),
456                            ty: Type::Defined {
457                                ident: QualifiedIdentifier {
458                                    module: Some("bar_proto".into()),
459                                    name: "Bar".into(),
460                                },
461                                args: vec![],
462                            },
463                        },
464                        NamedField {
465                            name: "baz".into(),
466                            ty: Type::Primitive(I8)
467                        },
468                    ]
469                }),
470            }
471        );
472    }
473
474    #[test]
475    fn test_enum_named_fields() {
476        use PrimitiveType::*;
477
478        let data = "enum Foo { Bar { x: u32, y: u8 }, Baz { bip: i8 } }";
479        let (_, parsed) = enum_def(data).unwrap();
480
481        assert_eq!(
482            parsed,
483            TypeDef {
484                name: "Foo".into(),
485                params: vec![],
486                body: TypeBody::Enum(Enum {
487                    variants: vec![
488                        (
489                            "Bar".into(),
490                            EnumVariant::NamedFields {
491                                fields: vec![
492                                    NamedField {
493                                        name: "x".into(),
494                                        ty: Type::Primitive(U32)
495                                    },
496                                    NamedField {
497                                        name: "y".into(),
498                                        ty: Type::Primitive(U8)
499                                    },
500                                ],
501                            }
502                        ),
503                        (
504                            "Baz".into(),
505                            EnumVariant::NamedFields {
506                                fields: vec![NamedField {
507                                    name: "bip".into(),
508                                    ty: Type::Primitive(I8)
509                                },],
510                            }
511                        ),
512                    ]
513                }),
514            }
515        );
516    }
517}