use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::parse::{Parse, ParseStream};
use syn::{Expr, ItemFn, Lit, Meta, Token, parse_macro_input};
struct ContractArgs {
contract_name: String,
equation_name: String,
}
impl Parse for ContractArgs {
fn parse(input: ParseStream) -> syn::Result<Self> {
let contract_lit: Lit = input.parse()?;
let contract_name = match &contract_lit {
Lit::Str(s) => s.value(),
_ => {
return Err(syn::Error::new_spanned(
contract_lit,
"expected string literal for contract name",
));
}
};
input.parse::<Token![,]>()?;
let meta: Meta = input.parse()?;
let equation_name = match &meta {
Meta::NameValue(nv) if nv.path.is_ident("equation") => match &nv.value {
Expr::Lit(expr_lit) => match &expr_lit.lit {
Lit::Str(s) => s.value(),
_ => {
return Err(syn::Error::new_spanned(
&nv.value,
"expected string literal for equation name",
));
}
},
_ => {
return Err(syn::Error::new_spanned(
&nv.value,
"expected string literal for equation name",
));
}
},
_ => {
return Err(syn::Error::new_spanned(
meta,
"expected `equation = \"...\"`",
));
}
};
Ok(ContractArgs {
contract_name,
equation_name,
})
}
}
#[proc_macro_attribute]
pub fn contract(attr: TokenStream, item: TokenStream) -> TokenStream {
let args = parse_macro_input!(attr as ContractArgs);
let input_fn = parse_macro_input!(item as ItemFn);
let env_key = make_env_key(&args.contract_name, &args.equation_name);
let const_name = format_ident!(
"_CONTRACT_CHECK_{}_{}",
args.contract_name.to_uppercase().replace(['-', '.'], "_"),
args.equation_name.to_uppercase().replace(['-', '.'], "_")
);
let contract_name = &args.contract_name;
let equation_name = &args.equation_name;
let fn_name = &input_fn.sig.ident;
let fn_name_str = fn_name.to_string();
let binding_const_name = format_ident!(
"_CONTRACT_BINDING_{}_{}",
args.contract_name.to_uppercase().replace(['-', '.'], "_"),
args.equation_name.to_uppercase().replace(['-', '.'], "_")
);
let precondition_asserts = read_contract_assertions(&env_key, "PRE", equation_name);
let postcondition_asserts = read_contract_assertions(&env_key, "POST", equation_name);
let has_postconditions = !postcondition_asserts.is_empty();
let fn_attrs = &input_fn.attrs;
let fn_vis = &input_fn.vis;
let fn_sig = &input_fn.sig;
let fn_stmts = &input_fn.block.stmts;
let body = if has_postconditions {
quote! {
#[allow(dead_code)]
const #const_name: Option<&str> = option_env!(#env_key);
#[allow(dead_code)]
const #binding_const_name: &str = concat!(
"contract=", #contract_name,
",equation=", #equation_name,
",module=", module_path!(),
",function=", #fn_name_str,
);
#(#precondition_asserts)*
let ret = { #(#fn_stmts)* };
#(#postcondition_asserts)*
ret
}
} else {
quote! {
#[allow(dead_code)]
const #const_name: Option<&str> = option_env!(#env_key);
#[allow(dead_code)]
const #binding_const_name: &str = concat!(
"contract=", #contract_name,
",equation=", #equation_name,
",module=", module_path!(),
",function=", #fn_name_str,
);
#(#precondition_asserts)*
#(#fn_stmts)*
}
};
let expanded = quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
#body
}
};
TokenStream::from(expanded)
}
fn read_contract_assertions(
env_key: &str,
kind: &str, equation_name: &str,
) -> Vec<proc_macro2::TokenStream> {
let mut asserts = Vec::new();
let count_key = format!("{env_key}_{kind}_COUNT");
let count: usize = std::env::var(&count_key)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let kind_label = if kind == "PRE" { "Pre" } else { "Post" };
for i in 0..count {
let var_key = format!("{env_key}_{kind}_{i}");
if let Ok(expr_str) = std::env::var(&var_key) {
if let Ok(expr) = expr_str.parse::<proc_macro2::TokenStream>() {
let msg = format!(
"Contract [{equation_name}] {kind_label}-condition violated: {expr_str}"
);
asserts.push(quote! {
debug_assert!(#expr, #msg);
});
}
}
}
asserts
}
#[proc_macro_attribute]
pub fn requires(attr: TokenStream, item: TokenStream) -> TokenStream {
let predicate: proc_macro2::TokenStream = attr.into();
let input_fn = parse_macro_input!(item as ItemFn);
let fn_attrs = &input_fn.attrs;
let fn_vis = &input_fn.vis;
let fn_sig = &input_fn.sig;
let fn_block = &input_fn.block;
let pred_str = predicate.to_string();
let expanded = quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
debug_assert!(#predicate, "Pre-condition violated: {}", #pred_str);
#fn_block
}
};
TokenStream::from(expanded)
}
#[proc_macro_attribute]
pub fn ensures(attr: TokenStream, item: TokenStream) -> TokenStream {
let predicate: proc_macro2::TokenStream = attr.into();
let input_fn = parse_macro_input!(item as ItemFn);
let fn_attrs = &input_fn.attrs;
let fn_vis = &input_fn.vis;
let fn_sig = &input_fn.sig;
let fn_block = &input_fn.block;
let pred_str = predicate.to_string();
let expanded = quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
let ret = #fn_block;
debug_assert!(#predicate, "Post-condition violated: {}", #pred_str);
ret
}
};
TokenStream::from(expanded)
}
#[proc_macro_attribute]
pub fn invariant(attr: TokenStream, item: TokenStream) -> TokenStream {
let predicate: proc_macro2::TokenStream = attr.into();
let input_fn = parse_macro_input!(item as ItemFn);
let fn_attrs = &input_fn.attrs;
let fn_vis = &input_fn.vis;
let fn_sig = &input_fn.sig;
let fn_block = &input_fn.block;
let pred_str = predicate.to_string();
let expanded = quote! {
#(#fn_attrs)*
#fn_vis #fn_sig {
debug_assert!(#predicate, "Invariant violated (pre): {}", #pred_str);
let ret = #fn_block;
debug_assert!(#predicate, "Invariant violated (post): {}", #pred_str);
ret
}
};
TokenStream::from(expanded)
}
#[proc_macro_attribute]
pub fn must_contract(_attr: TokenStream, item: TokenStream) -> TokenStream {
let input_fn = parse_macro_input!(item as ItemFn);
let fn_name = &input_fn.sig.ident;
let fn_name_upper = fn_name.to_string().to_uppercase();
let env_prefix = "CONTRACT_";
let has_binding =
std::env::vars().any(|(k, _)| k.starts_with(env_prefix) && k.ends_with(&fn_name_upper));
if has_binding {
quote! { #input_fn }.into()
} else {
let warning_msg = format!(
"Function `{fn_name}` has no contract binding. Add #[contract(\"...\", equation = \"...\")] or add a binding.yaml entry."
);
quote! {
#[deprecated(note = #warning_msg)]
#input_fn
}
.into()
}
}
fn make_env_key(contract: &str, equation: &str) -> String {
let contract_part = contract.to_uppercase().replace(['-', '.'], "_");
let equation_part = equation.to_uppercase().replace(['-', '.'], "_");
format!("CONTRACT_{contract_part}_{equation_part}")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_make_env_key() {
assert_eq!(
make_env_key("rmsnorm-kernel-v1", "rmsnorm"),
"CONTRACT_RMSNORM_KERNEL_V1_RMSNORM"
);
assert_eq!(
make_env_key("attention-kernel-v1", "scaled_dot_product"),
"CONTRACT_ATTENTION_KERNEL_V1_SCALED_DOT_PRODUCT"
);
assert_eq!(
make_env_key("gated-delta-net-v1", "decay"),
"CONTRACT_GATED_DELTA_NET_V1_DECAY"
);
}
#[test]
fn test_make_env_key_with_dots() {
assert_eq!(make_env_key("v1.0", "eq.1"), "CONTRACT_V1_0_EQ_1");
}
}