Iter 8a: closure-pair ABI for fn-values

Every fn-value is now a `ptr` to a closure pair `{ ptr thunk, ptr env }`,
not a raw fn-pointer. Top-level fns get an auto-generated adapter
`@ail_<m>_<f>_adapter(ptr %_env, params...)` that ignores env and
forwards to the real fn, plus a static closure constant
`@ail_<m>_<f>_clos = { @adapter, null }`. References to a fn as a
value return the closure-pair address.

Indirect calls now GEP the thunk + env slots, load both, and call
`thunk(env, args...)`. Direct calls to statically-known callees stay
on the original fast path (no adapter).

This is the ABI groundwork for Iter 8b lambdas: a captured-env closure
will reuse the same value shape, just with a non-null env produced by
malloc.

Tests: 48 still green. IR snapshots refreshed (every fn now carries an
adapter + static closure pair). Iter 7's hof.ail.json prints 42
unchanged — `apply(inc, 41)` now hands `@ail_hof_inc_clos` to apply,
which unpacks and indirect-calls as designed.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-05-07 12:56:12 +02:00
parent c6c0a10788
commit 99f68e89fa
6 changed files with 166 additions and 17 deletions
+7
View File
@@ -8,12 +8,19 @@ declare i32 @printf(ptr, ...)
declare i32 @puts(ptr)
declare ptr @malloc(i64)
@ail_hello_main_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_hello_main_adapter, ptr null }
define i8 @ail_hello_main() {
entry:
call i32 @puts(ptr @.str_hello_str_0)
ret i8 0
}
define i8 @ail_hello_main_adapter(ptr %_env) {
entry:
%r = call i8 @ail_hello_main()
ret i8 %r
}
define i32 @main() {
call i8 @ail_hello_main()
+14
View File
@@ -8,6 +8,8 @@ declare i32 @printf(ptr, ...)
declare i32 @puts(ptr)
declare ptr @malloc(i64)
@ail_list_sum_list_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_list_sum_list_adapter, ptr null }
@ail_list_main_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_list_main_adapter, ptr null }
define i64 @ail_list_sum_list(ptr %arg_xs) {
entry:
%v1 = load i64, ptr %arg_xs, align 8
@@ -32,6 +34,12 @@ mjoin.2:
ret i64 %v9
}
define i64 @ail_list_sum_list_adapter(ptr %_env, ptr %a0) {
entry:
%r = call i64 @ail_list_sum_list(ptr %a0)
ret i64 %r
}
define i8 @ail_list_main() {
entry:
%v1 = call ptr @malloc(i64 8)
@@ -59,6 +67,12 @@ entry:
ret i8 0
}
define i8 @ail_list_main_adapter(ptr %_env) {
entry:
%r = call i8 @ail_list_main()
ret i8 %r
}
define i32 @main() {
call i8 @ail_list_main()
+21
View File
@@ -8,6 +8,9 @@ declare i32 @printf(ptr, ...)
declare i32 @puts(ptr)
declare ptr @malloc(i64)
@ail_max3_max_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_max3_max_adapter, ptr null }
@ail_max3_max3_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_max3_max3_adapter, ptr null }
@ail_max3_main_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_max3_main_adapter, ptr null }
define i64 @ail_max3_max(i64 %arg_a, i64 %arg_b) {
entry:
%v1 = icmp sgt i64 %arg_a, %arg_b
@@ -21,6 +24,12 @@ join.2:
ret i64 %v3
}
define i64 @ail_max3_max_adapter(ptr %_env, i64 %a0, i64 %a1) {
entry:
%r = call i64 @ail_max3_max(i64 %a0, i64 %a1)
ret i64 %r
}
define i64 @ail_max3_max3(i64 %arg_a, i64 %arg_b, i64 %arg_c) {
entry:
%v1 = icmp sgt i64 %arg_a, %arg_b
@@ -50,6 +59,12 @@ join.2:
ret i64 %v9
}
define i64 @ail_max3_max3_adapter(ptr %_env, i64 %a0, i64 %a1, i64 %a2) {
entry:
%r = call i64 @ail_max3_max3(i64 %a0, i64 %a1, i64 %a2)
ret i64 %r
}
define i8 @ail_max3_main() {
entry:
%v1 = call i64 @ail_max3_max3(i64 3, i64 17, i64 9)
@@ -57,6 +72,12 @@ entry:
ret i8 0
}
define i8 @ail_max3_main_adapter(ptr %_env) {
entry:
%r = call i8 @ail_max3_main()
ret i8 %r
}
define i32 @main() {
call i8 @ail_max3_main()
+14
View File
@@ -8,6 +8,8 @@ declare i32 @printf(ptr, ...)
declare i32 @puts(ptr)
declare ptr @malloc(i64)
@ail_sum_sum_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_sum_sum_adapter, ptr null }
@ail_sum_main_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_sum_main_adapter, ptr null }
define i64 @ail_sum_sum(i64 %arg_n) {
entry:
%v1 = icmp eq i64 %arg_n, 0
@@ -24,6 +26,12 @@ join.2:
ret i64 %v6
}
define i64 @ail_sum_sum_adapter(ptr %_env, i64 %a0) {
entry:
%r = call i64 @ail_sum_sum(i64 %a0)
ret i64 %r
}
define i8 @ail_sum_main() {
entry:
%v1 = call i64 @ail_sum_sum(i64 10)
@@ -31,6 +39,12 @@ entry:
ret i8 0
}
define i8 @ail_sum_main_adapter(ptr %_env) {
entry:
%r = call i8 @ail_sum_main()
ret i8 %r
}
define i32 @main() {
call i8 @ail_sum_main()
+14
View File
@@ -8,12 +8,20 @@ declare i32 @printf(ptr, ...)
declare i32 @puts(ptr)
declare ptr @malloc(i64)
@ail_ws_lib_add_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_ws_lib_add_adapter, ptr null }
@ail_ws_main_main_clos = private unnamed_addr constant { ptr, ptr } { ptr @ail_ws_main_main_adapter, ptr null }
define i64 @ail_ws_lib_add(i64 %arg_a, i64 %arg_b) {
entry:
%v1 = add i64 %arg_a, %arg_b
ret i64 %v1
}
define i64 @ail_ws_lib_add_adapter(ptr %_env, i64 %a0, i64 %a1) {
entry:
%r = call i64 @ail_ws_lib_add(i64 %a0, i64 %a1)
ret i64 %r
}
define i8 @ail_ws_main_main() {
entry:
%v1 = call i64 @ail_ws_lib_add(i64 2, i64 3)
@@ -21,6 +29,12 @@ entry:
ret i8 0
}
define i8 @ail_ws_main_main_adapter(ptr %_env) {
entry:
%r = call i8 @ail_ws_main_main()
ret i8 %r
}
define i32 @main() {
call i8 @ail_ws_main_main()
+96 -17
View File
@@ -428,9 +428,58 @@ impl<'a> Emitter<'a> {
}
self.body
.push_str(&format!(" ret {val_ty} {val}\n}}\n\n"));
// Iter 8a: emit closure-pair scaffold (adapter + static closure)
// for this fn. The adapter takes an extra `ptr %_env` (ignored,
// null sentinel for top-level fns) and forwards to the real fn.
// The static closure pair `{ adapter_ptr, null }` is the value
// produced when this fn is referenced as a `Term::Var` value
// (closure-pair pointer ABI).
self.emit_adapter_and_static_closure(&f.name, &llvm_param_tys, &llvm_ret);
Ok(())
}
/// Iter 8a: closure-pair scaffold for a top-level fn. Always emitted
/// (one wrapper per fn), so cross-module references just use the
/// `<m>_<f>_clos` symbol without coordination.
fn emit_adapter_and_static_closure(
&mut self,
fn_name: &str,
param_tys: &[String],
ret_ty: &str,
) {
let m = self.module_name;
// Adapter: `(ptr %_env, params...) -> ret` calls the real fn,
// returning whatever it returned.
let mut adapter = format!(
"define {ret} @ail_{m}_{fn_name}_adapter(ptr %_env",
ret = ret_ty,
);
for (i, pty) in param_tys.iter().enumerate() {
adapter.push_str(&format!(", {pty} %a{i}"));
}
adapter.push_str(") {\nentry:\n");
let mut call_args = String::new();
for (i, pty) in param_tys.iter().enumerate() {
if i > 0 {
call_args.push_str(", ");
}
call_args.push_str(&format!("{pty} %a{i}"));
}
adapter.push_str(&format!(
" %r = call {ret} @ail_{m}_{fn_name}({call_args})\n",
ret = ret_ty,
));
adapter.push_str(&format!(" ret {ret} %r\n}}\n\n", ret = ret_ty));
self.body.push_str(&adapter);
// Static closure pair: `{ adapter_ptr, null }`. The address of
// this global IS the fn-value that escapes to other code.
self.header.push_str(&format!(
"@ail_{m}_{fn_name}_clos = private unnamed_addr constant {{ ptr, ptr }} {{ ptr @ail_{m}_{fn_name}_adapter, ptr null }}\n"
));
}
/// Lowers a term to (SSA value string, LLVM type).
fn lower_term(&mut self, t: &Term) -> Result<(String, String)> {
match t {
@@ -894,8 +943,11 @@ impl<'a> Emitter<'a> {
Ok((dst, sig.ret.clone()))
}
/// Iter 7: indirect call through a fn-pointer SSA value. Sig must
/// already be known (sidetable lookup happened in the caller).
/// Iter 8a: indirect call through a closure-pair pointer. The
/// callee SSA points at `{ ptr thunk, ptr env }`; we GEP+load both
/// halves and call `thunk(env, args...)`. The user-visible `sig`
/// describes only the user-level params/ret — the env_ptr is
/// inserted by codegen, transparent to the source language.
fn emit_indirect_call(
&mut self,
callee_ssa: &str,
@@ -919,16 +971,38 @@ impl<'a> Emitter<'a> {
}
compiled.push((v, vty));
}
let arglist = compiled
.iter()
.map(|(v, t)| format!("{t} {v}"))
.collect::<Vec<_>>()
.join(", ");
let dst = self.fresh_ssa();
// LLVM indirect-call form: `call <ret>(<param-tys>) <ptr>(<args>)`.
let param_tys = sig.params.join(", ");
// Unpack the closure pair: thunk pointer at offset 0, env pointer
// at offset 8. Use a typed GEP through `{ ptr, ptr }` so the
// offsets are computed correctly across targets.
let thunk_p = self.fresh_ssa();
let thunk = self.fresh_ssa();
let env_p = self.fresh_ssa();
let env = self.fresh_ssa();
self.body.push_str(&format!(
" {dst} = call {ret} ({ptys}) {callee_ssa}({arglist})\n",
" {thunk_p} = getelementptr inbounds {{ ptr, ptr }}, ptr {callee_ssa}, i64 0, i32 0\n"
));
self.body
.push_str(&format!(" {thunk} = load ptr, ptr {thunk_p}\n"));
self.body.push_str(&format!(
" {env_p} = getelementptr inbounds {{ ptr, ptr }}, ptr {callee_ssa}, i64 0, i32 1\n"
));
self.body
.push_str(&format!(" {env} = load ptr, ptr {env_p}\n"));
// Build the actual call. The thunk's signature is `(ptr, params...)`
// — env_ptr is the implicit first arg, transparent to the user.
let mut arglist = format!("ptr {env}");
for (v, t) in &compiled {
arglist.push_str(&format!(", {t} {v}"));
}
let mut param_tys = String::from("ptr");
for pt in &sig.params {
param_tys.push_str(", ");
param_tys.push_str(pt);
}
let dst = self.fresh_ssa();
self.body.push_str(&format!(
" {dst} = call {ret} ({ptys}) {thunk}({arglist})\n",
ret = sig.ret,
ptys = param_tys,
));
@@ -951,23 +1025,28 @@ impl<'a> Emitter<'a> {
.is_some_and(|m| m.contains_key(name))
}
/// Iter 7: resolve `name` to a top-level fn-pointer (`@ail_<m>_<def>`)
/// + its sig, or return None. Mirrors the dispatch in `lower_app` for
/// fn refs only — operators / `not` are not first-class values, so
/// `is_static_callee` excludes them here.
/// Iter 8a: resolve `name` to a top-level fn-value, i.e. the address
/// of its static closure pair `@ail_<m>_<def>_clos`, plus the user-
/// visible FnSig (params/ret WITHOUT the env_ptr — that's added at
/// the call site by the closure ABI). Returns None if the name does
/// not refer to a top-level fn. Operators / `not` are not first-
/// class values; `is_static_callee` filters them earlier.
fn resolve_top_level_fn(&self, name: &str) -> Option<(String, FnSig)> {
if name.matches('.').count() == 1 {
let (prefix, suffix) = name.split_once('.')?;
let target = self.import_map.get(prefix)?;
let sig = self.module_user_fns.get(target)?.get(suffix)?.clone();
return Some((format!("@ail_{target}_{suffix}"), sig));
return Some((format!("@ail_{target}_{suffix}_clos"), sig));
}
let sig = self
.module_user_fns
.get(self.module_name)?
.get(name)?
.clone();
Some((format!("@ail_{module}_{name}", module = self.module_name), sig))
Some((
format!("@ail_{module}_{name}_clos", module = self.module_name),
sig,
))
}
fn lower_effect_op(&mut self, op: &str, args: &[Term]) -> Result<(String, String)> {