use assura_parser::ast::{ClauseKind, Decl, ServiceItem};
use assura_resolve::{SymbolKind, SymbolTable};
use crate::clauses::{
collect_input_param_types, extract_output_type_from_body, register_input_clause_params,
};
use crate::convert::{enum_field_type_tokens, parse_type_tokens, resolve_type_opt, type_from_expr};
use crate::domain::StdlibTypes;
use crate::types::builtin_type;
use crate::{Type, TypeEnv};
pub(crate) fn build_type_env(
symbols: &SymbolTable,
source: &assura_parser::ast::SourceFile,
) -> TypeEnv {
let mut env = TypeEnv::new();
for sym in &symbols.symbols {
let ty = match sym.kind {
SymbolKind::BuiltinType => builtin_type(&sym.name).unwrap_or(Type::Unknown),
SymbolKind::TypeDef
| SymbolKind::ContractDef
| SymbolKind::ServiceDef
| SymbolKind::EnumDef => Type::Named(sym.name.clone()),
SymbolKind::FnDef | SymbolKind::ExternFn | SymbolKind::BindFn => Type::Fn {
params: Vec::new(),
ret: Box::new(Type::Unknown),
},
SymbolKind::Operation | SymbolKind::Query => Type::Fn {
params: Vec::new(),
ret: Box::new(Type::Unknown),
},
SymbolKind::TypeParam => Type::TypeParam(sym.name.clone()),
SymbolKind::Parameter | SymbolKind::Field => Type::Unknown,
SymbolKind::EnumVariant => Type::Named(sym.name.clone()),
SymbolKind::Prophecy => Type::Unknown,
SymbolKind::CodecRegistry => Type::Named(sym.name.clone()),
};
env.insert(sym.name.clone(), ty);
}
for decl in &source.decls {
match &decl.node {
Decl::FnDef(f) => {
for p in &f.params {
let ty = resolve_type_opt(p.ty.as_ref());
env.insert(p.name.clone(), ty);
}
let param_types: Vec<Type> = f
.params
.iter()
.map(|p| resolve_type_opt(p.ty.as_ref()))
.collect();
let ret = resolve_type_opt(f.return_ty.as_ref());
env.insert(
f.name.clone(),
Type::Fn {
params: param_types,
ret: Box::new(ret),
},
);
}
Decl::Extern(e) => {
for p in &e.params {
let ty = resolve_type_opt(p.ty.as_ref());
env.insert(p.name.clone(), ty);
}
let param_types: Vec<Type> = e
.params
.iter()
.map(|p| resolve_type_opt(p.ty.as_ref()))
.collect();
let ret = resolve_type_opt(e.return_ty.as_ref());
env.insert(
e.name.clone(),
Type::Fn {
params: param_types,
ret: Box::new(ret),
},
);
}
Decl::Contract(c) => {
for clause in &c.clauses {
if clause.kind == ClauseKind::Input {
register_input_clause_params(&clause.body, &mut env);
}
}
}
Decl::Service(s) => {
for item in &s.items {
let (name, clauses) = match item {
ServiceItem::Operation { name, clauses } => (name, clauses),
ServiceItem::Query { name, clauses } => (name, clauses),
_ => continue,
};
let mut param_types = Vec::new();
for clause in clauses {
if clause.kind == ClauseKind::Input {
collect_input_param_types(&clause.body, &mut param_types);
}
}
let mut ret = Type::Unit;
for clause in clauses {
if clause.kind == ClauseKind::Output {
let ty = extract_output_type_from_body(&clause.body);
if !ty.is_indeterminate() {
ret = ty;
break;
}
}
}
env.insert(
name.clone(),
Type::Fn {
params: param_types,
ret: Box::new(ret),
},
);
}
}
Decl::TypeDef(td) => {
if let assura_parser::ast::TypeBody::Struct(fields) = &td.body {
let field_types: Vec<(String, Type)> = fields
.iter()
.map(|f| (f.name.clone(), resolve_type_opt(f.ty.as_ref())))
.collect();
env.struct_fields.insert(td.name.clone(), field_types);
}
}
Decl::EnumDef(e) => {
for variant in &e.variants {
if !variant.fields.is_empty() {
let field_types: Vec<Type> = variant
.fields
.iter()
.map(|f| parse_type_tokens(&enum_field_type_tokens(f)))
.collect();
env.insert(
variant.name.clone(),
Type::Fn {
params: field_types,
ret: Box::new(Type::Named(e.name.clone())),
},
);
}
}
}
Decl::Prophecy(p) => {
if let Some(te) = &p.ty {
env.insert(p.name.clone(), type_from_expr(te));
}
}
Decl::Bind(b) => {
for p in &b.params {
let ty = resolve_type_opt(p.ty.as_ref());
env.insert(p.name.clone(), ty);
}
let param_types: Vec<Type> = b
.params
.iter()
.map(|p| resolve_type_opt(p.ty.as_ref()))
.collect();
let ret = resolve_type_opt(b.return_ty.as_ref());
env.insert(
b.name.clone(),
Type::Fn {
params: param_types,
ret: Box::new(ret),
},
);
}
Decl::CodecRegistry(_) | Decl::Block { .. } => {}
}
}
let stdlib = StdlibTypes::new();
for sdef in stdlib.all_types() {
if env.lookup(&sdef.name).is_none() {
env.insert(sdef.name.clone(), sdef.base_type.clone());
}
}
env
}
#[cfg(test)]
mod tests {
use super::*;
fn env_from_source(src: &str) -> TypeEnv {
let source = assura_parser::parse_unwrap(src);
let resolved = assura_resolve::resolve(&source).unwrap();
build_type_env(&resolved.symbols, &source)
}
#[test]
fn empty_source_has_stdlib_types() {
let env = env_from_source("");
env.lookup("Pos").unwrap();
env.lookup("NonNeg").unwrap();
}
#[test]
fn fndef_params_enriched() {
let env = env_from_source("fn add(a: Int, b: Int) -> Int { requires { a > 0 } }");
assert_eq!(env.lookup("a"), Some(&Type::Int));
assert_eq!(env.lookup("b"), Some(&Type::Int));
match env.lookup("add") {
Some(Type::Fn { params, ret }) => {
assert_eq!(params.len(), 2);
assert_eq!(params[0], Type::Int);
assert_eq!(**ret, Type::Int);
}
other => panic!("expected Fn type for add, got {other:?}"),
}
}
#[test]
fn fndef_no_return_type_defaults_unit() {
let env = env_from_source("fn noop() { ensures { true } }");
match env.lookup("noop") {
Some(Type::Fn { ret, .. }) => assert_eq!(**ret, Type::Unit),
other => panic!("expected Fn, got {other:?}"),
}
}
#[test]
fn extern_params_enriched() {
let env = env_from_source("extern fn ext(x: Bool) -> Nat");
assert_eq!(env.lookup("x"), Some(&Type::Bool));
match env.lookup("ext") {
Some(Type::Fn { params, ret }) => {
assert_eq!(params[0], Type::Bool);
assert_eq!(**ret, Type::Nat);
}
other => panic!("expected Fn, got {other:?}"),
}
}
#[test]
fn bind_params_enriched() {
let env = env_from_source("bind \"std::collections::HashMap\" as bd {\n input(n: Int)\n}");
assert_eq!(env.lookup("n"), Some(&Type::Int));
env.lookup("bd").unwrap();
}
#[test]
fn typedef_struct_fields_registered() {
let env = env_from_source("type Point { x: Float, y: Float }");
let fields = env.struct_fields.get("Point").unwrap();
assert_eq!(fields.len(), 2);
assert_eq!(fields[0].0, "x");
assert_eq!(fields[0].1, Type::Float);
}
#[test]
fn typedef_struct_fields_newline_without_separators() {
let env = env_from_source("type Point {\n x: Int\n y: Int\n}");
let fields = env.struct_fields.get("Point").expect("Point fields");
assert_eq!(
fields.len(),
2,
"newline-separated fields must both register, got {fields:?}"
);
assert_eq!(fields[0].0, "x");
assert_eq!(fields[0].1, Type::Int);
assert_eq!(fields[1].0, "y");
assert_eq!(fields[1].1, Type::Int);
}
#[test]
fn enumdef_variant_constructors() {
let env = env_from_source("enum Shape { Rect(Int, Int), Circle(Float) }");
match env.lookup("Rect") {
Some(Type::Fn { params, ret }) => {
assert_eq!(params.len(), 2);
assert_eq!(params[0], Type::Int);
assert_eq!(params[1], Type::Int);
assert_eq!(**ret, Type::Named("Shape".into()));
}
other => panic!("expected Fn constructor for Rect, got {other:?}"),
}
match env.lookup("Circle") {
Some(Type::Fn { params, ret }) => {
assert_eq!(params.len(), 1);
assert_eq!(params[0], Type::Float);
assert_eq!(**ret, Type::Named("Shape".into()));
}
other => panic!("expected Fn constructor for Circle, got {other:?}"),
}
}
#[test]
fn enumdef_multi_token_payload_constructors() {
let env = env_from_source(
"enum E { Box(List<Int>), Pair((Int, Bool)), Both(List<Int>, (Int,)) }",
);
match env.lookup("Box") {
Some(Type::Fn { params, ret }) => {
assert_eq!(params.len(), 1, "Box should be unary");
assert_eq!(params[0], Type::List(Box::new(Type::Int)));
assert_eq!(**ret, Type::Named("E".into()));
}
other => panic!("expected Fn constructor for Box, got {other:?}"),
}
match env.lookup("Pair") {
Some(Type::Fn { params, .. }) => {
assert_eq!(params.len(), 1);
assert_eq!(
params[0],
Type::Tuple(vec![Type::Int, Type::Bool]),
"Pair payload should be a 2-tuple type"
);
}
other => panic!("expected Fn constructor for Pair, got {other:?}"),
}
match env.lookup("Both") {
Some(Type::Fn { params, .. }) => {
assert_eq!(params.len(), 2);
assert_eq!(params[0], Type::List(Box::new(Type::Int)));
assert_eq!(params[1], Type::Tuple(vec![Type::Int]));
}
other => panic!("expected Fn constructor for Both, got {other:?}"),
}
}
#[test]
fn contract_input_params_registered() {
let env = env_from_source("contract C { input(n: Nat) ensures { n > 0 } }");
env.lookup("C").unwrap();
}
#[test]
fn prophecy_type_registered() {
let env = env_from_source("ghost prophecy p: Int");
assert_eq!(env.lookup("p"), Some(&Type::Int));
}
#[test]
fn prophecy_no_type_stays_unknown() {
let env = env_from_source("ghost prophecy q");
assert_eq!(env.lookup("q"), Some(&Type::Unknown));
}
#[test]
fn multiple_decls_all_registered() {
let env = env_from_source(
"contract A { ensures { true } }\n\
fn f(x: Int) -> Bool { ensures { true } }\n\
type T { val: Nat }",
);
env.lookup("A").unwrap();
env.lookup("f").unwrap();
env.lookup("T").unwrap();
assert_eq!(env.lookup("x"), Some(&Type::Int));
}
}