Skip to main content

rusteron_code_gen/
parser.rs

1use crate::generator::{parse_custom_methods, CBinding, CWrapper, Method};
2use crate::{Arg, ArgProcessing, CHandler};
3use itertools::Itertools;
4use quote::ToTokens;
5use std::collections::{BTreeMap, BTreeSet};
6use std::fs;
7use std::path::PathBuf;
8use syn::{Attribute, Item, ItemForeignMod, ItemStruct, ItemType, Lit, Meta, MetaNameValue};
9
10/// bindgen maps C `va_list` to platform-specific opaque types (e.g.
11/// `__va_list_tag` on Linux, `VaListTag` on some platforms).  These never
12/// need a real Rust repr — replace every occurrence with `*mut c_char`
13/// so the generated handler wrappers compile everywhere.
14fn normalise_va_list(c_type: &str) -> String {
15    for needle in &["va_list", "__va_list_tag", "VaListTag", "__builtin_va_list"] {
16        if c_type.contains(needle) {
17            return "*mut ::std::os::raw::c_char".to_string();
18        }
19    }
20    c_type.to_string()
21}
22
23pub fn parse_bindings(out: &PathBuf) -> CBinding {
24    parse_bindings_with_custom(out, &[])
25}
26
27/// Like [`parse_bindings`], additionally scanning per-crate custom code (appended after the
28/// common `aeron_custom.rs`) so its hand-written methods are skipped during generation.
29pub fn parse_bindings_with_custom(out: &PathBuf, extra_custom_code: &[&str]) -> CBinding {
30    let file_content = fs::read_to_string(out.clone()).expect("Unable to read file");
31    let syntax_tree = syn::parse_file(&file_content).expect("Unable to parse file");
32    let mut wrappers = BTreeMap::new();
33    let mut methods = Vec::new();
34    let mut handlers = Vec::new();
35
36    let items = syntax_tree.items;
37
38    for item in &items {
39        if let Item::Type(ty) = item {
40            process_type(&mut wrappers, &mut handlers, ty);
41        }
42    }
43
44    let handler_names = handlers
45        .iter()
46        .filter(|h| {
47            !["aeron_udp_channel", "aeron_udp_transport"]
48                .iter()
49                .any(|&filter| h.type_name.starts_with(filter))
50        })
51        .map(|handler| handler.type_name.clone())
52        .collect();
53
54    for item in &items {
55        if let Item::Struct(s) = item {
56            process_struct(&mut wrappers, s, &handler_names);
57        }
58    }
59
60    for item in &items {
61        if let Item::ForeignMod(fm) = item {
62            process_c_method(&mut wrappers, &mut methods, fm, &handler_names);
63        }
64    }
65
66    let mut bindings = CBinding {
67        wrappers: wrappers
68            .into_iter()
69            .filter(|(_, wrapper)| {
70                // these are from media driver and do not follow convention
71                ![
72                    "aeron_thread",
73                    "aeron_command",
74                    "aeron_executor",
75                    "aeron_name_resolver",
76                    "aeron_udp_channel_transport", // this one I have issues with handlers
77                    "aeron_udp_transport",         // this one I have issues with handlers
78                ]
79                .iter()
80                .any(|&filter| wrapper.type_name.starts_with(filter))
81            })
82            .collect(),
83        methods,
84        handlers: handlers
85            .into_iter()
86            .filter(|h| {
87                !["aeron_udp_channel", "aeron_udp_transport"]
88                    .iter()
89                    .any(|&filter| h.type_name.starts_with(filter))
90            })
91            .collect(),
92    };
93
94    let mismatched_types = bindings
95        .wrappers
96        .iter()
97        .filter(|(key, w)| key.as_str() != w.type_name)
98        .map(|(a, b)| (a.clone(), b.clone()))
99        .collect_vec();
100    assert_eq!(Vec::<(String, CWrapper)>::new(), mismatched_types);
101
102    let mut custom = parse_custom_methods(crate::CUSTOM_AERON_CODE);
103    for code in extra_custom_code {
104        for (class, methods) in parse_custom_methods(code) {
105            custom.entry(class).or_default().extend(methods);
106        }
107    }
108    for wrapper in bindings.wrappers.values_mut() {
109        if let Some(methods) = custom.get(&wrapper.class_name) {
110            wrapper.skipped_methods = methods.clone();
111        }
112    }
113
114    bindings
115}
116
117fn process_c_method(
118    wrappers: &mut BTreeMap<String, CWrapper>,
119    methods: &mut Vec<Method>,
120    fm: &ItemForeignMod,
121    handler_names: &BTreeSet<String>,
122) {
123    // Extract functions inside extern "C" blocks
124    if fm.abi.name.is_some() && fm.abi.name.as_ref().unwrap().value() == "C" {
125        for foreign_item in &fm.items {
126            if let syn::ForeignItem::Fn(f) = foreign_item {
127                let docs = get_doc_comments(&f.attrs);
128                let fn_name = f.sig.ident.to_string();
129
130                // aeronc.h has a deprecated typo alias:
131                //   aeron_async_add_exclusive_exclusive_publication_get_registration_id
132                // (note the doubled "exclusive"). It attaches to the same wrapper as the
133                // canonical aeron_async_add_exclusive_publication_get_registration_id and
134                // would either collide with it or emit a garbled method name. Upstream
135                // marks it @deprecated, so drop it outright.
136                if fn_name.contains("exclusive_exclusive") {
137                    continue;
138                }
139
140                // Get function arguments and return type as Rust code
141                let args = extract_function_arguments(&f.sig.inputs);
142                let ret = extract_return_type(&f.sig.output);
143
144                let option = if let Some(arg) = args
145                    .iter()
146                    .skip_while(|a| a.is_mut_pointer() && a.is_primitive())
147                    .next()
148                {
149                    let ty = &arg.c_type;
150                    let ty = ty.split(' ').last().map(|t| t.to_string()).unwrap();
151                    if wrappers.contains_key(&ty) {
152                        Some(ty)
153                    } else {
154                        find_closest_wrapper_from_method_name(wrappers, &fn_name)
155                    }
156                } else {
157                    find_closest_wrapper_from_method_name(wrappers, &fn_name)
158                };
159
160                match option {
161                    Some(key) => {
162                        let wrapper = wrappers.get_mut(&key).unwrap();
163                        wrapper.methods.push(Method {
164                            fn_name: fn_name.clone(),
165                            struct_method_name: fn_name
166                                .replace(&wrapper.type_name[..wrapper.type_name.len() - 1], "")
167                                .to_string(),
168                            return_type: Arg {
169                                name: "".to_string(),
170                                c_type: ret.clone(),
171                                processing: ArgProcessing::Default,
172                            },
173                            arguments: process_types(args.clone(), Some(handler_names)),
174                            docs: docs.clone(),
175                        });
176                    }
177                    None => methods.push(Method {
178                        fn_name: fn_name.clone(),
179                        struct_method_name: "".to_string(),
180                        return_type: Arg {
181                            name: "".to_string(),
182                            c_type: ret.clone(),
183                            processing: ArgProcessing::Default,
184                        },
185                        arguments: process_types(args.clone(), Some(handler_names)),
186                        docs: docs.clone(),
187                    }),
188                }
189            }
190        }
191    }
192}
193
194fn find_closest_wrapper_from_method_name(
195    wrappers: &mut BTreeMap<String, CWrapper>,
196    fn_name: &String,
197) -> Option<String> {
198    let type_names = get_possible_wrappers(&fn_name);
199
200    let mut value = None;
201    for ty in type_names {
202        if wrappers.contains_key(&ty) {
203            value = Some(ty);
204            break;
205        }
206    }
207
208    value
209}
210
211pub fn get_possible_wrappers(fn_name: &str) -> Vec<String> {
212    // Try all possible wrapper names by splitting at underscores
213    fn_name
214        .char_indices()
215        .filter(|(_, c)| *c == '_')
216        .map(|(i, _)| format!("{}_t", &fn_name[..i]))
217        .rev()
218        .collect_vec()
219}
220
221fn process_type(wrappers: &mut BTreeMap<String, CWrapper>, handlers: &mut Vec<CHandler>, ty: &ItemType) {
222    // Handle type definitions and get docs
223    let docs = get_doc_comments(&ty.attrs);
224
225    let type_name = ty.ident.to_string();
226    let class_name = snake_to_pascal_case(&type_name);
227
228    if is_struct_typedef(&ty.ty) {
229        wrappers
230            .entry(type_name.clone())
231            .or_insert(CWrapper {
232                class_name,
233                without_name: type_name[..type_name.len() - 2].to_string(),
234                type_name,
235                ..Default::default()
236            })
237            .docs
238            .extend(docs);
239    } else {
240        // Parse the function pointer type -> it is typically used for handlers/callbacks
241        if let syn::Type::Path(type_path) = &*ty.ty {
242            if let Some(segment) = type_path.path.segments.last() {
243                if segment.ident.to_string() == "Option" {
244                    if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
245                        if let Some(syn::GenericArgument::Type(syn::Type::BareFn(bare_fn))) = args.args.first() {
246                            let args: Vec<Arg> = bare_fn
247                                .inputs
248                                .iter()
249                                .map(|arg| {
250                                    let arg_name = match &arg.name {
251                                        Some((ident, _)) => ident.to_string(),
252                                        None => "".to_string(),
253                                    };
254                                    let arg_type = arg.ty.to_token_stream().to_string();
255                                    (arg_name, arg_type)
256                                })
257                                .map(|(field_name, field_type)| Arg {
258                                    name: field_name,
259                                    c_type: normalise_va_list(&field_type),
260                                    processing: ArgProcessing::Default,
261                                })
262                                .collect();
263                            let string = bare_fn.output.to_token_stream().to_string();
264                            let mut return_type = string.trim();
265
266                            if return_type.starts_with("-> ") {
267                                return_type = &return_type[3..];
268                            }
269
270                            if return_type.is_empty() {
271                                return_type = "()";
272                            }
273
274                            if is_handler_typedef(&args) {
275                                let value = CHandler {
276                                    type_name: ty.ident.to_string(),
277                                    args: process_types(args, None),
278                                    return_type: Arg {
279                                        name: "".to_string(),
280                                        c_type: return_type.to_string(),
281                                        processing: ArgProcessing::Default,
282                                    },
283                                    docs: docs.clone(),
284                                    fn_mut_signature: Default::default(),
285                                    closure_type_name: Default::default(),
286                                };
287                                handlers.push(value);
288                            }
289                        }
290                    }
291                }
292            }
293        }
294    }
295}
296
297fn is_handler_typedef(args: &[Arg]) -> bool {
298    args.iter().filter(|arg| arg.is_c_void()).count() == 1
299        || args
300            .first()
301            .map(|arg| arg.is_c_void() && is_client_data_arg(&arg.name))
302            .unwrap_or(false)
303}
304
305fn is_client_data_arg(name: &str) -> bool {
306    name == "clientd"
307        || name == "state"
308        || name == "task_clientd"
309        || name.ends_with("_clientd")
310        || name.ends_with("_state")
311}
312
313fn is_struct_typedef(ty: &syn::Type) -> bool {
314    if let syn::Type::Path(type_path) = ty {
315        if let Some(segment) = type_path.path.segments.last() {
316            return segment.ident.to_string().ends_with("_stct");
317        }
318    }
319
320    false
321}
322
323fn process_struct(wrappers: &mut BTreeMap<String, CWrapper>, s: &ItemStruct, handler_names: &BTreeSet<String>) {
324    if !matches!(s.fields, syn::Fields::Named(_)) {
325        return;
326    }
327
328    // Print the struct name and its doc comments
329    let docs = get_doc_comments(&s.attrs);
330    let type_name = s.ident.to_string().replace("_stct", "_t");
331    let class_name = snake_to_pascal_case(&type_name);
332
333    let fields: Vec<Arg> = s
334        .fields
335        .iter()
336        .map(|f| {
337            let field_name = f.ident.as_ref().unwrap().to_string();
338            let field_type = f.ty.to_token_stream().to_string();
339            (field_name, field_type)
340        })
341        .map(|(field_name, field_type)| Arg {
342            name: field_name,
343            c_type: field_type,
344            processing: ArgProcessing::Default,
345        })
346        .collect();
347
348    let w = wrappers.entry(type_name.to_string()).or_insert(CWrapper {
349        class_name,
350        without_name: type_name[..type_name.len() - 2].to_string(),
351        type_name,
352        ..Default::default()
353    });
354    w.docs.extend(docs);
355    w.fields = process_types(fields, Some(handler_names));
356}
357
358fn process_types(mut name_and_type: Vec<Arg>, handler_names: Option<&BTreeSet<String>>) -> Vec<Arg> {
359    // now mark arguments which can be reduced
360    for i in 1..name_and_type.len() {
361        let param1 = &name_and_type[i - 1];
362        let param2 = &name_and_type[i];
363
364        let is_int = param2.c_type == "usize" || param2.c_type == "i32";
365        // Length fields use varied names across Aeron headers; match the common ones so
366        // (buffer, length) pairs merge into a slice rather than leaking two raw params.
367        let length_field = matches!(param2.name.as_str(), "length" | "len" | "count" | "capacity")
368            || param2.name.ends_with("_length")
369            || param2.name.ends_with("_len")
370            || param2.name.ends_with("_size");
371        if param2.is_c_void()
372            && !param1.is_mut_pointer()
373            && param1.c_type.ends_with("_t")
374            && handler_names
375                .map(|handler_names| handler_names.contains(&param1.c_type))
376                .unwrap_or(false)
377        {
378            // closures
379            //         handler: aeron_on_available_counter_t,
380            //         clientd: *mut ::std::os::raw::c_void,
381            let processing = ArgProcessing::Handler(vec![param1.clone(), param2.clone()]);
382            name_and_type[i - 1].processing = processing.clone();
383            name_and_type[i].processing = processing.clone();
384        } else if param1.is_c_string_any() && is_int && length_field {
385            //     pub stripped_channel: *mut ::std::os::raw::c_char,
386            //     pub stripped_channel_length: usize,
387            // `*const` pairs become `&str` arguments; `*mut` pairs are C fill-buffers
388            // (snprintf-style) and become `&mut [u8]` arguments / `&str` field reads.
389            let processing = ArgProcessing::StringWithLength(vec![param1.clone(), param2.clone()]);
390            name_and_type[i - 1].processing = processing.clone();
391            name_and_type[i].processing = processing.clone();
392        } else if param1.is_byte_array() && is_int && length_field {
393            //         key_buffer: *const u8,
394            //         key_buffer_length: usize,
395            let processing = ArgProcessing::ByteArrayWithLength(vec![param1.clone(), param2.clone()]);
396            name_and_type[i - 1].processing = processing.clone();
397            name_and_type[i].processing = processing.clone();
398        }
399
400        //
401    }
402
403    name_and_type
404}
405
406// Helper function to extract doc comments
407fn get_doc_comments(attrs: &[Attribute]) -> BTreeSet<String> {
408    attrs
409        .iter()
410        .filter_map(|attr| {
411            // Parse the attribute meta to check if it is a `Meta::NameValue`
412            if let Meta::NameValue(MetaNameValue {
413                path,
414                value: syn::Expr::Lit(expr_lit),
415                ..
416            }) = &attr.meta
417            {
418                // Check if the path is "doc"
419                if path.is_ident("doc") {
420                    // Check if the literal is a string and return its value
421                    if let Lit::Str(lit_str) = &expr_lit.lit {
422                        return Some(lit_str.value().trim().to_string());
423                    }
424                }
425            }
426            None
427        })
428        .collect()
429}
430
431pub fn snake_to_pascal_case(mut snake: &str) -> String {
432    if snake.ends_with("_t") {
433        snake = &snake[..snake.len() - 2];
434    }
435    snake
436        .split('_')
437        .filter(|x| *x != "on") // Split the string by underscores
438        .map(|word| {
439            let mut chars = word.chars();
440            // Capitalize the first letter and collect the rest of the letters
441            match chars.next() {
442                Some(c) => c.to_uppercase().collect::<String>() + chars.as_str(),
443                None => String::new(),
444            }
445        })
446        .collect()
447}
448
449// Helper function to extract function arguments as Rust code
450fn extract_function_arguments(inputs: &syn::punctuated::Punctuated<syn::FnArg, syn::token::Comma>) -> Vec<Arg> {
451    inputs
452        .iter()
453        .map(|arg| match arg {
454            syn::FnArg::Receiver(_) => "self".to_string(), // Handle self receiver
455            syn::FnArg::Typed(pat_type) => pat_type.to_token_stream().to_string(), // Convert the pattern and type to Rust code
456        })
457        .map(|arg| {
458            arg.splitn(2, ':')
459                .map(|s| s.trim().to_string())
460                .collect_tuple()
461                .unwrap()
462        })
463        .map(|(name, ty)| Arg {
464            name,
465            c_type: ty,
466            processing: ArgProcessing::Default,
467        })
468        .collect_vec()
469}
470
471// Helper function to extract return type as Rust code
472fn extract_return_type(output: &syn::ReturnType) -> String {
473    match output {
474        syn::ReturnType::Default => "()".to_string(), // No return type, equivalent to ()
475        syn::ReturnType::Type(_, ty) => ty.to_token_stream().to_string(), // Convert the type to Rust code
476    }
477}
478
479#[cfg(test)]
480mod tests {
481    use crate::parser::parse_bindings;
482    use crate::ArgProcessing;
483    use std::path::PathBuf;
484
485    fn running_under_valgrind() -> bool {
486        std::env::var_os("RUSTERON_VALGRIND").is_some()
487    }
488
489    #[test]
490    fn media_driver() {
491        if running_under_valgrind() {
492            return;
493        }
494
495        let path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
496            .join("bindings")
497            .join("media-driver.rs");
498        let bindings = parse_bindings(&path);
499        assert_eq!(
500            "AeronImageFragmentAssembler",
501            bindings
502                .wrappers
503                .get("aeron_image_fragment_assembler_t")
504                .unwrap()
505                .class_name
506        );
507    }
508    #[test]
509    fn client() {
510        if running_under_valgrind() {
511            return;
512        }
513
514        let path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
515            .join("bindings")
516            .join("client.rs");
517        let bindings = parse_bindings(&path);
518        assert_eq!(
519            "AeronImageFragmentAssembler",
520            bindings
521                .wrappers
522                .get("aeron_image_fragment_assembler_t")
523                .unwrap()
524                .class_name
525        );
526        assert!(bindings.handlers.len() > 1);
527    }
528
529    /// `(buffer, frame_length)` must merge into one slice; old heuristic missed `frame_length`.
530    #[test]
531    fn reserved_value_supplier_buffer_merges_with_frame_length() {
532        if running_under_valgrind() {
533            return;
534        }
535
536        let path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
537            .join("bindings")
538            .join("client.rs");
539        let bindings = parse_bindings(&path);
540        let handler = bindings
541            .handlers
542            .iter()
543            .find(|h| h.type_name == "aeron_reserved_value_supplier_t")
544            .expect("reserved value supplier handler missing");
545        let buffer = handler.args.iter().find(|a| a.name == "buffer").unwrap();
546        assert!(
547            matches!(buffer.processing, ArgProcessing::ByteArrayWithLength(_)),
548            "buffer must merge with frame_length into a slice, got {:?}",
549            buffer.processing
550        );
551    }
552
553    /// `*mut c_char` fill-buffers must merge so the wrapper takes `&mut [u8]`.
554    #[test]
555    fn mut_string_fill_buffer_merges() {
556        if running_under_valgrind() {
557            return;
558        }
559
560        let path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
561            .join("bindings")
562            .join("client.rs");
563        let bindings = parse_bindings(&path);
564        let reader = bindings.wrappers.get("aeron_counters_reader_t").unwrap();
565        let method = reader
566            .methods
567            .iter()
568            .find(|m| m.fn_name == "aeron_counters_reader_counter_label")
569            .expect("counter_label method missing");
570        let buffer = method.arguments.iter().find(|a| a.name == "buffer").unwrap();
571        assert!(
572            matches!(buffer.processing, ArgProcessing::StringWithLength(_)),
573            "mut string buffer must merge with its length, got {:?}",
574            buffer.processing
575        );
576    }
577}