Skip to main content

web_rpc_macro/
lib.rs

1use std::collections::HashSet;
2
3use proc_macro::TokenStream;
4use proc_macro2::TokenStream as TokenStream2;
5use quote::{format_ident, quote, quote_spanned, ToTokens};
6use syn::{
7    braced,
8    ext::IdentExt,
9    parenthesized,
10    parse::{Parse, ParseStream},
11    parse_macro_input, parse_quote,
12    punctuated::Punctuated,
13    spanned::Spanned,
14    token::Comma,
15    Attribute, FnArg, Ident, Lifetime, Pat, PatType, ReturnType, Token, Type, Visibility,
16};
17
18macro_rules! extend_errors {
19    ($errors: ident, $e: expr) => {
20        match $errors {
21            Ok(_) => $errors = Err($e),
22            Err(ref mut errors) => errors.extend($e),
23        }
24    };
25}
26
27/// If `ty` is `impl Stream<Item = T>`, returns Some(T).
28fn stream_item_type(ty: &Type) -> Option<&Type> {
29    if let Type::ImplTrait(impl_trait) = ty {
30        for bound in &impl_trait.bounds {
31            if let syn::TypeParamBound::Trait(trait_bound) = bound {
32                let last_segment = trait_bound.path.segments.last()?;
33                if last_segment.ident == "Stream" {
34                    if let syn::PathArguments::AngleBracketed(args) = &last_segment.arguments {
35                        for arg in &args.args {
36                            if let syn::GenericArgument::Binding(binding) = arg {
37                                if binding.ident == "Item" {
38                                    return Some(&binding.ty);
39                                }
40                            }
41                        }
42                    }
43                }
44            }
45        }
46    }
47    None
48}
49
50/// If `ty` is `Option<T>`, returns Some(T).
51fn option_inner_type(ty: &Type) -> Option<&Type> {
52    if let Type::Path(type_path) = ty {
53        let last_seg = type_path.path.segments.last()?;
54        if last_seg.ident == "Option" {
55            if let syn::PathArguments::AngleBracketed(args) = &last_seg.arguments {
56                if args.args.len() == 1 {
57                    if let syn::GenericArgument::Type(inner) = &args.args[0] {
58                        return Some(inner);
59                    }
60                }
61            }
62        }
63    }
64    None
65}
66
67/// If `ty` is `Result<T, E>`, returns Some((T, E)).
68fn result_inner_types(ty: &Type) -> Option<(&Type, &Type)> {
69    if let Type::Path(type_path) = ty {
70        let last_seg = type_path.path.segments.last()?;
71        if last_seg.ident == "Result" {
72            if let syn::PathArguments::AngleBracketed(args) = &last_seg.arguments {
73                if args.args.len() == 2 {
74                    if let (syn::GenericArgument::Type(ok_ty), syn::GenericArgument::Type(err_ty)) =
75                        (&args.args[0], &args.args[1])
76                    {
77                        return Some((ok_ty, err_ty));
78                    }
79                }
80            }
81        }
82    }
83    None
84}
85
86/// True if `ty` is `&str` or `&[u8]` — the only reference shapes we route through
87/// the existing serde-borrowing path (zero-copy, with an `'a` lifetime injected
88/// into the request enum). Any other reference shape goes through the JS path.
89fn is_borrowed_serde_ref(ty: &Type) -> bool {
90    if let Type::Reference(r) = ty {
91        match &*r.elem {
92            Type::Path(p) if p.path.is_ident("str") => return true,
93            Type::Slice(s) => {
94                if let Type::Path(p) = &*s.elem {
95                    if p.path.is_ident("u8") {
96                        return true;
97                    }
98                }
99            }
100            _ => {}
101        }
102    }
103    false
104}
105
106/// True if `ty` is a reference to a presumed JS type (anything other than
107/// `&str`/`&[u8]`). The receiver side decodes these via `JsCast::dyn_ref`.
108fn is_js_ref(ty: &Type) -> bool {
109    matches!(ty, Type::Reference(_)) && !is_borrowed_serde_ref(ty)
110}
111
112/// True if `attr` is a cfg-style attribute (`#[cfg(...)]` or `#[cfg_attr(...)]`).
113/// These are propagated onto every generated artifact derived from a method so
114/// that rustc strips them in lockstep after macro expansion.
115fn is_cfg_attr(attr: &Attribute) -> bool {
116    attr.path.is_ident("cfg") || attr.path.is_ident("cfg_attr")
117}
118
119/// Recursively emit code that encodes a value of type `ty` into a `WireArg`,
120/// pushing JS values onto `post` as a side-effect.
121///
122/// Caller supplies `value` as a token-tree expression (typically an ident binding).
123/// The emitted code matches structure on `Option`/`Result` and recurses; bare
124/// leaves dispatch through the autoref encoder traits.
125///
126/// Reference-to-JS types should be handled by the caller before invoking this
127/// helper — they cannot be encoded as nested elements (no `Decoder<&T>` impl on
128/// the receiver side).
129fn emit_encode(ty: &Type, value: TokenStream2, post: &TokenStream2) -> TokenStream2 {
130    // Match against `&#value` so the original binding remains accessible to any
131    // transfer-side code emitted alongside the encoder. Match ergonomics binds
132    // `__inner` as a reference inside each arm.
133    if let Some(inner) = option_inner_type(ty) {
134        let inner_enc = emit_encode(inner, quote!(__inner), post);
135        quote_spanned! {ty.span()=>
136            match &#value {
137                ::core::option::Option::Some(__inner) =>
138                    web_rpc::codec::WireArg::Some(std::boxed::Box::new(#inner_enc)),
139                ::core::option::Option::None =>
140                    web_rpc::codec::WireArg::None,
141            }
142        }
143    } else if let Some((ok, err)) = result_inner_types(ty) {
144        let ok_enc = emit_encode(ok, quote!(__inner), post);
145        let err_enc = emit_encode(err, quote!(__inner), post);
146        quote_spanned! {ty.span()=>
147            match &#value {
148                ::core::result::Result::Ok(__inner) =>
149                    web_rpc::codec::WireArg::Ok(std::boxed::Box::new(#ok_enc)),
150                ::core::result::Result::Err(__inner) =>
151                    web_rpc::codec::WireArg::Err(std::boxed::Box::new(#err_enc)),
152            }
153        }
154    } else {
155        quote_spanned! {ty.span()=>
156            {
157                #[allow(unused_imports)]
158                use web_rpc::codec::{
159                    __RpcJsEncode as _,
160                    __RpcSerialEncode as _,
161                };
162                (&#value).__rpc_encode(#post)
163            }
164        }
165    }
166}
167
168/// Recursively emit code that decodes a `WireArg` of type `ty` into a Rust value,
169/// shifting JS values off `post` as needed.
170///
171/// Caller supplies `wire` as a token-tree expression evaluating to a `WireArg`.
172/// Reference-to-JS types should be handled by the caller — see `emit_encode`.
173fn emit_decode(ty: &Type, wire: TokenStream2, post: &TokenStream2) -> TokenStream2 {
174    if let Some(inner) = option_inner_type(ty) {
175        let inner_dec = emit_decode(inner, quote!(*__inner), post);
176        quote_spanned! {ty.span()=>
177            match #wire {
178                web_rpc::codec::WireArg::Some(__inner) =>
179                    ::core::option::Option::Some(#inner_dec),
180                web_rpc::codec::WireArg::None =>
181                    ::core::option::Option::None,
182                _ => panic!("web_rpc: wire/type mismatch — expected Some or None"),
183            }
184        }
185    } else if let Some((ok, err)) = result_inner_types(ty) {
186        let ok_dec = emit_decode(ok, quote!(*__inner), post);
187        let err_dec = emit_decode(err, quote!(*__inner), post);
188        quote_spanned! {ty.span()=>
189            match #wire {
190                web_rpc::codec::WireArg::Ok(__inner) =>
191                    ::core::result::Result::Ok(#ok_dec),
192                web_rpc::codec::WireArg::Err(__inner) =>
193                    ::core::result::Result::Err(#err_dec),
194                _ => panic!("web_rpc: wire/type mismatch — expected Ok or Err"),
195            }
196        }
197    } else {
198        quote_spanned! {ty.span()=>
199            {
200                #[allow(unused_imports)]
201                use web_rpc::codec::{
202                    __RpcJsDecode as _,
203                    __RpcSerialDecode as _,
204                };
205                (&web_rpc::codec::Decoder::<#ty>::default()).__rpc_decode(#wire, #post)
206            }
207        }
208    }
209}
210
211struct Service {
212    attrs: Vec<Attribute>,
213    vis: Visibility,
214    ident: Ident,
215    rpcs: Vec<RpcMethod>,
216}
217
218struct RpcMethod {
219    is_async: Option<Token![async]>,
220    attrs: Vec<Attribute>,
221    receiver: syn::Receiver,
222    ident: Ident,
223    args: Vec<PatType>,
224    transfer: Vec<TransferClause>,
225    output: ReturnType,
226}
227
228/// One entry inside a `#[transfer(...)]` attribute.
229#[allow(dead_code)]
230enum TransferClause {
231    /// `name` — push the parameter itself, unconditionally.
232    BareParam(Ident),
233    /// `data => data.buffer()` — push the expression's result, unconditionally.
234    ParamExpr { name: Ident, body: syn::Expr },
235    /// `data => |Some(d)| d.buffer()` (closure) or
236    /// `data => match { Some(d) => d.buffer(), ... }` (match-block).
237    /// Each `Gate` becomes one `if let pat = &name { __transfer.push(body) }`.
238    ParamGated { name: Ident, gates: Vec<Gate> },
239    /// `return` — push the response value itself, unconditionally.
240    BareReturn,
241    /// `return => |Ok(o)| o.buffer()` or `return => match { ... }`.
242    ReturnGated { gates: Vec<Gate> },
243}
244
245#[allow(dead_code)]
246struct Gate {
247    pat: syn::Pat,
248    body: syn::Expr,
249}
250
251struct ServiceGenerator<'a> {
252    trait_ident: &'a Ident,
253    service_ident: &'a Ident,
254    client_ident: &'a Ident,
255    request_ident: &'a Ident,
256    response_ident: &'a Ident,
257    vis: &'a Visibility,
258    attrs: &'a [Attribute],
259    rpcs: &'a [RpcMethod],
260    camel_case_idents: &'a [Ident],
261    has_streaming_methods: bool,
262}
263
264impl<'a> ServiceGenerator<'a> {
265    fn enum_request(&self) -> TokenStream2 {
266        let &Self {
267            vis,
268            request_ident,
269            camel_case_idents,
270            rpcs,
271            ..
272        } = self;
273        let variants = rpcs.iter().zip(camel_case_idents.iter()).map(
274            |(RpcMethod { attrs, args, .. }, camel_case_ident)| {
275                let cfg_attrs = attrs.iter().filter(|a| is_cfg_attr(a));
276                let fields = args.iter().map(|arg| {
277                    let pat = &arg.pat;
278                    if is_borrowed_serde_ref(&arg.ty) {
279                        // `&str` / `&[u8]` — keep zero-copy serde borrowing path.
280                        let mut type_ref = match &*arg.ty {
281                            Type::Reference(r) => r.clone(),
282                            _ => unreachable!("is_borrowed_serde_ref guarantees a reference"),
283                        };
284                        type_ref.lifetime =
285                            Some(Lifetime::new("'a", type_ref.and_token.span()));
286                        quote_spanned! {arg.ty.span()=> #pat: #type_ref }
287                    } else {
288                        // Everything else (including `&JsT`) uses the universal
289                        // recursive WireArg representation.
290                        quote_spanned! {arg.ty.span()=>
291                            #pat: web_rpc::codec::WireArg
292                        }
293                    }
294                });
295                quote! {
296                    #(#cfg_attrs)*
297                    #camel_case_ident { #( #fields ),* }
298                }
299            },
300        );
301        // `<'a>` is always emitted, with a hidden variant that uses it via
302        // `PhantomData`. This keeps the enum well-formed regardless of which
303        // methods rustc strips via cfg-evaluation after macro expansion: a
304        // service whose only borrowing methods get cfg'd out would otherwise
305        // hit E0392. The macro never constructs this variant on the wire; the
306        // server's match arm panics if it ever appears.
307        quote! {
308            #[derive(web_rpc::serde::Serialize, web_rpc::serde::Deserialize)]
309            #vis enum #request_ident<'a> {
310                #( #variants, )*
311                #[doc(hidden)]
312                __WebRpcPhantom(std::marker::PhantomData<&'a ()>),
313            }
314        }
315    }
316
317    fn enum_response(&self) -> TokenStream2 {
318        let &Self {
319            vis,
320            response_ident,
321            camel_case_idents,
322            rpcs,
323            ..
324        } = self;
325        let variants = rpcs.iter().zip(camel_case_idents.iter()).map(
326            |(RpcMethod { attrs, .. }, camel_case_ident)| {
327                let cfg_attrs = attrs.iter().filter(|a| is_cfg_attr(a));
328                // Every method's response variant carries a single uniform
329                // `WireArg`. Notification methods (no return) still get a
330                // variant — the macro fills it with a placeholder that the
331                // client never reads.
332                quote! {
333                    #(#cfg_attrs)*
334                    #camel_case_ident ( web_rpc::codec::WireArg )
335                }
336            },
337        );
338        quote! {
339            #[derive(web_rpc::serde::Serialize, web_rpc::serde::Deserialize)]
340            #vis enum #response_ident {
341                #( #variants ),*
342            }
343        }
344    }
345
346    fn trait_service(&self) -> TokenStream2 {
347        let &Self {
348            attrs,
349            rpcs,
350            vis,
351            trait_ident,
352            ..
353        } = self;
354
355        let unit_type: &Type = &parse_quote!(());
356        let rpc_fns = rpcs.iter().map(
357            |RpcMethod {
358                 attrs,
359                 args,
360                 receiver,
361                 ident,
362                 is_async,
363                 output,
364                 ..
365             }| {
366                if let ReturnType::Type(_, ref ty) = output {
367                    if let Some(item_ty) = stream_item_type(ty) {
368                        return quote_spanned! {ident.span()=>
369                            #( #attrs )*
370                            #is_async fn #ident(#receiver, #( #args ),*) -> impl web_rpc::futures_core::Stream<Item = #item_ty>;
371                        };
372                    }
373                }
374                let output = match output {
375                    ReturnType::Type(_, ref ty) => ty,
376                    ReturnType::Default => unit_type,
377                };
378                quote_spanned! {ident.span()=>
379                    #( #attrs )*
380                    #is_async fn #ident(#receiver, #( #args ),*) -> #output;
381                }
382            },
383        );
384
385        let forward_fns = rpcs
386            .iter()
387            .map(
388                |RpcMethod {
389                     attrs,
390                     args,
391                     receiver,
392                     ident,
393                     is_async,
394                     output,
395                     ..
396                 }| {
397                    {
398                        let output = if let ReturnType::Type(_, ref ty) = output {
399                            if let Some(item_ty) = stream_item_type(ty) {
400                                quote! { impl web_rpc::futures_core::Stream<Item = #item_ty> }
401                            } else {
402                                let ty: &Type = ty;
403                                quote! { #ty }
404                            }
405                        } else {
406                            let ty = unit_type;
407                            quote! { #ty }
408                        };
409                        let do_await = match is_async {
410                            Some(token) => quote_spanned!(token.span=> .await),
411                            None => quote!(),
412                        };
413                        let forward_args = args.iter().filter_map(|arg| match &*arg.pat {
414                            Pat::Ident(ident) => Some(&ident.ident),
415                            _ => None,
416                        });
417                        quote_spanned! {ident.span()=>
418                            #( #attrs )*
419                            #is_async fn #ident(#receiver, #( #args ),*) -> #output {
420                                T::#ident(self, #( #forward_args ),*)#do_await
421                            }
422                        }
423                    }
424                },
425            )
426            .collect::<Vec<_>>();
427
428        quote! {
429            #( #attrs )*
430            #[allow(async_fn_in_trait)]
431            #vis trait #trait_ident {
432                #( #rpc_fns )*
433            }
434
435            impl<T> #trait_ident for std::sync::Arc<T> where T: #trait_ident {
436                #( #forward_fns )*
437            }
438            impl<T> #trait_ident for std::boxed::Box<T> where T: #trait_ident {
439                #( #forward_fns )*
440            }
441            impl<T> #trait_ident for std::rc::Rc<T> where T: #trait_ident {
442                #( #forward_fns )*
443            }
444        }
445    }
446
447    fn struct_client(&self) -> TokenStream2 {
448        let &Self {
449            vis,
450            client_ident,
451            request_ident,
452            response_ident,
453            camel_case_idents,
454            rpcs,
455            has_streaming_methods,
456            ..
457        } = self;
458
459        let rpc_fns = rpcs
460            .iter()
461            .zip(camel_case_idents.iter())
462            .map(|(RpcMethod { attrs, args, transfer, ident, output, .. }, camel_case_ident)| {
463                // 1. Per-arg encoding: borrowed `&str`/`&[u8]` pass through inline;
464                // everything else routes through the autoref-dispatched `__rpc_encode`.
465                let mut arg_encodings = Vec::<TokenStream2>::new();
466                let mut request_struct_fields = Vec::<TokenStream2>::new();
467                for arg in args {
468                    let id = match &*arg.pat {
469                        Pat::Ident(p) => &p.ident,
470                        _ => continue,
471                    };
472                    if is_borrowed_serde_ref(&arg.ty) {
473                        request_struct_fields.push(quote! { #id });
474                    } else {
475                        let wire_ident = format_ident!("__wire_{}", id);
476                        let post = quote!(&__post);
477                        let enc = emit_encode(&arg.ty, quote!(#id), &post);
478                        arg_encodings.push(quote! { let #wire_ident = #enc; });
479                        request_struct_fields.push(quote! { #id: #wire_ident });
480                    }
481                }
482
483                // 2. Per-method transfer pushes (param-side only; return-side
484                // clauses are handled in struct_server).
485                let transfer_pushes = transfer.iter().filter_map(|c| match c {
486                    TransferClause::BareParam(name) => Some(quote! {
487                        __transfer.push(#name.as_ref());
488                    }),
489                    TransferClause::ParamExpr { name, body } => Some(quote_spanned! {body.span()=>
490                        {
491                            let _ = &#name; // ensure name is referenced
492                            __transfer.push((#body).as_ref());
493                        }
494                    }),
495                    TransferClause::ParamGated { name, gates } => {
496                        let arms = gates.iter().map(|g| {
497                            let pat = &g.pat;
498                            let body = &g.body;
499                            quote_spanned! {body.span()=>
500                                if let #pat = &#name {
501                                    __transfer.push((#body).as_ref());
502                                }
503                            }
504                        });
505                        Some(quote! { #( #arms )* })
506                    }
507                    TransferClause::BareReturn | TransferClause::ReturnGated { .. } => None,
508                });
509
510                let send_request = quote! {
511                    let __seq_id = self.seq_id.replace_with(|seq_id| seq_id.wrapping_add(1));
512                    let __post = web_rpc::js_sys::Array::new();
513                    let __transfer = web_rpc::js_sys::Array::new();
514                    #( #arg_encodings )*
515                    let __request = #request_ident::#camel_case_ident {
516                        #( #request_struct_fields ),*
517                    };
518                    let __header = web_rpc::MessageHeader::Request(__seq_id);
519                    let __header_bytes = web_rpc::bincode::serialize(&__header).unwrap();
520                    let __header_buffer = web_rpc::js_sys::Uint8Array::from(&__header_bytes[..]).buffer();
521                    let __payload_bytes = web_rpc::bincode::serialize(&__request).unwrap();
522                    let __payload_buffer = web_rpc::js_sys::Uint8Array::from(&__payload_bytes[..]).buffer();
523                    // Prepend [header, payload] in front of the encoded JS values.
524                    __post.unshift(&__payload_buffer);
525                    __post.unshift(&__header_buffer);
526                    __transfer.push(__header_buffer.as_ref());
527                    __transfer.push(__payload_buffer.as_ref());
528                    #( #transfer_pushes )*
529                    self.port.post_message(&__post, &__transfer).unwrap();
530                };
531
532                let is_streaming = matches!(
533                    output,
534                    ReturnType::Type(_, ref ty) if stream_item_type(ty).is_some()
535                );
536
537                if is_streaming {
538                    let item_ty = match output {
539                        ReturnType::Type(_, ref ty) => stream_item_type(ty).unwrap(),
540                        _ => unreachable!(),
541                    };
542                    let dec = emit_decode(item_ty, quote!(__wire), &quote!(&__post_array));
543
544                    let unpack_stream_item = quote! {
545                        |(__response, __post_array): (#response_ident, web_rpc::js_sys::Array)| {
546                            let #response_ident::#camel_case_ident(__wire) = __response else {
547                                panic!("web_rpc: received incorrect response variant")
548                            };
549                            #dec
550                        }
551                    };
552
553                    quote! {
554                        #( #attrs )*
555                        #vis fn #ident(
556                            &self,
557                            #( #args ),*
558                        ) -> web_rpc::client::StreamReceiver<#item_ty> {
559                            #send_request
560                            let (__item_tx, __item_rx) = web_rpc::futures_channel::mpsc::unbounded();
561                            self.stream_callback_map.borrow_mut().insert(__seq_id, __item_tx);
562                            let __mapped_rx = web_rpc::futures_util::StreamExt::map(
563                                __item_rx,
564                                #unpack_stream_item
565                            );
566                            let __abort_sender = self.abort_sender.clone();
567                            let __stream_callback_map = self.stream_callback_map.clone();
568                            let __dispatcher = self.dispatcher.clone();
569                            web_rpc::client::StreamReceiver::new(
570                                __mapped_rx,
571                                __dispatcher,
572                                std::boxed::Box::new(move || {
573                                    __stream_callback_map.borrow_mut().remove(&__seq_id);
574                                    (__abort_sender)(__seq_id);
575                                }),
576                            )
577                        }
578                    }
579                } else {
580                    let return_type = match output {
581                        ReturnType::Type(_, ref ty) => quote! {
582                            web_rpc::client::RequestFuture<#ty>
583                        },
584                        _ => quote!(()),
585                    };
586                    let maybe_register_callback = match output {
587                        ReturnType::Type(_, _) => quote! {
588                            let (__response_tx, __response_rx) =
589                                web_rpc::futures_channel::oneshot::channel();
590                            self.callback_map.borrow_mut().insert(__seq_id, __response_tx);
591                        },
592                        _ => Default::default(),
593                    };
594
595                    let maybe_unpack_and_return_future = match output {
596                        ReturnType::Type(_, ref ret_ty) => {
597                            let dec = emit_decode(ret_ty, quote!(__wire), &quote!(&__post_array));
598                            quote! {
599                                let __response_future = web_rpc::futures_util::FutureExt::map(
600                                    __response_rx,
601                                    |response| {
602                                        let (__serialize_response, __post_array) = response.unwrap();
603                                        let #response_ident::#camel_case_ident(__wire) = __serialize_response else {
604                                            panic!("web_rpc: received incorrect response variant")
605                                        };
606                                        #dec
607                                    }
608                                );
609                                let __abort_sender = self.abort_sender.clone();
610                                let __dispatcher = self.dispatcher.clone();
611                                web_rpc::client::RequestFuture::new(
612                                    __response_future,
613                                    __dispatcher,
614                                    std::boxed::Box::new(move || (__abort_sender)(__seq_id)))
615                            }
616                        }
617                        _ => Default::default(),
618                    };
619
620                    quote! {
621                        #( #attrs )*
622                        #vis fn #ident(
623                            &self,
624                            #( #args ),*
625                        ) -> #return_type {
626                            #send_request
627                            #maybe_register_callback
628                            #maybe_unpack_and_return_future
629                        }
630                    }
631                }
632            });
633
634        let stream_callback_map_field = if has_streaming_methods {
635            // `#[allow(dead_code)]` covers the case where every streaming method
636            // is stripped via cfg — the field is still bound by the
637            // `From<Configuration>` impl but no surviving method reads it.
638            quote! {
639                #[allow(dead_code)]
640                stream_callback_map: std::rc::Rc<
641                    std::cell::RefCell<
642                        web_rpc::client::StreamCallbackMap<#response_ident>
643                    >
644                >,
645            }
646        } else {
647            quote!()
648        };
649
650        let stream_callback_map_pat = if has_streaming_methods {
651            quote! { stream_callback_map, }
652        } else {
653            quote! { _, }
654        };
655
656        let stream_callback_map_init = if has_streaming_methods {
657            quote! { stream_callback_map, }
658        } else {
659            quote! {}
660        };
661
662        quote! {
663            #[derive(core::clone::Clone)]
664            #vis struct #client_ident {
665                callback_map: std::rc::Rc<
666                    std::cell::RefCell<
667                        web_rpc::client::CallbackMap<#response_ident>
668                    >
669                >,
670                #stream_callback_map_field
671                port: web_rpc::port::Port,
672                listener: std::rc::Rc<web_rpc::gloo_events::EventListener>,
673                dispatcher: web_rpc::futures_util::future::Shared<
674                    web_rpc::futures_core::future::LocalBoxFuture<'static, ()>
675                >,
676                abort_sender: std::rc::Rc<dyn std::ops::Fn(usize)>,
677                seq_id: std::rc::Rc<std::cell::RefCell<usize>>
678            }
679            impl std::fmt::Debug for #client_ident {
680                fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
681                    formatter.debug_struct(std::stringify!(#client_ident))
682                        .finish()
683                }
684            }
685            impl web_rpc::client::Client for #client_ident {
686                type Response = #response_ident;
687            }
688            impl From<web_rpc::client::Configuration<#response_ident>>
689                for #client_ident {
690                fn from((callback_map, #stream_callback_map_pat port, listener, dispatcher, abort_sender):
691                    web_rpc::client::Configuration<#response_ident>) -> Self {
692                    Self {
693                        callback_map,
694                        #stream_callback_map_init
695                        port,
696                        listener,
697                        dispatcher,
698                        abort_sender,
699                        seq_id: std::default::Default::default()
700                    }
701                }
702            }
703            impl #client_ident {
704                #( #rpc_fns )*
705            }
706        }
707    }
708
709    fn struct_server(&self) -> TokenStream2 {
710        let &Self {
711            vis,
712            trait_ident,
713            service_ident,
714            request_ident,
715            response_ident,
716            camel_case_idents,
717            rpcs,
718            ..
719        } = self;
720
721        let request_type = quote! { #request_ident<'_> };
722
723        let handlers = rpcs.iter()
724            .zip(camel_case_idents.iter())
725            .map(|(RpcMethod { is_async, ident, args, transfer, output, attrs, .. }, camel_case_ident)| {
726                let cfg_attrs: Vec<_> = attrs.iter().filter(|a| is_cfg_attr(a)).collect();
727                // 1. Destructure pattern for the request enum variant.
728                // Borrowed args use their own ident; non-borrowed args bind to __wire_<id>.
729                let destructure_fields: Vec<_> = args.iter()
730                    .filter_map(|arg| {
731                        let id = match &*arg.pat {
732                            Pat::Ident(p) => &p.ident,
733                            _ => return None,
734                        };
735                        Some(if is_borrowed_serde_ref(&arg.ty) {
736                            quote! { #id }
737                        } else {
738                            let wire_ident = format_ident!("__wire_{}", id);
739                            quote! { #id: #wire_ident }
740                        })
741                    })
742                    .collect();
743
744                // 2. Per-arg decoding statements.
745                let arg_decodes: Vec<_> = args.iter()
746                    .filter_map(|arg| {
747                        let id = match &*arg.pat {
748                            Pat::Ident(p) => &p.ident,
749                            _ => return None,
750                        };
751                        if is_borrowed_serde_ref(&arg.ty) {
752                            // Already bound by destructuring.
753                            None
754                        } else if is_js_ref(&arg.ty) {
755                            // `&T` where T: JsCast — no `Decoder<&T>` impl, so we
756                            // shift from the post-array and bind via dyn_ref locally.
757                            let inner_ty = match &*arg.ty {
758                                Type::Reference(r) => &*r.elem,
759                                _ => unreachable!(),
760                            };
761                            let tmp_ident = format_ident!("__tmp_{}", id);
762                            let wire_ident = format_ident!("__wire_{}", id);
763                            let arg_ty = &arg.ty;
764                            Some(quote! {
765                                let #tmp_ident = match #wire_ident {
766                                    web_rpc::codec::WireArg::Js => __js_args.shift(),
767                                    _ => panic!("web_rpc: expected Js wire variant for reference arg"),
768                                };
769                                let #id: #arg_ty = web_rpc::wasm_bindgen::JsCast::dyn_ref::<#inner_ty>(&#tmp_ident)
770                                    .unwrap();
771                            })
772                        } else {
773                            let wire_ident = format_ident!("__wire_{}", id);
774                            let dec = emit_decode(&arg.ty, quote!(#wire_ident), &quote!(&__js_args));
775                            Some(quote! { let #id = #dec; })
776                        }
777                    })
778                    .collect();
779
780                let call_args: Vec<_> = args.iter().filter_map(|arg| match &*arg.pat {
781                    Pat::Ident(ident) => Some(&ident.ident),
782                    _ => None,
783                }).collect();
784
785                // Return-side transfer clauses (BareReturn / ReturnGated).
786                // The scrutinee is `__response` for non-streaming and `__item` for streaming.
787                let make_return_transfer = |scrutinee_ident: &Ident| -> TokenStream2 {
788                    let pushes = transfer.iter().filter_map(|c| match c {
789                        TransferClause::BareReturn => Some(quote! {
790                            __transfer.push(#scrutinee_ident.as_ref());
791                        }),
792                        TransferClause::ReturnGated { gates } => {
793                            let arms = gates.iter().map(|g| {
794                                let pat = &g.pat;
795                                let body = &g.body;
796                                quote_spanned! {body.span()=>
797                                    if let #pat = &#scrutinee_ident {
798                                        __transfer.push((#body).as_ref());
799                                    }
800                                }
801                            });
802                            Some(quote! { #( #arms )* })
803                        }
804                        _ => None,
805                    });
806                    quote! { #( #pushes )* }
807                };
808
809                let is_streaming = matches!(
810                    output,
811                    ReturnType::Type(_, ref ty) if stream_item_type(ty).is_some()
812                );
813
814                if is_streaming {
815                    let item_ty = match output {
816                        ReturnType::Type(_, ref ty) => stream_item_type(ty).unwrap(),
817                        _ => unreachable!(),
818                    };
819                    let item_enc = emit_encode(item_ty, quote!(__item), &quote!(&__post));
820                    let item_ident = Ident::new("__item", proc_macro2::Span::call_site());
821                    let return_transfer = make_return_transfer(&item_ident);
822
823                    let wrap_item = quote! {
824                        let __post = web_rpc::js_sys::Array::new();
825                        let __transfer = web_rpc::js_sys::Array::new();
826                        let __wire_item = #item_enc;
827                        #return_transfer
828                        let __response = #response_ident::#camel_case_ident(__wire_item);
829                    };
830
831                    let fwd_body = quote! {
832                        let __stream_tx_clone = __stream_tx.clone();
833                        web_rpc::pin_utils::pin_mut!(__user_rx);
834                        let __fwd = async move {
835                            while let Some(__item) = web_rpc::futures_util::StreamExt::next(&mut __user_rx).await {
836                                #wrap_item
837                                if __stream_tx_clone.unbounded_send((__seq_id, Some((__response, __post, __transfer)))).is_err() {
838                                    break;
839                                }
840                            }
841                        };
842                        let __fwd = web_rpc::futures_util::FutureExt::fuse(__fwd);
843                        web_rpc::pin_utils::pin_mut!(__fwd);
844                        web_rpc::futures_util::select! {
845                            _ = __abort_rx => {},
846                            _ = __fwd => {},
847                        }
848                        let _ = __stream_tx.unbounded_send((__seq_id, None));
849                        web_rpc::service::ExecuteResult::StreamComplete
850                    };
851
852                    match is_async {
853                        Some(_) => quote! {
854                            #( #cfg_attrs )*
855                            #request_ident::#camel_case_ident { #( #destructure_fields ),* } => {
856                                #( #arg_decodes )*
857                                let __get_rx = web_rpc::futures_util::FutureExt::fuse(
858                                    self.server_impl.#ident(#( #call_args ),*)
859                                );
860                                web_rpc::pin_utils::pin_mut!(__get_rx);
861                                let __maybe_rx = web_rpc::futures_util::select! {
862                                    _ = __abort_rx => None,
863                                    __rx = __get_rx => Some(__rx),
864                                };
865                                if let Some(mut __user_rx) = __maybe_rx {
866                                    #fwd_body
867                                } else {
868                                    let _ = __stream_tx.unbounded_send((__seq_id, None));
869                                    web_rpc::service::ExecuteResult::StreamComplete
870                                }
871                            }
872                        },
873                        None => quote! {
874                            #( #cfg_attrs )*
875                            #request_ident::#camel_case_ident { #( #destructure_fields ),* } => {
876                                #( #arg_decodes )*
877                                let mut __user_rx = self.server_impl.#ident(#( #call_args ),*);
878                                #fwd_body
879                            }
880                        },
881                    }
882                } else {
883                    // Non-streaming.
884                    let resp_ident = Ident::new("__response", proc_macro2::Span::call_site());
885                    let return_transfer = make_return_transfer(&resp_ident);
886                    let return_response = match output {
887                        ReturnType::Type(_, ref ret_ty) => {
888                            let enc = emit_encode(ret_ty, quote!(__response), &quote!(&__post));
889                            quote! {
890                                let __post = web_rpc::js_sys::Array::new();
891                                let __transfer = web_rpc::js_sys::Array::new();
892                                let __wire = #enc;
893                                #return_transfer
894                                (#response_ident::#camel_case_ident(__wire), __post, __transfer)
895                            }
896                        }
897                        _ => {
898                            // Notification — emit a placeholder WireArg.
899                            quote! {
900                                let _ = __response;
901                                let __post = web_rpc::js_sys::Array::new();
902                                let __transfer = web_rpc::js_sys::Array::new();
903                                let __wire = web_rpc::codec::WireArg::Bytes(
904                                    web_rpc::bincode::serialize(&()).unwrap()
905                                );
906                                (#response_ident::#camel_case_ident(__wire), __post, __transfer)
907                            }
908                        }
909                    };
910
911                    match is_async {
912                        Some(_) => quote! {
913                            #( #cfg_attrs )*
914                            #request_ident::#camel_case_ident { #( #destructure_fields ),* } => {
915                                #( #arg_decodes )*
916                                let __task =
917                                    web_rpc::futures_util::FutureExt::fuse(self.server_impl.#ident(#( #call_args ),*));
918                                web_rpc::pin_utils::pin_mut!(__task);
919                                web_rpc::service::ExecuteResult::Response(
920                                    web_rpc::futures_util::select! {
921                                        _ = __abort_rx => None,
922                                        __response = __task => Some({
923                                            #return_response
924                                        })
925                                    }
926                                )
927                            }
928                        },
929                        None => quote! {
930                            #( #cfg_attrs )*
931                            #request_ident::#camel_case_ident { #( #destructure_fields ),* } => {
932                                #( #arg_decodes )*
933                                let __response = self.server_impl.#ident(#( #call_args ),*);
934                                web_rpc::service::ExecuteResult::Response(
935                                    Some({
936                                        #return_response
937                                    })
938                                )
939                            }
940                        }
941                    }
942                }
943            });
944
945        quote! {
946            #vis struct #service_ident<T> {
947                server_impl: T
948            }
949            impl<T: #trait_ident> web_rpc::service::Service for #service_ident<T> {
950                type Response = #response_ident;
951                async fn execute(
952                    &self,
953                    __seq_id: usize,
954                    mut __abort_rx: web_rpc::futures_channel::oneshot::Receiver<()>,
955                    __payload: std::vec::Vec<u8>,
956                    __js_args: web_rpc::js_sys::Array,
957                    __stream_tx: web_rpc::futures_channel::mpsc::UnboundedSender<
958                        web_rpc::service::StreamMessage<Self::Response>
959                    >,
960                ) -> (usize, web_rpc::service::ExecuteResult<Self::Response>) {
961                    let __request: #request_type = web_rpc::bincode::deserialize(&__payload).unwrap();
962                    let __result = match __request {
963                        #( #handlers )*
964                        #request_ident::__WebRpcPhantom(_) => {
965                            unreachable!("web_rpc: __WebRpcPhantom variant received on wire")
966                        }
967                    };
968                    (__seq_id, __result)
969                }
970            }
971            impl<T: #trait_ident> std::convert::From<T> for #service_ident<T> {
972                fn from(server_impl: T) -> Self {
973                    Self { server_impl }
974                }
975            }
976        }
977    }
978}
979
980impl<'a> ToTokens for ServiceGenerator<'a> {
981    fn to_tokens(&self, output: &mut TokenStream2) {
982        output.extend(vec![
983            self.enum_request(),
984            self.enum_response(),
985            self.trait_service(),
986            self.struct_client(),
987            self.struct_server(),
988        ])
989    }
990}
991
992impl Parse for Service {
993    fn parse(input: ParseStream) -> syn::Result<Self> {
994        let attrs = input.call(Attribute::parse_outer)?;
995        let vis = input.parse()?;
996        input.parse::<Token![trait]>()?;
997        let ident: Ident = input.parse()?;
998        let content;
999        braced!(content in input);
1000        let mut rpcs = Vec::<RpcMethod>::new();
1001        while !content.is_empty() {
1002            rpcs.push(content.parse()?);
1003        }
1004
1005        Ok(Self {
1006            attrs,
1007            vis,
1008            ident,
1009            rpcs,
1010        })
1011    }
1012}
1013
1014/// Parsed RHS of a `name => ...` or `return => ...` clause.
1015enum TransferRhs {
1016    Expr(syn::Expr),
1017    Gates(Vec<Gate>),
1018}
1019
1020fn parse_transfer_rhs(input: ParseStream) -> syn::Result<TransferRhs> {
1021    if input.peek(Token![|]) || input.peek(Token![||]) {
1022        // Closure form: `|pat| body` (or `|| body` — rejected).
1023        let closure: syn::ExprClosure = input.parse()?;
1024        if closure.inputs.len() != 1 {
1025            return Err(syn::Error::new_spanned(
1026                &closure,
1027                "transfer closure must have exactly one parameter",
1028            ));
1029        }
1030        let pat = closure.inputs.into_iter().next().unwrap();
1031        let body = *closure.body;
1032        Ok(TransferRhs::Gates(vec![Gate { pat, body }]))
1033    } else if input.peek(Token![match]) {
1034        // `match { arms }` — no scrutinee. Bespoke syntax.
1035        input.parse::<Token![match]>()?;
1036        let content;
1037        braced!(content in input);
1038        let arms: Punctuated<syn::Arm, Token![,]> =
1039            content.parse_terminated(syn::Arm::parse)?;
1040        let gates = arms
1041            .into_iter()
1042            .map(|a| Gate {
1043                pat: a.pat,
1044                body: *a.body,
1045            })
1046            .collect();
1047        Ok(TransferRhs::Gates(gates))
1048    } else {
1049        // Bare expression — only valid for params; the caller checks.
1050        let body: syn::Expr = input.parse()?;
1051        Ok(TransferRhs::Expr(body))
1052    }
1053}
1054
1055impl Parse for TransferClause {
1056    fn parse(input: ParseStream) -> syn::Result<Self> {
1057        let is_return = input.peek(Token![return]);
1058        let lhs_name: Option<Ident> = if is_return {
1059            input.parse::<Token![return]>()?;
1060            None
1061        } else {
1062            Some(input.parse()?)
1063        };
1064
1065        if input.peek(Token![=>]) {
1066            input.parse::<Token![=>]>()?;
1067            let rhs = parse_transfer_rhs(input)?;
1068            match (lhs_name, rhs) {
1069                (Some(name), TransferRhs::Expr(body)) => {
1070                    Ok(TransferClause::ParamExpr { name, body })
1071                }
1072                (Some(name), TransferRhs::Gates(gates)) => {
1073                    Ok(TransferClause::ParamGated { name, gates })
1074                }
1075                (None, TransferRhs::Gates(gates)) => {
1076                    Ok(TransferClause::ReturnGated { gates })
1077                }
1078                (None, TransferRhs::Expr(_)) => Err(syn::Error::new(
1079                    input.span(),
1080                    "`return =>` requires a closure (`|pat| body`) or `match { arms }` block",
1081                )),
1082            }
1083        } else {
1084            Ok(match lhs_name {
1085                Some(name) => TransferClause::BareParam(name),
1086                None => TransferClause::BareReturn,
1087            })
1088        }
1089    }
1090}
1091
1092impl Parse for RpcMethod {
1093    fn parse(input: ParseStream) -> syn::Result<Self> {
1094        let mut errors = Ok(());
1095        let attrs = input.call(Attribute::parse_outer)?;
1096
1097        // Reject the removed `#[post(...)]` attribute with a migration message.
1098        for attr in &attrs {
1099            if attr
1100                .path
1101                .segments
1102                .last()
1103                .is_some_and(|seg| seg.ident == "post")
1104            {
1105                extend_errors!(
1106                    errors,
1107                    syn::Error::new_spanned(
1108                        attr,
1109                        "`#[post(...)]` has been removed. JS-vs-serialize routing is now \
1110                         inferred from each argument and return type. For transfer semantics, \
1111                         use `#[transfer(...)]` (e.g. `#[transfer(canvas)]`, \
1112                         `#[transfer(data => data.buffer())]`, or \
1113                         `#[transfer(return => |Ok(o)| o.buffer())]`)."
1114                    )
1115                );
1116            }
1117        }
1118
1119        // Partition out the new `#[transfer(...)]` attribute(s).
1120        let (transfer_attrs, attrs): (Vec<_>, Vec<_>) = attrs.into_iter().partition(|attr| {
1121            attr.path
1122                .segments
1123                .last()
1124                .is_some_and(|last_segment| last_segment.ident == "transfer")
1125        });
1126        let mut transfer: Vec<TransferClause> = Vec::new();
1127        for transfer_attr in transfer_attrs {
1128            let parsed = transfer_attr
1129                .parse_args_with(Punctuated::<TransferClause, Token![,]>::parse_terminated)?;
1130            transfer.extend(parsed.into_iter());
1131        }
1132
1133        let is_async = input.parse::<Token![async]>().ok();
1134        input.parse::<Token![fn]>()?;
1135        let ident: Ident = input.parse()?;
1136
1137        // Reject generic methods up front — autoref dispatch needs concrete types.
1138        if input.peek(Token![<]) {
1139            let generics: syn::Generics = input.parse()?;
1140            extend_errors!(
1141                errors,
1142                syn::Error::new_spanned(
1143                    generics,
1144                    "web_rpc::service trait methods may not have generic parameters; \
1145                     concrete types are required so the macro can route each argument."
1146                )
1147            );
1148        }
1149
1150        let content;
1151        parenthesized!(content in input);
1152        let mut receiver: Option<syn::Receiver> = None;
1153        let mut args = Vec::new();
1154        for arg in content.parse_terminated::<FnArg, Comma>(FnArg::parse)? {
1155            match arg {
1156                FnArg::Typed(captured) => match &*captured.pat {
1157                    Pat::Ident(_) => {
1158                        // Reject reference args other than `&str`/`&[u8]`/`&JsT`.
1159                        // (The is_js_ref / is_borrowed_serde_ref classifiers will
1160                        // accept any reference; we let them through here and rely
1161                        // on the receiver-side decoder to fail on unsupported
1162                        // shapes. A dedicated diagnostic comes later.)
1163                        args.push(captured)
1164                    }
1165                    _ => extend_errors!(
1166                        errors,
1167                        syn::Error::new(
1168                            captured.pat.span(),
1169                            "patterns are not allowed in RPC arguments"
1170                        )
1171                    ),
1172                },
1173                FnArg::Receiver(ref recv) => {
1174                    if recv.reference.is_none() || recv.mutability.is_some() {
1175                        extend_errors!(
1176                            errors,
1177                            syn::Error::new(
1178                                arg.span(),
1179                                "RPC methods only support `&self` as a receiver"
1180                            )
1181                        );
1182                    }
1183                    receiver = Some(recv.clone());
1184                }
1185            }
1186        }
1187        let receiver = match receiver {
1188            Some(r) => r,
1189            None => {
1190                extend_errors!(
1191                    errors,
1192                    syn::Error::new(
1193                        ident.span(),
1194                        "RPC methods must include `&self` as the first parameter"
1195                    )
1196                );
1197                parse_quote!(&self)
1198            }
1199        };
1200        let output: ReturnType = input.parse()?;
1201        input.parse::<Token![;]>()?;
1202
1203        // Validate that every transfer clause references a real parameter
1204        // (or `return`, which has no name to check).
1205        let arg_names: HashSet<_> = args
1206            .iter()
1207            .filter_map(|arg| match &*arg.pat {
1208                Pat::Ident(pat_ident) => Some(pat_ident.ident.clone()),
1209                _ => None,
1210            })
1211            .collect();
1212        for clause in &transfer {
1213            let name_ref = match clause {
1214                TransferClause::BareParam(name)
1215                | TransferClause::ParamExpr { name, .. }
1216                | TransferClause::ParamGated { name, .. } => Some(name),
1217                TransferClause::BareReturn | TransferClause::ReturnGated { .. } => None,
1218            };
1219            if let Some(name) = name_ref {
1220                if !arg_names.contains(name) {
1221                    extend_errors!(
1222                        errors,
1223                        syn::Error::new(
1224                            name.span(),
1225                            format!(
1226                                "`{}` in #[transfer(...)] does not match any parameter",
1227                                name
1228                            )
1229                        )
1230                    );
1231                }
1232            }
1233        }
1234        errors?;
1235
1236        Ok(Self {
1237            is_async,
1238            attrs,
1239            receiver,
1240            ident,
1241            args,
1242            transfer,
1243            output,
1244        })
1245    }
1246}
1247
1248/// This attribute macro should applied to traits that need to be turned into RPCs. The
1249/// macro will consume the trait and output three items in its place. For example,
1250/// a trait `Calculator` will be replaced with two structs `CalculatorClient` and
1251/// `CalculatorService` and a new trait by the same name. All methods must include
1252/// `&self` as their first parameter.
1253#[proc_macro_attribute]
1254pub fn service(_attr: TokenStream, input: TokenStream) -> TokenStream {
1255    let Service {
1256        ref attrs,
1257        ref vis,
1258        ref ident,
1259        ref rpcs,
1260    } = parse_macro_input!(input as Service);
1261
1262    let camel_case_fn_names: &Vec<_> = &rpcs
1263        .iter()
1264        .map(|rpc| snake_to_camel(&rpc.ident.unraw().to_string()))
1265        .collect();
1266
1267    let has_streaming_methods = rpcs.iter().any(
1268        |rpc| matches!(&rpc.output, ReturnType::Type(_, ref ty) if stream_item_type(ty).is_some()),
1269    );
1270
1271    ServiceGenerator {
1272        trait_ident: ident,
1273        service_ident: &format_ident!("{}Service", ident),
1274        client_ident: &format_ident!("{}Client", ident),
1275        request_ident: &format_ident!("{}Request", ident),
1276        response_ident: &format_ident!("{}Response", ident),
1277        vis,
1278        attrs,
1279        rpcs,
1280        camel_case_idents: &rpcs
1281            .iter()
1282            .zip(camel_case_fn_names.iter())
1283            .map(|(rpc, name)| Ident::new(name, rpc.ident.span()))
1284            .collect::<Vec<_>>(),
1285        has_streaming_methods,
1286    }
1287    .into_token_stream()
1288    .into()
1289}
1290
1291fn snake_to_camel(ident_str: &str) -> String {
1292    let mut camel_ty = String::with_capacity(ident_str.len());
1293
1294    let mut last_char_was_underscore = true;
1295    for c in ident_str.chars() {
1296        match c {
1297            '_' => last_char_was_underscore = true,
1298            c if last_char_was_underscore => {
1299                camel_ty.extend(c.to_uppercase());
1300                last_char_was_underscore = false;
1301            }
1302            c => camel_ty.extend(c.to_lowercase()),
1303        }
1304    }
1305
1306    camel_ty.shrink_to_fit();
1307    camel_ty
1308}