use super::*;
#[derive(Debug, Clone, PartialEq)]
pub(crate) enum UsageGrade {
Erased,
Linear,
Exact(u32),
Unlimited,
}
impl std::fmt::Display for UsageGrade {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
UsageGrade::Erased => write!(f, "erased (grade 0)"),
UsageGrade::Linear => write!(f, "linear (grade 1)"),
UsageGrade::Exact(n) => write!(f, "exact (grade {n})"),
UsageGrade::Unlimited => write!(f, "unlimited (grade ω)"),
}
}
}
#[derive(Debug, Clone, Default)]
pub(crate) struct UsageTracker {
usages: HashMap<String, (UsageGrade, u32, Range<usize>)>,
}
impl UsageTracker {
pub fn new() -> Self {
Self {
usages: HashMap::new(),
}
}
pub fn declare(&mut self, name: String, grade: UsageGrade, span: Range<usize>) {
self.usages.insert(name, (grade, 0, span));
}
pub fn use_var(&mut self, name: &str) {
if let Some((_grade, count, _span)) = self.usages.get_mut(name) {
*count += 1;
}
}
pub fn check(&self) -> Vec<TypeError> {
let mut errors = Vec::new();
for (name, (grade, count, span)) in &self.usages {
match grade {
UsageGrade::Erased => {
if *count > 0 {
errors.push(TypeError {
code: "A05002".into(),
message: format!(
"erased variable `{name}` must not be used at runtime, \
but was used {count} time(s)"
),
span: span.clone(),
secondary: None,
suggestion: None,
});
}
}
UsageGrade::Linear => {
if *count == 0 {
errors.push(TypeError {
code: "A05002".into(),
message: format!("linear variable `{name}` was never used"),
span: span.clone(),
secondary: None,
suggestion: None,
});
} else if *count > 1 {
errors.push(TypeError {
code: "A05001".into(),
message: format!(
"linear variable `{name}` used {count} times, \
but must be used exactly once"
),
span: span.clone(),
secondary: None,
suggestion: None,
});
}
}
UsageGrade::Exact(expected) => {
if count != expected {
errors.push(TypeError {
code: "A05003".into(),
message: format!(
"variable `{name}` used {count} time(s), \
but must be used exactly {expected} time(s)"
),
span: span.clone(),
secondary: None,
suggestion: None,
});
}
}
UsageGrade::Unlimited => {
}
}
}
errors.sort_by_key(|e| e.span.start);
errors
}
pub fn get_count(&self, name: &str) -> Option<u32> {
self.usages.get(name).map(|(_, count, _)| *count)
}
pub fn set_count(&mut self, name: &str, count: u32) {
if let Some((_grade, c, _span)) = self.usages.get_mut(name) {
*c = count;
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct LinearContext {
tracker: UsageTracker,
}
impl LinearContext {
pub fn new(tracker: UsageTracker) -> Self {
Self { tracker }
}
pub fn use_var(&mut self, name: &str) {
self.tracker.use_var(name);
}
pub fn declare(&mut self, name: String, grade: UsageGrade, span: Range<usize>) {
self.tracker.declare(name, grade, span);
}
#[cfg(test)]
pub fn get_count(&self, name: &str) -> Option<u32> {
self.tracker.get_count(name)
}
pub fn fork(&self) -> (LinearContext, LinearContext) {
(self.clone(), self.clone())
}
pub fn merge(&mut self, branch_a: &LinearContext, branch_b: &LinearContext) -> Vec<TypeError> {
let mut errors = Vec::new();
let base_state: Vec<(String, UsageGrade, u32, Range<usize>)> = self
.tracker
.usages
.iter()
.map(|(name, (grade, count, span))| (name.clone(), grade.clone(), *count, span.clone()))
.collect();
for (name, grade, base_count, span) in &base_state {
let a_count = branch_a.tracker.get_count(name).unwrap_or(*base_count);
let b_count = branch_b.tracker.get_count(name).unwrap_or(*base_count);
let delta_a = a_count.saturating_sub(*base_count);
let delta_b = b_count.saturating_sub(*base_count);
if matches!(grade, UsageGrade::Linear | UsageGrade::Exact(_)) && delta_a != delta_b {
errors.push(TypeError {
code: "A05004".into(),
message: format!(
"linear variable `{name}` used inconsistently across branches: \
used {delta_a} time(s) in one branch, {delta_b} time(s) in the other"
),
span: span.clone(),
secondary: None,
suggestion: None,
});
}
let merged_count = base_count + std::cmp::max(delta_a, delta_b);
self.tracker.set_count(name, merged_count);
}
errors
}
pub fn merge_arms(&mut self, arm_contexts: &[LinearContext]) -> Vec<TypeError> {
if arm_contexts.is_empty() {
return Vec::new();
}
let mut errors = Vec::new();
let base_state: Vec<(String, UsageGrade, u32, Range<usize>)> = self
.tracker
.usages
.iter()
.map(|(name, (grade, count, span))| (name.clone(), grade.clone(), *count, span.clone()))
.collect();
for (name, grade, base_count, span) in &base_state {
let deltas: Vec<u32> = arm_contexts
.iter()
.map(|arm| {
arm.tracker
.get_count(name)
.unwrap_or(*base_count)
.saturating_sub(*base_count)
})
.collect();
if matches!(grade, UsageGrade::Linear | UsageGrade::Exact(_)) {
let first = deltas[0];
for (i, &delta) in deltas.iter().enumerate().skip(1) {
if delta != first {
errors.push(TypeError {
code: "A05004".into(),
message: format!(
"linear variable `{name}` used inconsistently across match arms: \
used {first} time(s) in arm 1, {delta} time(s) in arm {}",
i + 1
),
span: span.clone(),
secondary: None,
suggestion: None,
});
break; }
}
}
let max_delta = deltas.iter().copied().max().unwrap_or(0);
self.tracker.set_count(name, base_count + max_delta);
}
errors
}
pub fn check(&self) -> Vec<TypeError> {
self.tracker.check()
}
}
pub(crate) fn check_expr_linearity(expr: &SpExpr, ctx: &mut LinearContext) -> Vec<TypeError> {
let mut errors = Vec::new();
check_expr_linearity_inner(expr, ctx, &mut errors);
errors
}
fn check_expr_linearity_inner(expr: &SpExpr, ctx: &mut LinearContext, errors: &mut Vec<TypeError>) {
match &expr.node {
Expr::Ident(name) => {
ctx.use_var(name);
}
Expr::Literal(_) => {}
Expr::Field(receiver, _field) => {
check_expr_linearity_inner(receiver, ctx, errors);
}
Expr::MethodCall { receiver, args, .. } => {
check_expr_linearity_inner(receiver, ctx, errors);
for arg in args {
check_expr_linearity_inner(arg, ctx, errors);
}
}
Expr::Call { func, args } => {
check_expr_linearity_inner(func, ctx, errors);
for arg in args {
check_expr_linearity_inner(arg, ctx, errors);
}
}
Expr::Index { expr: base, index } => {
check_expr_linearity_inner(base, ctx, errors);
check_expr_linearity_inner(index, ctx, errors);
}
Expr::BinOp { lhs, rhs, .. } => {
check_expr_linearity_inner(lhs, ctx, errors);
check_expr_linearity_inner(rhs, ctx, errors);
}
Expr::UnaryOp { expr: inner, .. } => {
check_expr_linearity_inner(inner, ctx, errors);
}
Expr::Old(_inner) => {
}
Expr::Forall {
var: _,
domain: _,
body: _,
}
| Expr::Exists {
var: _,
domain: _,
body: _,
} => {
}
Expr::If {
cond,
then_branch,
else_branch,
} => {
check_expr_linearity_inner(cond, ctx, errors);
let (mut ctx_then, mut ctx_else) = ctx.fork();
check_expr_linearity_inner(then_branch, &mut ctx_then, errors);
if let Some(else_br) = else_branch {
check_expr_linearity_inner(else_br, &mut ctx_else, errors);
}
let merge_errors = ctx.merge(&ctx_then, &ctx_else);
errors.extend(merge_errors);
}
Expr::List(items) => {
for item in items {
check_expr_linearity_inner(item, ctx, errors);
}
}
Expr::Cast { expr: inner, .. } => {
check_expr_linearity_inner(inner, ctx, errors);
}
Expr::Block(exprs) => {
for e in exprs {
check_expr_linearity_inner(e, ctx, errors);
}
}
Expr::Ghost(_inner) => {
}
Expr::Apply { .. } => {
}
Expr::Match { scrutinee, arms } => {
check_expr_linearity_inner(scrutinee, ctx, errors);
if arms.is_empty() {
return;
}
let mut arm_contexts: Vec<LinearContext> = Vec::new();
for arm in arms {
let mut arm_ctx = ctx.clone();
check_expr_linearity_inner(&arm.body, &mut arm_ctx, errors);
arm_contexts.push(arm_ctx);
}
let merge_errs = ctx.merge_arms(&arm_contexts);
errors.extend(merge_errs);
}
Expr::Let { value, body, .. } => {
check_expr_linearity_inner(value, ctx, errors);
check_expr_linearity_inner(body, ctx, errors);
}
Expr::Tuple(elems) => {
for e in elems {
check_expr_linearity_inner(e, ctx, errors);
}
}
Expr::Raw(_) => {
}
}
}
fn infer_usage_grade(ty_tokens: &[String]) -> UsageGrade {
for (i, t) in ty_tokens.iter().enumerate() {
match t.as_str() {
"linear" => return UsageGrade::Linear,
"ghost" | "erased" => return UsageGrade::Erased,
"exact" => {
if let Some(n_str) = ty_tokens.get(i + 1)
&& let Ok(n) = n_str.parse::<u32>()
{
return UsageGrade::Exact(n);
}
return UsageGrade::Linear;
}
_ => {}
}
}
UsageGrade::Unlimited
}
pub(crate) fn declare_linear_params_from_expr(
expr: &SpExpr,
tracker: &mut UsageTracker,
span: &std::ops::Range<usize>,
) {
match &expr.node {
Expr::Raw(tokens) => {
declare_linear_params_from_raw(tokens, tracker, span);
}
Expr::Call { args, .. } => {
for arg in args {
declare_linear_single_param(arg, tracker, span);
}
}
Expr::Cast { expr: inner, ty } => {
if ty.contains("linear")
&& let Expr::Ident(name) = &inner.as_ref().node
{
tracker.declare(name.clone(), UsageGrade::Linear, span.clone());
}
}
Expr::Ident(_) => {}
Expr::Tuple(items) | Expr::Block(items) => {
for item in items {
declare_linear_single_param(item, tracker, span);
}
}
_ => {}
}
}
fn declare_linear_single_param(
expr: &SpExpr,
tracker: &mut UsageTracker,
span: &std::ops::Range<usize>,
) {
match &expr.node {
Expr::Cast { expr: inner, ty } => {
if ty.contains("linear")
&& let Expr::Ident(name) = &inner.as_ref().node
{
tracker.declare(name.clone(), UsageGrade::Linear, span.clone());
}
}
Expr::Raw(tokens) => {
declare_linear_params_from_raw(tokens, tracker, span);
}
_ => {}
}
}
fn declare_linear_params_from_raw(
tokens: &[String],
tracker: &mut UsageTracker,
span: &std::ops::Range<usize>,
) {
let mut i = 0;
while i < tokens.len() {
let sep = tokens.get(i + 1).map(|s| s.as_str());
if i + 2 < tokens.len()
&& matches!(sep, Some(":" | "as"))
&& tokens[i + 2..]
.iter()
.take_while(|t| *t != ",")
.any(|t| t == "linear")
{
let name = &tokens[i];
tracker.declare(name.clone(), UsageGrade::Linear, span.clone());
while i < tokens.len() && tokens[i] != "," {
i += 1;
}
}
i += 1;
}
}
pub(crate) fn run_linearity_checks_source(
source: &assura_parser::ast::SourceFile,
) -> Vec<TypeError> {
use assura_parser::ast::{ClauseKind, Decl, ServiceItem};
let mut errors = Vec::new();
for decl in &source.decls {
if let Decl::Contract(c) = &decl.node {
let mut tracker = UsageTracker::new();
for clause in &c.clauses {
if clause.kind == ClauseKind::Input {
declare_linear_params_from_expr(&clause.body, &mut tracker, &decl.span);
}
}
let mut ctx = LinearContext::new(tracker);
for clause in &c.clauses {
if matches!(
clause.kind,
ClauseKind::Requires | ClauseKind::Ensures | ClauseKind::Invariant
) {
errors.extend(check_expr_linearity(&clause.body, &mut ctx));
}
}
errors.extend(ctx.check());
} else if matches!(&decl.node, Decl::FnDef(_) | Decl::Extern(_)) {
let tracker = UsageTracker::new();
let mut ctx = LinearContext::new(tracker);
for param in decl.node.params() {
let p_tokens = param.ty.as_ref().map(|t| t.to_tokens()).unwrap_or_default();
let grade = infer_usage_grade(&p_tokens);
if grade != UsageGrade::Unlimited {
ctx.declare(param.name.clone(), grade, decl.span.clone());
}
}
for clause in decl.node.clauses() {
errors.extend(check_expr_linearity(&clause.body, &mut ctx));
}
errors.extend(ctx.check());
} else if let Decl::Service(s) = &decl.node {
for item in &s.items {
if let ServiceItem::Operation { clauses, .. } | ServiceItem::Query { clauses, .. } =
item
{
let tracker = UsageTracker::new();
let mut ctx = LinearContext::new(tracker);
for clause in clauses {
errors.extend(check_expr_linearity(&clause.body, &mut ctx));
}
errors.extend(ctx.check());
}
}
}
}
errors
}
#[cfg(test)]
mod tests {
use super::*;
use assura_parser::ast::Spanned;
fn span() -> Range<usize> {
0..10
}
fn ident(s: &str) -> SpExpr {
Spanned::no_span(Expr::Ident(s.to_string()))
}
fn int_lit(n: i64) -> SpExpr {
Spanned::no_span(Expr::Literal(Literal::Int(n.to_string())))
}
#[test]
fn tracker_linear_used_once_ok() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
t.use_var("x");
let errs = t.check();
assert!(errs.is_empty());
}
#[test]
fn tracker_linear_never_used() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let errs = t.check();
assert_eq!(errs.len(), 1);
assert_eq!(errs[0].code.as_ref(), "A05002");
}
#[test]
fn tracker_linear_used_twice() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
t.use_var("x");
t.use_var("x");
let errs = t.check();
assert_eq!(errs.len(), 1);
assert_eq!(errs[0].code.as_ref(), "A05001");
}
#[test]
fn tracker_erased_used_at_runtime() {
let mut t = UsageTracker::new();
t.declare("g".into(), UsageGrade::Erased, span());
t.use_var("g");
let errs = t.check();
assert_eq!(errs.len(), 1);
assert_eq!(errs[0].code.as_ref(), "A05002");
}
#[test]
fn tracker_erased_not_used_ok() {
let mut t = UsageTracker::new();
t.declare("g".into(), UsageGrade::Erased, span());
let errs = t.check();
assert!(errs.is_empty());
}
#[test]
fn tracker_exact_correct_count() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Exact(3), span());
t.use_var("x");
t.use_var("x");
t.use_var("x");
let errs = t.check();
assert!(errs.is_empty());
}
#[test]
fn tracker_exact_wrong_count() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Exact(2), span());
t.use_var("x");
let errs = t.check();
assert_eq!(errs.len(), 1);
assert_eq!(errs[0].code.as_ref(), "A05003");
}
#[test]
fn tracker_unlimited_any_count_ok() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Unlimited, span());
t.use_var("x");
t.use_var("x");
t.use_var("x");
t.use_var("x");
let errs = t.check();
assert!(errs.is_empty());
}
#[test]
fn tracker_get_count() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
assert_eq!(t.get_count("x"), Some(0));
t.use_var("x");
assert_eq!(t.get_count("x"), Some(1));
assert_eq!(t.get_count("unknown"), None);
}
#[test]
fn tracker_use_undeclared_is_noop() {
let mut t = UsageTracker::new();
t.use_var("unknown"); let errs = t.check();
assert!(errs.is_empty());
}
#[test]
fn ctx_fork_merge_consistent() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let (mut a, mut b) = ctx.fork();
a.use_var("x");
b.use_var("x");
let errs = ctx.merge(&a, &b);
assert!(errs.is_empty()); }
#[test]
fn ctx_fork_merge_inconsistent() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let (mut a, b) = ctx.fork();
a.use_var("x");
let errs = ctx.merge(&a, &b);
assert_eq!(errs.len(), 1);
assert_eq!(errs[0].code.as_ref(), "A05004");
}
#[test]
fn ctx_merge_arms_consistent() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let mut arm1 = ctx.clone();
let mut arm2 = ctx.clone();
let mut arm3 = ctx.clone();
arm1.use_var("x");
arm2.use_var("x");
arm3.use_var("x");
let errs = ctx.merge_arms(&[arm1, arm2, arm3]);
assert!(errs.is_empty());
}
#[test]
fn ctx_merge_arms_inconsistent() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let mut arm1 = ctx.clone();
let arm2 = ctx.clone(); arm1.use_var("x");
let errs = ctx.merge_arms(&[arm1, arm2]);
assert_eq!(errs.len(), 1);
assert_eq!(errs[0].code.as_ref(), "A05004");
}
#[test]
fn ctx_merge_arms_empty() {
let t = UsageTracker::new();
let mut ctx = LinearContext::new(t);
let errs = ctx.merge_arms(&[]);
assert!(errs.is_empty());
}
#[test]
fn linearity_ident_records_use() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let errs = check_expr_linearity(&ident("x"), &mut ctx);
assert!(errs.is_empty());
assert_eq!(ctx.get_count("x"), Some(1));
}
#[test]
fn linearity_if_forks_context() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let expr = Spanned::no_span(Expr::If {
cond: Box::new(Spanned::no_span(Expr::Literal(Literal::Bool(true)))),
then_branch: Box::new(ident("x")),
else_branch: Some(Box::new(ident("x"))),
});
let errs = check_expr_linearity(&expr, &mut ctx);
assert!(errs.is_empty()); }
#[test]
fn linearity_if_one_branch_only() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let expr = Spanned::no_span(Expr::If {
cond: Box::new(Spanned::no_span(Expr::Literal(Literal::Bool(true)))),
then_branch: Box::new(ident("x")),
else_branch: Some(Box::new(int_lit(0))),
});
let errs = check_expr_linearity(&expr, &mut ctx);
assert!(!errs.is_empty());
assert!(errs.iter().any(|e| e.code.as_ref() == "A05004"));
}
#[test]
fn linearity_old_does_not_count() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let expr = Spanned::no_span(Expr::Old(Box::new(ident("x"))));
let errs = check_expr_linearity(&expr, &mut ctx);
assert!(errs.is_empty());
assert_eq!(ctx.get_count("x"), Some(0)); }
#[test]
fn linearity_ghost_does_not_count() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let expr = Spanned::no_span(Expr::Ghost(Box::new(ident("x"))));
let errs = check_expr_linearity(&expr, &mut ctx);
assert!(errs.is_empty());
assert_eq!(ctx.get_count("x"), Some(0));
}
#[test]
fn linearity_quantifier_does_not_count() {
let mut t = UsageTracker::new();
t.declare("x".into(), UsageGrade::Linear, span());
let mut ctx = LinearContext::new(t);
let expr = Spanned::no_span(Expr::Forall {
var: "i".into(),
domain: Box::new(ident("x")),
body: Box::new(Spanned::no_span(Expr::Literal(Literal::Bool(true)))),
});
let errs = check_expr_linearity(&expr, &mut ctx);
assert!(errs.is_empty());
assert_eq!(ctx.get_count("x"), Some(0));
}
#[test]
fn usage_grade_display() {
assert_eq!(UsageGrade::Erased.to_string(), "erased (grade 0)");
assert_eq!(UsageGrade::Linear.to_string(), "linear (grade 1)");
assert_eq!(UsageGrade::Exact(3).to_string(), "exact (grade 3)");
assert!(UsageGrade::Unlimited.to_string().contains("unlimited"));
}
}