Skip to main content

extendr_macros/
extendr_module.rs

1use crate::wrappers;
2use proc_macro::TokenStream;
3use quote::{format_ident, quote};
4use syn::{parse::ParseStream, parse_macro_input, Ident, Token, Type};
5
6pub fn extendr_module(item: TokenStream) -> TokenStream {
7    let module = parse_macro_input!(item as Module);
8    let Module {
9        modname,
10        fnnames,
11        implnames,
12        usenames,
13    } = module;
14    let modname = modname.expect("cannot include unnamed modules");
15    let modname_string = modname.to_string();
16    let module_init_name = format_ident!("R_init_{}_extendr", modname);
17
18    let module_metadata_name = format_ident!("get_{}_metadata", modname);
19    let module_metadata_name_string = module_metadata_name.to_string();
20    let wrap_module_metadata_name =
21        format_ident!("{}get_{}_metadata", wrappers::WRAP_PREFIX, modname);
22    let wrap_module_metadata_name_str = wrap_module_metadata_name.to_string();
23
24    let make_module_wrappers_name = format_ident!("make_{}_wrappers", modname);
25    let make_module_wrappers_name_string = make_module_wrappers_name.to_string();
26    let wrap_make_module_wrappers =
27        format_ident!("{}make_{}_wrappers", wrappers::WRAP_PREFIX, modname);
28    let wrap_make_module_wrappers_string = wrap_make_module_wrappers.to_string();
29    let write_make_module_wrappers = format_ident!("write__make_{}_wrappers", modname);
30
31    let fnmetanames = fnnames
32        .iter()
33        .map(|id| format_ident!("{}{}", wrappers::META_PREFIX, id));
34    let implmetanames = implnames
35        .iter()
36        .map(|id| format_ident!("{}{}", wrappers::META_PREFIX, wrappers::type_name(id)));
37    let usemetanames = usenames
38        .iter()
39        .map(|id| format_ident!("get_{}_metadata", id))
40        .collect::<Vec<Ident>>();
41
42    TokenStream::from(quote! {
43        #[no_mangle]
44        #[allow(non_snake_case)]
45        pub fn #module_metadata_name() -> extendr_api::metadata::Metadata {
46            let mut functions = Vec::new();
47            let mut impls = Vec::new();
48
49            // Pushes metadata (eg. extendr_api::metadata::Func) to functions and impl vectors.
50            #( #fnmetanames(&mut functions); )*
51            #( #implmetanames(&mut impls); )*
52
53            // Extends functions and impls with the submodules metadata
54            #( functions.extend(#usenames::#usemetanames().functions); )*
55            #( impls.extend(#usenames::#usemetanames().impls); )*
56
57            // Add this function to the list, but set hidden: true.
58            functions.push(extendr_api::metadata::Func {
59                doc: "Metadata access function.",
60                rust_name: #module_metadata_name_string,
61                mod_name: #module_metadata_name_string,
62                r_name: #module_metadata_name_string,
63                c_name: #wrap_module_metadata_name_str,
64                args: Vec::new(),
65                return_type: "Metadata",
66                func_ptr: #wrap_module_metadata_name as * const u8,
67                hidden: true,
68                invisible: None,
69            });
70            let mut args = vec![
71                extendr_api::metadata::Arg { name: "use_symbols", arg_type: "bool", default: None },
72                extendr_api::metadata::Arg { name: "package_name", arg_type: "&str", default: None }
73            ];
74            let args = args;
75
76            // Add this function to the list, but set hidden: true.
77            functions.push(extendr_api::metadata::Func {
78                doc: "Wrapper generator.",
79                rust_name: #make_module_wrappers_name_string,
80                mod_name: #make_module_wrappers_name_string,
81                r_name: #make_module_wrappers_name_string,
82                c_name: #wrap_make_module_wrappers_string,
83                args,
84                return_type: "String",
85                func_ptr: #wrap_make_module_wrappers as * const u8,
86                hidden: true,
87                invisible: None,
88            });
89
90            extendr_api::metadata::Metadata {
91                name: #modname_string,
92                functions,
93                impls,
94            }
95        }
96
97        #[no_mangle]
98        #[allow(non_snake_case)]
99        pub extern "C" fn #wrap_module_metadata_name() -> extendr_api::SEXP {
100            use extendr_api::GetSexp;
101            unsafe { extendr_api::Robj::from(#module_metadata_name()).get() }
102        }
103
104        #[no_mangle]
105        #[allow(non_snake_case, clippy::not_unsafe_ptr_arg_deref)]
106        pub extern "C" fn #wrap_make_module_wrappers(
107            use_symbols_sexp: extendr_api::SEXP,
108            package_name_sexp: extendr_api::SEXP,
109        ) -> extendr_api::SEXP {
110            unsafe {
111                use extendr_api::robj::*;
112                use extendr_api::GetSexp;
113                let robj = Robj::from_sexp(use_symbols_sexp);
114                let use_symbols: bool = <bool>::try_from(&robj).unwrap();
115
116                let robj = Robj::from_sexp(package_name_sexp);
117                let package_name: &str = <&str>::try_from(&robj).unwrap();
118
119                extendr_api::Robj::from(
120                    #module_metadata_name()
121                        .make_r_wrappers(
122                            use_symbols,
123                            package_name,
124                        ).unwrap()
125                ).get()
126            }
127        }
128
129        #[no_mangle]
130        #[allow(non_snake_case, clippy::not_unsafe_ptr_arg_deref)]
131        pub extern "C" fn #write_make_module_wrappers(
132            package_name: *const std::os::raw::c_char,
133            out_path: *const std::os::raw::c_char,
134        ) -> i32 {
135            let pkg = match unsafe { std::ffi::CStr::from_ptr(package_name) }.to_str() {
136                Ok(s) => s,
137                Err(e) => {
138                    eprintln!("extendr: package_name is not valid UTF-8: {}", e);
139                    return 2;
140                }
141            };
142            let path = match unsafe { std::ffi::CStr::from_ptr(out_path) }.to_str() {
143                Ok(s) => s,
144                Err(e) => {
145                    eprintln!("extendr: out_path is not valid UTF-8: {}", e);
146                    return 2;
147                }
148            };
149            match #module_metadata_name().write_r_wrappers(pkg, path) {
150                Ok(()) => 0,
151                Err(e) => {
152                    eprintln!("extendr: writing wrappers for '{pkg}' failed: {e}");
153                    1
154                }
155            }
156        }
157
158        #[no_mangle]
159        #[allow(non_snake_case, clippy::not_unsafe_ptr_arg_deref)]
160        pub extern "C" fn #module_init_name(info: * mut extendr_api::DllInfo) {
161            unsafe { extendr_api::register_call_methods(info, #module_metadata_name()) };
162        }
163    })
164}
165
166#[derive(Debug)]
167struct Module {
168    modname: Option<Ident>,
169    fnnames: Vec<Ident>,
170    implnames: Vec<Type>,
171    usenames: Vec<Ident>,
172}
173
174// Custom parser for the module.
175impl syn::parse::Parse for Module {
176    fn parse(input: ParseStream) -> syn::Result<Self> {
177        use syn::spanned::Spanned;
178        let mut res = Self {
179            modname: None,
180            fnnames: Vec::new(),
181            implnames: Vec::new(),
182            usenames: Vec::new(),
183        };
184        while !input.is_empty() {
185            if let Ok(kmod) = input.parse::<Token![mod]>() {
186                let name: Ident = input.parse()?;
187                if res.modname.is_some() {
188                    return Err(syn::Error::new(kmod.span(), "only one mod allowed"));
189                }
190                res.modname = Some(name);
191            } else if input.parse::<Token![fn]>().is_ok() {
192                res.fnnames.push(input.parse()?);
193            } else if input.parse::<Token![impl]>().is_ok() {
194                res.implnames.push(input.parse()?);
195            } else if input.parse::<Token![use]>().is_ok() {
196                res.usenames.push(input.parse()?);
197            } else {
198                return Err(syn::Error::new(input.span(), "expected mod, fn or impl"));
199            }
200
201            input.parse::<Token![;]>()?;
202        }
203        if res.modname.is_none() {
204            return Err(syn::Error::new(input.span(), "expected one 'mod name'"));
205        }
206        Ok(res)
207    }
208}