use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{
parse::{self, Parse, ParseStream, discouraged::Speculative},
parse_macro_input,
spanned::Spanned,
};
struct UnzipN {
range: Range,
trait_name: syn::Ident,
explicit_name: bool,
visibility: syn::Visibility,
}
fn int_from_expr(expr: &syn::Expr) -> parse::Result<usize> {
match expr {
syn::Expr::Lit(expr_lit) => match &expr_lit.lit {
syn::Lit::Int(lit_int) => lit_int.base10_parse(),
_ => {
let span = expr_lit.span();
Err(syn::Error::new(span, "Only integer literals are supported"))
}
},
_ => {
let span = expr.span();
Err(syn::Error::new(span, "Only integer literals are supported"))
}
}
}
fn extract_range(range: syn::ExprRange) -> parse::Result<Range> {
let start = range
.start
.as_ref()
.map(|start| int_from_expr(start))
.transpose()?
.unwrap_or(2)
.max(2);
let end = range.end.as_ref().ok_or_else(|| {
let span = range.span();
syn::Error::new(span, "Range must be terminated explicitly")
})?;
let end = int_from_expr(end)?;
Ok(match range.limits {
syn::RangeLimits::HalfOpen(..) => Range::Exclusive(start..end),
syn::RangeLimits::Closed(..) => Range::Inclusive(start..=end),
})
}
#[derive(Clone, Debug)]
enum Range {
Inclusive(std::ops::RangeInclusive<usize>),
Exclusive(std::ops::Range<usize>),
}
enum RangeIter {
Inclusive(std::ops::RangeInclusive<usize>),
Exclusive(std::ops::Range<usize>),
}
impl IntoIterator for &Range {
type Item = usize;
type IntoIter = RangeIter;
fn into_iter(self) -> Self::IntoIter {
match self {
Range::Inclusive(range_inclusive) => RangeIter::Inclusive(range_inclusive.clone()),
Range::Exclusive(range) => RangeIter::Exclusive(range.clone()),
}
}
}
impl Iterator for RangeIter {
type Item = usize;
fn next(&mut self) -> Option<Self::Item> {
match self {
RangeIter::Inclusive(range_inclusive) => range_inclusive.next(),
RangeIter::Exclusive(range) => range.next(),
}
}
}
impl Range {
fn single(&self) -> Option<usize> {
match self {
Range::Inclusive(range_inclusive) => {
if range_inclusive.start() == range_inclusive.end() {
Some(*range_inclusive.start())
} else {
None
}
}
Range::Exclusive(range) => {
if range.start + 1 == range.end {
Some(range.start)
} else {
None
}
}
}
}
}
impl Parse for UnzipN {
fn parse(input: ParseStream) -> parse::Result<Self> {
let visibility = input.parse()?;
let maybe_name = if input.peek(syn::Ident) {
Some(input.parse::<syn::Ident>()?)
} else {
None
};
let range = {
let fork = input.fork();
if let Ok(range_expr) = fork.parse::<syn::ExprRange>() {
input.advance_to(&fork);
extract_range(range_expr)?
} else {
let n: usize = input.parse::<syn::LitInt>()?.base10_parse()?;
Range::Inclusive(n..=n)
}
};
let (trait_name, explicit_name) = match maybe_name {
Some(name) => (name, true),
None => (format_ident!("Unzip"), false),
};
Ok(UnzipN {
range,
trait_name,
explicit_name,
visibility,
})
}
}
impl UnzipN {
fn make_trait(&self, n: usize, no_std: bool, trait_name: &syn::Ident) -> TokenStream {
let trait_doc =
format!("Extension trait for unzipping iterators over tuples of size {n}.",);
let generic_types = (0..n)
.map(|id| format_ident!("Type_{id}"))
.collect::<Vec<_>>();
let collections = (0..n)
.map(|id| format_ident!("Collection_{id}"))
.collect::<Vec<_>>();
let visibility = &self.visibility;
let unzip_n_vec = if no_std {
TokenStream::new()
} else {
quote!(
fn unzip_n_vec(self) -> ( #( Vec< #generic_types >, )* )
where
Self: Sized,
{
self.unzip_n()
}
)
};
quote!(
#[doc = #trait_doc]
#visibility trait #trait_name < #( #generic_types, )* > {
fn unzip_n<#(#collections,)*>(self) -> ( #(#collections,)* )
where
#( #collections: Default + Extend< #generic_types >, )*
;
#unzip_n_vec
}
)
}
fn make_implementation(&self, n: usize, trait_name: &syn::Ident) -> TokenStream {
let generic_types = (0..n)
.map(|id| format_ident!("Type_{id}"))
.collect::<Vec<_>>();
let collections = (0..n)
.map(|id| format_ident!("Collection_{id}"))
.collect::<Vec<_>>();
let containers: Vec<_> = (0..n).map(|id| format_ident!("container_{}", id)).collect();
let values: Vec<_> = (0..n).map(|id| format_ident!("val_{}", id)).collect();
quote!(
impl<Iter, #(#generic_types,)*> #trait_name <#(#generic_types,)*> for Iter
where
Iter: Iterator<Item = (#(#generic_types,)*)>,
{
fn unzip_n<#(#collections,)*>(self) -> ( #(#collections,)* )
where
#( #collections: Default + Extend< #generic_types >, )*
{
#( let mut #containers = #collections :: default() ;)*
self.for_each(|( #( #values, )* )| {
#( #containers.extend(std::iter::once(#values)) ;)*
});
(#( #containers, )* )
}
}
)
}
pub fn generate(&self, no_std: bool) -> TokenStream {
if let Some(single) = self.range.single() {
let trait_name = if self.explicit_name {
self.trait_name.clone()
} else {
format_ident!("{}{single}", self.trait_name)
};
let trait_decl = self.make_trait(single, no_std, &trait_name);
let impl_block = self.make_implementation(single, &trait_name);
quote!( #trait_decl #impl_block )
} else {
let mut tokens = TokenStream::new();
for n in &self.range {
let trait_name = format_ident!("{}{n}", self.trait_name);
tokens.extend(self.make_trait(n, no_std, &trait_name));
tokens.extend(self.make_implementation(n, &trait_name));
}
tokens
}
}
}
#[proc_macro]
pub fn unzip_n(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
parse_macro_input!(input as UnzipN).generate(false).into()
}
#[proc_macro]
pub fn unzip_n_nostd(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
parse_macro_input!(input as UnzipN).generate(true).into()
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn test_half_open_range() {
let range: syn::ExprRange = syn::parse2(quote! { 4..6 }).unwrap();
let result = extract_range(range).unwrap();
assert!(matches!(result, Range::Exclusive(r) if r == (4..6)));
}
#[test]
fn test_closed_range() {
let range: syn::ExprRange = syn::parse2(quote! { 2..=5 }).unwrap();
let result = extract_range(range).unwrap();
assert!(matches!(result, Range::Inclusive(r) if r == (2..=5)));
}
#[test]
fn test_unbounded_start_defaults_to_2() {
let range: syn::ExprRange = syn::parse2(quote! { ..10 }).unwrap();
let result = extract_range(range).unwrap();
assert!(matches!(result, Range::Exclusive(r) if r.start == 2 && r.end == 10));
}
#[test]
fn test_start_below_2_clamps_to_2() {
let range: syn::ExprRange = syn::parse2(quote! { 0..5 }).unwrap();
let result = extract_range(range).unwrap();
assert!(matches!(result, Range::Exclusive(r) if r.start == 2));
}
#[test]
fn test_unbounded_end_errors() {
let range: syn::ExprRange = syn::parse2(quote! { 4.. }).unwrap();
let result = extract_range(range);
assert!(result.is_err());
}
#[test]
fn check_parse_unzip_n() {
let _unzip_n: UnzipN = syn::parse2(quote! { pub 5 }).unwrap();
}
#[test]
fn parse_single_number_no_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { 5 }).unwrap();
assert!(!unzip_n.explicit_name);
assert_eq!(unzip_n.trait_name, "Unzip");
assert!(matches!(unzip_n.range, Range::Inclusive(r) if r == (5..=5)));
}
#[test]
fn parse_single_number_with_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { MyTrait 5 }).unwrap();
assert!(unzip_n.explicit_name);
assert_eq!(unzip_n.trait_name, "MyTrait");
assert!(matches!(unzip_n.range, Range::Inclusive(r) if r == (5..=5)));
}
#[test]
fn parse_range_no_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { 2..5 }).unwrap();
assert!(!unzip_n.explicit_name);
assert_eq!(unzip_n.trait_name, "Unzip");
assert!(matches!(unzip_n.range, Range::Exclusive(r) if r == (2..5)));
}
#[test]
fn parse_range_with_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { MyTrait 2..5 }).unwrap();
assert!(unzip_n.explicit_name);
assert_eq!(unzip_n.trait_name, "MyTrait");
assert!(matches!(unzip_n.range, Range::Exclusive(r) if r == (2..5)));
}
#[test]
fn parse_pub_visibility() {
let unzip_n: UnzipN = syn::parse2(quote! { pub 5 }).unwrap();
assert!(matches!(unzip_n.visibility, syn::Visibility::Public(_)));
}
#[test]
fn parse_pub_crate_visibility() {
let unzip_n: UnzipN = syn::parse2(quote! { pub(crate) 5 }).unwrap();
assert!(matches!(unzip_n.visibility, syn::Visibility::Restricted(_)));
}
#[test]
fn parse_inherited_visibility() {
let unzip_n: UnzipN = syn::parse2(quote! { 5 }).unwrap();
assert!(matches!(unzip_n.visibility, syn::Visibility::Inherited));
}
#[test]
fn parse_pub_with_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { pub MyTrait 5 }).unwrap();
assert!(matches!(unzip_n.visibility, syn::Visibility::Public(_)));
assert!(unzip_n.explicit_name);
assert_eq!(unzip_n.trait_name, "MyTrait");
}
#[test]
fn generate_single_no_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { 3 }).unwrap();
let output = unzip_n.generate(false).to_string();
assert!(
output.contains("trait Unzip3"),
"Expected 'trait Unzip3' in output: {}",
output
);
}
#[test]
fn generate_single_with_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { MyTrait 3 }).unwrap();
let output = unzip_n.generate(false).to_string();
assert!(
output.contains("trait MyTrait <"),
"Expected 'trait MyTrait' in output: {}",
output
);
assert!(
!output.contains("trait MyTrait3"),
"Should not contain 'trait MyTrait3' in output: {}",
output
);
}
#[test]
fn generate_range_no_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { 2..4 }).unwrap();
let output = unzip_n.generate(false).to_string();
assert!(
output.contains("trait Unzip2"),
"Expected 'trait Unzip2' in output: {}",
output
);
assert!(
output.contains("trait Unzip3"),
"Expected 'trait Unzip3' in output: {}",
output
);
assert!(
!output.contains("trait Unzip4"),
"Should not contain 'trait Unzip4' in output: {}",
output
);
}
#[test]
fn generate_range_with_explicit_name() {
let unzip_n: UnzipN = syn::parse2(quote! { MyTrait 2..4 }).unwrap();
let output = unzip_n.generate(false).to_string();
assert!(
output.contains("trait MyTrait2"),
"Expected 'trait MyTrait2' in output: {}",
output
);
assert!(
output.contains("trait MyTrait3"),
"Expected 'trait MyTrait3' in output: {}",
output
);
assert!(
!output.contains("trait MyTrait4"),
"Should not contain 'trait MyTrait4' in output: {}",
output
);
}
#[test]
fn generate_inclusive_range() {
let unzip_n: UnzipN = syn::parse2(quote! { 2..=3 }).unwrap();
let output = unzip_n.generate(false).to_string();
assert!(
output.contains("trait Unzip2"),
"Expected 'trait Unzip2' in output: {}",
output
);
assert!(
output.contains("trait Unzip3"),
"Expected 'trait Unzip3' in output: {}",
output
);
}
#[test]
fn generate_pub_visibility() {
let unzip_n: UnzipN = syn::parse2(quote! { pub 3 }).unwrap();
let output = unzip_n.generate(false).to_string();
assert!(
output.contains("pub trait Unzip3"),
"Expected 'pub trait Unzip3' in output: {}",
output
);
}
}