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}