Refactor type checker for series and overloads

Introduce `unify_matched_overload` to handle type variable propagation
for overloaded functions. This ensures that when an overload is chosen,
its concrete parameter types are unified with the actual argument types.
This is crucial for resolving type variables nested within closures
called through overloaded functions.

Additionally, this commit:
- Allows `series` to infer schema type from `push` calls, simplifying
  its signature.
- Adds a fallback to `ValueSeries` in `rtl::series` if schema injection
  fails.
- Updates `StaticType::TypeVar` to return `Any` when accessed,
  simplifying field access logic.
- Adjusts tests to reflect the simplified `series` signature.
This commit is contained in:
2026-03-28 14:03:48 +01:00
parent d2cba2d55d
commit d7d1aef8ed
5 changed files with 150 additions and 67 deletions
+135 -63
View File
@@ -243,6 +243,15 @@ impl TypeChecker {
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);
}
}
// Any and Error are already handled by is_assignable_from — silently succeed
(StaticType::Any, _) | (_, StaticType::Any) => {}
(StaticType::Error, _) | (_, StaticType::Error) => {}
@@ -257,6 +266,71 @@ impl TypeChecker {
matches!(&node.kind, NodeKind::Identifier { symbol, .. } if symbol.name.as_ref() == name)
}
/// 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.
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);
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`.
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` (or any element of a Tuple) contains a TypeVar.
fn has_typevar_component(ty: &StaticType) -> bool {
match ty {
StaticType::TypeVar(_) => true,
StaticType::Tuple(elems) => elems.iter().any(|e| matches!(e, StaticType::TypeVar(_))),
_ => false,
}
}
fn is_concrete(ty: &StaticType) -> bool {
!matches!(ty, StaticType::Any | StaticType::TypeVar(_) | StaticType::Error)
}
/// Converts a resolved series element type to the schema `Value` passed to the
/// `series` runtime: a type keyword (`:float`, `:int`, …) for scalar types, or a
/// schema record (`{:price :float …}`) for record types.
@@ -313,41 +387,30 @@ impl TypeChecker {
// Call: detect `(series n)` and inject the resolved schema arg
NodeKind::Call { callee, args } => {
let callee_fin = Rc::new(Self::finalize_node((*callee).clone(), subst));
if Self::is_identifier_named("series", &callee_fin) {
if let StaticType::Series(inner) = node_ty {
if let NodeKind::Tuple { elements } = &args.kind {
if elements.len() == 1 {
if let Some(schema_val) =
Self::series_element_to_schema_value(inner)
{
let orig_arg = Rc::new(Self::finalize_node(
(*elements[0]).clone(),
subst,
));
let schema_ty = schema_val.static_type();
let schema_node = Rc::new(Node {
kind: NodeKind::Constant(schema_val),
ty: schema_ty,
identity: identity.clone(),
comments: Rc::from([]),
});
let new_elems = vec![orig_arg, schema_node];
let elem_types: Vec<StaticType> =
new_elems.iter().map(|e| e.ty.clone()).collect();
let new_args = Node {
kind: NodeKind::Tuple { elements: new_elems },
ty: StaticType::Tuple(elem_types),
identity: args.identity.clone(),
comments: args.comments.clone(),
};
return NodeKind::Call {
callee: callee_fin,
args: Rc::new(new_args),
};
}
}
}
}
if Self::is_identifier_named("series", &callee_fin)
&& let StaticType::Series(inner) = node_ty
&& let NodeKind::Tuple { elements } = &args.kind
&& elements.len() == 1
&& let Some(schema_val) = Self::series_element_to_schema_value(inner)
{
let orig_arg = Rc::new(Self::finalize_node((*elements[0]).clone(), subst));
let schema_ty = schema_val.static_type();
let schema_node = Rc::new(Node {
kind: NodeKind::Constant(schema_val),
ty: schema_ty,
identity: identity.clone(),
comments: Rc::from([]),
});
let new_elems = vec![orig_arg, schema_node];
let elem_types: Vec<StaticType> =
new_elems.iter().map(|e| e.ty.clone()).collect();
let new_args = Node {
kind: NodeKind::Tuple { elements: new_elems },
ty: StaticType::Tuple(elem_types),
identity: args.identity.clone(),
comments: args.comments.clone(),
};
return NodeKind::Call { callee: callee_fin, args: Rc::new(new_args) };
}
let args_fin = Rc::new(Self::finalize_node((*args).clone(), subst));
NodeKind::Call { callee: callee_fin, args: args_fin }
@@ -701,6 +764,7 @@ impl TypeChecker {
) -> TypedNode {
match specialized_ty {
StaticType::Any
| StaticType::TypeVar(_) // TypeVar may resolve to any destructurable type
| StaticType::Tuple(_)
| StaticType::Vector(_, _)
| StaticType::Matrix(_, _)
@@ -1000,9 +1064,17 @@ impl TypeChecker {
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(),
&StaticType::Any,
&param_hint_ty,
&mut lambda_ctx,
diag,
);
@@ -1104,39 +1176,39 @@ impl TypeChecker {
}
};
// 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: (series n) returns Series(Any) — replace with a fresh TypeVar so
// subsequent push calls can unify the element type.
if Self::is_identifier_named("series", &callee_typed) {
if let StaticType::Series(inner) = &ret_ty {
if **inner == StaticType::Any {
ret_ty = StaticType::Series(Box::new(self.fresh_var()));
}
}
if Self::is_identifier_named("series", &callee_typed)
&& let StaticType::Series(inner) = &ret_ty
&& **inner == StaticType::Any
{
ret_ty = StaticType::Series(Box::new(self.fresh_var()));
}
// HM: (push series val) — unify the series element TypeVar with the value
// type and propagate the resolved type back to the binding in the context.
if Self::is_identifier_named("push", &callee_typed) {
if let NodeKind::Tuple { elements } = &args_typed.kind {
if elements.len() >= 2 {
let series_arg = &elements[0];
let value_arg = &elements[1];
if let StaticType::Series(inner) = &series_arg.ty {
let inner_ty = (**inner).clone();
let val_ty = value_arg.ty.clone();
self.unify(inner_ty, val_ty, diag);
if let NodeKind::Identifier {
binding: IdentifierBinding::Reference(addr),
..
} = &series_arg.kind
{
let resolved = Self::apply_subst(
series_arg.ty.clone(),
&self.subst.borrow(),
);
ctx.set_type(*addr, resolved);
}
}
if Self::is_identifier_named("push", &callee_typed)
&& let NodeKind::Tuple { elements } = &args_typed.kind
&& elements.len() >= 2
{
let series_arg = &elements[0];
let value_arg = &elements[1];
if let StaticType::Series(inner) = &series_arg.ty {
let inner_ty = (**inner).clone();
let val_ty = value_arg.ty.clone();
self.unify(inner_ty, val_ty, diag);
if let NodeKind::Identifier {
binding: IdentifierBinding::Reference(addr),
..
} = &series_arg.kind
{
let resolved =
Self::apply_subst(series_arg.ty.clone(), &self.subst.borrow());
ctx.set_type(*addr, resolved);
}
}
}