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
27fn 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
50fn 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
67fn 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
86fn 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
106fn is_js_ref(ty: &Type) -> bool {
109 matches!(ty, Type::Reference(_)) && !is_borrowed_serde_ref(ty)
110}
111
112fn is_cfg_attr(attr: &Attribute) -> bool {
116 attr.path.is_ident("cfg") || attr.path.is_ident("cfg_attr")
117}
118
119fn emit_encode(ty: &Type, value: TokenStream2, post: &TokenStream2) -> TokenStream2 {
130 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
168fn 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#[allow(dead_code)]
230enum TransferClause {
231 BareParam(Ident),
233 ParamExpr { name: Ident, body: syn::Expr },
235 ParamGated { name: Ident, gates: Vec<Gate> },
239 BareReturn,
241 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 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 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 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 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 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 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; __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 __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), "e!(&__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), "e!(&__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 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 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 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 None
754 } else if is_js_ref(&arg.ty) {
755 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), "e!(&__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 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), "e!(&__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 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), "e!(&__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 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
1014enum 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 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 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 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 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 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 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 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 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#[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}