use proc_macro2::{Group, Ident, Span, TokenStream, TokenTree};
use quote::quote;
use std::collections::HashMap;
use crate::preprocess::consts_ctx::{ConstCtx, UserConsts};
use crate::util::bracket_is_passthrough;
use crate::util::{compile_err, compile_err_at, compile_error_str};
fn builtin_named(name: &str) -> Option<Vec<&'static str>> {
match name {
"uint" => Some(vec!["u8", "u16", "u32", "u64", "u128", "usize"]),
"int" => Some(vec!["i8", "i16", "i32", "i64", "i128", "isize"]),
"float" => Some(vec!["f32", "f64"]),
"num" => Some(vec![
"u8", "u16", "u32", "u64", "u128", "usize", "i8", "i16", "i32", "i64",
"i128", "isize", "f32", "f64",
]),
"scalar" => Some(vec![
"u8", "u16", "u32", "u64", "u128", "usize", "i8", "i16", "i32", "i64",
"i128", "isize", "f32", "f64", "bool", "char",
]),
_ => None,
}
}
fn split_range_endpoint(s: &str) -> Option<(char, u32)> {
let (fam, width_str) = s.split_at(1);
let fam = fam.chars().next()?;
let width: u32 = width_str.parse().ok()?;
let legal: &[u32] = match fam {
'u' | 'i' => &[8, 16, 32, 64, 128],
'f' => &[32, 64],
_ => return None,
};
legal.contains(&width).then_some((fam, width))
}
fn builtin_range(start: &str, end: &str) -> Result<Vec<String>, String> {
let Some((fam1, w1)) = split_range_endpoint(start) else {
return Err(format!(
"`@{}` has an invalid width (legal: u/i are 8/16/32/64/128, \
f is 32/64)",
start
));
};
let Some((fam2, w2)) = split_range_endpoint(end) else {
return Err(format!(
"`@{}` has an invalid width (legal: u/i are 8/16/32/64/128, \
f is 32/64)",
end
));
};
if fam1 != fam2 {
return Err(format!(
"range endpoint families mismatch: `{}` and `{}`",
start, end
));
}
if w1 > w2 {
return Err(format!(
"range start is greater than end: `{}..{}`",
start, end
));
}
let widths: &[u32] = match fam1 {
'u' | 'i' => &[8, 16, 32, 64, 128],
_ => &[32, 64],
};
Ok(widths
.iter()
.filter(|w| **w >= w1 && **w <= w2)
.map(|w| format!("{}{}", fam1, w))
.collect())
}
fn render_list<'a>(names: impl IntoIterator<Item = &'a str>) -> TokenTree {
let idents: Vec<Ident> =
names.into_iter().map(|s| Ident::new(s, Span::call_site())).collect();
Group::new(delimiter![[]], quote!(#(#idents),*)).into()
}
fn render_list_strings(names: impl IntoIterator<Item = String>) -> TokenTree {
let idents: Vec<Ident> =
names.into_iter().map(|s| Ident::new(&s, Span::call_site())).collect();
Group::new(delimiter![[]], quote!(#(#idents),*)).into()
}
fn try_expand_at(
tokens: &[TokenTree], ctx: ConstCtx,
) -> Result<Option<(Vec<TokenTree>, usize)>, TokenStream> {
let Some(TokenTree::Ident(name)) = tokens.get(1) else {
let sp = tokens
.first()
.map(|t| t.span())
.unwrap_or_else(proc_macro2::Span::call_site);
return Err(compile_error_str(
"batch-impl: `@` must be followed by a constant name (e.g. `@uint`, \
`@u8..u128`)",
sp,
));
};
let name_str = name.to_string();
if let Some(TokenTree::Punct(eq)) = tokens.get(2)
&& eq.as_char() == '='
{
let msg = if ctx.user_table().is_some() {
format!(
"batch-impl: constant definition `@{}=...` must appear before all \
`batch_trait!` trait segments (only the leading position can \
define)",
name_str
)
} else {
"batch-impl: `#[batch_impl]` / `#[batch_impl_only]` do not support \
custom constant definitions; custom constants are supported only \
by `batch_trait!` (leading `@name=value;` segment)"
.to_string()
};
let sp = tokens
.first()
.map(|t| t.span())
.unwrap_or_else(proc_macro2::Span::call_site);
return Err(compile_error_str(&msg, sp));
}
if let Some(TokenTree::Punct(d1)) = tokens.get(2)
&& d1.as_char() == '.'
&& d1.spacing() == proc_macro2::Spacing::Joint
&& let Some(TokenTree::Punct(d2)) = tokens.get(3)
&& d2.as_char() == '.'
{
let end_idx = if let Some(TokenTree::Punct(eq)) = tokens.get(4)
&& eq.as_char() == '='
{
5
} else {
4
};
let Some(TokenTree::Ident(end)) = tokens.get(end_idx) else {
return Err(compile_err!(
"batch-impl: range constant `@{}{}..` is missing an end point \
(e.g. `@u8..u128`)",
name_str,
".."
));
};
let types = builtin_range(&name_str, &end.to_string())
.map_err(|msg| compile_err!("batch-impl: {}", msg))?;
return Ok(Some((
vec![render_list(types.iter().map(|s| s.as_str()))],
end_idx + 1,
)));
}
if name_str == "trait" {
return match ctx.trait_full_path() {
Some(path) => Ok(Some((path.clone().into_iter().collect(), 2))),
None => Ok(None),
};
}
if let Some((kinds, default, receiver)) =
crate::preprocess::resolve_all_marker(&name_str)
{
return match ctx.trait_def() {
Some(td) => {
let ids = crate::preprocess::get_trait_item_names(
td, kinds.0, kinds.1, kinds.2, default, receiver,
);
Ok(Some((
vec![render_list_strings(ids.iter().map(|i| i.to_string()))],
2,
)))
}
None => Err(compile_err!(
"batch-impl: `@{}` is supported only by `#[batch_impl]` / \
`#[batch_impl_only]` (needs a trait definition to select \
items; `batch_trait!` is a function-like macro without one)",
name_str
)),
};
}
if let Some(expanded) = ctx.user_table().and_then(|t| t.get(&name_str)) {
return Ok(Some((expanded.clone(), 2)));
}
match builtin_named(&name_str) {
Some(types) => Ok(Some((vec![render_list(types.iter().copied())], 2))),
None => Err(compile_err_at!(
tokens[0].span(),
"batch-impl: unknown @ constant `@{}`; built-ins: `@uint` `@int` \
`@float` `@num` `@scalar` and ranges `@u8..u128` `@i8..i128` \
`@f32..f64`\
{}",
name_str,
if ctx.user_table().is_some() {
"; batch_trait! user constants must be defined before the \
reference (defining them later has no effect)"
} else {
""
}
)),
}
}
fn check_value_refs(
tokens: &[TokenTree], table: &HashMap<String, Vec<TokenTree>>, def_name: &str,
) -> Result<(), TokenStream> {
let mut i = 0;
while i < tokens.len() {
match &tokens[i] {
TokenTree::Punct(p) if p.as_char() == '@' => {
let Some(TokenTree::Ident(name)) = tokens.get(i + 1) else {
return Err(compile_error_str(
"batch-impl: inside a constant value, `@` must be followed \
by a constant name (e.g. `@uint`, `@u8..u128`)",
tokens[i].span(),
));
};
let name_str = name.to_string();
let known = name_str == "trait"
|| builtin_named(&name_str).is_some()
|| split_range_endpoint(&name_str).is_some()
|| table.contains_key(&name_str);
if !known {
return Err(compile_err!(
"batch-impl: constant `@{}` references unknown `@{}` \
(undefined or defined later; inside a constant \
definition, only built-in constants or previously \
defined constants can be referenced)",
def_name,
name_str
));
}
i += 2;
}
TokenTree::Group(g) => {
check_value_refs(
&g.stream().into_iter().collect::<Vec<_>>(),
table,
def_name,
)?;
i += 1;
}
_ => i += 1,
}
}
Ok(())
}
pub(crate) fn expand_consts(
tokens: &[TokenTree], ctx: ConstCtx,
) -> Result<Vec<TokenTree>, TokenStream> {
let mut result = vec![];
let mut i = 0;
while i < tokens.len() {
match &tokens[i] {
TokenTree::Group(g)
if g.delimiter() == delimiter![()]
|| g.delimiter() == delimiter![[]]
|| g.delimiter() == delimiter![none] =>
{
if g.delimiter() == delimiter![[]]
&& bracket_is_passthrough(tokens, i)
{
result.push(tokens[i].clone());
} else {
let inner: Vec<_> = g.stream().into_iter().collect();
result.push(
Group::new(
g.delimiter(),
expand_consts(&inner, ctx)?.into_iter().collect(),
)
.into(),
);
}
i += 1;
}
TokenTree::Punct(p) if p.as_char() == '@' => {
match try_expand_at(&tokens[i..], ctx)? {
Some((expanded, consumed)) => {
let expanded = expand_consts(&expanded, ctx)?;
result.extend(expanded);
i += consumed;
}
None => {
result.push(tokens[i].clone());
i += 1;
}
}
}
_ => {
result.push(tokens[i].clone());
i += 1;
}
}
}
Ok(result)
}
pub(crate) fn collect_user_consts(
tokens: &[TokenTree],
) -> Result<(Vec<TokenTree>, UserConsts), TokenStream> {
let mut i = 0;
let mut table = UserConsts::new();
while let Some(TokenTree::Punct(at)) = tokens.get(i) {
if at.as_char() != '@' {
break;
}
let Some(TokenTree::Ident(name)) = tokens.get(i + 1) else { break };
let Some(TokenTree::Punct(eq)) = tokens.get(i + 2) else { break };
if eq.as_char() != '=' {
break;
}
let name_str = name.to_string();
if name_str == "trait" {
return Err(compile_err!(
"batch-impl: constant name `@trait` is a reserved marker \
(segment-level substitution into a trait path); please rename"
));
}
if builtin_named(&name_str).is_some() {
return Err(compile_err!(
"batch-impl: user constant `@{}` collides with a built-in \
constant name; please rename",
name_str
));
}
let mut j = i + 3;
let mut end = None;
while j < tokens.len() {
if let TokenTree::Punct(p) = &tokens[j]
&& p.as_char() == ';'
{
end = Some(j);
break;
}
j += 1;
}
let Some(end) = end else {
return Err(compile_err!(
"batch-impl: constant definition `@{}=...` is missing the \
trailing `;`",
name_str
));
};
let value: Vec<TokenTree> = tokens[i + 3..end].to_vec();
check_value_refs(&value, &table, &name_str)?;
table.insert(name_str, value);
i = end + 1;
}
Ok((tokens[i..].to_vec(), table))
}