use proc_macro2::{Group, Ident, Literal, Punct, Spacing, Span, TokenStream, TokenTree};
use crate::codegen::VarSeg;
use crate::util::{MAX_NEST_DEPTH, compile_error_str, depth_err, is_punct_at};
pub(crate) fn expand_repeat_blocks(
tokens: TokenStream, segs: &[VarSeg],
) -> Result<TokenStream, TokenStream> {
let v = fix_literal_at(tokens.into_iter().collect::<Vec<_>>());
expand_stream(&v, segs, 0).map(|out| out.into_iter().collect())
}
fn fix_literal_at(tokens: Vec<TokenTree>) -> Vec<TokenTree> {
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
if let TokenTree::Literal(lit) = &tokens[i] {
let s = lit.to_string();
if s.ends_with('.')
&& is_punct_at(&tokens, i + 1, '@')
&& let Ok(n) = s[..s.len() - 1].parse::<u64>()
{
out.push(TokenTree::Literal(Literal::u64_unsuffixed(n)));
out.push(TokenTree::Punct(Punct::new('.', Spacing::Alone)));
i += 1;
continue;
}
}
if let TokenTree::Group(g) = &tokens[i] {
let inner = fix_literal_at(g.stream().into_iter().collect::<Vec<_>>());
let mut ng = Group::new(g.delimiter(), inner.into_iter().collect());
ng.set_span(g.span());
out.push(TokenTree::Group(ng));
i += 1;
continue;
}
out.push(tokens[i].clone());
i += 1;
}
out
}
fn expand_stream(
tokens: &[TokenTree], segs: &[VarSeg], depth: usize,
) -> Result<Vec<TokenTree>, TokenStream> {
if depth > MAX_NEST_DEPTH {
return Err(depth_err(tokens, ""));
}
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
if is_punct_at(tokens, i, '@') {
if let Some(TokenTree::Ident(id)) = tokens.get(i + 1)
&& let Some(TokenTree::Group(g)) = tokens.get(i + 2)
&& g.delimiter() == delimiter![()]
&& is_punct_at(tokens, i + 3, '.')
&& is_punct_at(tokens, i + 4, '.')
{
let body = g.stream().into_iter().collect::<Vec<_>>();
out.extend(expand_block(&body, segs, depth + 1, Some(id.clone()))?);
i += 5;
continue;
}
if let Some(TokenTree::Group(g)) = tokens.get(i + 1)
&& g.delimiter() == delimiter![()]
&& is_punct_at(tokens, i + 2, '.')
&& is_punct_at(tokens, i + 3, '.')
{
let body = g.stream().into_iter().collect::<Vec<_>>();
out.extend(expand_block(&body, segs, depth + 1, None)?);
i += 4;
continue;
}
return Err(compile_error_str(
"batch-impl: `@` inside an impl body must start a repeat block \
`@(...)..` (or `@ident(...)..` with the driving segment declared)",
tokens[i].span(),
));
}
if let TokenTree::Group(g) = &tokens[i] {
if depth + 1 > MAX_NEST_DEPTH {
return Err(depth_err(&tokens[i..i + 1], ""));
}
let inner = g.stream().into_iter().collect::<Vec<_>>();
let expanded = expand_stream(&inner, segs, depth + 1)?;
let mut ng = Group::new(g.delimiter(), expanded.into_iter().collect());
ng.set_span(g.span());
out.push(TokenTree::Group(ng));
i += 1;
continue;
}
out.push(tokens[i].clone());
i += 1;
}
Ok(out)
}
fn expand_block(
body: &[TokenTree], segs: &[VarSeg], depth: usize, driver: Option<Ident>,
) -> Result<Vec<TokenTree>, TokenStream> {
if depth > MAX_NEST_DEPTH {
return Err(depth_err(body, " in a repeat block"));
}
let body = expand_nested(body, segs, depth)?;
let (inner_prefixes, inner_len) = collect_drivers(&body, segs)?;
let len = match driver {
Some(id) => {
let prefix = id.to_string();
let Some(seg) = segs.iter().find(|s| s.prefix == prefix) else {
return Err(compile_error_str(
&format!(
"batch-impl: repeat block driver `@{}` is not a variadic \
segment (the `impl{{...}}` template declares no `{}@..`)",
prefix, prefix,
),
id.span(),
));
};
for p in &inner_prefixes {
if *p != prefix {
return Err(compile_error_str(
&format!(
"batch-impl: repeat block driver `@{}` conflicts with the \
inner segment reference `@{}` (they must be the same)",
prefix, p,
),
id.span(),
));
}
}
seg.len
}
None => match inner_len {
Some(l) => l,
None if segs.len() == 1 => segs[0].len,
None => {
return Err(compile_error_str(
"batch-impl: a repeat block needs a driving segment to determine \
its length — write `@ident(...)..` with the segment declared, or \
reference a segment inside",
body.first().map_or_else(Span::call_site, |t| t.span()),
));
}
},
};
let mut out = vec![];
for round in 0..len {
out.extend(substitute(&body, segs, round, depth + 1)?);
}
Ok(out)
}
fn expand_nested(
tokens: &[TokenTree], segs: &[VarSeg], depth: usize,
) -> Result<Vec<TokenTree>, TokenStream> {
if depth > MAX_NEST_DEPTH {
return Err(depth_err(tokens, " in a repeat block"));
}
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
if is_punct_at(tokens, i, '@')
&& let Some(TokenTree::Ident(id)) = tokens.get(i + 1)
&& let Some(TokenTree::Group(g)) = tokens.get(i + 2)
&& g.delimiter() == delimiter![()]
&& is_punct_at(tokens, i + 3, '.')
&& is_punct_at(tokens, i + 4, '.')
{
let body = g.stream().into_iter().collect::<Vec<_>>();
out.extend(expand_block(&body, segs, depth + 1, Some(id.clone()))?);
i += 5;
continue;
}
if is_punct_at(tokens, i, '@')
&& let Some(TokenTree::Group(g)) = tokens.get(i + 1)
&& g.delimiter() == delimiter![()]
&& is_punct_at(tokens, i + 2, '.')
&& is_punct_at(tokens, i + 3, '.')
{
let body = g.stream().into_iter().collect::<Vec<_>>();
out.extend(expand_block(&body, segs, depth + 1, None)?);
i += 4;
continue;
}
if let TokenTree::Group(g) = &tokens[i] {
if depth + 1 > MAX_NEST_DEPTH {
return Err(depth_err(&tokens[i..i + 1], ""));
}
let inner = g.stream().into_iter().collect::<Vec<_>>();
let expanded = expand_nested(&inner, segs, depth + 1)?;
let mut ng = Group::new(g.delimiter(), expanded.into_iter().collect());
ng.set_span(g.span());
out.push(TokenTree::Group(ng));
i += 1;
continue;
}
out.push(tokens[i].clone());
i += 1;
}
Ok(out)
}
fn collect_drivers(
tokens: &[TokenTree], segs: &[VarSeg],
) -> Result<(Vec<String>, Option<usize>), TokenStream> {
let mut prefixes: Vec<String> = vec![];
let mut len: Option<usize> = None;
let mut i = 0;
while i < tokens.len() {
if is_punct_at(tokens, i, '@') {
if matches!(tokens.get(i + 1), Some(TokenTree::Literal(_))) {
i += 2;
continue;
}
let Some(TokenTree::Ident(id)) = tokens.get(i + 1) else {
return Err(compile_error_str(
"batch-impl: `@` inside a repeat block must be followed by a \
segment name (`@ident`) or an index (`@N`)",
tokens[i].span(),
));
};
let prefix = id.to_string();
let Some(seg) = segs.iter().find(|s| s.prefix == prefix) else {
return Err(compile_error_str(
&format!(
"batch-impl: repeat block references unknown variadic segment \
`@{}` (the `impl{{...}}` template declares no `{}@..`)",
prefix, prefix,
),
id.span(),
));
};
if !prefixes.contains(&prefix) {
prefixes.push(prefix);
}
match len {
None => len = Some(seg.len),
Some(l) if l != seg.len => {
return Err(compile_error_str(
&format!(
"batch-impl: repeat block segments have different lengths \
({} vs {}); all referenced segments must be equal-length",
l, seg.len,
),
id.span(),
));
}
_ => {}
}
i += 2;
continue;
}
if let TokenTree::Group(g) = &tokens[i] {
let inner = g.stream().into_iter().collect::<Vec<_>>();
let (p, l) = collect_drivers(&inner, segs)?;
for p in p {
if !prefixes.contains(&p) {
prefixes.push(p);
}
}
match (len, l) {
(None, _) => len = l,
(Some(a), Some(b)) if a != b => {
return Err(compile_error_str(
"batch-impl: repeat block segments have different lengths; all \
referenced segments must be equal-length",
tokens[i].span(),
));
}
_ => {}
}
i += 1;
continue;
}
i += 1;
}
Ok((prefixes, len))
}
fn substitute(
tokens: &[TokenTree], segs: &[VarSeg], round: usize, depth: usize,
) -> Result<Vec<TokenTree>, TokenStream> {
if depth > MAX_NEST_DEPTH {
return Err(depth_err(tokens, ""));
}
let mut out = vec![];
let mut i = 0;
while i < tokens.len() {
if is_punct_at(tokens, i, '@') {
match tokens.get(i + 1) {
Some(TokenTree::Ident(id)) => {
let prefix = id.to_string();
let Some(seg) = segs.iter().find(|s| s.prefix == prefix) else {
return Err(compile_error_str(
&format!("batch-impl: unknown variadic segment `@{}`", prefix),
id.span(),
));
};
let name = Ident::new(&format!("{}{}", prefix, seg.start + round), id.span());
out.push(TokenTree::Ident(name));
i += 2;
continue;
}
Some(TokenTree::Literal(lit)) => {
let Ok(n) = lit.to_string().parse::<usize>() else {
return Err(compile_error_str(
"batch-impl: `@` inside a repeat block must be followed by a \
segment name (`@ident`) or a number (`@0`)",
lit.span(),
));
};
let val = Literal::u64_unsuffixed((n + round) as u64);
out.push(TokenTree::Literal(val));
i += 2;
continue;
}
_ => {
return Err(compile_error_str(
"batch-impl: `@` inside a repeat block must be followed by a \
segment name (`@ident`) or an index (`@N`)",
tokens[i].span(),
));
}
}
}
if let TokenTree::Group(g) = &tokens[i] {
if depth + 1 > MAX_NEST_DEPTH {
return Err(depth_err(&tokens[i..i + 1], ""));
}
let inner = g.stream().into_iter().collect::<Vec<_>>();
let substituted = substitute(&inner, segs, round, depth + 1)?;
let mut ng = Group::new(g.delimiter(), substituted.into_iter().collect());
ng.set_span(g.span());
out.push(TokenTree::Group(ng));
i += 1;
continue;
}
out.push(tokens[i].clone());
i += 1;
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn segs() -> Vec<VarSeg> {
vec![
VarSeg { prefix: "A".into(), start: 0, len: 3 },
VarSeg { prefix: "B".into(), start: 1, len: 2 },
]
}
fn expand(s: &str) -> Result<String, String> {
let ts = s.parse::<TokenStream>().map_err(|e| e.to_string())?;
expand_repeat_blocks(ts, &segs()).map(|o| o.to_string()).map_err(|e| e.to_string())
}
#[test]
fn single_segment_rounds() {
assert_eq!(
expand("@(@A::f(&self.@0),)..").unwrap(),
"A0 :: f (& self .0) , A1 :: f (& self .1) , A2 :: f (& self .2) ,"
);
}
#[test]
fn offset_start_name_numbering() {
assert_eq!(
expand("@(@B::f(&self.@1),)..").unwrap(),
"B1 :: f (& self .1) , B2 :: f (& self .2) ,"
);
}
#[test]
fn multi_segment_parallel_rounds() {
let segs = vec![
VarSeg { prefix: "A".into(), start: 0, len: 2 },
VarSeg { prefix: "B".into(), start: 2, len: 2 },
];
let ts = "@(@A + @B,)..".parse::<TokenStream>().unwrap();
let out = expand_repeat_blocks(ts, &segs).unwrap().to_string();
assert_eq!(out, "A0 + B2 , A1 + B3 ,");
}
#[test]
fn unequal_segment_lengths_error() {
let segs = vec![
VarSeg { prefix: "A".into(), start: 0, len: 3 },
VarSeg { prefix: "B".into(), start: 1, len: 2 },
];
let ts = "@(@A + @B,)..".parse::<TokenStream>().unwrap();
assert!(expand_repeat_blocks(ts, &segs).is_err());
}
#[test]
fn nested_cartesian() {
let out = expand("@(@A::f(&self.@0) @(@B::g(&self.@1),)..)..").unwrap();
assert_eq!(
out,
"A0 :: f (& self .0) B1 :: g (& self .1) , B2 :: g (& self .2) , \
A1 :: f (& self .1) B1 :: g (& self .1) , B2 :: g (& self .2) , \
A2 :: f (& self .2) B1 :: g (& self .1) , B2 :: g (& self .2) ,"
);
}
#[test]
fn no_trailing_separator_concatenates() {
assert_eq!(expand("@(@A)..").unwrap(), "A0 A1 A2");
}
#[test]
fn float_literal_at_path_fixed() {
let segs = vec![VarSeg { prefix: "A".into(), start: 0, len: 2 }];
let ts = "@(@A::from(self.0.@0),)..".parse::<TokenStream>().unwrap();
let out = expand_repeat_blocks(ts, &segs).unwrap().to_string();
assert_eq!(out, "A0 :: from (self . 0 . 0) , A1 :: from (self . 0 . 1) ,");
}
#[test]
fn plain_body_passthrough() {
let s = "fn combine (& self , rhs : & Self) -> Self { todo ! () }";
assert_eq!(expand(s).unwrap(), s);
}
#[test]
fn declared_driver_cursor_only() {
let segs = vec![VarSeg { prefix: "A".into(), start: 0, len: 3 }];
let ts = "@A(self.@0,)..".parse::<TokenStream>().unwrap();
let out = expand_repeat_blocks(ts, &segs).unwrap().to_string();
assert_eq!(out, "self .0 , self .1 , self .2 ,");
}
#[test]
fn cursor_only_single_segment() {
let segs = vec![VarSeg { prefix: "A".into(), start: 0, len: 2 }];
let ts = "@(self.@0,)..".parse::<TokenStream>().unwrap();
let out = expand_repeat_blocks(ts, &segs).unwrap().to_string();
assert_eq!(out, "self .0 , self .1 ,");
}
#[test]
fn cursor_only_multi_segment_errors() {
let segs = vec![
VarSeg { prefix: "A".into(), start: 0, len: 2 },
VarSeg { prefix: "B".into(), start: 2, len: 2 },
];
let ts = "@(self.@0,)..".parse::<TokenStream>().unwrap();
assert!(expand_repeat_blocks(ts, &segs).is_err());
}
#[test]
fn declared_driver_conflict_errors() {
let segs = vec![
VarSeg { prefix: "A".into(), start: 0, len: 2 },
VarSeg { prefix: "B".into(), start: 2, len: 2 },
];
let ts = "@A(@B::f(),)..".parse::<TokenStream>().unwrap();
assert!(expand_repeat_blocks(ts, &segs).is_err());
}
#[test]
fn bare_at_errors() {
assert!(expand("x @ 0").is_err());
}
#[test]
fn unknown_segment_errors() {
assert!(expand("@(@X::f(),)..").is_err());
}
#[test]
fn no_driver_errors() {
assert!(expand("@(@0,)..").is_err());
}
}