Skip to main content

xacto_derive/
lib.rs

1use proc_macro::TokenStream;
2use quote::{format_ident, quote};
3use syn::{
4    Data, DataEnum, DeriveInput, Error, Ident, Type, TypePath, parse_macro_input, spanned::Spanned,
5};
6
7struct MessageVariant {
8    ident: Ident,
9    variant: Ident,
10    request_fields: Vec<Type>,
11    response_type: Option<Type>,
12}
13
14impl MessageVariant {
15    fn get_request_variant(&self) -> proc_macro2::TokenStream {
16        let ident = &self.variant;
17        if self.request_fields.is_empty() {
18            quote! { #ident }
19        } else {
20            let fields = &self.request_fields;
21            quote! { #ident ( #(#fields),* ) }
22        }
23    }
24
25    fn get_response_variant(&self) -> Option<proc_macro2::TokenStream> {
26        if let Some(response_type) = &self.response_type {
27            let ident = &self.variant;
28            Some(quote! { #ident ( #response_type ) })
29        } else {
30            None
31        }
32    }
33
34    fn get_original_arm(&self) -> proc_macro2::TokenStream {
35        let o_ident = &self.ident;
36        let v_ident = &self.variant;
37        let mut fields = vec![];
38        for i in 0..self.request_fields.len() {
39            let field_ident = format_ident!("a{}", i);
40            fields.push(quote! { #field_ident });
41        }
42        if self.response_type.is_some() {
43            let reply = quote! { reply };
44            fields.push(reply);
45        }
46        if fields.is_empty() {
47            quote! { #o_ident::#v_ident }
48        } else {
49            quote! { #o_ident::#v_ident ( #(#fields),* ) }
50        }
51    }
52
53    fn get_request_arm(&self) -> proc_macro2::TokenStream {
54        let v_ident = &self.variant;
55        let mut fields = vec![];
56        for i in 0..self.request_fields.len() {
57            let field_ident = format_ident!("a{}", i);
58            fields.push(quote! { #field_ident });
59        }
60        if fields.is_empty() {
61            quote! { Self::Request::#v_ident }
62        } else {
63            quote! { Self::Request::#v_ident ( #(#fields),* ) }
64        }
65    }
66
67    fn get_response_arm(&self) -> proc_macro2::TokenStream {
68        let v_ident = &self.variant;
69        quote! { Self::Response::#v_ident ( response ) }
70    }
71}
72
73#[proc_macro_derive(RpcMessage)]
74pub fn rpc_message(input: TokenStream) -> TokenStream {
75    let input = parse_macro_input!(input as DeriveInput);
76    match parse_rpc_message(input) {
77        Ok(output) => output.into(),
78        Err(error) => error.to_compile_error().into(),
79    }
80}
81
82fn parse_rpc_message(input: DeriveInput) -> Result<TokenStream, Error> {
83    let Data::Enum(DataEnum { variants, .. }) = &input.data else {
84        return Err(Error::new(
85            input.span(),
86            "RpcMessage can only be derived for enums",
87        ));
88    };
89
90    let mut new_variants = vec![];
91
92    for v in variants {
93        match &v.fields {
94            syn::Fields::Named(fields) => {
95                return Err(Error::new(fields.span(), "Named fields are not supported"));
96            }
97            syn::Fields::Unnamed(fields) => {
98                let mut mv = MessageVariant {
99                    ident: input.ident.clone(),
100                    variant: v.ident.clone(),
101                    request_fields: vec![],
102                    response_type: None,
103                };
104
105                for field in &fields.unnamed {
106                    if is_reply_type(&field.ty) {
107                        if mv.response_type.is_some() {
108                            return Err(Error::new(
109                                field.span(),
110                                "Only one reply type is allowed per variant",
111                            ));
112                        }
113                        let inner_ty = get_inner_reply_type(&field.ty)?;
114                        mv.response_type = Some(inner_ty);
115                    } else {
116                        mv.request_fields.push(field.ty.clone());
117                    }
118                }
119
120                new_variants.push(mv);
121            }
122            syn::Fields::Unit => {
123                new_variants.push(MessageVariant {
124                    ident: input.ident.clone(),
125                    variant: v.ident.clone(),
126                    request_fields: vec![],
127                    response_type: None,
128                });
129            }
130        }
131    }
132
133    let ident = input.ident.clone();
134    let request_ident = format_ident!("{}Request", ident);
135    let response_ident = format_ident!("{}Response", ident);
136
137    let request_variants = new_variants
138        .iter()
139        .map(|mv| mv.get_request_variant())
140        .collect::<Vec<_>>();
141
142    let response_variants = new_variants
143        .iter()
144        .filter_map(|mv| mv.get_response_variant())
145        .collect::<Vec<_>>();
146
147    let request_enum = quote! {
148        #[derive(Debug,Clone, Serialize, Deserialize)]
149        pub enum #request_ident {
150            #(#request_variants),*
151        }
152    };
153
154    let response_enum = if !response_variants.is_empty() {
155        Some(quote! {
156            #[derive(Debug, Clone, Serialize, Deserialize)]
157            pub enum #response_ident {
158                #(#response_variants),*
159            }
160        })
161    } else {
162        None
163    };
164
165    let mut into_request_arms = vec![];
166    let mut proxy_request_arms = vec![];
167    let mut proxy_response_arms = vec![];
168
169    for mv in &new_variants {
170        let original_arm = mv.get_original_arm();
171        let request_arm = mv.get_request_arm();
172        let response_arm = mv.get_response_arm();
173
174        if mv.response_type.is_none() {
175            into_request_arms.push(quote! {
176                #original_arm => RpcEnvelope {
177                    id: 0,
178                    payload: #request_arm,
179                }
180            });
181
182            proxy_request_arms.push(quote! {
183                #request_arm => {
184                    let msg = #original_arm;
185                    let (msg, act) = f(msg).ok_or(())?;
186                    act.cast(msg).await.map_err(|_| ())?;
187                    Ok(None)
188                }
189            });
190        } else {
191            into_request_arms.push(quote! {
192                #original_arm => {
193                    let id = replies.insert_reply(reply);
194                    RpcEnvelope {
195                        id,
196                        payload: #request_arm,
197                    }
198                }
199            });
200
201            proxy_request_arms.push(quote! {
202                #request_arm => {
203                    let (tx, rx) = oneshot::channel();
204                    let reply = Reply::new(tx);
205                    let msg = #original_arm;
206                    let (msg, act) = f(msg).ok_or(())?;
207                    act.cast(msg).await.map_err(|_| ())?;
208                    let response = rx.await.unwrap();
209                    let env = RpcEnvelope {
210                        id: env.id,
211                        payload: #response_arm,
212                    };
213                    Ok(Some(env))
214                }
215            });
216
217            proxy_response_arms.push(quote! {
218                #response_arm => {
219                    let reply = replies.get_reply(env.id).ok_or(())?;
220                    reply.send(response).map_err(|_| ())?;
221                    Ok(())
222                }
223            });
224        }
225    }
226
227    let response_assoc_type = if response_variants.is_empty() {
228        quote! { () }
229    } else {
230        quote! { #response_ident}
231    };
232
233    let proxy_response_impl = if response_variants.is_empty() {
234        quote! {
235            unreachable!()
236        }
237    } else {
238        quote! {
239            match env.payload {
240                #(#proxy_response_arms),*
241            }
242        }
243    };
244
245    let rpc_message_impl = quote! {
246        impl RpcMessage for #ident {
247            type Request = #request_ident;
248            type Response = #response_assoc_type;
249
250            fn into_request(self, replies: &mut ReplyMap) -> RpcEnvelope<Self::Request> {
251                match self {
252                    #(#into_request_arms),*
253                }
254            }
255
256            async fn proxy_request<F>(
257                env: RpcEnvelope<Self::Request>,
258                f: F,
259            ) -> Result<Option<RpcEnvelope<Self::Response>>, ()>
260            where
261                F: FnOnce(Self) -> Option<(Self, Act<Self>)>,
262                Self: Sized,
263            {
264                match env.payload {
265                    #(#proxy_request_arms),*
266                }
267            }
268
269            async fn proxy_response(
270                env: RpcEnvelope<Self::Response>,
271                replies: &mut ReplyMap,
272            ) -> Result<(), ()> {
273                #proxy_response_impl
274            }
275        }
276    };
277
278    let mut out = request_enum;
279
280    if let Some(response_enum) = response_enum {
281        out.extend(response_enum);
282    }
283
284    out.extend(rpc_message_impl);
285
286    Ok(out.into())
287}
288
289fn is_reply_type(ty: &Type) -> bool {
290    match ty {
291        Type::Path(TypePath { qself: None, path }) => path
292            .segments
293            .last()
294            .map_or(false, |seg| seg.ident == "Reply"),
295        _ => false,
296    }
297}
298
299fn get_inner_reply_type(ty: &Type) -> Result<Type, Error> {
300    match ty {
301        Type::Path(TypePath { qself: None, path }) => {
302            if let Some(segment) = path.segments.last() {
303                if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
304                    if let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() {
305                        return Ok(inner_ty.clone());
306                    }
307                }
308            }
309        }
310        _ => {}
311    };
312    Err(Error::new(ty.span(), "Expected Reply<T> type"))
313}