Add type inference infrastructure
Introduces a `TypeChecker` struct with logic for Hindley-Milner type inference. This includes: - `var_counter` for generating unique type variable IDs. - `subst` for the global substitution map. - `fresh_var` to create new type variables. - `apply_subst` to resolve type variables recursively. - `occurs` to check for cycles during unification. - `unify` to merge types and extend the substitution. The `StaticType` enum is extended with a `TypeVar(u32)` variant to represent unresolved type variables. The `Display` and `is_assignable_from` implementations for `StaticType` are updated to handle `TypeVar`.
This commit is contained in:
@@ -4,6 +4,7 @@ use crate::ast::nodes::{
|
||||
};
|
||||
use crate::ast::diagnostics::Diagnostics;
|
||||
use crate::ast::types::{Keyword, RecordLayout, Signature, StaticType};
|
||||
use std::cell::Cell;
|
||||
use std::collections::HashMap;
|
||||
use std::rc::Rc;
|
||||
|
||||
@@ -131,11 +132,128 @@ fn extract_lambda_param_hints(
|
||||
|
||||
pub struct TypeChecker {
|
||||
root_types: Rc<std::cell::RefCell<Vec<StaticType>>>,
|
||||
/// Monotonic counter for generating unique type variable IDs.
|
||||
var_counter: Cell<u32>,
|
||||
/// Global substitution map: TypeVar ID → resolved StaticType.
|
||||
/// Shared across all scopes within a single type-checking pass.
|
||||
subst: std::cell::RefCell<HashMap<u32, StaticType>>,
|
||||
}
|
||||
|
||||
impl TypeChecker {
|
||||
pub fn new(root_types: Rc<std::cell::RefCell<Vec<StaticType>>>) -> Self {
|
||||
Self { root_types }
|
||||
Self {
|
||||
root_types,
|
||||
var_counter: Cell::new(0),
|
||||
subst: std::cell::RefCell::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Creates a fresh, unique type variable.
|
||||
#[allow(dead_code)]
|
||||
fn fresh_var(&self) -> StaticType {
|
||||
let id = self.var_counter.get();
|
||||
self.var_counter.set(id + 1);
|
||||
StaticType::TypeVar(id)
|
||||
}
|
||||
|
||||
/// Recursively applies the substitution map to a type, replacing all
|
||||
/// resolved `TypeVar`s with their concrete types.
|
||||
#[allow(dead_code)]
|
||||
fn apply_subst(ty: StaticType, subst: &HashMap<u32, StaticType>) -> StaticType {
|
||||
match ty {
|
||||
StaticType::TypeVar(n) => {
|
||||
if let Some(resolved) = subst.get(&n) {
|
||||
// Follow the chain (handles transitive substitutions)
|
||||
Self::apply_subst(resolved.clone(), subst)
|
||||
} else {
|
||||
StaticType::TypeVar(n)
|
||||
}
|
||||
}
|
||||
StaticType::Series(inner) => {
|
||||
StaticType::Series(Box::new(Self::apply_subst(*inner, subst)))
|
||||
}
|
||||
StaticType::Stream(inner) => {
|
||||
StaticType::Stream(Box::new(Self::apply_subst(*inner, subst)))
|
||||
}
|
||||
StaticType::Optional(inner) => {
|
||||
StaticType::Optional(Box::new(Self::apply_subst(*inner, subst)))
|
||||
}
|
||||
StaticType::Function(sig) => StaticType::Function(Box::new(Signature {
|
||||
params: Self::apply_subst(sig.params, subst),
|
||||
ret: Self::apply_subst(sig.ret, subst),
|
||||
})),
|
||||
StaticType::Tuple(elems) => {
|
||||
StaticType::Tuple(elems.into_iter().map(|t| Self::apply_subst(t, subst)).collect())
|
||||
}
|
||||
other => other,
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns true if `TypeVar(var_id)` appears anywhere in `ty` under the
|
||||
/// current substitution. Used to prevent infinite types (occurs check).
|
||||
#[allow(dead_code)]
|
||||
fn occurs(var_id: u32, ty: &StaticType, subst: &HashMap<u32, StaticType>) -> bool {
|
||||
match ty {
|
||||
StaticType::TypeVar(n) => {
|
||||
if *n == var_id {
|
||||
return true;
|
||||
}
|
||||
// Follow chain in substitution
|
||||
if let Some(resolved) = subst.get(n) {
|
||||
Self::occurs(var_id, resolved, subst)
|
||||
} else {
|
||||
false
|
||||
}
|
||||
}
|
||||
StaticType::Series(inner)
|
||||
| StaticType::Stream(inner)
|
||||
| StaticType::Optional(inner) => Self::occurs(var_id, inner, subst),
|
||||
StaticType::Function(sig) => {
|
||||
Self::occurs(var_id, &sig.params, subst)
|
||||
|| Self::occurs(var_id, &sig.ret, subst)
|
||||
}
|
||||
StaticType::Tuple(elems) => elems.iter().any(|t| Self::occurs(var_id, t, subst)),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Unifies two types under the current substitution.
|
||||
/// On success, the substitution is extended so that `ty1` and `ty2` become equal.
|
||||
/// On failure (type mismatch or occurs check), a diagnostic error is emitted.
|
||||
#[allow(dead_code)]
|
||||
fn unify(&self, ty1: StaticType, ty2: StaticType, diag: &mut Diagnostics) {
|
||||
let mut subst = self.subst.borrow_mut();
|
||||
let ty1 = Self::apply_subst(ty1, &subst);
|
||||
let ty2 = Self::apply_subst(ty2, &subst);
|
||||
|
||||
match (ty1, ty2) {
|
||||
(a, b) if a == b => {}
|
||||
(StaticType::TypeVar(n), ty) | (ty, StaticType::TypeVar(n)) => {
|
||||
if Self::occurs(n, &ty, &subst) {
|
||||
diag.push_error(format!("Infinite type: ?{} = {}", n, ty), None);
|
||||
return;
|
||||
}
|
||||
subst.insert(n, ty);
|
||||
}
|
||||
(StaticType::Series(a), StaticType::Series(b)) => {
|
||||
drop(subst);
|
||||
self.unify(*a, *b, diag);
|
||||
}
|
||||
(StaticType::Stream(a), StaticType::Stream(b)) => {
|
||||
drop(subst);
|
||||
self.unify(*a, *b, diag);
|
||||
}
|
||||
(StaticType::Optional(a), StaticType::Optional(b)) => {
|
||||
drop(subst);
|
||||
self.unify(*a, *b, diag);
|
||||
}
|
||||
// Any and Error are already handled by is_assignable_from — silently succeed
|
||||
(StaticType::Any, _) | (_, StaticType::Any) => {}
|
||||
(StaticType::Error, _) | (_, StaticType::Error) => {}
|
||||
(a, b) => {
|
||||
diag.push_error(format!("Type mismatch: expected {}, got {}", a, b), None);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn check(
|
||||
|
||||
@@ -373,6 +373,9 @@ pub enum StaticType {
|
||||
resolve_return: fn(&StaticType) -> Option<StaticType>,
|
||||
resolve_arg_hints: Option<ArgHintResolver>,
|
||||
},
|
||||
/// An unresolved type variable used during Hindley-Milner type inference.
|
||||
/// The `u32` is a unique ID assigned by the type checker. Resolved via substitution.
|
||||
TypeVar(u32),
|
||||
/// A diagnostic poison type, allowing type-checking to continue after an error.
|
||||
Error,
|
||||
}
|
||||
@@ -432,6 +435,7 @@ impl fmt::Display for StaticType {
|
||||
}
|
||||
StaticType::Object(name) => write!(f, "{}", name),
|
||||
StaticType::PolymorphicFn { .. } => write!(f, "<polymorphic-fn>"),
|
||||
StaticType::TypeVar(n) => write!(f, "?{}", n),
|
||||
StaticType::Error => write!(f, "<error>"),
|
||||
}
|
||||
}
|
||||
@@ -479,6 +483,8 @@ impl StaticType {
|
||||
|| matches!(other, StaticType::Any)
|
||||
|| matches!(self, StaticType::Error)
|
||||
|| matches!(other, StaticType::Error)
|
||||
|| matches!(self, StaticType::TypeVar(_))
|
||||
|| matches!(other, StaticType::TypeVar(_))
|
||||
{
|
||||
return true;
|
||||
}
|
||||
@@ -589,6 +595,7 @@ impl StaticType {
|
||||
.or_else(|| sigs.iter().find(|sig| sig.params.is_assignable_from(args_ty))) // 2. Fallback to implicit coercion
|
||||
.map(|sig| sig.ret.clone()),
|
||||
StaticType::PolymorphicFn { resolve_return, .. } => resolve_return(args_ty),
|
||||
StaticType::TypeVar(_) => Some(StaticType::Any),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user