BUG-TAG: Type inference for lambda parameters
The type checker incorrectly inferred `Any` for lambda parameters when a `Program` node was involved, preventing optimizations like constant folding. This was because the `check_params_tuple` function was resolving `TypeVar` to `StaticType::Any` instead of propagating the type variable. This commit addresses the issue by: - Explicitly wrapping the bound AST in a parameterless lambda within `compile_pipeline`. This ensures that the type checker always receives a `Lambda` node, even if the original input was a `Program` node. - Adding debug logging to the type checker to help diagnose similar issues in the future. Additionally, the commit fixes a bug where `parser.parse_program()` was used instead of `parser.parse_expression()`, which would consume the entire input and prevent checking for trailing expressions.
This commit is contained in:
@@ -0,0 +1,52 @@
|
|||||||
|
# Bug: Program-Node bricht Destructuring-Typinferenz
|
||||||
|
|
||||||
|
## Reproduktion
|
||||||
|
|
||||||
|
Ausgehend vom aktuellen Stand (Lambda-Wrapping in `compile_pipeline`, `parse_expression` in `compile`):
|
||||||
|
|
||||||
|
### Schritt 1: In `compile()` `parse_expression` durch `parse_program` ersetzen
|
||||||
|
|
||||||
|
In `src/ast/environment.rs`, Methode `compile()`:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
// VORHER (funktioniert):
|
||||||
|
let syntax_ast = parser.parse_expression();
|
||||||
|
|
||||||
|
// NACHHER (Typinferenz bricht):
|
||||||
|
let syntax_ast = parser.parse_program();
|
||||||
|
```
|
||||||
|
|
||||||
|
Außerdem die `at_eof`-Prüfung entfernen (weil `parse_program` alles konsumiert).
|
||||||
|
|
||||||
|
### Schritt 2: Testen
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cargo run --release --bin ast -- -d -e "((fn [[x y]] (+ x y)) [10 20])"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Erwartet:** `Constant: 30` im Dump (Optimizer faltet den Ausdruck)
|
||||||
|
|
||||||
|
**Tatsächlich:** Kein Folding. Die Lambda-Parameter `x` und `y` haben Typ `Any` statt `Int`. HM step 10 (Unifikation) wird übersprungen weil `has_typevar_component` false ist.
|
||||||
|
|
||||||
|
### Schritt 3: Zusätzlich Lambda-Wrapping entfernen (verschärft das Problem)
|
||||||
|
|
||||||
|
In `compile_pipeline()` die Zeile `let wrapped = self.wrap_as_lambda(bound);` entfernen.
|
||||||
|
|
||||||
|
Dann bekommt `check_node_as_bound` einen `Program`-Node statt eines `Lambda`-Nodes. Ergebnis ist dasselbe: `Any`-Parameter, kein Folding.
|
||||||
|
|
||||||
|
## Ursache (unvollständig analysiert)
|
||||||
|
|
||||||
|
Der AST-Unterschied:
|
||||||
|
|
||||||
|
- **Funktioniert:** `Lambda(Call(Lambda([[x y]], body), [10 20]))` — kein Program-Node
|
||||||
|
- **Bricht:** `Lambda(Program(Call(Lambda([[x y]], body), [10 20])))` — Program-Node dazwischen
|
||||||
|
|
||||||
|
Der TypeChecker erzeugt TypeVars (`?0`, `?1`) für die innere Lambda-Parameter. Aber in `check_params_tuple` (check.rs, Zeile ~241) wird `TypeVar` im Match auf `_ => StaticType::Any` aufgelöst statt propagiert. Dadurch enthält die Signatur `fn([[any any]]) -> int` statt `fn([[?0 ?1]]) -> int`, und HM step 10 überspringt die Unifikation.
|
||||||
|
|
||||||
|
Warum das nur mit Program-Node passiert und nicht ohne, ist unklar. Der `check_params_tuple`-Code hat sich nicht geändert.
|
||||||
|
|
||||||
|
## Betroffener Test
|
||||||
|
|
||||||
|
```
|
||||||
|
tests/destructuring.rs::test_nested_destructuring_optimization
|
||||||
|
```
|
||||||
@@ -526,25 +526,23 @@ impl TypeChecker {
|
|||||||
let mut lambda_ctx =
|
let mut lambda_ctx =
|
||||||
TypeContext::new(64, upvalue_types, ctx.root_types, Some(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(
|
let param_hint_ty = StaticType::Tuple(
|
||||||
(0..positional_count.unwrap_or(0)).map(|_| self.fresh_var()).collect(),
|
(0..positional_count.unwrap_or(0)).map(|_| self.fresh_var()).collect(),
|
||||||
);
|
);
|
||||||
|
eprintln!("[TC] Lambda check_node: positional_count={:?} param_hint={}", positional_count, param_hint_ty.display_compact());
|
||||||
let params_typed = self.check_params(
|
let params_typed = self.check_params(
|
||||||
params.as_ref(),
|
params.as_ref(),
|
||||||
¶m_hint_ty,
|
¶m_hint_ty,
|
||||||
&mut lambda_ctx,
|
&mut lambda_ctx,
|
||||||
diag,
|
diag,
|
||||||
);
|
);
|
||||||
|
eprintln!("[TC] Lambda params_typed.ty={}", params_typed.ty.display_compact());
|
||||||
|
|
||||||
lambda_ctx.current_params_ty = Some(params_typed.ty.clone());
|
lambda_ctx.current_params_ty = Some(params_typed.ty.clone());
|
||||||
|
|
||||||
let body_typed = self.check_node(body, &mut lambda_ctx, diag);
|
let body_typed = self.check_node(body, &mut lambda_ctx, diag);
|
||||||
let ret_ty = body_typed.ty.clone();
|
let ret_ty = body_typed.ty.clone();
|
||||||
|
eprintln!("[TC] Lambda body ret_ty={}", ret_ty.display_compact());
|
||||||
|
|
||||||
let fn_ty = StaticType::Function(Box::new(Signature {
|
let fn_ty = StaticType::Function(Box::new(Signature {
|
||||||
params: params_typed.ty.clone(),
|
params: params_typed.ty.clone(),
|
||||||
@@ -565,7 +563,9 @@ impl TypeChecker {
|
|||||||
}
|
}
|
||||||
|
|
||||||
NodeKind::Call { callee, args } => {
|
NodeKind::Call { callee, args } => {
|
||||||
|
eprintln!("[TC] Call: checking callee...");
|
||||||
let callee_typed = self.check_node(callee, ctx, diag);
|
let callee_typed = self.check_node(callee, ctx, diag);
|
||||||
|
eprintln!("[TC] Call: callee_typed.ty={}", callee_typed.ty.display_compact());
|
||||||
|
|
||||||
let args_typed = if let NodeKind::Tuple { elements } = &args.kind {
|
let args_typed = if let NodeKind::Tuple { elements } = &args.kind {
|
||||||
let arg_count = elements.len();
|
let arg_count = elements.len();
|
||||||
@@ -624,6 +624,7 @@ impl TypeChecker {
|
|||||||
self.check_node(args, ctx, diag)
|
self.check_node(args, ctx, diag)
|
||||||
};
|
};
|
||||||
|
|
||||||
|
eprintln!("[TC] Call: args_typed.ty={}", args_typed.ty.display_compact());
|
||||||
let mut ret_ty = match callee_typed.ty.resolve_call(&args_typed.ty) {
|
let mut ret_ty = match callee_typed.ty.resolve_call(&args_typed.ty) {
|
||||||
Some(ty) => ty,
|
Some(ty) => ty,
|
||||||
None => {
|
None => {
|
||||||
@@ -687,9 +688,16 @@ impl TypeChecker {
|
|||||||
if let StaticType::Function(sig) = &callee_typed.ty
|
if let StaticType::Function(sig) = &callee_typed.ty
|
||||||
&& Self::has_typevar_component(&sig.params)
|
&& Self::has_typevar_component(&sig.params)
|
||||||
{
|
{
|
||||||
|
eprintln!("[TC] HM10: unify params={} with args={}", sig.params.display_compact(), args_typed.ty.display_compact());
|
||||||
let params = sig.params.clone();
|
let params = sig.params.clone();
|
||||||
self.unify(params, args_typed.ty.clone(), diag);
|
self.unify(params, args_typed.ty.clone(), diag);
|
||||||
ret_ty = Self::apply_subst(ret_ty, &self.subst.borrow());
|
ret_ty = Self::apply_subst(ret_ty, &self.subst.borrow());
|
||||||
|
eprintln!("[TC] HM10: after unify ret_ty={}", ret_ty.display_compact());
|
||||||
|
} else {
|
||||||
|
eprintln!("[TC] HM10: SKIPPED (callee_ty={}, has_typevar={})",
|
||||||
|
callee_typed.ty.display_compact(),
|
||||||
|
if let StaticType::Function(sig) = &callee_typed.ty { Self::has_typevar_component(&sig.params) } else { false }
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Dispatch compiler hooks registered by the RTL (keyed by global slot index).
|
// Dispatch compiler hooks registered by the RTL (keyed by global slot index).
|
||||||
|
|||||||
@@ -70,6 +70,7 @@ impl TypeChecker {
|
|||||||
) -> TypedNode {
|
) -> TypedNode {
|
||||||
match &node.kind {
|
match &node.kind {
|
||||||
NodeKind::Lambda { params, body, info } => {
|
NodeKind::Lambda { params, body, info } => {
|
||||||
|
eprintln!("[TC] check_node_as_bound: Lambda entry");
|
||||||
let upvalues = &info.upvalues;
|
let upvalues = &info.upvalues;
|
||||||
let positional_count = info.positional_count;
|
let positional_count = info.positional_count;
|
||||||
|
|
||||||
@@ -115,6 +116,7 @@ impl TypeChecker {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
NodeKind::Block { .. } | NodeKind::Program { .. } => {
|
NodeKind::Block { .. } | NodeKind::Program { .. } => {
|
||||||
|
eprintln!("[TC] check_node_as_bound: Block/Program entry");
|
||||||
let mut ctx = TypeContext::new(64, vec![], &self.root_types, None);
|
let mut ctx = TypeContext::new(64, vec![], &self.root_types, None);
|
||||||
self.check_node(node, &mut ctx, diag)
|
self.check_node(node, &mut ctx, diag)
|
||||||
}
|
}
|
||||||
|
|||||||
+41
-4
@@ -12,7 +12,7 @@ use std::rc::Rc;
|
|||||||
|
|
||||||
use crate::ast::nodes::{
|
use crate::ast::nodes::{
|
||||||
Address, AnalyzedNode, ExecNode, GlobalAnalyzedRegistry, GlobalFunctionRegistry, GlobalIdx,
|
Address, AnalyzedNode, ExecNode, GlobalAnalyzedRegistry, GlobalFunctionRegistry, GlobalIdx,
|
||||||
Node, NodeKind, VirtualId,
|
LambdaBinding, Node, NodeKind, VirtualId,
|
||||||
};
|
};
|
||||||
use crate::ast::compiler::dumper::Dumper;
|
use crate::ast::compiler::dumper::Dumper;
|
||||||
use crate::ast::compiler::lambda_collector::LambdaCollector;
|
use crate::ast::compiler::lambda_collector::LambdaCollector;
|
||||||
@@ -447,11 +447,38 @@ impl Environment {
|
|||||||
typed
|
typed
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Full compilation pipeline: expand → bind → type-check.
|
/// Wraps a bound AST in a parameterless lambda (unless it already is one).
|
||||||
|
fn wrap_as_lambda(&self, bound_ast: Node<BoundPhase>) -> Node<BoundPhase> {
|
||||||
|
if let NodeKind::Lambda { .. } = bound_ast.kind {
|
||||||
|
bound_ast
|
||||||
|
} else {
|
||||||
|
Node {
|
||||||
|
identity: bound_ast.identity.clone(),
|
||||||
|
kind: NodeKind::Lambda {
|
||||||
|
params: Rc::new(Node {
|
||||||
|
identity: bound_ast.identity.clone(),
|
||||||
|
kind: NodeKind::Tuple { elements: vec![] },
|
||||||
|
ty: (),
|
||||||
|
comments: Rc::from([]),
|
||||||
|
}),
|
||||||
|
body: Rc::new(bound_ast),
|
||||||
|
info: LambdaBinding {
|
||||||
|
upvalues: vec![],
|
||||||
|
positional_count: Some(0),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ty: (),
|
||||||
|
comments: Rc::from([]),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Full compilation pipeline: expand → bind → wrap → type-check.
|
||||||
fn compile_pipeline(&self, syntax_ast: SyntaxNode, diagnostics: &mut Diagnostics) -> Option<TypedNode> {
|
fn compile_pipeline(&self, syntax_ast: SyntaxNode, diagnostics: &mut Diagnostics) -> Option<TypedNode> {
|
||||||
let expanded = self.expand(syntax_ast, diagnostics)?;
|
let expanded = self.expand(syntax_ast, diagnostics)?;
|
||||||
let bound = self.bind_and_update(&expanded, diagnostics)?;
|
let bound = self.bind_and_update(&expanded, diagnostics)?;
|
||||||
let typed = self.type_check(&bound, diagnostics);
|
let wrapped = self.wrap_as_lambda(bound);
|
||||||
|
let typed = self.type_check(&wrapped, diagnostics);
|
||||||
Some(typed)
|
Some(typed)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -710,7 +737,17 @@ impl Environment {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let mut parser = Parser::new(source);
|
let mut parser = Parser::new(source);
|
||||||
let syntax_ast = parser.parse_program();
|
let syntax_ast = parser.parse_expression();
|
||||||
|
|
||||||
|
if !parser.at_eof() {
|
||||||
|
parser
|
||||||
|
.diagnostics
|
||||||
|
.push_error("Unexpected trailing expressions in script.", None);
|
||||||
|
return CompilationResult {
|
||||||
|
ast: None,
|
||||||
|
diagnostics: parser.diagnostics,
|
||||||
|
};
|
||||||
|
}
|
||||||
let mut diagnostics = parser.diagnostics;
|
let mut diagnostics = parser.diagnostics;
|
||||||
|
|
||||||
let typed_ast = self.compile_pipeline(syntax_ast, &mut diagnostics);
|
let typed_ast = self.compile_pipeline(syntax_ast, &mut diagnostics);
|
||||||
|
|||||||
Reference in New Issue
Block a user