1use proc_macro::TokenStream;
2use quote::{format_ident, quote};
3use syn::{
4 Error, FnArg, GenericArgument, Ident, Item, ItemImpl, ItemType, PathArguments, ReturnType,
5 Type, parse_macro_input,
6};
7
8#[macro_use]
9mod ts_type;
10mod ts_macro;
11
12use crate::ts_type::ToTsType;
13
14#[proc_macro_attribute]
53pub fn ts(attr: TokenStream, input: TokenStream) -> TokenStream {
54 let item = parse_macro_input!(input as Item);
55 ts_internal_dispatcher(attr.into(), item).into()
56}
57
58fn ts_internal_dispatcher(attr: proc_macro2::TokenStream, item: Item) -> proc_macro2::TokenStream {
59 let attr_args = attr.clone();
60
61 match &item {
62 Item::Struct(item_struct) => {
63 let args = match syn::parse2::<ts_macro::TsArgs>(attr_args) {
64 Ok(args) => args,
65 Err(err) => return err.to_compile_error(),
66 };
67 ts_macro::ts_internal(args, item_struct.clone())
68 }
69 Item::Enum(item_enum) => {
70 let enum_name = &item_enum.ident;
71 let variants = item_enum
72 .variants
73 .iter()
74 .map(|variant| {
75 if !matches!(variant.fields, syn::Fields::Unit) {
76 return Err(Error::new_spanned(
77 variant,
78 "#[ts] enums must only contain unit variants",
79 ));
80 }
81 Ok(&variant.ident)
82 })
83 .collect::<syn::Result<Vec<_>>>();
84 let variants = match variants {
85 Ok(variants) => variants,
86 Err(err) => return err.to_compile_error(),
87 };
88 let variant_constants = variants
89 .iter()
90 .map(|variant| {
91 let name = format_ident!(
92 "__TS_FUNCTION_VARIANT_{}",
93 variant.to_string().to_uppercase()
94 );
95 quote! { const #name: u32 = #enum_name::#variant as u32; }
96 })
97 .collect::<Vec<_>>();
98 let conversion_arms = variants
99 .iter()
100 .map(|variant| {
101 let name = format_ident!(
102 "__TS_FUNCTION_VARIANT_{}",
103 variant.to_string().to_uppercase()
104 );
105 quote! { #name => ::std::result::Result::Ok(Self::#variant), }
106 })
107 .collect::<Vec<_>>();
108 quote! {
109 #[::wasm_bindgen::prelude::wasm_bindgen]
110 #item_enum
111
112 impl ::std::convert::TryFrom<::wasm_bindgen::JsValue> for #enum_name {
113 type Error = ::wasm_bindgen::JsValue;
114
115 #[inline]
116 fn try_from(value: ::wasm_bindgen::JsValue) -> ::std::result::Result<Self, Self::Error> {
117 let value = match value.as_f64() {
118 ::std::option::Option::Some(value) => value,
119 ::std::option::Option::None => {
120 return ::std::result::Result::Err(::wasm_bindgen::JsValue::from_str(
121 concat!("Expected a number for enum ", stringify!(#enum_name)),
122 ));
123 }
124 };
125 if !value.is_finite()
126 || value.fract() != 0.0
127 || !(0.0..=u32::MAX as f64).contains(&value)
128 {
129 return ::std::result::Result::Err(::wasm_bindgen::JsValue::from_str(&format!(
130 "Invalid {} variant: {}", stringify!(#enum_name), value
131 )));
132 }
133 #(#variant_constants)*
134 match value as u32 {
135 #(#conversion_arms)*
136 _ => ::std::result::Result::Err(::wasm_bindgen::JsValue::from_str(&format!(
137 "Invalid {} variant: {}", stringify!(#enum_name), value
138 ))),
139 }
140 }
141 }
142 }
143 }
144 Item::Type(item_type) => match parse_item_type(item_type) {
145 Ok(tokens) => tokens,
146 Err(err) => err.to_compile_error(),
147 },
148 Item::Impl(item_impl) => match parse_item_impl(item_impl) {
149 Ok(tokens) => tokens,
150 Err(err) => err.to_compile_error(),
151 },
152 _ => Error::new_spanned(
153 item,
154 "#[ts] can only be applied to a struct, enum, type alias, or impl block",
155 )
156 .to_compile_error(),
157 }
158}
159
160struct ParsedSignature<'a> {
161 struct_ident: &'a Ident,
162 args: Vec<(Ident, &'a Type)>,
163 output: &'a ReturnType,
164}
165
166pub(crate) fn generate_try_convert_support(struct_ident: &syn::Ident) -> proc_macro2::TokenStream {
167 let try_convert_name = format_ident!("try_convert_{}", struct_ident);
168 let trait_name = format_ident!("IntoJsValue_{}", struct_ident);
169 quote! {
170 #[allow(non_camel_case_types)]
171 trait #trait_name {
172 fn into_js_value(self) -> ::wasm_bindgen::JsValue;
173 }
174
175 impl #trait_name for ::wasm_bindgen::JsValue {
176 #[inline]
177 fn into_js_value(self) -> ::wasm_bindgen::JsValue {
178 self
179 }
180 }
181
182 impl #trait_name for ::std::convert::Infallible {
183 #[inline]
184 fn into_js_value(self) -> ::wasm_bindgen::JsValue {
185 match self {}
186 }
187 }
188
189 #[inline]
190 #[allow(non_snake_case)]
191 fn #try_convert_name<T, E>(res: ::wasm_bindgen::JsValue) -> ::std::result::Result<T, ::wasm_bindgen::JsValue>
192 where
193 T: ::std::convert::TryFrom<::wasm_bindgen::JsValue, Error = E>,
194 E: #trait_name,
195 {
196 ::std::convert::TryInto::<T>::try_into(res).map_err(#trait_name::into_js_value)
197 }
198 }
199}
200
201pub(crate) fn generate_return_conversion(
202 struct_ident: &syn::Ident,
203 ty: &Type,
204) -> syn::Result<proc_macro2::TokenStream> {
205 let try_convert_name = format_ident!("try_convert_{}", struct_ident);
206 match ty {
207 Type::Path(type_path) => {
208 let segment = type_path
209 .path
210 .segments
211 .last()
212 .ok_or_else(|| Error::new_spanned(ty, "Expected a type segment"))?;
213 let ident = &segment.ident;
214 let ident_str = ident.to_string();
215
216 if let Some(inner_ty) = get_slice_element_type(ty)
217 && let Some(arr_type) = get_typed_array_ident(inner_ty)
218 {
219 return Ok(quote! {
220 let arr: ::js_sys::#arr_type = ::wasm_bindgen::JsCast::dyn_into(res)
221 .map_err(|_| ::wasm_bindgen::JsValue::from_str(concat!("Expected a ", stringify!(#arr_type))))?;
222 ::std::result::Result::Ok::<_, ::wasm_bindgen::JsValue>(::std::convert::Into::<#ty>::into(arr.to_vec()))
223 });
224 }
225
226 match ident_str.as_str() {
227 "f32" | "f64" | "i8" | "i16" | "i32" | "u8" | "u16" | "u32" => Ok(quote! {
228 res.as_f64().map(|v| v as #ty).ok_or_else(|| ::wasm_bindgen::JsValue::from_str("Expected a number"))
229 }),
230 "i64" | "u64" => Ok(quote! {
231 ::std::convert::TryInto::<#ty>::try_into(res).map_err(|_| ::wasm_bindgen::JsValue::from_str("Expected a BigInt"))
232 }),
233 "bool" => Ok(quote! {
234 res.as_bool().ok_or_else(|| ::wasm_bindgen::JsValue::from_str("Expected a boolean"))
235 }),
236 "String" => Ok(quote! {
237 res.as_string().ok_or_else(|| ::wasm_bindgen::JsValue::from_str("Expected a string"))
238 }),
239 "JsValue" => Ok(quote! {
240 ::std::result::Result::Ok::<_, ::wasm_bindgen::JsValue>(res)
241 }),
242 "Option" => {
243 let PathArguments::AngleBracketed(args) = &segment.arguments else {
244 return Err(Error::new_spanned(
245 ty,
246 "Expected generic argument for Option",
247 ));
248 };
249 let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() else {
250 return Err(Error::new_spanned(ty, "Expected type argument for Option"));
251 };
252 let inner_conversion = generate_return_conversion(struct_ident, inner_ty)?;
253 Ok(quote! {
254 if res.is_null() || res.is_undefined() {
255 ::std::result::Result::Ok::<_, ::wasm_bindgen::JsValue>(None)
256 } else {
257 let res = { #inner_conversion };
258 res.map(Some)
259 }
260 })
261 }
262 _ => Ok(quote! {
263 #try_convert_name::<#ty, _>(res)
264 }),
265 }
266 }
267 _ => Err(Error::new_spanned(
268 ty,
269 "Unsupported return type in type alias pattern. Use the `impl` escape hatch instead.",
270 )),
271 }
272}
273
274fn parse_item_type(item_type: &ItemType) -> syn::Result<proc_macro2::TokenStream> {
275 item_type.modifiers.require_empty()?;
276
277 let Type::FnPtr(bare_fn) = &*item_type.ty else {
278 return Err(Error::new_spanned(
279 &item_type.ty,
280 "Expected a function pointer type (e.g., `fn(x: f64)`)",
281 ));
282 };
283
284 let struct_ident = &item_type.ident;
285 let mut args = Vec::new();
286
287 for (i, arg) in bare_fn.inputs.iter().enumerate() {
288 let ident = match &arg.name {
289 Some((ident, _)) => ident.clone(),
290 None => format_ident!("arg{}", i),
291 };
292 args.push((ident, &arg.ty));
293 }
294
295 let parsed = ParsedSignature {
296 struct_ident,
297 args: args.clone(),
298 output: &bare_fn.output,
299 };
300
301 let abi_traits = generate_abi_traits(&parsed)?;
302
303 let mut fn_args = Vec::new();
304 let mut arg_conversions = Vec::new();
305 let mut call_args = Vec::new();
306 for (ident, ty) in &args {
307 fn_args.push(quote! { #ident: #ty });
308 let conversion = generate_conversion(ident, ty)?;
309 arg_conversions.push(conversion);
310 call_args.push(quote! { &#ident });
311 }
312
313 let args_len = call_args.len();
314 if args_len > 9 {
315 return Err(Error::new_spanned(
316 item_type,
317 "Functions with more than 9 arguments are not supported yet",
318 ));
319 }
320 let call_method_name = format_ident!("call{}", args_len);
321 let call_method = quote! { #call_method_name(&::wasm_bindgen::JsValue::NULL, #(#call_args),*) };
322
323 let output = parsed.output;
324 let (ret_type, ret_stmt) = match output {
325 ReturnType::Default => (quote! { () }, quote! { self.0.#call_method.map(|_| ()) }),
326 ReturnType::Type(_, ty) => {
327 let conversion = generate_return_conversion(struct_ident, ty)?;
328 (
329 quote! { #ty },
330 quote! {
331 let res = self.0.#call_method?;
332 #conversion
333 },
334 )
335 }
336 };
337
338 Ok(quote! {
339 pub struct #struct_ident(pub ::js_sys::Function);
340
341 const _: () = {
342 #abi_traits
343
344 impl #struct_ident {
345 pub fn call(&self, #(#fn_args),*) -> Result<#ret_type, ::wasm_bindgen::JsValue> {
346 #(#arg_conversions)*
347 #ret_stmt
348 }
349 }
350 };
351 })
352}
353
354fn generate_conversion(ident: &Ident, ty: &Type) -> syn::Result<proc_macro2::TokenStream> {
355 if let Type::ImplTrait(type_impl) = ty {
356 for bound in &type_impl.bounds {
357 if let syn::TypeParamBound::Trait(trait_bound) = bound
358 && let Some(segment) = trait_bound.path.segments.last()
359 && let PathArguments::AngleBracketed(args) = &segment.arguments
360 && let Some(GenericArgument::Type(inner_ty)) = args.args.first()
361 {
362 match segment.ident.to_string().as_str() {
363 "Into" => {
364 let inner_conversion = generate_conversion(ident, inner_ty)?;
365 return Ok(quote! {
366 let #ident = ::std::convert::Into::<#inner_ty>::into(#ident);
367 #inner_conversion
368 });
369 }
370 "AsRef" => {
371 if let Type::Slice(slice) = inner_ty {
372 return Ok(generate_typed_array_conversion(ident, &slice.elem));
373 }
374 }
375 _ => {}
376 }
377 }
378 }
379 return Err(Error::new_spanned(
380 ty,
381 "Unsupported `impl Trait`. Only `impl Into<T>` and `impl AsRef<[T]>` are supported.",
382 ));
383 }
384
385 if let Some(inner_ty) = get_slice_element_type(ty) {
386 Ok(generate_typed_array_conversion(ident, inner_ty))
387 } else {
388 Ok(quote! {
389 let #ident = ::std::convert::Into::<::wasm_bindgen::JsValue>::into(#ident);
390 })
391 }
392}
393
394fn generate_typed_array_conversion(ident: &Ident, inner_ty: &Type) -> proc_macro2::TokenStream {
395 if let Some(arr_type) = get_typed_array_ident(inner_ty) {
396 quote! {
397 let #ident = ::wasm_bindgen::JsValue::from(::js_sys::#arr_type::from(::std::convert::AsRef::<[#inner_ty]>::as_ref(&#ident)));
398 }
399 } else {
400 quote! {
401 let #ident = ::wasm_bindgen::JsValue::from(
402 ::std::convert::AsRef::<[#inner_ty]>::as_ref(&#ident)
403 .iter()
404 .map(::wasm_bindgen::JsValue::from)
405 .collect::<::js_sys::Array>()
406 );
407 }
408 }
409}
410
411fn get_typed_array_ident(inner_ty: &Type) -> Option<proc_macro2::TokenStream> {
412 let inner_str = match inner_ty {
413 Type::Path(p) => p.path.segments.last().map(|s| s.ident.to_string()),
414 _ => None,
415 };
416
417 match inner_str.as_deref() {
418 Some("u8") => Some(quote! { Uint8Array }),
419 Some("i8") => Some(quote! { Int8Array }),
420 Some("u16") => Some(quote! { Uint16Array }),
421 Some("i16") => Some(quote! { Int16Array }),
422 Some("u32") => Some(quote! { Uint32Array }),
423 Some("i32") => Some(quote! { Int32Array }),
424 Some("f32") => Some(quote! { Float32Array }),
425 Some("f64") => Some(quote! { Float64Array }),
426 Some("u64") => Some(quote! { BigUint64Array }),
427 Some("i64") => Some(quote! { BigInt64Array }),
428 _ => None,
429 }
430}
431
432fn get_slice_element_type(ty: &Type) -> Option<&Type> {
433 match ty {
434 Type::Path(type_path) => {
435 let segment = type_path.path.segments.last()?;
436 if matches!(
438 segment.ident.to_string().as_str(),
439 "Vec" | "Box" | "Arc" | "Rc"
440 ) && let PathArguments::AngleBracketed(args) = &segment.arguments
441 && let Some(syn::GenericArgument::Type(inner)) = args.args.first()
442 {
443 if let Type::Slice(slice) = inner {
444 return Some(&*slice.elem);
445 }
446 return Some(inner);
447 }
448 }
449 Type::Reference(type_ref) => {
450 if let Type::Slice(type_slice) = &*type_ref.elem {
451 return Some(&*type_slice.elem);
452 }
453 return get_slice_element_type(&type_ref.elem);
454 }
455 _ => {}
456 }
457 None
458}
459
460fn parse_item_impl(item_impl: &ItemImpl) -> syn::Result<proc_macro2::TokenStream> {
461 item_impl.modifiers.require_empty()?;
462
463 if item_impl.trait_.is_some() {
464 return Err(Error::new_spanned(
465 item_impl,
466 "#[ts_function] cannot be applied to trait impls",
467 ));
468 }
469
470 let Type::Path(type_path) = &*item_impl.self_ty else {
471 return Err(Error::new_spanned(
472 &item_impl.self_ty,
473 "Expected a simple path for the struct",
474 ));
475 };
476
477 let struct_ident = type_path.path.get_ident().ok_or_else(|| {
478 Error::new_spanned(
479 &type_path.path,
480 "Expected a single identifier for the struct",
481 )
482 })?;
483
484 let method = item_impl
485 .items
486 .iter()
487 .find_map(|item| {
488 if let syn::ImplItem::Fn(method) = item
489 && method.sig.ident == "call"
490 {
491 return Some(method);
492 }
493 None
494 })
495 .ok_or_else(|| Error::new_spanned(item_impl, "Missing `call` method in impl block"))?;
496
497 let mut args = Vec::new();
498 let mut inputs_iter = method.sig.inputs.iter();
499
500 match inputs_iter.next() {
502 Some(FnArg::Receiver(_)) => {}
503 _ => {
504 return Err(Error::new_spanned(
505 &method.sig,
506 "The `call` method must take `&self` or `&mut self` as its first parameter",
507 ));
508 }
509 }
510
511 for (i, arg) in inputs_iter.enumerate() {
512 let FnArg::Typed(pat_type) = arg else {
513 return Err(Error::new_spanned(arg, "Expected a typed argument"));
514 };
515
516 let ident = if let syn::Pat::Ident(pat_ident) = &*pat_type.pat {
517 pat_ident.ident.clone()
518 } else {
519 format_ident!("arg{}", i)
520 };
521
522 args.push((ident, &*pat_type.ty));
523 }
524
525 let parsed = ParsedSignature {
526 struct_ident,
527 args,
528 output: &method.sig.output,
529 };
530
531 let abi_traits = generate_abi_traits(&parsed)?;
532
533 Ok(quote! {
534 #item_impl
535 #abi_traits
536 })
537}
538
539fn generate_abi_traits(parsed: &ParsedSignature) -> syn::Result<proc_macro2::TokenStream> {
540 let struct_ident = parsed.struct_ident;
541 let mut ts_args = Vec::new();
542
543 for (ident, ty) in &parsed.args {
544 let ts_ty = ty
545 .to_ts_type()
546 .map_err(|e| Error::new_spanned(ty, e.message))?
547 .to_string();
548 ts_args.push(format!("{}: {}", ident, ts_ty));
549 }
550
551 let ts_output = match parsed.output {
552 ReturnType::Default => "void".to_string(),
553 ReturnType::Type(_, ty) => ty
554 .to_ts_type()
555 .map_err(|e| Error::new_spanned(ty, e.message))?
556 .to_string(),
557 };
558
559 let ts_string = format!(
560 "type {} = ({}) => {};",
561 struct_ident,
562 ts_args.join(", "),
563 ts_output
564 );
565
566 let try_convert_support = generate_try_convert_support(struct_ident);
567
568 let generated = quote! {
569 #[::wasm_bindgen::prelude::wasm_bindgen(typescript_custom_section)]
570 const _: &'static str = #ts_string;
571
572 #try_convert_support
573
574 impl ::wasm_bindgen::describe::WasmDescribe for #struct_ident {
575 fn describe() {
576 <::js_sys::Function as ::wasm_bindgen::describe::WasmDescribe>::describe()
577 }
578 }
579
580 impl ::wasm_bindgen::convert::FromWasmAbi for #struct_ident {
581 type Abi = <::js_sys::Function as ::wasm_bindgen::convert::FromWasmAbi>::Abi;
582
583 unsafe fn from_abi(js: Self::Abi) -> Self {
584 Self(::js_sys::Function::from_abi(js))
585 }
586 }
587
588 impl ::wasm_bindgen::convert::OptionFromWasmAbi for #struct_ident {
589 fn is_none(abi: &Self::Abi) -> bool {
590 <::js_sys::Function as ::wasm_bindgen::convert::OptionFromWasmAbi>::is_none(abi)
591 }
592 }
593
594 impl From<::js_sys::Function> for #struct_ident {
595 fn from(f: ::js_sys::Function) -> Self {
596 Self(f)
597 }
598 }
599
600 impl ::std::convert::TryFrom<::wasm_bindgen::JsValue> for #struct_ident {
601 type Error = ::wasm_bindgen::JsValue;
602
603 #[inline]
604 fn try_from(value: ::wasm_bindgen::JsValue) -> ::std::result::Result<Self, Self::Error> {
605 use ::wasm_bindgen::JsCast;
606 let f = value.dyn_into::<::js_sys::Function>()?;
607 ::std::result::Result::Ok(Self(f))
608 }
609 }
610
611 impl From<#struct_ident> for ::wasm_bindgen::JsValue {
612 fn from(f: #struct_ident) -> Self {
613 ::wasm_bindgen::JsValue::from(f.0)
614 }
615 }
616 };
617
618 Ok(generated)
619}
620
621#[cfg(test)]
622mod tests {
623 use super::*;
624 use syn::parse_quote;
625
626 #[test]
627 fn test_item_type() {
628 let item_type: ItemType = parse_quote! {
629 pub type OnClick = fn(x: f64, y: impl Into<f64>, arr: js_sys::Float64Array);
630 };
631 let result = parse_item_type(&item_type).unwrap();
632 let result_str = result.to_string();
633
634 assert!(
635 result_str
636 .contains("type OnClick = (x: number, y: number, arr: Float64Array) => void;")
637 );
638 assert!(result_str.contains("pub struct OnClick (pub :: js_sys :: Function) ;"));
639 assert!(result_str.contains(
640 "pub fn call (& self , x : f64 , y : impl Into < f64 > , arr : js_sys :: Float64Array)"
641 ));
642 }
643
644 #[test]
645 fn test_item_impl() {
646 let item_impl: ItemImpl = parse_quote! {
647 impl OnScroll {
648 pub fn call(&self, y: f64) {
649 }
651 }
652 };
653 let result = parse_item_impl(&item_impl).unwrap();
654 let result_str = result.to_string();
655
656 assert!(result_str.contains("type OnScroll = (y: number) => void;"));
657 assert!(
658 result_str.contains("impl :: wasm_bindgen :: describe :: WasmDescribe for OnScroll")
659 );
660 }
661
662 #[test]
663 fn test_dispatcher_item_struct() {
664 let input: Item = parse_quote! {
665 pub struct MyStruct {
666 pub field: f64,
667 }
668 };
669 let attr = quote! {};
670 let result = ts_internal_dispatcher(attr, input);
671 let result_str = result.to_string();
672
673 assert!(result_str.contains("export interface MyStruct"));
674 assert!(result_str.contains("field: number;"));
675 }
676
677 #[test]
678 fn test_dispatcher_item_type() {
679 let input: Item = parse_quote! {
680 pub type OnClick = fn(x: f64);
681 };
682 let attr = quote! {};
683 let result = ts_internal_dispatcher(attr, input);
684 let result_str = result.to_string();
685
686 assert!(result_str.contains("type OnClick = (x: number) => void;"));
687 assert!(result_str.contains("pub struct OnClick (pub :: js_sys :: Function) ;"));
688 }
689
690 #[test]
691 fn test_dispatcher_item_impl() {
692 let input: Item = parse_quote! {
693 impl OnScroll {
694 pub fn call(&self, y: f64) {}
695 }
696 };
697 let attr = quote! {};
698 let result = ts_internal_dispatcher(attr, input);
699 let result_str = result.to_string();
700
701 assert!(result_str.contains("type OnScroll = (y: number) => void;"));
702 assert!(
703 result_str.contains("impl :: wasm_bindgen :: describe :: WasmDescribe for OnScroll")
704 );
705 }
706
707 #[test]
708 fn test_enum_item() {
709 let input: Item = parse_quote! {
710 pub enum Status { Active, Inactive }
711 };
712 let attr = quote! {};
713 let result = ts_internal_dispatcher(attr, input);
714 let result_str = result.to_string();
715
716 assert!(result_str.contains("# [:: wasm_bindgen :: prelude :: wasm_bindgen]"));
717 assert!(result_str.contains("pub enum Status { Active , Inactive }"));
718 }
719
720 #[test]
721 fn test_recursive_generics() {
722 let item_type: ItemType = parse_quote! {
723 pub type ResultFn = fn(res: Result<String, i32>);
724 };
725 let result = parse_item_type(&item_type).unwrap();
726 let result_str = result.to_string();
727
728 assert!(result_str.contains("type ResultFn = (res: Result<string, number>) => void;"));
729
730 let item_type: ItemType = parse_quote! {
731 pub type NestedVecFn = fn(args: Vec<Vec<f64>>);
732 };
733 let result = parse_item_type(&item_type).unwrap();
734 let result_str = result.to_string();
735
736 assert!(result_str.contains("type NestedVecFn = (args: Float64Array[]) => void;"));
737 }
738
739 #[test]
740 fn test_item_impl_rejects_modifiers() {
741 let item_impl: ItemImpl = parse_quote! {
742 default impl Callback {
743 fn call(&self) {}
744 }
745 };
746
747 assert!(parse_item_impl(&item_impl).is_err());
748 }
749}