use crate::ast::{CallArg, Expr, FieldDef, ParamDef, TypeExpr};
use crate::token::TokenKind;
use crate::value::Value;
use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone)]
pub struct TypeCheckError {
pub message: String,
}
impl std::fmt::Display for TypeCheckError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "type error: {}", self.message)
}
}
#[derive(Clone)]
struct FnSig {
#[allow(dead_code)]
type_params: Vec<String>,
params: Vec<ParamDef>,
return_type: Option<TypeExpr>,
}
#[derive(Clone)]
struct ClassInfo {
#[allow(dead_code)]
type_params: Vec<String>,
fields: Vec<FieldDef>,
methods: HashMap<String, FnSig>,
}
pub struct TypeChecker {
functions: HashMap<String, FnSig>,
classes: HashMap<String, ClassInfo>,
type_aliases: HashMap<String, TypeExpr>,
errors: Vec<TypeCheckError>,
current_return_type: Option<TypeExpr>,
var_scopes: Vec<HashMap<String, TypeExpr>>,
type_vars: Vec<HashSet<String>>,
}
impl TypeChecker {
fn new() -> Self {
Self {
functions: HashMap::new(),
classes: HashMap::new(),
type_aliases: HashMap::new(),
errors: Vec::new(),
current_return_type: None,
var_scopes: vec![HashMap::new()],
type_vars: Vec::new(),
}
}
pub fn check(exprs: &[Expr]) -> Vec<TypeCheckError> {
let mut tc = Self::new();
for e in exprs {
tc.collect_def(e);
}
for e in exprs {
tc.check_expr(e);
}
tc.errors
}
fn resolve_type(&self, te: TypeExpr) -> TypeExpr {
match te {
TypeExpr::Named(ref n) => {
if let Some(expanded) = self.type_aliases.get(n) {
self.resolve_type(expanded.clone())
} else {
te
}
}
TypeExpr::Apply(name, args) => {
TypeExpr::Apply(name, args.into_iter().map(|a| self.resolve_type(a)).collect())
}
TypeExpr::Union(arms) => {
let resolved: Vec<TypeExpr> = arms.into_iter().map(|a| self.resolve_type(a)).collect();
let mut flat = Vec::new();
for arm in resolved {
match arm {
TypeExpr::Union(inner) => flat.extend(inner),
other => flat.push(other),
}
}
if flat.len() == 1 {
flat.remove(0)
} else {
TypeExpr::Union(flat)
}
}
TypeExpr::Literal(_) => te,
TypeExpr::Any => TypeExpr::Any,
}
}
fn push_type_vars(&mut self, params: &[String]) {
self.type_vars.push(params.iter().cloned().collect());
}
fn pop_type_vars(&mut self) {
self.type_vars.pop();
}
fn is_type_var(&self, name: &str) -> bool {
self.type_vars.iter().rev().any(|scope| scope.contains(name))
}
fn types_compat(&self, actual: &TypeExpr, expected: &TypeExpr) -> bool {
if let TypeExpr::Named(n) = expected
&& self.is_type_var(n)
{
return true;
}
if let TypeExpr::Named(n) = actual
&& self.is_type_var(n)
{
return true;
}
let a = self.resolve_type(actual.clone());
let e = self.resolve_type(expected.clone());
types_compatible(&a, &e)
}
fn collect_def(&mut self, expr: &Expr) {
match expr {
Expr::TypeAlias { name, type_expr } => {
self.type_aliases.insert(name.clone(), type_expr.clone());
}
Expr::Function {
name,
type_params,
params,
return_type,
..
} => {
self.functions.insert(
name.clone(),
FnSig {
type_params: type_params.clone(),
params: params.clone(),
return_type: return_type.clone(),
},
);
}
Expr::Class {
name,
type_params,
fields,
methods,
..
} => {
let method_sigs = methods
.iter()
.map(|m| {
(
m.name.clone(),
FnSig {
type_params: m.type_params.clone(),
params: m.params.clone(),
return_type: m.return_type.clone(),
},
)
})
.collect();
self.classes.insert(
name.clone(),
ClassInfo {
type_params: type_params.clone(),
fields: fields.clone(),
methods: method_sigs,
},
);
}
_ => {}
}
}
fn push_scope(&mut self) {
self.var_scopes.push(HashMap::new());
}
fn pop_scope(&mut self) {
self.var_scopes.pop();
}
fn validate_type_ann(&mut self, te: &TypeExpr) {
match te {
TypeExpr::Apply(_, args) => {
for arg in args {
self.validate_type_ann(arg);
}
}
_ => {
let resolved = self.resolve_type(te.clone());
if let Some(msg) = check_union_duplicates(&resolved) {
self.errors.push(TypeCheckError { message: msg });
}
}
}
}
fn set_var(&mut self, name: &str, ty: TypeExpr) {
if let Some(scope) = self.var_scopes.last_mut() {
scope.insert(name.to_string(), ty);
}
}
fn get_var(&self, name: &str) -> Option<TypeExpr> {
for scope in self.var_scopes.iter().rev() {
if let Some(ty) = scope.get(name) {
return Some(ty.clone());
}
}
None
}
fn check_expr(&mut self, expr: &Expr) {
match expr {
Expr::Return(inner) => {
if let Some(rt) = self.current_return_type.clone()
&& let Some(actual) = self.infer_type(inner)
&& !self.types_compat(&actual, &rt)
{
self.errors.push(TypeCheckError {
message: format!(
"return value expected {}, got {}",
te_name(&rt),
te_name(&actual)
),
});
}
self.check_expr(inner);
}
Expr::While { condition, body } => {
self.check_expr(condition);
self.push_scope();
for s in body {
self.check_expr(s);
}
self.pop_scope();
}
Expr::Lambda { body, .. } => {
self.push_scope();
for s in body {
self.check_expr(s);
}
self.pop_scope();
}
Expr::Raise(inner) => self.check_expr(inner),
Expr::Break(inner) | Expr::Next(inner) => self.check_expr(inner),
Expr::MultiAssign { names, values } => {
for (name, ve) in names.iter().zip(values.iter()) {
if let Some(ty) = self.infer_type(ve) {
self.set_var(name, ty);
}
self.check_expr(ve);
}
}
Expr::Call { callee, args, .. } => self.check_call(callee, args),
Expr::Assign { name, value } => {
let ty = self.infer_type(value).or_else(|| {
if let Expr::Call { callee, .. } = value.as_ref()
&& let Expr::Get { object, name: m } = callee.as_ref()
&& m == "new"
&& let Expr::Variable(cn) = object.as_ref()
&& self.classes.contains_key(cn)
{
return Some(TypeExpr::Named(cn.clone()));
}
None
});
if let Some(ty) = ty {
self.set_var(name, ty);
}
self.check_expr(value);
}
Expr::Binary { left, right, .. } => {
self.check_expr(left);
self.check_expr(right);
}
Expr::Unary { right, .. } => self.check_expr(right),
Expr::Get { object, .. } | Expr::SafeGet { object, .. } => self.check_expr(object),
Expr::Set {
object,
value,
name,
} => {
if let Some(TypeExpr::Named(class_name)) = self.infer_type(object)
&& let Some(cls) = self.classes.get(&class_name).cloned()
&& let Some(fd) = cls.fields.iter().find(|f| &f.name == name)
&& let Some(te) = &fd.type_ann
&& let Some(actual) = self.infer_type(value)
&& !self.types_compat(&actual, te)
{
self.errors.push(TypeCheckError {
message: format!(
"field '{}' expected {}, got {}",
name,
te_name(te),
te_name(&actual)
),
});
}
self.check_expr(object);
self.check_expr(value);
}
Expr::Index { object, index } => {
self.check_expr(object);
self.check_expr(index);
}
Expr::IndexSet {
object,
index,
value,
} => {
self.check_expr(object);
self.check_expr(index);
self.check_expr(value);
}
Expr::Range { from, to } => {
self.check_expr(from);
self.check_expr(to);
}
Expr::ListLit(elems) => {
for e in elems {
self.check_expr(e);
}
}
Expr::MapLit(pairs) => {
for (_, v) in pairs {
self.check_expr(v);
}
}
Expr::Grouping(inner) => self.check_expr(inner),
Expr::If {
condition,
then_branch,
else_branch,
} => {
self.check_expr(condition);
self.push_scope();
for s in then_branch {
self.check_expr(s);
}
self.pop_scope();
if let Some(branch) = else_branch {
self.push_scope();
for s in branch {
self.check_expr(s);
}
self.pop_scope();
}
}
Expr::Begin {
body,
rescue_body,
else_body,
..
} => {
for s in body {
self.check_expr(s);
}
for s in rescue_body {
self.check_expr(s);
}
for s in else_body {
self.check_expr(s);
}
}
Expr::Print(inner) => self.check_expr(inner),
Expr::Class { type_params, methods, .. } => {
self.push_type_vars(type_params);
for method in methods {
let saved = self.current_return_type.take();
if let Some(rt) = &method.return_type {
self.validate_type_ann(rt);
}
self.current_return_type = method.return_type.clone();
self.push_scope();
self.push_type_vars(&method.type_params);
for p in &method.params {
if let Some(te) = &p.type_ann {
self.validate_type_ann(te);
self.set_var(&p.name, te.clone());
}
}
for s in &method.body {
self.check_expr(s);
}
if let Some(rt) = &method.return_type.clone()
&& let Some(last_expr) = method.body.last()
&& let Some(actual) = self.infer_type(last_expr)
&& !self.types_compat(&actual, rt)
{
self.errors.push(TypeCheckError {
message: format!(
"return value expected {}, got {}",
te_name(rt),
te_name(&actual)
),
});
}
self.pop_type_vars();
self.pop_scope();
self.current_return_type = saved;
}
self.pop_type_vars();
}
Expr::Function {
name,
type_params,
params,
return_type,
body,
} => {
self.functions.insert(
name.clone(),
FnSig {
type_params: type_params.clone(),
params: params.clone(),
return_type: return_type.clone(),
},
);
let saved = self.current_return_type.take();
if let Some(rt) = return_type {
self.validate_type_ann(rt);
}
self.current_return_type = return_type.clone();
self.push_scope();
self.push_type_vars(type_params);
for p in params {
if let Some(te) = &p.type_ann {
self.validate_type_ann(te);
self.set_var(&p.name, te.clone());
}
}
for s in body {
self.check_expr(s);
}
if let Some(rt) = return_type
&& let Some(last_expr) = body.last()
&& let Some(actual) = self.infer_type(last_expr)
&& !self.types_compat(&actual, rt)
{
self.errors.push(TypeCheckError {
message: format!(
"return value expected {}, got {}",
te_name(rt),
te_name(&actual)
),
});
}
self.pop_type_vars();
self.pop_scope();
self.current_return_type = saved;
}
Expr::TypeAlias { .. } => {}
_ => {}
}
}
fn check_call(&mut self, callee: &Expr, args: &[CallArg]) {
for arg in args {
self.check_expr(&arg.value);
}
match callee {
Expr::Variable(name) => {
if let Some(sig) = self.functions.get(name).cloned() {
self.push_type_vars(&sig.type_params.clone());
self.check_args(&sig.params, args, name);
self.pop_type_vars();
}
}
Expr::Get {
object,
name: method_name,
} => {
self.check_expr(object);
if method_name == "new" {
if let Expr::Variable(class_name) = object.as_ref()
&& let Some(cls) = self.classes.get(class_name).cloned()
{
self.push_type_vars(&cls.type_params.clone());
for arg in args {
if let Some(fname) = &arg.name
&& let Some(fd) = cls.fields.iter().find(|f| &f.name == fname)
&& let Some(te) = &fd.type_ann
&& let Some(actual) = self.infer_type(&arg.value)
&& !self.types_compat(&actual, te)
{
self.errors.push(TypeCheckError {
message: format!(
"field '{}' expected {}, got {}",
fname,
te_name(te),
te_name(&actual)
),
});
}
}
self.pop_type_vars();
}
} else if let Some(TypeExpr::Named(class_name)) = self.infer_type(object)
&& let Some(cls) = self.classes.get(&class_name).cloned()
&& let Some(sig) = cls.methods.get(method_name).cloned()
{
let combined: Vec<String> = cls.type_params.iter()
.chain(sig.type_params.iter())
.cloned()
.collect();
self.push_type_vars(&combined);
self.check_args(&sig.params, args, method_name);
self.pop_type_vars();
}
}
_ => self.check_expr(callee),
}
}
fn check_args(&mut self, params: &[ParamDef], args: &[CallArg], fn_name: &str) {
for (param, arg) in params.iter().zip(args.iter()) {
if let Some(te) = ¶m.type_ann
&& let Some(actual) = self.infer_type(&arg.value)
&& !self.types_compat(&actual, te)
{
self.errors.push(TypeCheckError {
message: format!(
"argument '{}' to '{}' expected {}, got {}",
param.name,
fn_name,
te_name(te),
te_name(&actual)
),
});
}
}
}
fn infer_type(&self, expr: &Expr) -> Option<TypeExpr> {
match expr {
Expr::Literal(v) => match v {
Value::Int(_) => Some(TypeExpr::Named("Int".into())),
Value::Float(_) => Some(TypeExpr::Named("Float".into())),
Value::Str(_) => Some(TypeExpr::Named("String".into())),
Value::Bool(_) => Some(TypeExpr::Named("Bool".into())),
Value::Nil => Some(TypeExpr::Named("Nil".into())),
},
Expr::Variable(name) => self.get_var(name),
Expr::Grouping(inner) => self.infer_type(inner),
Expr::StringInterp(_) => Some(TypeExpr::Named("String".into())),
Expr::ListLit(_) => Some(TypeExpr::Named("List".into())),
Expr::MapLit(_) => Some(TypeExpr::Named("Map".into())),
Expr::Range { .. } => Some(TypeExpr::Named("Range".into())),
Expr::Binary { left, op, right } => match &op.kind {
TokenKind::Plus
| TokenKind::Minus
| TokenKind::Star
| TokenKind::Slash
| TokenKind::Percent => {
let l = self.infer_type(left);
let r = self.infer_type(right);
match (&l, &r) {
(Some(TypeExpr::Named(a)), Some(TypeExpr::Named(b))) => {
if a == "Float" || b == "Float" {
Some(TypeExpr::Named("Float".into()))
} else if a == "Int" && b == "Int" {
Some(TypeExpr::Named("Int".into()))
} else {
None
}
}
_ => None,
}
}
TokenKind::EqEq
| TokenKind::BangEq
| TokenKind::Less
| TokenKind::LessEq
| TokenKind::Greater
| TokenKind::GreaterEq
| TokenKind::AmpAmp
| TokenKind::PipePipe => Some(TypeExpr::Named("Bool".into())),
_ => None,
},
Expr::Print(inner) => self.infer_type(inner),
Expr::If { .. }
| Expr::Begin { .. }
| Expr::While { .. }
| Expr::MultiAssign { .. }
| Expr::Return(_)
| Expr::Break(_)
| Expr::Next(_)
| Expr::Raise(_) => None,
Expr::Lambda { .. } => None,
Expr::Class { name, .. } => Some(TypeExpr::Named(name.clone())),
Expr::Function { .. } => Some(TypeExpr::Named("String".into())),
Expr::Call { callee, .. } => match callee.as_ref() {
Expr::Variable(name) => {
self.functions.get(name).and_then(|s| s.return_type.clone())
}
Expr::Get {
object,
name: method_name,
} => {
if method_name == "new"
&& let Expr::Variable(cn) = object.as_ref()
&& self.classes.contains_key(cn)
{
return Some(TypeExpr::Named(cn.clone()));
}
if let Some(TypeExpr::Named(cn)) = self.infer_type(object)
&& let Some(cls) = self.classes.get(&cn)
{
return cls
.methods
.get(method_name)
.and_then(|s| s.return_type.clone());
}
None
}
_ => None,
},
_ => None,
}
}
}
fn types_compatible(actual: &TypeExpr, expected: &TypeExpr) -> bool {
let literal_base_named = |v: &Value| match v {
Value::Int(_) => TypeExpr::Named("Int".to_string()),
Value::Float(_) => TypeExpr::Named("Float".to_string()),
Value::Str(_) => TypeExpr::Named("String".to_string()),
Value::Bool(_) => TypeExpr::Named("Bool".to_string()),
Value::Nil => TypeExpr::Named("Nil".to_string()),
};
match (actual, expected) {
(_, TypeExpr::Any) | (TypeExpr::Any, _) => true,
(TypeExpr::Apply(an, a_args), TypeExpr::Apply(en, e_args)) => {
an == en
&& a_args.len() == e_args.len()
&& a_args.iter().zip(e_args.iter()).all(|(a, e)| types_compatible(a, e))
}
(TypeExpr::Named(a), TypeExpr::Apply(e, _)) | (TypeExpr::Apply(a, _), TypeExpr::Named(e)) => a == e,
(TypeExpr::Union(arms), _) => arms.iter().all(|a| types_compatible(a, expected)),
(_, TypeExpr::Union(arms)) => arms.iter().any(|e| types_compatible(actual, e)),
(TypeExpr::Literal(a), TypeExpr::Literal(e)) => a == e,
(TypeExpr::Literal(a), TypeExpr::Named(e)) => {
let base = literal_base_named(a);
types_compatible(&base, &TypeExpr::Named(e.clone()))
}
(TypeExpr::Named(a), TypeExpr::Named(e)) => {
a == e || (e == "Num" && (a == "Int" || a == "Float"))
}
(TypeExpr::Named(_), TypeExpr::Literal(_))
| (TypeExpr::Apply(_, _), TypeExpr::Literal(_))
| (TypeExpr::Literal(_), TypeExpr::Apply(_, _)) => false,
}
}
fn te_name(te: &TypeExpr) -> String {
match te {
TypeExpr::Named(n) => n.clone(),
TypeExpr::Apply(n, args) => {
format!("{}[{}]", n, args.iter().map(te_name).collect::<Vec<_>>().join(", "))
}
TypeExpr::Literal(Value::Int(n)) => n.to_string(),
TypeExpr::Literal(Value::Float(n)) => n.to_string(),
TypeExpr::Literal(Value::Str(s)) => format!("{:?}", s),
TypeExpr::Literal(Value::Bool(b)) => b.to_string(),
TypeExpr::Literal(Value::Nil) => "Nil".to_string(),
TypeExpr::Any => "Any".to_string(),
TypeExpr::Union(arms) => arms.iter().map(te_name).collect::<Vec<_>>().join(" | "),
}
}
fn check_union_duplicates(te: &TypeExpr) -> Option<String> {
if let TypeExpr::Union(arms) = te {
let mut seen = std::collections::HashSet::new();
for arm in arms {
let key = te_name(arm);
if !seen.insert(key.clone()) {
return Some(format!("duplicate type '{}' in union", key));
}
}
}
None
}