Add type inference tests

Adds a suite of tests to verify type inference for constants, variable
propagation,
block expressions, lambda return types, and assignment updates.
This commit is contained in:
Michael Schimmel
2026-02-19 12:34:20 +01:00
parent d93727e198
commit a611c50a92
+69
View File
@@ -225,3 +225,72 @@ impl TypeChecker {
}) })
} }
} }
#[cfg(test)]
mod tests {
use super::*;
use crate::ast::parser::Parser;
use crate::ast::compiler::binder::Binder;
use crate::ast::types::StaticType;
use std::cell::RefCell;
fn check_source(source: &str) -> TypedNode {
let mut parser = Parser::new(source).unwrap();
let untyped = parser.parse_expression().unwrap();
let globals = Rc::new(RefCell::new(HashMap::new()));
let bound = Binder::bind_root(globals, &untyped).unwrap();
let global_types = Rc::new(RefCell::new(HashMap::new()));
let checker = TypeChecker::new(global_types);
checker.check(bound).unwrap()
}
#[test]
fn test_inference_constants() {
assert_eq!(check_source("10").ty, StaticType::Int);
assert_eq!(check_source("10.5").ty, StaticType::Float);
assert_eq!(check_source("true").ty, StaticType::Bool);
assert_eq!(check_source("\"hello\"").ty, StaticType::Text);
}
#[test]
fn test_inference_variable_propagation() {
// (do (def x 10) x) -> The last 'x' must be Int
let typed = check_source("(do (def x 10) x)");
if let BoundKind::Block { exprs } = typed.kind {
let last_expr = exprs.last().unwrap();
assert_eq!(last_expr.ty, StaticType::Int, "Variable 'x' should be inferred as Int");
} else { panic!("Expected block"); }
}
#[test]
fn test_inference_block_type() {
// Block type = last expression type
assert_eq!(check_source("(do 1 2.5)").ty, StaticType::Float);
assert_eq!(check_source("(do 1.5 \"test\")").ty, StaticType::Text);
}
#[test]
fn test_inference_lambda_return() {
// (fn [a] 10) -> fn(any) -> Int
let typed = check_source("(fn [a] 10)");
if let StaticType::Function { ret, .. } = typed.ty {
assert_eq!(*ret, StaticType::Int);
} else { panic!("Expected function type, got {:?}", typed.ty); }
// Nested: (fn [] (do 1 2.5)) -> fn() -> Float
let typed_nested = check_source("(fn [] (do 1 2.5))");
if let StaticType::Function { ret, .. } = typed_nested.ty {
assert_eq!(*ret, StaticType::Float);
} else { panic!("Expected function type"); }
}
#[test]
fn test_inference_assignment_updates_type() {
// (do (def x 10) (assign x 20.5) x) -> x becomes Float after assignment
let typed = check_source("(do (def x 10) (assign x 20.5) x)");
if let BoundKind::Block { exprs } = typed.kind {
let last_expr = exprs.last().unwrap();
assert_eq!(last_expr.ty, StaticType::Float, "Variable 'x' should be specialized to Float after assignment");
} else { panic!("Expected block"); }
}
}