use proc_macro2::{Group, TokenStream, TokenTree};
use crate::ast::fresh::{FreshRange, parse_range_fresh};
use crate::ast::{MAX_EXPAND, parse_grouped_fresh};
use crate::util::compile_error_str;
fn sorted_fresh(impl_names: &[TokenStream]) -> Vec<(usize, usize, &TokenStream)> {
let mut fresh_sorted: Vec<(usize, usize, &TokenStream)> = impl_names
.iter()
.filter_map(|n| {
let (g, i) = parse_grouped_fresh(&n.to_string())?;
Some((g, i, n))
})
.collect();
fresh_sorted.sort_by_key(|&(g, i, _)| (g, i));
fresh_sorted
}
#[allow(clippy::needless_lifetimes)]
pub(crate) fn group_fresh<'a>(
fresh: &'a [(usize, usize, &'a TokenStream)], group: usize, span: proc_macro2::Span,
) -> Result<&'a [(usize, usize, &'a TokenStream)], TokenStream> {
let start = fresh.iter().position(|&(g, _, _)| g == group).ok_or_else(|| {
compile_error_str(
&format!(
"batch-impl: `@{}_..` group {} does not exist — this impl has \
no generator group {}",
group, group, group,
),
span,
)
})?;
let end =
fresh[start..].iter().position(|&(g, _, _)| g != group).map_or(fresh.len(), |p| start + p);
Ok(&fresh[start..end])
}
pub(crate) fn range_count(
range: FreshRange, scope_len: usize, span: proc_macro2::Span,
) -> Result<usize, TokenStream> {
match range.end {
Some(end) => {
if end >= scope_len || range.start > end {
return Err(compile_error_str(
&format!(
"batch-impl: `@{}..={}` out of range — this scope has {} fresh \
generics (numbered from 0 in document order)",
range.start, end, scope_len,
),
span,
));
}
let count = end - range.start + 1;
if count > MAX_EXPAND {
return Err(compile_error_str(
&format!(
"batch-impl: `@{}..={}` expands to {} elements (max {})",
range.start, end, count, MAX_EXPAND,
),
span,
));
}
Ok(count)
}
None => Ok(scope_len.saturating_sub(range.start)),
}
}
fn range_entries<'a>(
range: FreshRange, fresh: &'a [(usize, usize, &'a TokenStream)],
) -> Result<Vec<&'a TokenStream>, TokenStream> {
let slice: &[(usize, usize, &TokenStream)] = match range.group {
Some(l) => group_fresh(fresh, l, proc_macro2::Span::call_site())?,
None => fresh,
};
let count = range_count(range, slice.len(), proc_macro2::Span::call_site())?;
Ok(slice[range.start..range.start + count].iter().map(|&(_, _, n)| n).collect())
}
pub(crate) fn expand_range_refs(
tokens: TokenStream, impl_names: &[TokenStream],
) -> Result<TokenStream, TokenStream> {
let fresh = sorted_fresh(impl_names);
let v = tokens.into_iter().collect::<Vec<_>>();
let out = expand_at(&v, &fresh, 0)?;
Ok(out.into_iter().collect())
}
pub(crate) fn expand_range_decls(
impl_generics: &mut Vec<(TokenStream, Option<crate::ast::Ty>)>, impl_names: &[TokenStream],
) -> Result<(), TokenStream> {
let fresh = sorted_fresh(impl_names);
let mut out: Vec<(TokenStream, Option<crate::ast::Ty>)> = vec![];
for (name, bound) in impl_generics.iter() {
let s = name.to_string();
if let Some(range) = parse_range_fresh(&s) {
for n in range_entries(range, &fresh)? {
out.push(((*n).clone(), None));
}
} else {
out.push((name.clone(), bound.clone()));
}
}
*impl_generics = out;
Ok(())
}
fn expand_at(
tokens: &[TokenTree], fresh: &[(usize, usize, &TokenStream)], depth: usize,
) -> Result<Vec<TokenTree>, TokenStream> {
if depth > crate::util::MAX_NEST_DEPTH {
return Err(crate::util::depth_err(tokens, ""));
}
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
match &tokens[i] {
TokenTree::Ident(id) => {
let s = id.to_string();
if let Some(range) = parse_range_fresh(&s) {
let mut first = true;
for n in range_entries(range, fresh)? {
if !first {
out.push(TokenTree::Punct(proc_macro2::Punct::new(
',',
proc_macro2::Spacing::Alone,
)));
}
first = false;
out.extend(n.clone());
}
i += 1;
continue;
}
out.push(tokens[i].clone());
i += 1;
}
TokenTree::Group(g) => {
if depth + 1 > crate::util::MAX_NEST_DEPTH {
return Err(crate::util::depth_err(&tokens[i..i + 1], ""));
}
let inner = g.stream().into_iter().collect::<Vec<_>>();
let mut ng = Group::new(
g.delimiter(),
expand_at(&inner, fresh, depth + 1)?.into_iter().collect(),
);
ng.set_span(g.span());
out.push(TokenTree::Group(ng));
i += 1;
}
other => {
out.push(other.clone());
i += 1;
}
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use quote::quote;
fn names() -> Vec<TokenStream> {
vec![
quote!(_Param_0_0_BatchGen_),
quote!(_Param_1_0_BatchGen_),
quote!(_Param_1_1_BatchGen_),
]
}
#[test]
fn open_range_in_generic_args() {
let ts: TokenStream = "Wrapper < _Param_0_With_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(
out.to_string(),
"Wrapper < _Param_0_0_BatchGen_ , _Param_1_0_BatchGen_ , _Param_1_1_BatchGen_ >"
);
}
#[test]
fn closed_range() {
let ts: TokenStream = "Wrapper < _Param_1_With_2_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(out.to_string(), "Wrapper < _Param_1_0_BatchGen_ , _Param_1_1_BatchGen_ >");
}
#[test]
fn open_range_with_offset() {
let ts: TokenStream = "Wrapper < _Param_1_With_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(out.to_string(), "Wrapper < _Param_1_0_BatchGen_ , _Param_1_1_BatchGen_ >");
}
#[test]
fn tuple_range() {
let ts: TokenStream = "( _Param_0_With_BatchGen_ , u8 )".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(
out.to_string(),
"(_Param_0_0_BatchGen_ , _Param_1_0_BatchGen_ , _Param_1_1_BatchGen_ , u8)"
);
}
#[test]
fn closed_range_out_of_bounds_errors() {
let ts: TokenStream = "Wrapper < _Param_1_With_5_BatchGen_ >".parse().unwrap();
assert!(expand_range_refs(ts, &names()).is_err());
}
#[test]
fn plain_fresh_names_untouched() {
let ts: TokenStream = "Wrapper < _Param_0_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(out.to_string(), "Wrapper < _Param_0_BatchGen_ >");
}
#[test]
fn decl_position_open_range() {
let mut gens: Vec<(TokenStream, Option<crate::ast::Ty>)> =
vec![("_Param_0_With_BatchGen_".parse().unwrap(), None)];
expand_range_decls(&mut gens, &names()).unwrap();
let got: Vec<String> = gens.iter().map(|(n, _)| n.to_string()).collect();
assert_eq!(got, ["_Param_0_0_BatchGen_", "_Param_1_0_BatchGen_", "_Param_1_1_BatchGen_"]);
}
#[test]
fn decl_position_closed_range() {
let mut gens: Vec<(TokenStream, Option<crate::ast::Ty>)> =
vec![("_Param_1_With_2_BatchGen_".parse().unwrap(), None)];
expand_range_decls(&mut gens, &names()).unwrap();
let got: Vec<String> = gens.iter().map(|(n, _)| n.to_string()).collect();
assert_eq!(got, ["_Param_1_0_BatchGen_", "_Param_1_1_BatchGen_"]);
}
#[test]
fn decl_position_mixed_with_plain() {
let mut gens: Vec<(TokenStream, Option<crate::ast::Ty>)> =
vec![("X".parse().unwrap(), None), ("_Param_0_With_BatchGen_".parse().unwrap(), None)];
expand_range_decls(&mut gens, &names()).unwrap();
let got: Vec<String> = gens.iter().map(|(n, _)| n.to_string()).collect();
assert_eq!(
got,
["X", "_Param_0_0_BatchGen_", "_Param_1_0_BatchGen_", "_Param_1_1_BatchGen_"]
);
}
#[test]
fn decl_position_closed_out_of_bounds_errors() {
let mut gens: Vec<(TokenStream, Option<crate::ast::Ty>)> =
vec![("_Param_0_With_5_BatchGen_".parse().unwrap(), None)];
assert!(expand_range_decls(&mut gens, &names()).is_err());
}
#[test]
fn grouped_range_open_in_generic_args() {
let ts: TokenStream = "Wrapper < _Param_0_0_With_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(out.to_string(), "Wrapper < _Param_0_0_BatchGen_ >");
}
#[test]
fn grouped_range_open_group1() {
let ts: TokenStream = "Wrapper < _Param_1_0_With_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(out.to_string(), "Wrapper < _Param_1_0_BatchGen_ , _Param_1_1_BatchGen_ >");
}
#[test]
fn grouped_range_closed_in_generic_args() {
let ts: TokenStream = "Wrapper < _Param_1_0_With_0_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(out.to_string(), "Wrapper < _Param_1_0_BatchGen_ >");
}
#[test]
fn grouped_range_second_group_tail() {
let ts: TokenStream = "Wrapper < _Param_1_1_With_BatchGen_ >".parse().unwrap();
let out = expand_range_refs(ts, &names()).unwrap();
assert_eq!(out.to_string(), "Wrapper < _Param_1_1_BatchGen_ >");
}
#[test]
fn grouped_range_unknown_group_errors() {
let ts: TokenStream = "Wrapper < _Param_3_0_With_BatchGen_ >".parse().unwrap();
assert!(expand_range_refs(ts, &names()).is_err());
}
#[test]
fn grouped_range_out_of_group_errors() {
let ts: TokenStream = "Wrapper < _Param_0_2_With_3_BatchGen_ >".parse().unwrap();
assert!(expand_range_refs(ts, &names()).is_err());
}
}