use proc_macro2::{Punct, Spacing, TokenStream, TokenTree};
use crate::ast::{MAX_EXPAND, parse_grouped_fresh};
use crate::util::{compile_err, compile_error_str};
pub(crate) fn resolve_where_at(
pred: &TokenStream, impl_names: &[TokenStream],
) -> Result<TokenStream, TokenStream> {
let mut fresh_sorted: Vec<&TokenStream> = impl_names
.iter()
.filter(|n| parse_grouped_fresh(&n.to_string()).is_some())
.collect();
fresh_sorted.sort_by_key(|n| parse_grouped_fresh(&n.to_string()).unwrap());
let tokens = pred.clone().into_iter().collect::<Vec<_>>();
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
if let TokenTree::Punct(p) = &tokens[i]
&& p.as_char() == '@'
{
match tokens.get(i + 1) {
Some(TokenTree::Ident(id)) if id == "all_fresh" => {
if fresh_sorted.is_empty() {
return Err(compile_error_str(
"batch-impl: `@all_fresh` in a where predicate but this impl has no fresh generics",
tokens[i].span(),
));
}
if fresh_sorted.len() > MAX_EXPAND {
return Err(compile_err!(
"batch-impl: `@all_fresh` expands to {} predicates (max {}); use `@N..M` for a subset",
fresh_sorted.len(),
MAX_EXPAND
));
}
let tail = tokens[i + 2..].to_vec();
let comma = TokenTree::Punct(Punct::new(',', Spacing::Alone));
for (k, &name) in fresh_sorted.iter().enumerate() {
if k > 0 {
out.push(comma.clone());
}
out.extend(name.clone());
out.extend(tail.iter().cloned());
}
i = tokens.len();
continue;
}
Some(TokenTree::Literal(lit)) => {
let s = lit.to_string();
if let Ok(start) = s.parse::<usize>()
&& matches!(tokens.get(i + 2), Some(TokenTree::Punct(p)) if p.as_char() == '.')
&& matches!(tokens.get(i + 3), Some(TokenTree::Punct(p)) if p.as_char() == '.')
{
let inclusive = matches!(tokens.get(i + 4), Some(TokenTree::Punct(p)) if p.as_char() == '=');
let end_idx = if inclusive { i + 5 } else { i + 4 };
let Some(TokenTree::Literal(end_lit)) = tokens.get(end_idx)
else {
return Err(compile_error_str(
"batch-impl: a `@N..M` range in a where predicate must end with a number (e.g. `@0..=2`)",
tokens[i].span(),
));
};
let Ok(end) = end_lit.to_string().parse::<usize>() else {
return Err(compile_error_str(
"batch-impl: a `@N..M` range in a where predicate must end with a number (e.g. `@0..=2`)",
end_lit.span(),
));
};
let count = if inclusive {
end.saturating_sub(start) + 1
} else {
end.saturating_sub(start)
};
if count == 0 {
return Err(compile_err!(
"batch-impl: `@{}..{}` is an empty range (start \
not below end); no predicates will be generated",
start,
end
));
}
if end >= fresh_sorted.len() || start > end {
return Err(compile_err!(
"batch-impl: `@{}..{}` out of range in a where \
predicate (impl has {} fresh generics, numbered \
from 0 in document order)",
start,
end,
fresh_sorted.len()
));
}
if count > MAX_EXPAND {
return Err(compile_err!(
"batch-impl: `@{}..{}` expands to {} predicates (max {})",
start,
end,
count,
MAX_EXPAND
));
}
let tail = tokens[end_idx + 1..].to_vec();
let comma = TokenTree::Punct(Punct::new(',', Spacing::Alone));
for (offset, &name) in
fresh_sorted[start..start + count].iter().enumerate()
{
if offset > 0 {
out.push(comma.clone());
}
out.extend(name.clone());
out.extend(tail.iter().cloned());
}
i = tokens.len();
continue;
}
if let Ok(idx) = s.parse::<usize>() {
let Some(&name) = fresh_sorted.get(idx) else {
return Err(compile_err!(
"batch-impl: `@{}` out of range in a where predicate \
(impl has {} fresh generics, numbered from 0 in \
document order; user-written params are addressed \
by name)",
idx,
fresh_sorted.len()
));
};
out.extend(name.clone());
i += 2;
continue;
}
if let Some((g, pos)) = s.split_once('_')
&& let (Ok(g), Ok(pos)) =
(g.parse::<usize>(), pos.parse::<usize>())
{
let target = format!("_Param_{}_{}_BatchGen_", g, pos);
let Some(name) =
impl_names.iter().find(|n| n.to_string() == target)
else {
return Err(compile_err!(
"batch-impl: `@{}` in a where predicate — this \
impl has no group {} position {} (grouped \
fresh names are `_Param_{{g}}_{{i}}_BatchGen_`; \
use `@N` for the impl's document-order fresh)",
s,
g,
pos
));
};
out.extend(name.clone());
i += 2;
continue;
}
return Err(compile_error_str(
"batch-impl: `@` in a where predicate must be followed by \
a position digit (e.g. `@0` or `@0_1`)",
tokens[i].span(),
));
}
_ => {
return Err(compile_error_str(
"batch-impl: `@` in a where predicate must be a position digit (e.g. `@0` or `@0_1`)",
tokens[i].span(),
));
}
}
} else {
out.push(tokens[i].clone());
i += 1;
}
}
Ok(out.into_iter().collect())
}
#[cfg(test)]
mod tests {
use crate::analyze::extract_trait_bounds;
use crate::ast::*;
use crate::codegen::generate_impl;
use quote::quote;
use syn::parse_quote;
#[test]
fn const_param_where_predicate_no_error() {
let trait_def: syn::ItemTrait = parse_quote!(
trait WhereArr<T, const N: usize>
where
[T; N]: Sized,
{
}
);
let tb = extract_trait_bounds(&trait_def);
let target = TyTuple(vec![]).to_ty();
let trait_ty = TyTrait(
quote!(WhereArr),
TyTypeParam {
params: vec![
(Box::new(TyPrimitive(quote!(T)).to_ty()), None),
(Box::new(TyPrimitive(quote!(N)).to_ty()), None),
],
bindings: vec![],
},
);
let wrapped = TyWithTrait(trait_ty, target.into());
let impl_ty = TyWithType(
TyTypeParam {
params: vec![
(Box::new(TyPrimitive(quote!(T)).to_ty()), None),
(
Box::new(TyPrimitive(quote!(const N)).to_ty()),
Some(TyPrimitive(quote!(usize)).to_ty()),
),
],
bindings: vec![],
},
wrapped.into(),
)
.into();
let out =
generate_impl(impl_ty, "e!(WhereArr), false, &tb, &[]).to_string();
assert!(
!out.contains("compile_error"),
"expansion must not contain compile_error: {out}"
);
assert!(
out.contains("where [T ; N] : Sized"),
"missing where predicate: {out}"
);
assert!(
out.contains("impl < T , const N : usize > WhereArr < T , N >"),
"unexpected impl generics: {out}"
);
}
}