use proc_macro2::{Group, Ident, Span, TokenStream, TokenTree};
use quote::quote;
use std::collections::HashMap;
use crate::diagnostic::compile_error_str;
use crate::scan::bracket_is_passthrough;
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!(
"`@{}` 宽度非法(合法:u/i 为 8/16/32/64/128,f 为 32/64)",
start
));
};
let Some((fam2, w2)) = split_range_endpoint(end) else {
return Err(format!(
"`@{}` 宽度非法(合法:u/i 为 8/16/32/64/128,f 为 32/64)",
end
));
};
if fam1 != fam2 {
return Err(format!("范围端点族不一致:`{}` 与 `{}`", start, end));
}
if w1 > w2 {
return Err(format!("范围起点大于终点:`{}..{}`", 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 try_expand_at(
tokens: &[TokenTree], user_table: Option<&HashMap<String, Vec<TokenTree>>>,
) -> Result<(Vec<TokenTree>, usize), TokenStream> {
let Some(TokenTree::Ident(name)) = tokens.get(1) else {
return Err(compile_error_str(
"batch-impl: `@` 后必须跟常量名(如 `@uint`、`@u8..u128`)",
));
};
let name_str = name.to_string();
if let Some(TokenTree::Punct(eq)) = tokens.get(2)
&& eq.as_char() == '='
{
let msg = if user_table.is_some() {
format!(
"batch-impl: 常量定义 `@{}=...` 必须位于 `batch_trait!` 的\
所有 trait 段之前(仅前导位置可定义)",
name_str
)
} else {
"batch-impl: `#[batch_impl]` / `#[batch_impl_only]` 不支持自定义\
常量定义;自定义常量仅 `batch_trait!` 支持(前导 `@name=值;` 段)"
.to_string()
};
return Err(compile_error_str(&msg));
}
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_error_str(&format!(
"batch-impl: 范围常量 `@{}{}..` 后缺少终点(如 `@u8..u128`)",
name_str, ".."
)));
};
let types = builtin_range(&name_str, &end.to_string())
.map_err(|msg| compile_error_str(&format!("batch-impl: {}", msg)))?;
return Ok((
vec![render_list(types.iter().map(|s| s.as_str()))],
end_idx + 1,
));
}
if let Some(expanded) = user_table.and_then(|t| t.get(&name_str)) {
return Ok((expanded.clone(), 2));
}
match builtin_named(&name_str) {
Some(types) => Ok((vec![render_list(types.iter().copied())], 2)),
None => Err(compile_error_str(&format!(
"batch-impl: 未知的 @ 常量 `@{}`;内置:`@uint` `@int` `@float` `@num` \
`@scalar` 与范围 `@u8..u128` `@i8..i128` `@f32..f64`\
{}",
name_str,
if user_table.is_some() {
";batch_trait! 用户常量须在引用前定义(定义在其后不生效)"
} 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: 常量值中 `@` 后必须跟常量名(如 `@uint`、`@u8..u128`)",
));
};
let name_str = name.to_string();
let known = builtin_named(&name_str).is_some()
|| split_range_endpoint(&name_str).is_some()
|| table.contains_key(&name_str);
if !known {
return Err(compile_error_str(&format!(
"batch-impl: 常量 `@{}` 引用未知的 `@{}`(未定义或定义在其后;\
常量定义内只能引用内置常量或此前已定义的常量)",
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], user_table: Option<&HashMap<String, Vec<TokenTree>>>,
) -> 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![[]] =>
{
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, user_table)?.into_iter().collect(),
)
.into(),
);
}
i += 1;
}
TokenTree::Punct(p) if p.as_char() == '@' => {
let (expanded, consumed) = try_expand_at(&tokens[i..], user_table)?;
let expanded = expand_consts(&expanded, user_table)?;
result.extend(expanded);
i += consumed;
}
_ => {
result.push(tokens[i].clone());
i += 1;
}
}
}
Ok(result)
}
pub(crate) type UserConsts = HashMap<String, Vec<TokenTree>>;
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 builtin_named(&name_str).is_some() {
return Err(compile_error_str(&format!(
"batch-impl: 用户常量 `@{}` 与内置常量重名;请换名",
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_error_str(&format!(
"batch-impl: 常量定义 `@{}=...` 缺少结尾 `;`",
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))
}