use anyhow::{Context, Result};
use heck::ToLowerCamelCase;
use serde_json::{json, Value};
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use syn::spanned::Spanned;
const MARKER: &str = "pinoc:constant";
const INTEGER_TYPES: &[&str] = &[
"u8", "u16", "u32", "u64", "u128", "i8", "i16", "i32", "i64", "i128",
];
const MAX_SAFE_INTEGER: i128 = (1 << 53) - 1;
#[derive(Debug, Clone, PartialEq)]
pub struct Constant {
pub name: String,
pub docs: Vec<String>,
pub ty: String,
pub value: i128,
}
impl Constant {
pub fn to_shank(&self) -> Value {
json!({ "name": self.name, "type": self.ty, "value": self.value.to_string() })
}
pub fn to_codama(&self) -> Value {
let mut node = json!({
"kind": "constantNode",
"name": self.name.to_lower_camel_case(),
"type": { "kind": "numberTypeNode", "format": self.ty, "endian": "le" },
"value": { "kind": "numberValueNode", "number": self.value as i64 },
});
if !self.docs.is_empty() {
node["docs"] = json!(self.docs);
}
node
}
}
pub fn find_constants(src_dir: &Path) -> Result<Vec<Constant>> {
let mut files = Vec::new();
collect_files(src_dir, &mut files)?;
let mut parsed = Vec::new();
let mut all: HashMap<String, syn::Expr> = HashMap::new();
for path in files {
let Ok(src) = std::fs::read_to_string(&path) else {
continue;
};
let Ok(file) = syn::parse_file(&src) else {
continue;
};
let mut consts = Vec::new();
collect_consts(&file.items, &mut consts);
for item in &consts {
all.insert(item.ident.to_string(), (*item.expr).clone());
}
parsed.push((path, src, consts));
}
let mut out = Vec::new();
for (path, src, consts) in &parsed {
let lines: Vec<&str> = src.lines().collect();
let mut claimed = Vec::new();
for item in consts {
let Some((marker_line, type_override)) = marker_above(item, &lines) else {
continue;
};
claimed.push(marker_line);
let name = item.ident.to_string();
let at = format!(
"{}:{}",
path.display(),
item.const_token.span().start().line
);
let ty = constant_type(item, type_override.as_deref())
.with_context(|| format!("constant `{name}` at {at}"))?;
let value = eval(&item.expr, &all, 0).with_context(|| {
format!("constant `{name}` at {at}: its value is not an integer expression pinoc can evaluate (literals, `+ - * / <<`, casts, and other consts)")
})?;
check_range(&ty, value).with_context(|| format!("constant `{name}` at {at}"))?;
out.push(Constant {
name,
docs: doc_lines(&item.attrs),
ty,
value,
});
}
for (index, line) in lines.iter().enumerate() {
if marker_of(line).is_some() && !claimed.contains(&index) {
anyhow::bail!(
"`// {MARKER}` at {}:{} is not directly above a `const` item",
path.display(),
index + 1
);
}
}
}
Ok(out)
}
fn collect_consts(items: &[syn::Item], out: &mut Vec<syn::ItemConst>) {
for item in items {
match item {
syn::Item::Const(c) => out.push(c.clone()),
syn::Item::Mod(m) => {
if let Some((_, inner)) = &m.content {
collect_consts(inner, out);
}
}
_ => {}
}
}
}
fn marker_of(line: &str) -> Option<Option<String>> {
let rest = line
.trim()
.strip_prefix("//")?
.trim()
.strip_prefix(MARKER)?;
let rest = rest.trim();
if rest.is_empty() {
return Some(None);
}
let ty = rest.strip_prefix('(')?.split(')').next()?.trim();
Some(Some(ty.to_string()))
}
fn marker_above(item: &syn::ItemConst, lines: &[&str]) -> Option<(usize, Option<String>)> {
let keyword = item.const_token.span().start().line.checked_sub(1)?;
let mut top = item.span().start().line.checked_sub(1)?.min(keyword);
while top > 0 && lines.get(top - 1)?.trim().starts_with("//") {
top -= 1;
}
(top..keyword).find_map(|index| Some((index, marker_of(lines.get(index)?)?)))
}
fn constant_type(item: &syn::ItemConst, type_override: Option<&str>) -> Result<String> {
let declared = match item.ty.as_ref() {
syn::Type::Path(p) => p.path.get_ident().map(|i| i.to_string()),
_ => None,
};
let chosen = type_override.map(str::to_string).or(declared.clone());
match chosen {
Some(ty) if INTEGER_TYPES.contains(&ty.as_str()) => Ok(ty),
Some(ty) if type_override.is_some() => anyhow::bail!(
"`// {MARKER}({ty})` is not an integer type; use one of {}",
INTEGER_TYPES.join(", ")
),
_ => anyhow::bail!(
"its type `{}` is not a fixed-width integer, so it has no wire type; state one in the marker, e.g. `// {MARKER}(u8)`",
declared.unwrap_or_else(|| "?".to_string())
),
}
}
fn check_range(ty: &str, value: i128) -> Result<()> {
let bits: u32 = ty[1..].parse()?;
let (min, max) = if ty.starts_with('u') {
(
0,
if bits == 128 {
i128::MAX
} else {
(1i128 << bits) - 1
},
)
} else if bits == 128 {
(i128::MIN, i128::MAX)
} else {
(-(1i128 << (bits - 1)), (1i128 << (bits - 1)) - 1)
};
if value < min || value > max {
anyhow::bail!("value {value} does not fit in `{ty}`");
}
if value.abs() > MAX_SAFE_INTEGER {
anyhow::bail!(
"value {value} is beyond 2^53, which an IDL number cannot carry exactly; a JavaScript client would read a rounded value"
);
}
Ok(())
}
fn eval(expr: &syn::Expr, consts: &HashMap<String, syn::Expr>, depth: usize) -> Result<i128> {
if depth > 16 {
anyhow::bail!("the expression refers to itself");
}
let next = |e: &syn::Expr| eval(e, consts, depth + 1);
match expr {
syn::Expr::Lit(syn::ExprLit {
lit: syn::Lit::Int(int),
..
}) => Ok(int.base10_parse::<i128>()?),
syn::Expr::Paren(e) => next(&e.expr),
syn::Expr::Group(e) => next(&e.expr),
syn::Expr::Cast(e) => next(&e.expr),
syn::Expr::Unary(syn::ExprUnary {
op: syn::UnOp::Neg(_),
expr,
..
}) => next(expr)?
.checked_neg()
.ok_or_else(|| anyhow::anyhow!("overflow")),
syn::Expr::Binary(bin) => {
let (left, right) = (next(&bin.left)?, next(&bin.right)?);
let result = match bin.op {
syn::BinOp::Add(_) => left.checked_add(right),
syn::BinOp::Sub(_) => left.checked_sub(right),
syn::BinOp::Mul(_) => left.checked_mul(right),
syn::BinOp::Div(_) => left.checked_div(right),
syn::BinOp::Shl(_) => u32::try_from(right).ok().and_then(|r| left.checked_shl(r)),
_ => anyhow::bail!("unsupported operator"),
};
result.ok_or_else(|| anyhow::anyhow!("overflow"))
}
syn::Expr::Path(path) => {
let name = path
.path
.segments
.last()
.map(|s| s.ident.to_string())
.unwrap_or_default();
let referenced = consts
.get(&name)
.ok_or_else(|| anyhow::anyhow!("`{name}` is not a const in this crate"))?;
next(referenced)
}
_ => anyhow::bail!("unsupported expression"),
}
}
fn doc_lines(attrs: &[syn::Attribute]) -> Vec<String> {
attrs
.iter()
.filter(|a| a.path().is_ident("doc"))
.filter_map(|a| match &a.meta {
syn::Meta::NameValue(nv) => match &nv.value {
syn::Expr::Lit(syn::ExprLit {
lit: syn::Lit::Str(s),
..
}) => Some(s.value().trim().to_string()),
_ => None,
},
_ => None,
})
.collect()
}
fn collect_files(dir: &Path, out: &mut Vec<PathBuf>) -> Result<()> {
for entry in std::fs::read_dir(dir)? {
let path = entry?.path();
if path.is_dir() {
collect_files(&path, out)?;
} else if path.extension().and_then(|e| e.to_str()) == Some("rs") {
out.push(path);
}
}
out.sort();
Ok(())
}