use assura_parser::ast::{ClauseKind, Decl, Expr, ServiceItem, SpExpr};
use crate::{
Type, TypeEnv, TypeError, check_ghost_fn_effects, check_lemma_fn_effects, infer_expr,
infer_expr_spanned, parse_type_tokens,
};
pub(crate) fn register_input_clause_params(body: &SpExpr, env: &mut TypeEnv) {
use assura_parser::ast::extract_clause_params;
for param in extract_clause_params(body) {
if param.ty.is_none() {
if env.lookup(¶m.name).is_none() {
env.insert(param.name, Type::Unknown);
}
} else {
let parsed = crate::convert::resolve_type_opt(param.ty.as_ref());
env.insert(param.name, parsed);
}
}
}
pub(crate) fn collect_input_param_types(body: &SpExpr, out: &mut Vec<Type>) {
use assura_parser::ast::extract_clause_params;
for param in extract_clause_params(body) {
if param.ty.is_none() {
out.push(Type::Unknown);
} else {
out.push(crate::convert::resolve_type_opt(param.ty.as_ref()));
}
}
}
pub(crate) fn bind_pattern_vars(
pattern: &assura_parser::ast::Pattern,
scrutinee_ty: &Type,
env: &mut TypeEnv,
) {
match pattern {
assura_parser::ast::Pattern::Ident(name) => {
env.insert(name.clone(), scrutinee_ty.clone());
}
assura_parser::ast::Pattern::Constructor { name, fields } => {
let param_types: Vec<Type> = match env.lookup(name) {
Some(Type::Fn { params, .. }) => params.clone(),
_ => Vec::new(),
};
for (i, field) in fields.iter().enumerate() {
let field_ty = param_types.get(i).cloned().unwrap_or(Type::Unknown);
bind_pattern_vars(field, &field_ty, env);
}
}
assura_parser::ast::Pattern::Tuple(pats) => {
if let Type::Tuple(elem_tys) = scrutinee_ty {
for (i, pat) in pats.iter().enumerate() {
let elem_ty = elem_tys.get(i).cloned().unwrap_or(Type::Unknown);
bind_pattern_vars(pat, &elem_ty, env);
}
} else {
for pat in pats {
bind_pattern_vars(pat, &Type::Unknown, env);
}
}
}
assura_parser::ast::Pattern::Wildcard | assura_parser::ast::Pattern::Literal(_) => {}
}
}
pub(crate) fn env_with_result(env: &TypeEnv, result_ty: &Type) -> TypeEnv {
let mut new_env = env.clone();
new_env.insert("result".to_string(), result_ty.clone());
new_env
}
pub(crate) fn extract_output_type_from_body(body: &SpExpr) -> Type {
match &body.node {
Expr::Cast { ty, .. } => parse_type_tokens(std::slice::from_ref(ty)),
Expr::Raw(tokens) => {
if let Some(colon_pos) = tokens.iter().position(|t| t == ":") {
let type_tokens: Vec<String> = tokens[colon_pos + 1..].to_vec();
if !type_tokens.is_empty() {
let ty = parse_type_tokens(&type_tokens);
if !ty.is_indeterminate() {
return ty;
}
}
}
Type::Unknown
}
Expr::Call { args, .. } => {
for arg in args {
let ty = extract_output_type_from_body(arg);
if !ty.is_indeterminate() {
return ty;
}
}
Type::Unknown
}
_ => {
let env = TypeEnv::new();
if let Ok(ty) = infer_expr(body, &env) {
ty
} else {
Type::Unknown
}
}
}
}
pub(crate) fn extract_contract_output_type(c: &assura_parser::ast::ContractDecl) -> Type {
for clause in &c.clauses {
if clause.kind == ClauseKind::Output {
let ty = extract_output_type_from_body(&clause.body);
if !ty.is_indeterminate() {
return ty;
}
}
}
Type::Unknown
}
pub(crate) fn check_clause_bodies(
source: &assura_parser::ast::SourceFile,
env: &TypeEnv,
) -> Vec<TypeError> {
let mut errors = Vec::new();
for decl in &source.decls {
let span = &decl.span;
match &decl.node {
Decl::Contract(c) => {
let output_ty = extract_contract_output_type(c);
let mut contract_env = env.clone();
for clause in &c.clauses {
if clause.kind == ClauseKind::Requires
|| clause.kind == ClauseKind::Input
|| clause.kind == ClauseKind::Ensures
{
register_input_clause_params(&clause.body, &mut contract_env);
}
}
for p in &c.fn_params {
if let Some(te) = &p.ty {
contract_env.insert(p.name.clone(), crate::convert::type_from_expr(te));
}
}
let ensures_env = env_with_result(&contract_env, &output_ty);
for clause in &c.clauses {
let clause_env = if clause.kind == ClauseKind::Ensures {
&ensures_env
} else {
&contract_env
};
check_clause_expr(&clause.kind, &clause.body, clause_env, &mut errors, span);
}
}
Decl::FnDef(f) => {
if f.is_ghost {
check_ghost_fn_effects(f, span, &mut errors);
}
if f.is_lemma {
check_lemma_fn_effects(f, span, &mut errors);
}
let ret_ty = crate::convert::resolve_type_opt(f.return_ty.as_ref());
let fn_env = env_with_result(env, &ret_ty);
for clause in &f.clauses {
let clause_env = if clause.kind == ClauseKind::Ensures {
&fn_env
} else {
env
};
check_clause_expr(&clause.kind, &clause.body, clause_env, &mut errors, span);
}
}
Decl::Extern(ex) => {
let ret_ty = crate::convert::resolve_type_opt(ex.return_ty.as_ref());
let ext_env = env_with_result(env, &ret_ty);
for clause in &ex.clauses {
let clause_env = if clause.kind == ClauseKind::Ensures {
&ext_env
} else {
env
};
check_clause_expr(&clause.kind, &clause.body, clause_env, &mut errors, span);
}
}
Decl::Bind(b) => {
let ret_ty = crate::convert::resolve_type_opt(b.return_ty.as_ref());
let bind_env = env_with_result(env, &ret_ty);
for clause in &b.clauses {
let clause_env = if clause.kind == ClauseKind::Ensures {
&bind_env
} else {
env
};
check_clause_expr(&clause.kind, &clause.body, clause_env, &mut errors, span);
}
}
Decl::Service(s) => {
let mut svc_env = env.clone();
svc_env.insert("self".to_string(), Type::Named(s.name.clone()));
for item in &s.items {
let clauses = match item {
ServiceItem::Operation { clauses, .. }
| ServiceItem::Query { clauses, .. } => clauses.as_slice(),
ServiceItem::Invariant(expr) => {
check_clause_expr(
&ClauseKind::Invariant,
expr,
&svc_env,
&mut errors,
span,
);
continue;
}
ServiceItem::Other { body, .. } => {
collect_expr_errors(body, &svc_env, &mut errors, span);
continue;
}
_ => continue,
};
let mut op_env = svc_env.clone();
let mut output_ty = Type::Unit;
for clause in clauses {
if clause.kind == ClauseKind::Input {
register_input_clause_params(&clause.body, &mut op_env);
}
if clause.kind == ClauseKind::Output {
let ty = extract_output_type_from_body(&clause.body);
if !ty.is_indeterminate() {
output_ty = ty;
}
}
}
let ensures_env = env_with_result(&op_env, &output_ty);
for clause in clauses {
let clause_env = if clause.kind == ClauseKind::Ensures {
&ensures_env
} else {
&op_env
};
check_clause_expr(
&clause.kind,
&clause.body,
clause_env,
&mut errors,
span,
);
}
}
}
Decl::Block { body, .. } => {
for clause in body {
check_clause_expr(&clause.kind, &clause.body, env, &mut errors, span);
}
}
Decl::TypeDef(_) | Decl::EnumDef(_) | Decl::Prophecy(_) | Decl::CodecRegistry(_) => {}
}
}
errors
}
fn collect_expr_errors(
expr: &SpExpr,
env: &TypeEnv,
errors: &mut Vec<TypeError>,
ctx_span: &std::ops::Range<usize>,
) {
match infer_expr_spanned(expr, env, ctx_span.clone()) {
Ok(_) => {}
Err(e) => {
errors.push(e);
}
}
}
fn clause_requires_bool(kind: &ClauseKind) -> bool {
matches!(
kind,
ClauseKind::Requires | ClauseKind::Ensures | ClauseKind::Invariant | ClauseKind::Rule
)
}
fn clause_kind_label(kind: &ClauseKind) -> &'static str {
match kind {
ClauseKind::Requires => "requires",
ClauseKind::Ensures => "ensures",
ClauseKind::Invariant => "invariant",
ClauseKind::Rule => "rule",
_ => "clause",
}
}
pub(crate) fn check_clause_expr(
kind: &ClauseKind,
body: &SpExpr,
env: &TypeEnv,
errors: &mut Vec<TypeError>,
ctx_span: &std::ops::Range<usize>,
) {
let body_span = if body.span != (0..0) {
body.span.clone()
} else {
ctx_span.clone()
};
match infer_expr_spanned(body, env, body_span.clone()) {
Ok(ty) => {
if clause_requires_bool(kind) && !ty.is_indeterminate() && ty != Type::Bool {
errors.push(TypeError {
code: "A03006".into(),
message: format!(
"{} clause must be Bool, found `{ty}`",
clause_kind_label(kind),
),
span: body_span.clone(),
secondary: None,
suggestion: None,
});
}
}
Err(e) => {
errors.push(e);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use assura_parser::ast::{Expr, Literal, Pattern, Spanned};
#[test]
fn register_input_params_typed() {
let body = Spanned::no_span(Expr::Call {
func: Box::new(Spanned::no_span(Expr::Ident("input".into()))),
args: vec![
Spanned::no_span(Expr::Cast {
expr: Box::new(Spanned::no_span(Expr::Ident("n".into()))),
ty: "Int".into(),
}),
Spanned::no_span(Expr::Cast {
expr: Box::new(Spanned::no_span(Expr::Ident("s".into()))),
ty: "String".into(),
}),
],
});
let mut env = TypeEnv::new();
register_input_clause_params(&body, &mut env);
assert_eq!(env.lookup("n"), Some(&Type::Int));
assert_eq!(env.lookup("s"), Some(&Type::String));
}
#[test]
fn register_input_params_untyped() {
let body = Spanned::no_span(Expr::Call {
func: Box::new(Spanned::no_span(Expr::Ident("input".into()))),
args: vec![Spanned::no_span(Expr::Ident("x".into()))],
});
let mut env = TypeEnv::new();
register_input_clause_params(&body, &mut env);
assert_eq!(env.lookup("x"), Some(&Type::Unknown));
}
#[test]
fn collect_input_types_typed() {
let body = Spanned::no_span(Expr::Call {
func: Box::new(Spanned::no_span(Expr::Ident("input".into()))),
args: vec![Spanned::no_span(Expr::Cast {
expr: Box::new(Spanned::no_span(Expr::Ident("n".into()))),
ty: "Int".into(),
})],
});
let mut types = Vec::new();
collect_input_param_types(&body, &mut types);
assert_eq!(types, vec![Type::Int]);
}
#[test]
fn bind_pattern_ident() {
let mut env = TypeEnv::new();
bind_pattern_vars(&Pattern::Ident("x".into()), &Type::Int, &mut env);
assert_eq!(env.lookup("x"), Some(&Type::Int));
}
#[test]
fn bind_pattern_wildcard_no_bind() {
let mut env = TypeEnv::new();
bind_pattern_vars(&Pattern::Wildcard, &Type::Int, &mut env);
assert!(env.lookup("_").is_none());
}
#[test]
fn bind_pattern_tuple() {
let mut env = TypeEnv::new();
let pat = Pattern::Tuple(vec![Pattern::Ident("a".into()), Pattern::Ident("b".into())]);
let ty = Type::Tuple(vec![Type::Int, Type::Bool]);
bind_pattern_vars(&pat, &ty, &mut env);
assert_eq!(env.lookup("a"), Some(&Type::Int));
assert_eq!(env.lookup("b"), Some(&Type::Bool));
}
#[test]
fn bind_pattern_constructor_with_fn_env() {
let mut env = TypeEnv::new();
env.insert(
"Some".into(),
Type::Fn {
params: vec![Type::Int],
ret: Box::new(Type::Named("Option".into())),
},
);
let pat = Pattern::Constructor {
name: "Some".into(),
fields: vec![Pattern::Ident("val".into())],
};
bind_pattern_vars(&pat, &Type::Named("Option".into()), &mut env);
assert_eq!(env.lookup("val"), Some(&Type::Int));
}
#[test]
fn env_with_result_adds_binding() {
let env = TypeEnv::new();
let new_env = env_with_result(&env, &Type::Int);
assert_eq!(new_env.lookup("result"), Some(&Type::Int));
assert!(env.lookup("result").is_none());
}
#[test]
fn requires_clause_is_bool() {
assert!(clause_requires_bool(&ClauseKind::Requires));
assert!(clause_requires_bool(&ClauseKind::Ensures));
assert!(clause_requires_bool(&ClauseKind::Invariant));
}
#[test]
fn non_predicate_clause_not_bool() {
assert!(!clause_requires_bool(&ClauseKind::Input));
assert!(!clause_requires_bool(&ClauseKind::Output));
assert!(!clause_requires_bool(&ClauseKind::Effects));
}
#[test]
fn clause_kind_labels() {
assert_eq!(clause_kind_label(&ClauseKind::Requires), "requires");
assert_eq!(clause_kind_label(&ClauseKind::Ensures), "ensures");
assert_eq!(clause_kind_label(&ClauseKind::Invariant), "invariant");
}
#[test]
fn check_clause_body_bool_ok() {
let env = TypeEnv::new();
let body = Spanned::no_span(Expr::Literal(Literal::Bool(true)));
let mut errors = Vec::new();
check_clause_expr(&ClauseKind::Requires, &body, &env, &mut errors, &(0..1));
assert!(errors.is_empty());
}
#[test]
fn check_clause_body_non_bool_error() {
let env = TypeEnv::new();
let body = Spanned::no_span(Expr::Literal(Literal::Int("42".into())));
let mut errors = Vec::new();
check_clause_expr(&ClauseKind::Requires, &body, &env, &mut errors, &(0..1));
assert_eq!(errors.len(), 1);
assert_eq!(errors[0].code, "A03006");
}
#[test]
fn check_clause_body_input_not_checked_for_bool() {
let env = TypeEnv::new();
let body = Spanned::no_span(Expr::Literal(Literal::Int("42".into())));
let mut errors = Vec::new();
check_clause_expr(&ClauseKind::Input, &body, &env, &mut errors, &(0..1));
assert!(errors.is_empty(), "input clauses should not require Bool");
}
}