Refactor TypeChecker into modules

The TypeChecker implementation was becoming quite large, so it has been
refactored into several modules:
- `context`: Handles type checking contexts and inference access.
- `inference`: Contains the core Hindley-Milner inference logic
  (unification, generalization, etc.).
- `finalize`: Manages the finalization step, applying substitutions and
  dispatching hooks.
- `check`: Implements the main type checking logic for AST nodes.

This modularization improves code organization and maintainability.
This commit is contained in:
2026-03-30 14:47:38 +02:00
parent 2a5b87f45f
commit 8ef93e2af5
7 changed files with 1875 additions and 1832 deletions
File diff suppressed because it is too large Load Diff
+845
View File
@@ -0,0 +1,845 @@
use crate::ast::diagnostics::Diagnostics;
use crate::ast::nodes::{
Address, AssignBinding, BoundLike, DefBinding, IdentifierBinding, LambdaBinding,
Node, NodeKind, TypedNode, TypedPhase,
};
use crate::ast::types::{Keyword, RecordLayout, Signature, StaticType};
use std::rc::Rc;
use super::context::{CheckerInferenceAccess, TypeContext, extract_lambda_param_hints};
use super::TypeChecker;
impl TypeChecker {
/// Types a lambda node using externally provided parameter type hints,
/// while preserving the current scope's upvalue types.
/// Unlike `check_node_as_bound`, this keeps the enclosing `TypeContext` as parent,
/// so captured variables retain their inferred types.
pub(super) fn check_lambda_with_param_hints<P: BoundLike>(
&self,
node: &Node<P>,
param_hints: &[StaticType],
ctx: &mut TypeContext,
diag: &mut Diagnostics,
) -> TypedNode {
let NodeKind::Lambda {
params,
body,
info,
} = &node.kind
else {
return self.check_node(node, ctx, diag);
};
let upvalues = &info.upvalues;
let positional_count = info.positional_count;
let mut upvalue_types = Vec::with_capacity(upvalues.len());
for &addr in upvalues {
upvalue_types.push(ctx.get_type(addr));
}
let mut lambda_ctx = TypeContext::new(64, upvalue_types, ctx.root_types, Some(ctx));
let hint_ty = StaticType::Tuple(param_hints.to_vec());
let params_typed = self.check_params(params.as_ref(), &hint_ty, &mut lambda_ctx, diag);
lambda_ctx.current_params_ty = Some(params_typed.ty.clone());
let body_typed = self.check_node(body, &mut lambda_ctx, diag);
let ret_ty = body_typed.ty.clone();
let fn_ty = StaticType::Function(Box::new(Signature {
params: params_typed.ty.clone(),
ret: ret_ty,
}));
Node {
identity: node.identity.clone(),
kind: NodeKind::Lambda {
params: Rc::new(params_typed),
body: Rc::new(body_typed),
info: LambdaBinding {
upvalues: upvalues.clone(),
positional_count,
},
},
ty: fn_ty,
comments: node.comments.clone(),
}
}
pub(super) fn check_params<P: BoundLike>(
&self,
node: &Node<P>,
specialized_ty: &StaticType,
ctx: &mut TypeContext,
diag: &mut Diagnostics,
) -> TypedNode {
let (kind, ty): (NodeKind<TypedPhase>, StaticType) = match &node.kind {
NodeKind::Def {
pattern,
info,
..
} => {
if let NodeKind::Identifier {
symbol,
binding: IdentifierBinding::Declaration { addr, kind: decl_kind },
} = &pattern.kind
{
ctx.set_type(*addr, specialized_ty.clone());
(
NodeKind::Def {
pattern: Rc::new(Node {
identity: pattern.identity.clone(),
kind: NodeKind::Identifier {
symbol: symbol.clone(),
binding: IdentifierBinding::Declaration {
addr: *addr,
kind: *decl_kind,
},
},
ty: specialized_ty.clone(),
comments: pattern.comments.clone(),
}),
value: Rc::new(Node {
identity: node.identity.clone(),
kind: NodeKind::Nop,
ty: specialized_ty.clone(),
comments: Rc::from([]),
}),
info: DefBinding {
captured_by: info.captured_by.clone(),
},
},
specialized_ty.clone(),
)
} else {
// Destructuring def in params — fall through to tuple handling
// if pattern is a Tuple, handle elements
if let NodeKind::Tuple { elements } = &pattern.kind {
return self.check_params_tuple(node, elements, specialized_ty, ctx, diag);
}
diag.push_error(
"Invalid pattern in parameter definition",
Some(node.identity.clone()),
);
(NodeKind::Error, StaticType::Error)
}
}
NodeKind::Assign {
info,
..
} => {
(
NodeKind::Assign {
target: Rc::new(Node {
identity: node.identity.clone(),
kind: NodeKind::Nop,
ty: specialized_ty.clone(),
comments: Rc::from([]),
}),
value: Rc::new(Node {
identity: node.identity.clone(),
kind: NodeKind::Nop,
ty: specialized_ty.clone(),
comments: Rc::from([]),
}),
info: AssignBinding {
addr: info.addr,
},
},
specialized_ty.clone(),
)
}
NodeKind::Identifier {
symbol,
binding: IdentifierBinding::Declaration { addr, kind: decl_kind },
} => {
ctx.set_type(*addr, specialized_ty.clone());
(
NodeKind::Identifier {
symbol: symbol.clone(),
binding: IdentifierBinding::Declaration {
addr: *addr,
kind: *decl_kind,
},
},
specialized_ty.clone(),
)
}
NodeKind::Tuple { elements } => {
return self.check_params_tuple(node, elements, specialized_ty, ctx, diag);
}
NodeKind::Nop => (NodeKind::Nop, StaticType::Void),
NodeKind::Error => (NodeKind::Error, StaticType::Error),
_ => {
diag.push_error(
"Invalid node in parameter list",
Some(node.identity.clone()),
);
(NodeKind::Error, StaticType::Error)
}
};
Node {
identity: node.identity.clone(),
kind,
ty,
comments: node.comments.clone(),
}
}
fn check_params_tuple<P: BoundLike>(
&self,
node: &Node<P>,
elements: &[Rc<Node<P>>],
specialized_ty: &StaticType,
ctx: &mut TypeContext,
diag: &mut Diagnostics,
) -> TypedNode {
match specialized_ty {
StaticType::Any
| StaticType::TypeVar(_) // TypeVar may resolve to any destructurable type
| StaticType::Tuple(_)
| StaticType::Vector(_, _)
| StaticType::Matrix(_, _)
| StaticType::List(_)
| StaticType::Record(_)
| StaticType::Error => {}
_ => {
diag.push_error(
format!(
"Cannot destructure type {} as a tuple/vector",
specialized_ty.display_compact()
),
Some(node.identity.clone()),
);
return Node {
identity: node.identity.clone(),
kind: NodeKind::Error,
ty: StaticType::Error,
comments: node.comments.clone(),
};
}
}
let mut typed_elements = Vec::new();
let mut elem_types = Vec::new();
for (i, el) in elements.iter().enumerate() {
let sub_ty = match specialized_ty {
StaticType::Tuple(t) => t.get(i).cloned().unwrap_or(StaticType::Any),
StaticType::Vector(inner, _) => (**inner).clone(),
StaticType::Matrix(inner, _) => (**inner).clone(),
StaticType::List(inner) => (**inner).clone(),
StaticType::Record(layout) => layout
.fields
.get(i)
.map(|(_, ty): &(Keyword, StaticType)| ty.clone())
.unwrap_or(StaticType::Any),
StaticType::Error => StaticType::Error,
_ => StaticType::Any,
};
let t = self.check_params(el.as_ref(), &sub_ty, ctx, diag);
elem_types.push(t.ty.clone());
typed_elements.push(Rc::new(t));
}
Node {
identity: node.identity.clone(),
kind: NodeKind::Tuple {
elements: typed_elements,
},
ty: StaticType::Tuple(elem_types),
comments: node.comments.clone(),
}
}
pub(super) fn check_node<P: BoundLike>(
&self,
node: &Node<P>,
ctx: &mut TypeContext,
diag: &mut Diagnostics,
) -> TypedNode {
let (kind, ty): (NodeKind<TypedPhase>, StaticType) = match &node.kind {
NodeKind::Nop => (NodeKind::Nop, StaticType::Void),
NodeKind::Constant(v) => {
let ty = v.static_type();
(NodeKind::Constant(v.clone()), ty)
}
NodeKind::Def {
pattern,
value,
info,
} => {
let val_typed = self.check_node(value, ctx, diag);
let ty = val_typed.ty.clone();
// Extract addr from pattern to register the type.
// Value restriction: only Function-typed values are generalized to Forall.
// Mutable state (Series, scalars) must remain monomorphic.
if let NodeKind::Identifier {
binding: IdentifierBinding::Declaration { addr, .. },
..
} = &pattern.kind
{
let stored_ty = if matches!(ty, StaticType::Function(..)) {
self.generalize(ty.clone(), ctx)
} else {
ty.clone()
};
ctx.set_type(*addr, stored_ty);
}
// For destructuring defs, check params on the pattern
if let NodeKind::Tuple { .. } = &pattern.kind {
let pat_typed = self.check_params(pattern.as_ref(), &val_typed.ty, ctx, diag);
let ty = val_typed.ty.clone();
return Node {
identity: node.identity.clone(),
kind: NodeKind::Def {
pattern: Rc::new(pat_typed),
value: Rc::new(val_typed),
info: DefBinding {
captured_by: info.captured_by.clone(),
},
},
ty,
comments: node.comments.clone(),
};
}
// Simple def — reconstruct the pattern node with the new type
let new_pattern: TypedNode = Node {
identity: pattern.identity.clone(),
kind: match &pattern.kind {
NodeKind::Identifier { symbol, binding } => NodeKind::Identifier {
symbol: symbol.clone(),
binding: match binding {
IdentifierBinding::Declaration { addr, kind: decl_kind } => {
IdentifierBinding::Declaration {
addr: *addr,
kind: *decl_kind,
}
}
IdentifierBinding::Reference(addr) => {
IdentifierBinding::Reference(*addr)
}
},
},
_ => NodeKind::Error,
},
ty: ty.clone(),
comments: pattern.comments.clone(),
};
(
NodeKind::Def {
pattern: Rc::new(new_pattern),
value: Rc::new(val_typed),
info: DefBinding {
captured_by: info.captured_by.clone(),
},
},
ty,
)
}
NodeKind::Identifier { symbol, binding } => {
if let IdentifierBinding::Reference(addr) = binding {
// Apply the current HM substitution so that TypeVars resolved in
// nested scopes (e.g. inside a `while` body) are visible here even
// when ctx.set_type could not propagate back through an upvalue address.
// Instantiate Forall types: each use site gets fresh TypeVars so that
// calls with different argument types remain independent.
let ty = Self::apply_subst(ctx.get_type(*addr), &self.subst.borrow());
let ty = self.instantiate(ty);
(
NodeKind::Identifier {
symbol: symbol.clone(),
binding: IdentifierBinding::Reference(*addr),
},
ty,
)
} else if let IdentifierBinding::Declaration { addr, kind: decl_kind } = binding {
let ty = Self::apply_subst(ctx.get_type(*addr), &self.subst.borrow());
(
NodeKind::Identifier {
symbol: symbol.clone(),
binding: IdentifierBinding::Declaration {
addr: *addr,
kind: *decl_kind,
},
},
ty,
)
} else {
(NodeKind::Error, StaticType::Error)
}
}
NodeKind::FieldAccessor(k) => {
(NodeKind::FieldAccessor(*k), StaticType::FieldAccessor(*k))
}
NodeKind::GetField { rec, field } => {
let rec_typed = self.check_node(rec, ctx, diag);
let field_ty = match &rec_typed.ty {
StaticType::Record(layout) => {
if let Some(idx) = layout.index_of(*field) {
layout.fields[idx].1.clone()
} else {
diag.push_error(
format!("Record does not have field :{}", field.name()),
Some(rec_typed.identity.clone()),
);
StaticType::Error
}
}
StaticType::Any => StaticType::Any,
StaticType::Error => StaticType::Error,
_ => {
diag.push_error(
format!(
"Cannot access field :{} on non-record type {}",
field.name(),
rec_typed.ty.display_compact()
),
Some(rec_typed.identity.clone()),
);
StaticType::Error
}
};
(
NodeKind::GetField {
rec: Rc::new(rec_typed),
field: *field,
},
field_ty,
)
}
NodeKind::Assign { target, value, info } => {
let val_typed = self.check_node(value, ctx, diag);
let ty = val_typed.ty.clone();
if let Some(addr) = info.addr {
ctx.set_type(addr, ty.clone());
}
// For destructuring assigns (addr = None), preserve the target pattern
// so the VM can unpack values. For simple assigns, the target is unused.
let target_typed = if info.addr.is_none() {
Rc::new(self.check_node(target, ctx, diag))
} else {
Rc::new(Node {
identity: node.identity.clone(),
kind: NodeKind::Nop,
ty: ty.clone(),
comments: Rc::from([]),
})
};
(
NodeKind::Assign {
target: target_typed,
value: Rc::new(val_typed),
info: AssignBinding {
addr: info.addr,
},
},
ty,
)
}
NodeKind::If {
cond,
then_br,
else_br,
} => {
let cond_typed = self.check_node(cond, ctx, diag);
let then_typed = self.check_node(then_br, ctx, diag);
let mut else_typed = None;
let mut final_ty = then_typed.ty.clone();
if let Some(e) = else_br {
let et = self.check_node(e, ctx, diag);
if et.ty != final_ty {
final_ty = Self::numeric_widen(&final_ty, &et.ty)
.or_else(|| Self::record_promote(&final_ty, &et.ty))
.unwrap_or(StaticType::Any);
}
else_typed = Some(Rc::new(et));
} else {
final_ty = StaticType::Optional(Box::new(then_typed.ty.clone()));
}
(
NodeKind::If {
cond: Rc::new(cond_typed),
then_br: Rc::new(then_typed),
else_br: else_typed,
},
final_ty,
)
}
NodeKind::Block { exprs } => {
let mut typed_exprs = Vec::new();
let mut last_ty = StaticType::Void;
for e in exprs {
let t = self.check_node(e, ctx, diag);
last_ty = t.ty.clone();
typed_exprs.push(Rc::new(t));
}
(NodeKind::Block { exprs: typed_exprs }, last_ty)
}
NodeKind::Lambda {
params,
body,
info,
} => {
let upvalues = &info.upvalues;
let positional_count = info.positional_count;
let mut upvalue_types = Vec::with_capacity(upvalues.len());
for &addr in upvalues {
upvalue_types.push(ctx.get_type(addr));
}
let mut lambda_ctx =
TypeContext::new(64, upvalue_types, ctx.root_types, Some(ctx));
// Generate a fresh TypeVar per positional parameter so that HM
// constraint propagation works across nested closures.
// `check_lambda_with_hints` (used at call sites) overrides these with
// concrete types; this path only fires for lambdas typed as values
// (e.g. returned from another lambda, stored in a def).
let param_hint_ty = StaticType::Tuple(
(0..positional_count.unwrap_or(0)).map(|_| self.fresh_var()).collect(),
);
let params_typed = self.check_params(
params.as_ref(),
&param_hint_ty,
&mut lambda_ctx,
diag,
);
lambda_ctx.current_params_ty = Some(params_typed.ty.clone());
let body_typed = self.check_node(body, &mut lambda_ctx, diag);
let ret_ty = body_typed.ty.clone();
let fn_ty = StaticType::Function(Box::new(Signature {
params: params_typed.ty.clone(),
ret: ret_ty,
}));
(
NodeKind::Lambda {
params: Rc::new(params_typed),
body: Rc::new(body_typed),
info: LambdaBinding {
upvalues: upvalues.clone(),
positional_count,
},
},
fn_ty,
)
}
NodeKind::Call { callee, args } => {
let callee_typed = self.check_node(callee, ctx, diag);
let args_typed = if let NodeKind::Tuple { elements } = &args.kind {
let arg_count = elements.len();
let mut typed_elements: Vec<Option<Rc<TypedNode>>> = vec![None; arg_count];
let mut known_types: Vec<Option<StaticType>> = vec![None; arg_count];
let mut lambda_indices = Vec::new();
// Phase 1: Type non-lambda arguments first
for (i, e) in elements.iter().enumerate() {
if matches!(e.kind, NodeKind::Lambda { .. }) {
lambda_indices.push(i);
} else {
let t = self.check_node(e, ctx, diag);
known_types[i] = Some(t.ty.clone());
typed_elements[i] = Some(Rc::new(t));
}
}
// Phase 2: Type lambda arguments with parameter hints (if available)
for i in lambda_indices {
let hints = extract_lambda_param_hints(
&callee_typed.ty,
i,
&known_types,
);
let t = if let Some(param_types) = hints {
self.check_lambda_with_param_hints(
&elements[i],
&param_types,
ctx,
diag,
)
} else {
self.check_node(&elements[i], ctx, diag)
};
known_types[i] = Some(t.ty.clone());
typed_elements[i] = Some(Rc::new(t));
}
let final_elements: Vec<Rc<TypedNode>> = typed_elements
.into_iter()
.map(|e| e.expect("all args should be typed"))
.collect();
let elem_types: Vec<StaticType> =
final_elements.iter().map(|e| e.ty.clone()).collect();
Node {
identity: args.identity.clone(),
kind: NodeKind::Tuple {
elements: final_elements,
},
ty: StaticType::Tuple(elem_types),
comments: args.comments.clone(),
}
} else {
self.check_node(args, ctx, diag)
};
let mut ret_ty = match callee_typed.ty.resolve_call(&args_typed.ty) {
Some(ty) => ty,
None => {
let callee_name = match &callee_typed.kind {
NodeKind::Identifier { symbol, .. } => format!("'{}'", symbol.name),
_ => "function".to_string(),
};
diag.push_error(
format!(
"{}: no matching overload for ({})",
callee_name, args_typed.ty.display_compact()
),
Some(node.identity.clone()),
);
StaticType::Error
}
};
// HM: propagate TypeVar constraints through overloaded calls so that
// e.g. `(+ Float TypeVar(1))` resolves TypeVar(1) = Float.
self.unify_matched_overload(&callee_typed.ty, &args_typed.ty, diag);
// HM step 9: when a TypeVar is called with a single Int argument
// (series lookback indexing pattern), record a lazy index-call constraint
// instead of eagerly unifying `TypeVar = Series(elem)`.
//
// The constraint pair (callee_var → result_var) is stored in
// `index_constraints`. When `bind_var` later resolves callee_var to
// `Series(inner)`, it automatically binds result_var to `inner`.
//
// This avoids over-constraining the function to Series-only: passing
// any other callable (e.g. a Function) simply leaves result_var
// unresolved and the call returns `Any` — no spurious type error.
if let StaticType::TypeVar(n) = &callee_typed.ty {
let is_index_call = matches!(&args_typed.ty,
StaticType::Tuple(elems) if elems.len() == 1
&& matches!(&elems[0], StaticType::Int)
);
if is_index_call && matches!(ret_ty, StaticType::Any) {
let existing = self.index_constraints.borrow().get(n).copied();
if let Some(existing_ret_id) = existing {
// Same TypeVar indexed again — all index results on the same
// parameter must share one element TypeVar. Unify the fresh
// var with the existing one so they resolve together.
let elem_var = self.fresh_var();
self.unify(elem_var.clone(), StaticType::TypeVar(existing_ret_id), diag);
ret_ty = Self::apply_subst(elem_var, &self.subst.borrow());
} else {
let elem_var = self.fresh_var();
let StaticType::TypeVar(elem_id) = &elem_var else { unreachable!() };
self.index_constraints.borrow_mut().insert(*n, *elem_id);
ret_ty = elem_var;
}
}
}
// HM step 10: unify Function parameter types with actual argument types,
// but only when the signature still contains TypeVars to resolve.
// Skip for fully concrete signatures (e.g. fn([any any])) to avoid
// false conflicts between Tuple and Vector representations.
if let StaticType::Function(sig) = &callee_typed.ty
&& Self::has_typevar_component(&sig.params)
{
let params = sig.params.clone();
self.unify(params, args_typed.ty.clone(), diag);
ret_ty = Self::apply_subst(ret_ty, &self.subst.borrow());
}
// Dispatch compiler hooks registered by the RTL (keyed by global slot index).
// Hooks handle type-inference extensions such as:
// - series: inject a fresh TypeVar for the element type
// - push: unify the series element TypeVar with the pushed value type
if let NodeKind::Identifier {
binding: IdentifierBinding::Reference(Address::Global(idx)),
..
} = &callee_typed.kind
&& let Some(hook) = self.compiler_hooks.get(&idx.0) {
let hook_ctx = CheckerInferenceAccess { checker: self, ctx };
ret_ty = hook.post_call(&args_typed, ret_ty, &hook_ctx, diag);
}
(
NodeKind::Call {
callee: Rc::new(callee_typed),
args: Rc::new(args_typed),
},
ret_ty,
)
}
NodeKind::Again { args } => {
let args_typed = if let NodeKind::Tuple { elements } = &args.kind {
let mut typed_elements = Vec::new();
let mut elem_types = Vec::new();
for e in elements {
let t = self.check_node(e, ctx, diag);
elem_types.push(t.ty.clone());
typed_elements.push(Rc::new(t));
}
Node {
identity: args.identity.clone(),
kind: NodeKind::Tuple {
elements: typed_elements,
},
ty: StaticType::Tuple(elem_types),
comments: args.comments.clone(),
}
} else {
self.check_node(args, ctx, diag)
};
if let Some(expected_ty) = &ctx.current_params_ty {
self.unify(expected_ty.clone(), args_typed.ty.clone(), diag);
}
(
NodeKind::Again {
args: Rc::new(args_typed),
},
StaticType::Any,
)
}
NodeKind::Tuple { elements } => {
let mut typed_elements = Vec::new();
for e in elements {
typed_elements.push(Rc::new(self.check_node(e, ctx, diag)));
}
let ty = if typed_elements.is_empty() {
StaticType::Vector(Box::new(StaticType::Any), 0)
} else {
let first_ty = &typed_elements[0].ty;
let all_same = typed_elements.iter().all(|e| e.ty == *first_ty);
if all_same {
match first_ty {
StaticType::Vector(inner, len) => {
StaticType::Matrix(inner.clone(), vec![typed_elements.len(), *len])
}
StaticType::Matrix(inner, shape) => {
let mut new_shape = vec![typed_elements.len()];
new_shape.extend(shape);
StaticType::Matrix(inner.clone(), new_shape)
}
_ => {
StaticType::Vector(Box::new(first_ty.clone()), typed_elements.len())
}
}
} else {
StaticType::Tuple(typed_elements.iter().map(|e| e.ty.clone()).collect())
}
};
(
NodeKind::Tuple {
elements: typed_elements,
},
ty,
)
}
NodeKind::Record { fields, layout } => {
let mut typed_fields = Vec::with_capacity(fields.len());
let mut fields_ty = Vec::with_capacity(fields.len());
for (i, (key_node, val_node)) in fields.iter().enumerate() {
let kt = self.check_node(key_node, ctx, diag);
let vt = self.check_node(val_node, ctx, diag);
fields_ty.push((layout.fields[i].0, vt.ty.clone()));
typed_fields.push((Rc::new(kt), Rc::new(vt)));
}
let new_layout = RecordLayout::get_or_create(fields_ty);
(
NodeKind::Record {
fields: typed_fields,
layout: new_layout.clone(),
},
StaticType::Record(new_layout),
)
}
NodeKind::Expansion {
original_call,
expanded,
} => {
let expanded_typed = self.check_node(expanded, ctx, diag);
let ty = expanded_typed.ty.clone();
(
NodeKind::Expansion {
original_call: original_call.clone(),
expanded: Rc::new(expanded_typed),
},
ty,
)
}
NodeKind::Extension(_ext) => {
diag.push_error(
format!(
"TypeChecking for extension '{}' not implemented",
_ext.display_name()
),
Some(node.identity.clone()),
);
(NodeKind::Error, StaticType::Error)
}
NodeKind::Error => (NodeKind::Error, StaticType::Error),
// Syntax-only variants should not appear in bound phases
NodeKind::MacroDecl { .. }
| NodeKind::Template(_)
| NodeKind::Placeholder(_)
| NodeKind::Splice(_) => {
diag.push_error(
"Unexpected syntax-only node in type checking",
Some(node.identity.clone()),
);
(NodeKind::Error, StaticType::Error)
}
};
Node {
identity: node.identity.clone(),
kind,
ty,
comments: node.comments.clone(),
}
}
}
+167
View File
@@ -0,0 +1,167 @@
use crate::ast::compiler::call_hooks::InferenceAccess;
use crate::ast::diagnostics::Diagnostics;
use crate::ast::nodes::{Address, VirtualId};
use crate::ast::types::{Signature, StaticType};
use std::collections::HashMap;
use super::TypeChecker;
/// Manages the types of locals and upvalues during a single type-checking pass.
pub(super) struct TypeContext<'a> {
pub(super) _parent: Option<&'a TypeContext<'a>>,
/// Maps slot index -> Inferred Type
pub(super) slots: HashMap<u32, StaticType>,
/// Types of captured variables (passed from outer scope)
pub(super) upvalue_types: Vec<StaticType>,
/// Access to root types for unified resolution
pub(super) root_types: &'a std::cell::RefCell<Vec<StaticType>>,
/// The expected parameters of the current function (for 'again' validation)
pub(super) current_params_ty: Option<StaticType>,
}
impl<'a> TypeContext<'a> {
pub(super) fn new(
_slot_count: u32,
upvalue_types: Vec<StaticType>,
root_types: &'a std::cell::RefCell<Vec<StaticType>>,
parent: Option<&'a TypeContext<'a>>,
) -> Self {
Self {
_parent: parent,
slots: HashMap::new(),
upvalue_types,
root_types,
current_params_ty: None,
}
}
pub(super) fn get_type(&self, addr: Address<VirtualId>) -> StaticType {
match addr {
Address::Local(slot) => self
.slots
.get(&slot.0)
.cloned()
.unwrap_or(StaticType::Any),
Address::Global(idx) => self
.root_types
.borrow()
.get(idx.0 as usize)
.cloned()
.unwrap_or(StaticType::Any),
Address::Upvalue(idx) => self
.upvalue_types
.get(idx.0 as usize)
.cloned()
.unwrap_or(StaticType::Any),
}
}
pub(super) fn set_type(&mut self, addr: Address<VirtualId>, ty: StaticType) {
match addr {
Address::Local(slot) => {
self.slots.insert(slot.0, ty);
}
Address::Global(idx) => {
let mut rt = self.root_types.borrow_mut();
if (idx.0 as usize) < rt.len() {
rt[idx.0 as usize] = ty;
}
}
_ => {}
}
}
}
/// Temporary wrapper that gives call-hooks unified access to both the
/// type-checker's inference state and the current scope's slot types.
/// Created on the stack at each hook dispatch site; zero allocation cost.
pub(super) struct CheckerInferenceAccess<'a, 'b> {
pub(super) checker: &'a TypeChecker,
pub(super) ctx: &'b TypeContext<'a>,
}
impl InferenceAccess for CheckerInferenceAccess<'_, '_> {
fn fresh_var(&self) -> StaticType {
self.checker.fresh_var()
}
fn unify(&self, a: StaticType, b: StaticType, diag: &mut Diagnostics) {
self.checker.unify(a, b, diag);
}
fn bind_typevar(&self, id: u32, ty: StaticType) {
self.checker.bind_var(id, ty);
}
fn apply_subst_ty(&self, ty: StaticType) -> StaticType {
TypeChecker::apply_subst(ty, &self.checker.subst.borrow())
}
fn try_numeric_widen(&self, a: &StaticType, b: &StaticType) -> Option<StaticType> {
TypeChecker::numeric_widen(a, b)
}
fn try_record_promote(&self, a: &StaticType, b: &StaticType) -> Option<StaticType> {
TypeChecker::record_promote(a, b)
}
fn get_slot_type(&self, addr: Address<VirtualId>) -> StaticType {
self.ctx.get_type(addr)
}
}
/// Extracts expected lambda parameter types from the callee type for a specific argument position.
/// Used for bidirectional type inference: when a Call's callee expects a function at `arg_index`,
/// this returns the expected parameter types for that function, derived from the callee's signature
/// and the already-known types of non-lambda arguments.
pub(super) fn extract_lambda_param_hints(
callee_ty: &StaticType,
arg_index: usize,
known_arg_types: &[Option<StaticType>],
) -> Option<Vec<StaticType>> {
/// Helper: given the expected type at a specific parameter position in a signature,
/// extract the lambda parameter types if it expects a function.
fn hints_from_param_type(param_ty: &StaticType) -> Option<Vec<StaticType>> {
if let StaticType::Function(sig) = param_ty {
if let StaticType::Tuple(params) = &sig.params {
return Some(params.clone());
}
return Some(vec![sig.params.clone()]);
}
None
}
/// Helper: extract the expected type at `arg_index` from a signature's params tuple.
fn param_at(sig: &Signature, arg_index: usize) -> Option<&StaticType> {
if let StaticType::Tuple(params) = &sig.params {
params.get(arg_index)
} else if arg_index == 0 {
Some(&sig.params)
} else {
None
}
}
match callee_ty {
StaticType::Function(sig) => {
let expected = param_at(sig, arg_index)?;
hints_from_param_type(expected)
}
StaticType::FunctionOverloads(sigs) => {
// Try each overload — return hints from the first one that has a function at this position
for sig in sigs {
if let Some(expected) = param_at(sig, arg_index)
&& let Some(hints) = hints_from_param_type(expected)
{
return Some(hints);
}
}
None
}
StaticType::PolymorphicFn {
resolve_arg_hints: Some(resolver),
..
} => resolver(arg_index, known_arg_types),
_ => None,
}
}
+135
View File
@@ -0,0 +1,135 @@
use crate::ast::nodes::{
Address, IdentifierBinding, Node, NodeKind, TypedNode, TypedPhase,
};
use crate::ast::types::{NodeIdentity, StaticType};
use std::collections::HashMap;
use std::rc::Rc;
use super::TypeChecker;
impl TypeChecker {
/// Applies the final HM substitution to all type annotations in the tree and
/// elaborates `(series n)` calls with their inferred schema arguments.
/// Must be called after `check()` or `check_node_as_bound()` completes.
pub fn finalize(&self, node: TypedNode) -> TypedNode {
let subst = self.subst.borrow().clone();
self.finalize_node(node, &subst)
}
/// Walks the typed AST, applies the HM substitution to every type annotation,
/// and dispatches finalization hooks (e.g. schema injection for `series`).
fn finalize_node(&self, node: TypedNode, subst: &HashMap<u32, StaticType>) -> TypedNode {
let new_ty = Self::apply_subst(node.ty, subst);
let new_kind = self.finalize_kind(node.kind, subst, &new_ty, &node.identity);
Node { kind: new_kind, ty: new_ty, identity: node.identity, comments: node.comments }
}
fn finalize_kind(
&self,
kind: NodeKind<TypedPhase>,
subst: &HashMap<u32, StaticType>,
node_ty: &StaticType,
_identity: &Rc<NodeIdentity>,
) -> NodeKind<TypedPhase> {
match kind {
// Leaf nodes — nothing to recurse into
NodeKind::Nop => NodeKind::Nop,
NodeKind::Constant(v) => NodeKind::Constant(v),
NodeKind::FieldAccessor(k) => NodeKind::FieldAccessor(k),
NodeKind::Error => NodeKind::Error,
NodeKind::Identifier { symbol, binding } => NodeKind::Identifier { symbol, binding },
NodeKind::Extension(ext) => NodeKind::Extension(ext),
// Call: dispatch finalize hook registered by RTL (keyed by global slot index).
// Hooks may rewrite the call node — e.g. series injects a schema argument.
NodeKind::Call { callee, args } => {
let callee_fin = Rc::new(self.finalize_node((*callee).clone(), subst));
let args_fin = Rc::new(self.finalize_node((*args).clone(), subst));
if let NodeKind::Identifier {
binding: IdentifierBinding::Reference(Address::Global(idx)),
..
} = &callee_fin.kind
&& let Some(hook) = self.compiler_hooks.get(&idx.0)
&& let Some(new_kind) = hook.finalize(
Rc::clone(&callee_fin),
Rc::clone(&args_fin),
node_ty,
subst,
) {
return new_kind;
}
NodeKind::Call { callee: callee_fin, args: args_fin }
}
NodeKind::If { cond, then_br, else_br } => NodeKind::If {
cond: Rc::new(self.finalize_node((*cond).clone(), subst)),
then_br: Rc::new(self.finalize_node((*then_br).clone(), subst)),
else_br: else_br.map(|e| Rc::new(self.finalize_node((*e).clone(), subst))),
},
NodeKind::Def { pattern, value, info } => NodeKind::Def {
pattern: Rc::new(self.finalize_node((*pattern).clone(), subst)),
value: Rc::new(self.finalize_node((*value).clone(), subst)),
info,
},
NodeKind::Assign { target, value, info } => NodeKind::Assign {
target: Rc::new(self.finalize_node((*target).clone(), subst)),
value: Rc::new(self.finalize_node((*value).clone(), subst)),
info,
},
NodeKind::Lambda { params, body, info } => NodeKind::Lambda {
params: Rc::new(self.finalize_node((*params).clone(), subst)),
body: Rc::new(self.finalize_node((*body).clone(), subst)),
info,
},
NodeKind::Again { args } => NodeKind::Again {
args: Rc::new(self.finalize_node((*args).clone(), subst)),
},
NodeKind::Block { exprs } => NodeKind::Block {
exprs: exprs
.into_iter()
.map(|e| Rc::new(self.finalize_node((*e).clone(), subst)))
.collect(),
},
NodeKind::Tuple { elements } => NodeKind::Tuple {
elements: elements
.into_iter()
.map(|e| Rc::new(self.finalize_node((*e).clone(), subst)))
.collect(),
},
NodeKind::Record { fields, layout } => NodeKind::Record {
fields: fields
.into_iter()
.map(|(k, v)| {
(
Rc::new(self.finalize_node((*k).clone(), subst)),
Rc::new(self.finalize_node((*v).clone(), subst)),
)
})
.collect(),
layout,
},
NodeKind::GetField { rec, field } => NodeKind::GetField {
rec: Rc::new(self.finalize_node((*rec).clone(), subst)),
field,
},
NodeKind::Expansion { original_call, expanded } => NodeKind::Expansion {
original_call,
expanded: Rc::new(self.finalize_node((*expanded).clone(), subst)),
},
NodeKind::MacroDecl { name, params, body } => NodeKind::MacroDecl {
name,
params: Rc::new(self.finalize_node((*params).clone(), subst)),
body: Rc::new(self.finalize_node((*body).clone(), subst)),
},
NodeKind::Template(inner) => {
NodeKind::Template(Rc::new(self.finalize_node((*inner).clone(), subst)))
}
NodeKind::Placeholder(inner) => {
NodeKind::Placeholder(Rc::new(self.finalize_node((*inner).clone(), subst)))
}
NodeKind::Splice(inner) => {
NodeKind::Splice(Rc::new(self.finalize_node((*inner).clone(), subst)))
}
}
}
}
+401
View File
@@ -0,0 +1,401 @@
use crate::ast::diagnostics::Diagnostics;
use crate::ast::types::{RecordLayout, Signature, StaticType};
use std::collections::HashMap;
use super::context::TypeContext;
use super::TypeChecker;
impl TypeChecker {
/// Creates a fresh, unique type variable.
pub(super) fn fresh_var(&self) -> StaticType {
let id = self.var_counter.get();
self.var_counter.set(id + 1);
StaticType::TypeVar(id)
}
/// Binds a TypeVar to a concrete type in the substitution, then propagates
/// any pending index-call constraints registered for the series lookback pattern.
///
/// When `index_constraints[id]` is set and `ty` is a `Series`, the result TypeVar
/// is bound to the element type — connecting the callee TypeVar to its element type
/// without the eager over-constraint of unification. Non-indexable types leave the
/// result TypeVar unresolved.
pub(super) fn bind_var(&self, id: u32, ty: StaticType) {
self.subst.borrow_mut().insert(id, ty.clone());
let ret_id = self.index_constraints.borrow().get(&id).copied();
if let (Some(ret_id), StaticType::Series(elem)) = (ret_id, &ty) {
self.subst.borrow_mut().insert(ret_id, *elem.clone());
}
}
/// Recursively applies the substitution map to a type, replacing all
/// resolved `TypeVar`s with their concrete types.
pub(super) 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())
}
// Substitute into the body but skip over the bound vars: they are local to
// this schema and must not be replaced by the global substitution.
StaticType::Forall(vars, body) => {
let filtered: HashMap<u32, StaticType> =
subst.iter().filter(|(k, _)| !vars.contains(k)).map(|(k, v)| (*k, v.clone())).collect();
StaticType::Forall(vars, Box::new(Self::apply_subst(*body, &filtered)))
}
other => other,
}
}
/// Returns true if `TypeVar(var_id)` appears anywhere in `ty` under the
/// current substitution. Used to prevent infinite types (occurs check).
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)),
// A bound var inside Forall does not count as a free occurrence.
StaticType::Forall(vars, body) => {
!vars.contains(&var_id) && Self::occurs(var_id, body, 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.
pub(super) fn unify(&self, ty1: StaticType, ty2: StaticType, diag: &mut Diagnostics) {
let 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.display_compact()), None);
return;
}
// Release the borrow before routing through bind_var so that
// constraint propagation can re-borrow subst without a panic.
drop(subst);
self.bind_var(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);
}
// Unify element-wise so that overload resolution can propagate TypeVar
// constraints: e.g. `(+ Float TypeVar(1))` matched against `(Float, Float)`
// triggers `unify(Float, TypeVar(1))` → `subst[1] = Float`.
(StaticType::Tuple(a), StaticType::Tuple(b)) => {
drop(subst);
for (ta, tb) in a.into_iter().zip(b.into_iter()) {
self.unify(ta, tb, diag);
}
}
// A homogeneous Vector is assignable from a Tuple — unify each element
// with the vector's inner type. Mirrors the `is_assignable_from` coercion.
(StaticType::Tuple(elems), StaticType::Vector(inner, len))
| (StaticType::Vector(inner, len), StaticType::Tuple(elems)) => {
if elems.len() != len {
diag.push_error(
format!("Type mismatch: expected tuple of length {}, got {}", len, elems.len()),
None,
);
return;
}
drop(subst);
for e in elems {
self.unify(e, (*inner).clone(), diag);
}
}
// Variadic(T) unifies with a Tuple by unifying each element with T.
(StaticType::Variadic(inner), StaticType::Tuple(elems))
| (StaticType::Tuple(elems), StaticType::Variadic(inner)) => {
drop(subst);
for e in elems {
self.unify(e, (*inner).clone(), diag);
}
}
// Any and Error are already handled by is_assignable_from — silently succeed
(StaticType::Any, _) | (_, StaticType::Any) => {}
(StaticType::Error, _) | (_, StaticType::Error) => {}
// Forall should be instantiated before unification; delegate to the body.
(StaticType::Forall(_, body), other) | (other, StaticType::Forall(_, body)) => {
drop(subst);
self.unify(*body, other, diag);
}
(a, b) => {
diag.push_error(format!("Type mismatch: expected {}, got {}", a.display_compact(), b.display_compact()), None);
}
}
}
/// After a `FunctionOverloads` call resolves successfully, unify the matched
/// overload's concrete parameter types with the actual argument types.
///
/// This propagates TypeVar constraints through overloaded calls — for example,
/// `(+ Float TypeVar(1))` matched against `(Float, Float) → Float` binds
/// `TypeVar(1) = Float` so that nested closures can resolve series element types.
///
/// The "concrete anchor" guard ensures we only unify when at least one argument
/// is a known concrete type. Without it, `(+ TypeVar TypeVar)` would
/// spuriously pick the first overload (e.g. Int) and bind both TypeVars to Int.
pub(super) fn unify_matched_overload(
&self,
callee_ty: &StaticType,
args_ty: &StaticType,
diag: &mut Diagnostics,
) {
let StaticType::FunctionOverloads(sigs) = callee_ty else { return };
// Require both a concrete anchor (to pin the overload choice) and at
// least one TypeVar (something to actually bind). Pure concrete calls
// like `(- DateTime DateTime)` have nothing to unify and would
// incorrectly trigger errors from coercion-only compatible types.
if !Self::has_concrete_component(args_ty) || !Self::has_typevar_component(args_ty) {
return;
}
// Only unify when EXACTLY ONE non-variadic overload matches.
// If multiple overloads match (e.g. `(* TypeVar Int)` matches both
// `(Int,Int)→Int` and `(Float,Float)→Float`), the choice is ambiguous
// and binding the TypeVar to the first match would be incorrect.
let is_variadic = |sig: &Signature| matches!(sig.params, StaticType::Any | StaticType::Variadic(_));
let mut unique_match: Option<&Signature> = None;
for sig in sigs {
if !is_variadic(sig) && sig.params.is_assignable_from(args_ty) {
if unique_match.is_some() {
return; // Ambiguous — more than one specific overload matches
}
unique_match = Some(sig);
}
}
if let Some(sig) = unique_match {
self.unify(sig.params.clone(), args_ty.clone(), diag);
}
}
/// Returns true if `ty` (or any element of a Tuple) is a concrete type —
/// i.e. not `Any`, `TypeVar`, or `Error`.
pub(super) fn has_concrete_component(ty: &StaticType) -> bool {
match ty {
StaticType::Tuple(elems) => elems.iter().any(Self::is_concrete),
other => Self::is_concrete(other),
}
}
/// Returns true if `ty` contains a TypeVar anywhere in its structure.
pub(super) fn has_typevar_component(ty: &StaticType) -> bool {
match ty {
StaticType::TypeVar(_) => true,
StaticType::Tuple(elems) => elems.iter().any(Self::has_typevar_component),
StaticType::Series(inner) | StaticType::Stream(inner) | StaticType::Optional(inner) => {
Self::has_typevar_component(inner)
}
StaticType::Function(sig) => {
Self::has_typevar_component(&sig.params) || Self::has_typevar_component(&sig.ret)
}
_ => false,
}
}
fn is_concrete(ty: &StaticType) -> bool {
!matches!(ty, StaticType::Any | StaticType::TypeVar(_) | StaticType::Error)
}
/// Returns the widened numeric type when one side is `Int` and the other `Float`.
/// This is the only implicit numeric promotion in Myc: Int is a subtype of Float.
pub(super) fn numeric_widen(a: &StaticType, b: &StaticType) -> Option<StaticType> {
match (a, b) {
(StaticType::Int, StaticType::Float) | (StaticType::Float, StaticType::Int) => {
Some(StaticType::Float)
}
_ => None,
}
}
/// If both types are Records with the same field names and compatible types
/// (same type or Int→Float widening), returns the promoted Record type.
pub(super) fn record_promote(a: &StaticType, b: &StaticType) -> Option<StaticType> {
let (StaticType::Record(la), StaticType::Record(lb)) = (a, b) else {
return None;
};
if la == lb {
return None;
}
if la.fields.len() != lb.fields.len() {
return None;
}
let mut promoted_fields = Vec::with_capacity(la.fields.len());
for ((ka, ta), (kb, tb)) in la.fields.iter().zip(lb.fields.iter()) {
if ka != kb {
return None;
}
let t = if ta == tb {
ta.clone()
} else if let Some(w) = Self::numeric_widen(ta, tb) {
w
} else {
return None;
};
promoted_fields.push((*ka, t));
}
Some(StaticType::Record(RecordLayout::get_or_create(promoted_fields)))
}
/// Collects all free TypeVar IDs that appear in `ty`, following the substitution
/// chain and skipping variables bound by `Forall`. Deduplicates via `out`.
fn collect_free_tvars(ty: &StaticType, subst: &HashMap<u32, StaticType>, out: &mut Vec<u32>) {
match ty {
StaticType::TypeVar(n) => {
if let Some(resolved) = subst.get(n) {
Self::collect_free_tvars(resolved, subst, out);
} else if !out.contains(n) {
out.push(*n);
}
}
StaticType::Series(inner) | StaticType::Stream(inner) | StaticType::Optional(inner) => {
Self::collect_free_tvars(inner, subst, out);
}
StaticType::Function(sig) => {
Self::collect_free_tvars(&sig.params, subst, out);
Self::collect_free_tvars(&sig.ret, subst, out);
}
StaticType::Tuple(elems) => {
for e in elems {
Self::collect_free_tvars(e, subst, out);
}
}
// Recurse into the body but skip the locally bound vars.
StaticType::Forall(vars, body) => {
let mut body_free = Vec::new();
Self::collect_free_tvars(body, subst, &mut body_free);
for v in body_free {
if !vars.contains(&v) && !out.contains(&v) {
out.push(v);
}
}
}
_ => {}
}
}
/// Collects all free TypeVar IDs that are currently visible in `ctx`.
/// Used by `generalize` to avoid quantifying TypeVars that are still shared
/// with other live bindings (series, scalars, outer function params).
fn ctx_free_tvars(ctx: &TypeContext, subst: &HashMap<u32, StaticType>) -> Vec<u32> {
let mut out = Vec::new();
for ty in ctx.slots.values() {
Self::collect_free_tvars(ty, subst, &mut out);
}
for ty in &ctx.upvalue_types {
Self::collect_free_tvars(ty, subst, &mut out);
}
for ty in ctx.root_types.borrow().iter() {
Self::collect_free_tvars(ty, subst, &mut out);
}
out
}
/// Generalizes a type at a `def` boundary (Algorithm W `gen` step).
/// Wraps all TypeVars that are free in `ty` but NOT free in `ctx` into a
/// `Forall`. Value restriction: only call this for `Function`-typed values.
pub(super) fn generalize(&self, ty: StaticType, ctx: &TypeContext) -> StaticType {
let subst = self.subst.borrow();
let resolved = Self::apply_subst(ty, &subst);
let mut tvars_in_ty = Vec::new();
Self::collect_free_tvars(&resolved, &subst, &mut tvars_in_ty);
let ctx_tvars = Self::ctx_free_tvars(ctx, &subst);
let quantified: Vec<u32> = tvars_in_ty.into_iter()
.filter(|v| !ctx_tvars.contains(v))
.collect();
if quantified.is_empty() {
resolved
} else {
StaticType::Forall(quantified, Box::new(resolved))
}
}
/// Instantiates a `Forall` type by replacing each bound TypeVar with a
/// fresh one (Algorithm W `inst` step). No-op for non-`Forall` types.
///
/// Also remaps any index-call constraints: if `index_constraints[old] = ret`
/// and both `old` and `ret` are bound vars, the fresh copies inherit the
/// same pairing so that `bind_var` propagation keeps working at each call site.
pub(super) fn instantiate(&self, ty: StaticType) -> StaticType {
let StaticType::Forall(vars, body) = ty else { return ty };
let mut local_subst: HashMap<u32, StaticType> = HashMap::new();
let mut var_mapping: HashMap<u32, u32> = HashMap::new();
for v in &vars {
let fresh = self.fresh_var();
if let StaticType::TypeVar(fresh_id) = &fresh {
var_mapping.insert(*v, *fresh_id);
}
local_subst.insert(*v, fresh);
}
// Copy index-call constraints for the freshly created TypeVars.
// Both the callee-var and its result-var must be in vars for the
// mapping to apply; partial remaps are dropped (they can't occur in
// a well-formed Forall, but the guard is cheap insurance).
let new_constraints: Vec<(u32, u32)> = {
let constraints = self.index_constraints.borrow();
vars.iter()
.filter_map(|v| {
let new_v = *var_mapping.get(v)?;
let old_ret = *constraints.get(v)?;
let new_ret = *var_mapping.get(&old_ret)?;
Some((new_v, new_ret))
})
.collect()
};
for (new_v, new_ret) in new_constraints {
self.index_constraints.borrow_mut().insert(new_v, new_ret);
}
Self::apply_subst(*body, &local_subst)
}
}
+134
View File
@@ -0,0 +1,134 @@
mod context;
mod inference;
mod finalize;
mod check;
#[cfg(test)]
mod tests;
use crate::ast::compiler::call_hooks::RtlCompilerHook;
use crate::ast::diagnostics::Diagnostics;
use crate::ast::nodes::{
BoundLike, BoundPhase, LambdaBinding,
Node, NodeKind, TypedNode,
};
use crate::ast::types::{Signature, StaticType};
use std::cell::Cell;
use std::collections::HashMap;
use std::rc::Rc;
use context::TypeContext;
pub struct TypeChecker {
pub(crate) root_types: Rc<std::cell::RefCell<Vec<StaticType>>>,
/// Monotonic counter for generating unique type variable IDs.
pub(crate) var_counter: Cell<u32>,
/// Global substitution map: TypeVar ID → resolved StaticType.
/// Shared across all scopes within a single type-checking pass.
pub(crate) subst: std::cell::RefCell<HashMap<u32, StaticType>>,
/// Index-call constraints: callee_var → result_var.
///
/// When a `TypeVar` is used as a callable with a single `Int` argument —
/// the series lookback pattern `(s 0)` — we record the pairing here instead
/// of eagerly unifying `TypeVar = Series(elem)`. The constraint is evaluated
/// lazily in `bind_var`: when the callee TypeVar is bound to a concrete type,
/// the element type is extracted directly from the `Series` variant.
pub(crate) index_constraints: std::cell::RefCell<HashMap<u32, u32>>,
/// Compiler hooks keyed by global slot index.
/// Populated from the frozen RTL snapshot; empty for non-RTL call sites.
pub(crate) compiler_hooks: Rc<HashMap<u32, Rc<dyn RtlCompilerHook>>>,
}
impl TypeChecker {
pub fn new(
root_types: Rc<std::cell::RefCell<Vec<StaticType>>>,
compiler_hooks: Rc<HashMap<u32, Rc<dyn RtlCompilerHook>>>,
) -> Self {
Self {
root_types,
var_counter: Cell::new(0),
subst: std::cell::RefCell::new(HashMap::new()),
index_constraints: std::cell::RefCell::new(HashMap::new()),
compiler_hooks,
}
}
pub fn check(
&self,
node: &Node<BoundPhase>,
arg_types: &[StaticType],
diag: &mut Diagnostics,
) -> TypedNode {
let typed = self.check_node_as_bound(node, arg_types, diag);
self.finalize(typed)
}
/// Allows re-checking a node from any phase as if it were a bound node.
/// This is useful for specialization where we re-type a TypedNode with more specific info.
pub fn check_node_as_bound<P: BoundLike>(
&self,
node: &Node<P>,
arg_types: &[StaticType],
diag: &mut Diagnostics,
) -> TypedNode {
match &node.kind {
NodeKind::Lambda {
params,
body,
info,
} => {
let upvalues = &info.upvalues;
let positional_count = info.positional_count;
let mut upvalue_types = Vec::with_capacity(upvalues.len());
for &_addr in upvalues {
upvalue_types.push(StaticType::Any);
}
let root_ctx = TypeContext::new(0, vec![], &self.root_types, None);
let mut lambda_ctx =
TypeContext::new(64, upvalue_types, &self.root_types, Some(&root_ctx));
let arg_tuple_ty = if arg_types.is_empty() {
StaticType::Any
} else {
StaticType::Tuple(arg_types.to_vec())
};
let params_typed = self.check_params(
params.as_ref(),
&arg_tuple_ty,
&mut lambda_ctx,
diag,
);
let body_typed = self.check_node(body, &mut lambda_ctx, diag);
let ret_ty = body_typed.ty.clone();
let final_params_ty = params_typed.ty.clone();
let fn_ty = StaticType::Function(Box::new(Signature {
params: final_params_ty,
ret: ret_ty,
}));
Node {
identity: node.identity.clone(),
kind: NodeKind::Lambda {
params: Rc::new(params_typed),
body: Rc::new(body_typed),
info: LambdaBinding {
upvalues: upvalues.clone(),
positional_count,
},
},
ty: fn_ty,
comments: node.comments.clone(),
}
}
_ => {
let mut root_ctx = TypeContext::new(0, vec![], &self.root_types, None);
self.check_node(node, &mut root_ctx, diag)
}
}
}
}
+193
View File
@@ -0,0 +1,193 @@
use super::*;
use crate::ast::environment::Environment;
use crate::ast::types::StaticType;
fn check_source(source: &str) -> TypedNode {
let env = Environment::new();
env.compile(source).into_result().unwrap()
}
fn get_ret_type(node: &TypedNode) -> StaticType {
if let StaticType::Function(sig) = &node.ty {
sig.ret.clone()
} else {
node.ty.clone()
}
}
#[test]
fn test_inference_constants() {
assert_eq!(get_ret_type(&check_source("10")), StaticType::Int);
assert_eq!(get_ret_type(&check_source("10.5")), StaticType::Float);
assert_eq!(get_ret_type(&check_source("true")), StaticType::Bool);
assert_eq!(get_ret_type(&check_source("\"hello\"")), 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)");
// Outer is Lambda, Body is Block
if let NodeKind::Lambda { body, .. } = &typed.kind {
if let NodeKind::Block { exprs } = &body.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 in lambda body");
}
} else {
panic!("Expected Lambda wrapper");
}
}
#[test]
fn test_inference_block_type() {
// Block type = last expression type
assert_eq!(get_ret_type(&check_source("(do 1 2.5)")), StaticType::Float);
assert_eq!(
get_ret_type(&check_source("(do 1.5 \"test\")")),
StaticType::Text
);
}
#[test]
fn test_inference_lambda_return() {
// (fn [a] 10) -> fn(any) -> Int
// Since it's already a Lambda, it's NOT wrapped further.
let typed = check_source("(fn [a] 10)");
if let StaticType::Function(sig) = &typed.ty {
assert_eq!(sig.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(sig) = &typed_nested.ty {
assert_eq!(sig.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 NodeKind::Lambda { body, .. } = &typed.kind {
if let NodeKind::Block { exprs } = &body.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");
}
} else {
panic!("Expected Lambda");
}
}
#[test]
fn test_operator_overloading_inference() {
assert_eq!(get_ret_type(&check_source("(+ 1 2)")), StaticType::Int);
assert_eq!(
get_ret_type(&check_source("(+ 1.0 2.0)")),
StaticType::Float
);
assert_eq!(
get_ret_type(&check_source("(+ \"a\" \"b\")")),
StaticType::Text
);
assert_eq!(get_ret_type(&check_source("(/ 1 2)")), StaticType::Float);
}
#[test]
fn test_datetime_inference() {
// date("2023-01-01") -> DateTime
assert_eq!(
get_ret_type(&check_source("(date \"2023-01-01\")")),
StaticType::DateTime
);
// DateTime + Int -> DateTime
assert_eq!(
get_ret_type(&check_source("(+ (date \"2023-01-01\") 86400000)")),
StaticType::DateTime
);
// DateTime - DateTime -> Int (Duration)
assert_eq!(
get_ret_type(&check_source(
"(- (date \"2023-01-02\") (date \"2023-01-01\"))"
)),
StaticType::Int
);
// DateTime comparison -> Bool
assert_eq!(
get_ret_type(&check_source(
"(> (date \"2023-01-02\") (date \"2023-01-01\"))"
)),
StaticType::Bool
);
}
#[test]
fn test_inference_tuple_vector_matrix() {
// Heterogeneous -> Tuple
assert_eq!(
get_ret_type(&check_source("[1 3.14 \"text\"]")),
StaticType::Tuple(vec![StaticType::Int, StaticType::Float, StaticType::Text])
);
// Homogeneous -> Vector
assert_eq!(
get_ret_type(&check_source("[10 20 30]")),
StaticType::Vector(Box::new(StaticType::Int), 3)
);
// Nested Homogeneous -> Matrix
assert_eq!(
get_ret_type(&check_source("[[1 2] [3 4]]")),
StaticType::Matrix(Box::new(StaticType::Int), vec![2, 2])
);
// Deep Matrix
assert_eq!(
get_ret_type(&check_source("[[[1 2]] [[3 4]]]")),
StaticType::Matrix(Box::new(StaticType::Int), vec![2, 1, 2])
);
// Shape mismatch -> Tuple of Vectors
let mixed = get_ret_type(&check_source("[[1 2] [3 4 5]]"));
if let StaticType::Tuple(elements) = mixed {
assert_eq!(
elements[0],
StaticType::Vector(Box::new(StaticType::Int), 2)
);
assert_eq!(
elements[1],
StaticType::Vector(Box::new(StaticType::Int), 3)
);
} else {
panic!("Expected Tuple for shape mismatch, got {:?}", mixed);
}
}
#[test]
fn test_inference_record() {
use crate::ast::types::Keyword;
let typed = check_source("{:x 1 :y 0.3}");
let ty = get_ret_type(&typed);
if let StaticType::Record(layout) = ty {
assert_eq!(layout.fields.len(), 2);
assert_eq!(layout.fields[0], (Keyword::intern("x"), StaticType::Int));
assert_eq!(layout.fields[1], (Keyword::intern("y"), StaticType::Float));
} else {
panic!("Expected Record, got {:?}", ty);
}
}