use proc_macro2::{Group, TokenStream, TokenTree};
use crate::util::compile_error_str;
use crate::util::{bracket_is_passthrough, is_arrow};
pub(crate) const MAX_NEST_DEPTH: usize = 128;
pub(crate) fn angle_collect(
tokens: &[TokenTree],
) -> Result<Vec<TokenTree>, TokenStream> {
angle_collect_at(tokens, 0)
}
fn angle_collect_at(
tokens: &[TokenTree], depth: usize,
) -> Result<Vec<TokenTree>, TokenStream> {
if depth > MAX_NEST_DEPTH {
let sp = tokens
.first()
.map(|t| t.span())
.unwrap_or_else(proc_macro2::Span::call_site);
return Err(compile_error_str(
&format!(
"batch-impl: nesting depth exceeds {} levels (perhaps an accidental extra bracket)",
MAX_NEST_DEPTH
),
sp,
));
}
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
match &tokens[i] {
TokenTree::Group(g) if g.delimiter() == delimiter![none] => {
let inner: Vec<_> = g.stream().into_iter().collect();
out.extend(angle_collect_at(&inner, depth + 1)?);
i += 1;
}
TokenTree::Group(g) if g.delimiter() == delimiter![()] => {
let inner: Vec<_> = g.stream().into_iter().collect();
let mut new_g = Group::new(
g.delimiter(),
angle_collect_at(&inner, depth + 1)?.into_iter().collect(),
);
new_g.set_span(g.span());
out.push(new_g.into());
i += 1;
}
TokenTree::Group(g) if g.delimiter() == delimiter![[]] => {
if bracket_is_passthrough(tokens, i) {
out.push(tokens[i].clone());
} else {
let inner: Vec<_> = g.stream().into_iter().collect();
let mut new_g = Group::new(
g.delimiter(),
angle_collect_at(&inner, depth + 1)?.into_iter().collect(),
);
new_g.set_span(g.span());
out.push(new_g.into());
}
i += 1;
}
TokenTree::Group(_) => {
out.push(tokens[i].clone());
i += 1;
}
TokenTree::Punct(p) if p.as_char() == '<' => {
let Some(close) = find_angle_close(tokens, i) else {
return Err(compile_error_str(
"batch-impl: unclosed `<` (missing matching `>`)",
tokens[i].span(),
));
};
let inner: Vec<_> = tokens[i + 1..close].to_vec();
out.push(
Group::new(
delimiter![<>],
angle_collect_at(&inner, depth + 1)?.into_iter().collect(),
)
.into(),
);
i = close + 1;
}
TokenTree::Punct(p) if p.as_char() == '>' && !is_arrow(tokens, i) => {
return Err(compile_error_str(
"batch-impl: extra `>` (missing matching `<`)",
tokens[i].span(),
));
}
_ => {
out.push(tokens[i].clone());
i += 1;
}
}
}
Ok(out)
}
fn find_angle_close(tokens: &[TokenTree], open: usize) -> Option<usize> {
let mut depth = 0usize;
for (idx, token) in tokens.iter().enumerate().skip(open + 1) {
if is_punct(token, '<') {
depth += 1;
} else if is_punct(token, '>') && !is_arrow(tokens, idx) {
if depth == 0 {
return Some(idx);
}
depth -= 1;
}
}
None
}
fn is_punct(token: &TokenTree, ch: char) -> bool {
matches!(token, TokenTree::Punct(p) if p.as_char() == ch)
}
pub(crate) fn render_angles(stream: TokenStream) -> TokenStream {
let mut out = TokenStream::new();
for tt in stream {
match tt {
TokenTree::Group(g) if g.delimiter() == delimiter![<>] => {
let inner = render_angles(g.stream());
out.extend([TokenTree::from(proc_macro2::Punct::new(
'<',
proc_macro2::Spacing::Alone,
))]);
out.extend(inner);
out.extend([TokenTree::from(proc_macro2::Punct::new(
'>',
proc_macro2::Spacing::Alone,
))]);
}
TokenTree::Group(g)
if matches!(g.delimiter(), delimiter![()] | delimiter![[]]) =>
{
let inner = render_angles(g.stream());
let mut new_g = Group::new(g.delimiter(), inner);
new_g.set_span(g.span());
out.extend([TokenTree::Group(new_g)]);
}
other => out.extend([other]),
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use proc_macro2::TokenStream as TS2;
use std::str::FromStr;
fn roundtrip(s: &str) -> String {
let ts: TS2 = FromStr::from_str(s).unwrap();
let v: Vec<_> = ts.into_iter().collect();
let collected = angle_collect(&v).unwrap();
render_angles(collected.into_iter().collect()).to_string()
}
#[test]
fn angle_nesting_limit() {
let ts: TS2 = FromStr::from_str(&format!(
"{}0{}",
"[".repeat(MAX_NEST_DEPTH + 1),
"]".repeat(MAX_NEST_DEPTH + 1)
))
.unwrap();
let v: Vec<_> = ts.into_iter().collect();
let err = angle_collect(&v).unwrap_err().to_string();
assert!(
err.contains("nesting depth exceeds"),
"expected depth-limit diagnostic, got: {err}"
);
}
#[test]
fn angle_roundtrip() {
assert_eq!(roundtrip("Vec<T>"), "Vec < T >");
assert_eq!(roundtrip("A<B<C>>"), "A < B < C > >");
assert_eq!(
roundtrip("Box<dyn Fn() + Send>"),
"Box < dyn Fn () + Send >"
);
assert_eq!(roundtrip("<T: Clone> A<T>"), "< T : Clone > A < T >");
assert_eq!(roundtrip("A<Item=T>"), "A < Item = T >");
assert_eq!(roundtrip("fn(A) -> B"), "fn (A) -> B");
}
#[test]
fn angle_unmatched_errors() {
let ts: TS2 = FromStr::from_str("A <").unwrap();
assert!(angle_collect(&ts.into_iter().collect::<Vec<_>>()).is_err());
let ts: TS2 = FromStr::from_str("A >").unwrap();
assert!(angle_collect(&ts.into_iter().collect::<Vec<_>>()).is_err());
let ts: TS2 = FromStr::from_str("m![a < b]").unwrap();
assert!(angle_collect(&ts.into_iter().collect::<Vec<_>>()).is_ok());
}
#[test]
fn bracket_passthrough_guards() {
for s in ["m![a < b]", "#[a < b]", "#[#zzz{1}]"] {
let ts: TS2 = FromStr::from_str(s).unwrap();
assert!(
angle_collect(&ts.into_iter().collect::<Vec<_>>()).is_ok(),
"input {s} should passthrough"
);
}
}
#[test]
fn none_group_flattened() {
let inner: TS2 = FromStr::from_str("Vec<T>").unwrap();
let none = proc_macro2::Group::new(delimiter![none], inner);
let collected = angle_collect(&[none.into()]).unwrap();
let rendered = render_angles(collected.into_iter().collect());
assert_eq!(rendered.to_string(), "Vec < T >");
}
#[test]
fn render_rebuilds_nested_groups() {
assert_eq!(
roundtrip("[Vec<T>, (U, W<X>)]"),
"[Vec < T > , (U , W < X >)]"
);
assert_eq!(roundtrip("{ a < b }"), "{ a < b }");
}
}