1#![allow(clippy::let_and_return)]
152#![deny(unused_variables, dead_code, unused_must_use, unused_imports)]
153
154use proc_macro2::TokenStream;
155use quote::{format_ident, quote, ToTokens};
156use std::convert::{TryFrom, TryInto};
157use syn::{
158 parse::{self, Parse, ParseStream},
159 parse_macro_input, parse_quote,
160 punctuated::Punctuated,
161 spanned::Spanned,
162 token::{Plus, Semi},
163 Ident, ImplItem, ItemImpl, Result, Token, TraitItemConst, TraitItemFn, Type, TypeArray,
164 TypeBareFn, TypeGroup, TypeNever, TypeParamBound, TypeParen, TypePath, TypePtr, TypeReference,
165 TypeSlice, TypeTraitObject, TypeTuple, Visibility,
166};
167
168#[derive(Debug)]
169struct Input {
170 item_impl: ItemImpl,
171 vis: Option<Visibility>,
172}
173
174impl Parse for Input {
175 fn parse(input: ParseStream) -> syn::Result<Self> {
176 let mut attributes = Vec::new();
177 if input.peek(syn::Token![#]) {
178 attributes.extend(syn::Attribute::parse_outer(input)?);
179 }
180
181 let vis = input
182 .parse::<Visibility>()
183 .ok()
184 .filter(|vis| vis != &Visibility::Inherited);
185
186 let mut item_impl = input.parse::<ItemImpl>()?;
187 item_impl.attrs.extend(attributes);
188
189 Ok(Self { item_impl, vis })
190 }
191}
192
193#[proc_macro_attribute]
195#[allow(clippy::unneeded_field_pattern)]
196pub fn ext(
197 attr: proc_macro::TokenStream,
198 item: proc_macro::TokenStream,
199) -> proc_macro::TokenStream {
200 let item = parse_macro_input!(item as Input);
201 let config = parse_macro_input!(attr as Config);
202 match go(item, config) {
203 Ok(tokens) => tokens,
204 Err(err) => err.into_compile_error().into(),
205 }
206}
207
208#[proc_macro_attribute]
224#[allow(clippy::unneeded_field_pattern)]
225pub fn ext_sized(
226 attr: proc_macro::TokenStream,
227 item: proc_macro::TokenStream,
228) -> proc_macro::TokenStream {
229 let item = parse_macro_input!(item as Input);
230 let mut config: Config = parse_macro_input!(attr as Config);
231
232 config.supertraits = if let Some(supertraits) = config.supertraits.take() {
233 Some(parse_quote!(#supertraits + Sized))
234 } else {
235 Some(parse_quote!(Sized))
236 };
237
238 match go(item, config) {
239 Ok(tokens) => tokens,
240 Err(err) => err.into_compile_error().into(),
241 }
242}
243
244fn go(item: Input, mut config: Config) -> Result<proc_macro::TokenStream> {
245 if let Some(vis) = item.vis {
246 if config.visibility != Visibility::Inherited {
247 return Err(syn::Error::new(
248 config.visibility.span(),
249 "Cannot set visibility on `#[ext]` and `impl` block",
250 ));
251 }
252
253 config.visibility = vis;
254 }
255
256 let ItemImpl {
257 attrs,
258 unsafety,
259 generics,
260 trait_,
261 self_ty,
262 items,
263 defaultness: _,
265 impl_token: _,
266 brace_token: _,
267 } = item.item_impl;
268
269 if let Some((_, path, _)) = trait_ {
270 return Err(syn::Error::new(
271 path.span(),
272 "Trait impls cannot be used for #[ext]",
273 ));
274 }
275
276 let self_ty = parse_self_ty(&self_ty)?;
277
278 let ext_trait_name = if let Some(ext_trait_name) = config.ext_trait_name {
279 ext_trait_name
280 } else {
281 ext_trait_name(&self_ty)?
282 };
283
284 let MethodsAndConsts {
285 trait_methods,
286 trait_consts,
287 } = extract_allowed_items(&items)?;
288
289 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
290
291 let visibility = &config.visibility;
292
293 let mut all_supertraits = Vec::<TypeParamBound>::new();
294
295 if let Some(supertraits_from_config) = config.supertraits {
296 all_supertraits.extend(supertraits_from_config);
297 }
298
299 let supertraits_quoted = if all_supertraits.is_empty() {
300 quote! {}
301 } else {
302 let supertraits_quoted = punctuated_from_iter::<_, _, Plus>(all_supertraits);
303 quote! { : #supertraits_quoted }
304 };
305
306 let code = (quote! {
307 #[allow(non_camel_case_types)]
308 #(#attrs)*
309 #visibility
310 #unsafety
311 trait #ext_trait_name #impl_generics #supertraits_quoted #where_clause {
312 #(
313 #trait_consts
314 )*
315
316 #(
317 #[allow(
318 patterns_in_fns_without_body,
319 clippy::inline_fn_without_body,
320 unused_attributes
321 )]
322 #trait_methods
323 )*
324 }
325
326 #(#attrs)*
327 impl #impl_generics #ext_trait_name #ty_generics for #self_ty #where_clause {
328 #(#items)*
329 }
330 })
331 .into();
332
333 Ok(code)
334}
335
336#[derive(Debug, Clone)]
337enum ExtType<'a> {
338 Array(&'a TypeArray),
339 Group(&'a TypeGroup),
340 Never(&'a TypeNever),
341 Paren(&'a TypeParen),
342 Path(&'a TypePath),
343 Ptr(&'a TypePtr),
344 Reference(&'a TypeReference),
345 Slice(&'a TypeSlice),
346 Tuple(&'a TypeTuple),
347 BareFn(&'a TypeBareFn),
348 TraitObject(&'a TypeTraitObject),
349}
350
351#[allow(clippy::wildcard_in_or_patterns)]
352fn parse_self_ty(self_ty: &Type) -> Result<ExtType> {
353 let ty = match self_ty {
354 Type::Array(inner) => ExtType::Array(inner),
355 Type::Group(inner) => ExtType::Group(inner),
356 Type::Never(inner) => ExtType::Never(inner),
357 Type::Paren(inner) => ExtType::Paren(inner),
358 Type::Path(inner) => ExtType::Path(inner),
359 Type::Ptr(inner) => ExtType::Ptr(inner),
360 Type::Reference(inner) => ExtType::Reference(inner),
361 Type::Slice(inner) => ExtType::Slice(inner),
362 Type::Tuple(inner) => ExtType::Tuple(inner),
363 Type::BareFn(inner) => ExtType::BareFn(inner),
364 Type::TraitObject(inner) => ExtType::TraitObject(inner),
365
366 Type::ImplTrait(_) | Type::Infer(_) | Type::Macro(_) | Type::Verbatim(_) | _ => {
367 return Err(syn::Error::new(
368 self_ty.span(),
369 "#[ext] is not supported for this kind of type",
370 ))
371 }
372 };
373 Ok(ty)
374}
375
376impl<'a> TryFrom<&'a Type> for ExtType<'a> {
377 type Error = syn::Error;
378
379 fn try_from(inner: &'a Type) -> Result<ExtType<'a>> {
380 parse_self_ty(inner)
381 }
382}
383
384impl<'a> ToTokens for ExtType<'a> {
385 fn to_tokens(&self, tokens: &mut TokenStream) {
386 match self {
387 ExtType::Array(inner) => inner.to_tokens(tokens),
388 ExtType::Group(inner) => inner.to_tokens(tokens),
389 ExtType::Never(inner) => inner.to_tokens(tokens),
390 ExtType::Paren(inner) => inner.to_tokens(tokens),
391 ExtType::Path(inner) => inner.to_tokens(tokens),
392 ExtType::Ptr(inner) => inner.to_tokens(tokens),
393 ExtType::Reference(inner) => inner.to_tokens(tokens),
394 ExtType::Slice(inner) => inner.to_tokens(tokens),
395 ExtType::Tuple(inner) => inner.to_tokens(tokens),
396 ExtType::BareFn(inner) => inner.to_tokens(tokens),
397 ExtType::TraitObject(inner) => inner.to_tokens(tokens),
398 }
399 }
400}
401
402fn ext_trait_name(self_ty: &ExtType) -> Result<Ident> {
403 fn inner_self_ty(self_ty: &ExtType) -> Result<Ident> {
404 match self_ty {
405 ExtType::Path(inner) => find_and_combine_idents(inner),
406 ExtType::Reference(inner) => {
407 let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
408 if inner.mutability.is_some() {
409 Ok(format_ident!("RefMut{}", name))
410 } else {
411 Ok(format_ident!("Ref{}", name))
412 }
413 }
414 ExtType::Array(inner) => {
415 let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
416 Ok(format_ident!("ListOf{}", name))
417 }
418 ExtType::Group(inner) => {
419 let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
420 Ok(format_ident!("Group{}", name))
421 }
422 ExtType::Paren(inner) => {
423 let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
424 Ok(format_ident!("Paren{}", name))
425 }
426 ExtType::Ptr(inner) => {
427 let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
428 Ok(format_ident!("PointerTo{}", name))
429 }
430 ExtType::Slice(inner) => {
431 let name = inner_self_ty(&(&*inner.elem).try_into()?)?;
432 Ok(format_ident!("SliceOf{}", name))
433 }
434 ExtType::Tuple(inner) => {
435 let mut name = format_ident!("TupleOf");
436 for elem in &inner.elems {
437 name = format_ident!("{}{}", name, inner_self_ty(&elem.try_into()?)?);
438 }
439 Ok(name)
440 }
441 ExtType::Never(_) => Ok(format_ident!("Never")),
442 ExtType::BareFn(inner) => {
443 let mut name = format_ident!("BareFn");
444 for input in inner.inputs.iter() {
445 name = format_ident!("{}{}", name, inner_self_ty(&(&input.ty).try_into()?)?);
446 }
447 match &inner.output {
448 syn::ReturnType::Default => {
449 name = format_ident!("{}Unit", name);
450 }
451 syn::ReturnType::Type(_, ty) => {
452 name = format_ident!("{}{}", name, inner_self_ty(&(&**ty).try_into()?)?);
453 }
454 }
455 Ok(name)
456 }
457 ExtType::TraitObject(inner) => {
458 let mut name = format_ident!("TraitObject");
459 for bound in inner.bounds.iter() {
460 match bound {
461 TypeParamBound::Trait(bound) => {
462 for segment in bound.path.segments.iter() {
463 name = format_ident!("{}{}", name, segment.ident);
464 }
465 }
466 TypeParamBound::Lifetime(lifetime) => {
467 name = format_ident!("{}{}", name, lifetime.ident);
468 }
469 other => {
470 return Err(syn::Error::new(other.span(), "unsupported bound"));
471 }
472 }
473 }
474 Ok(name)
475 }
476 }
477 }
478
479 Ok(format_ident!("{}Ext", inner_self_ty(self_ty)?))
480}
481
482fn find_and_combine_idents(type_path: &TypePath) -> Result<Ident> {
483 use syn::visit::{self, Visit};
484
485 struct IdentVisitor<'a>(Vec<&'a Ident>);
486
487 impl<'a> Visit<'a> for IdentVisitor<'a> {
488 fn visit_ident(&mut self, i: &'a Ident) {
489 self.0.push(i);
490 }
491 }
492
493 let mut visitor = IdentVisitor(Vec::new());
494 visit::visit_type_path(&mut visitor, type_path);
495 let idents = visitor.0;
496
497 if idents.is_empty() {
498 Err(syn::Error::new(type_path.span(), "Empty type path"))
499 } else {
500 let start = &idents[0].span();
501 let combined_span = idents
502 .iter()
503 .map(|i| i.span())
504 .fold(*start, |a, b| a.join(b).unwrap_or(a));
505
506 let combined_name = idents.iter().map(|i| i.to_string()).collect::<String>();
507
508 Ok(Ident::new(&combined_name, combined_span))
509 }
510}
511
512#[derive(Debug, Default)]
513struct MethodsAndConsts {
514 trait_methods: Vec<TraitItemFn>,
515 trait_consts: Vec<TraitItemConst>,
516}
517
518#[allow(clippy::wildcard_in_or_patterns)]
519fn extract_allowed_items(items: &[ImplItem]) -> Result<MethodsAndConsts> {
520 let mut acc = MethodsAndConsts::default();
521 for item in items {
522 match item {
523 ImplItem::Fn(method) => acc.trait_methods.push(TraitItemFn {
524 attrs: method.attrs.clone(),
525 sig: {
526 let mut sig = method.sig.clone();
527 sig.inputs = sig
528 .inputs
529 .into_iter()
530 .map(|fn_arg| match fn_arg {
531 syn::FnArg::Receiver(recv) => syn::FnArg::Receiver(recv),
532 syn::FnArg::Typed(mut pat_type) => {
533 pat_type.pat = Box::new(match *pat_type.pat {
534 syn::Pat::Ident(pat_ident) => syn::Pat::Ident(pat_ident),
535 _ => {
536 parse_quote!(_)
537 }
538 });
539 syn::FnArg::Typed(pat_type)
540 }
541 })
542 .collect();
543 sig
544 },
545 default: None,
546 semi_token: Some(Semi::default()),
547 }),
548 ImplItem::Const(const_) => acc.trait_consts.push(TraitItemConst {
549 attrs: const_.attrs.clone(),
550 generics: const_.generics.clone(),
551 const_token: Default::default(),
552 ident: const_.ident.clone(),
553 colon_token: Default::default(),
554 ty: const_.ty.clone(),
555 default: None,
556 semi_token: Default::default(),
557 }),
558 ImplItem::Type(_) => {
559 return Err(syn::Error::new(
560 item.span(),
561 "Associated types are not allowed in #[ext] impls",
562 ))
563 }
564 ImplItem::Macro(_) => {
565 return Err(syn::Error::new(
566 item.span(),
567 "Macros are not allowed in #[ext] impls",
568 ))
569 }
570 ImplItem::Verbatim(_) | _ => {
571 return Err(syn::Error::new(item.span(), "Not allowed in #[ext] impls"))
572 }
573 }
574 }
575 Ok(acc)
576}
577
578#[derive(Debug)]
579struct Config {
580 ext_trait_name: Option<Ident>,
581 visibility: Visibility,
582 supertraits: Option<Punctuated<TypeParamBound, Plus>>,
583}
584
585impl Parse for Config {
586 fn parse(input: ParseStream) -> parse::Result<Self> {
587 let mut config = Config::default();
588
589 if let Ok(visibility) = input.parse::<Visibility>() {
590 config.visibility = visibility;
591 }
592
593 input.parse::<Token![,]>().ok();
594
595 while !input.is_empty() {
596 let ident = input.parse::<Ident>()?;
597 input.parse::<Token![=]>()?;
598
599 match &*ident.to_string() {
600 "name" => {
601 config.ext_trait_name = Some(input.parse()?);
602 }
603 "supertraits" => {
604 config.supertraits =
605 Some(Punctuated::<TypeParamBound, Plus>::parse_terminated(input)?);
606 }
607 _ => return Err(syn::Error::new(ident.span(), "Unknown configuration name")),
608 }
609
610 input.parse::<Token![,]>().ok();
611 }
612
613 Ok(config)
614 }
615}
616
617impl Default for Config {
618 fn default() -> Self {
619 Self {
620 ext_trait_name: None,
621 visibility: Visibility::Inherited,
622 supertraits: None,
623 }
624 }
625}
626
627fn punctuated_from_iter<I, T, P>(i: I) -> Punctuated<T, P>
628where
629 P: Default,
630 I: IntoIterator<Item = T>,
631{
632 let mut iter = i.into_iter().peekable();
633 let mut acc = Punctuated::default();
634
635 while let Some(item) = iter.next() {
636 acc.push_value(item);
637
638 if iter.peek().is_some() {
639 acc.push_punct(P::default());
640 }
641 }
642
643 acc
644}
645
646#[cfg(test)]
647mod test {
648 #[allow(unused_imports)]
649 use super::*;
650
651 #[test]
652 fn test_ui() {
653 let t = trybuild::TestCases::new();
654 t.pass("tests/compile_pass/*.rs");
655 t.compile_fail("tests/compile_fail/*.rs");
656 }
657}