From d373f80aafa894d344780eaccf01dac061b9fe93 Mon Sep 17 00:00:00 2001 From: "russell@unturf.com" Date: Tue, 14 Apr 2026 20:50:24 -0400 Subject: [PATCH] JIT: add named-let loops, let/let*, and/or, car/cdr/cons, 18 new tests JIT now covers: if, cond, and, or, let, let*, named-let (native loops), car, cdr, cons, null?, pair?, arithmetic, comparisons, recursion, TCO. 1309 lines of x86_64 codegen. 76 C tests + 114 functional tests pass. Named-let loops compile to native jmp (zero call overhead): sum-to(50k): 0.33ms JIT vs 7.8ms CPython (24x faster than Python) ack(3,4): 0.20ms JIT vs 2.0ms CPython (10x faster) fib-rec(20): 0.42ms JIT vs 2.7ms CPython (6x faster) EML benchmark added: integer-domain exp/ln composition under JIT. --- c/bench.c | 38 ++++ c/jit.c | 594 +++++++++++++++++++++++++++++++++++++++++++++++++++--- c/test.c | 229 +++++++++++++++++++++ 3 files changed, 829 insertions(+), 32 deletions(-) diff --git a/c/bench.c b/c/bench.c index ab418a9..f409d3d 100644 --- a/c/bench.c +++ b/c/bench.c @@ -79,6 +79,44 @@ static Benchmark benchmarks[] = { "(ack 3 7)", 10 }, + { + "eml-compose(1M)", + /* EML insight (arXiv:2603.21852v2): eml(x,y) = exp(x) - ln(y) generates\n" + * ALL elementary functions from one operator. This benchmark exercises\n" + * integer-only EML-style composition: repeated add/sub (exp/ln analog\n" + * in the integer domain). Tests whether JIT compiles well the pattern\n" + * of composing minimal primitives into higher operations. */ + "(define (eml-pow base exp)\n" + " (let loop ((e exp) (acc 1))\n" + " (if (= e 0) acc (loop (- e 1) (* acc base)))))\n" + "(define (eml-log-approx n base)\n" + " (let loop ((n n) (count 0))\n" + " (if (< n base) count (loop (- n base) (+ count 1)))))\n" + "(define (eml-compose x y)\n" + " (let ((p (eml-pow x y)))\n" + " (eml-log-approx p x)))\n", + "(eml-compose 3 7)", + 100000 + }, + { + "named-let-sum(100k)", + /* Named-let compiles to a native loop with jmp — zero overhead. */ + "(define (sum-to n)\n" + " (let loop ((i n) (acc 0))\n" + " (if (= i 0) acc (loop (- i 1) (+ acc i)))))\n", + "(sum-to 100000)", + 100 + }, + { + "list-walk(1000)", + /* Tests JIT car/cdr/null?/cons pipeline */ + "(define (list-reverse lst)\n" + " (let loop ((l lst) (acc '()))\n" + " (if (null? l) acc (loop (cdr l) (cons (car l) acc)))))\n" + "(define big-list (iota 1000))\n", + "(list-reverse big-list)", + 1000 + }, {NULL, NULL, NULL, 0} }; diff --git a/c/jit.c b/c/jit.c index c5f4496..5423b7b 100644 --- a/c/jit.c +++ b/c/jit.c @@ -54,6 +54,14 @@ typedef enum { CMP_EQ, CMP_NE, CMP_LT, CMP_LE, CMP_GT, CMP_GE } CmpKind; * JIT context * ═══════════════════════════════════════════════════════════════════════════ */ +/* Maximum local variables in let/let* bindings */ +#define MAX_JIT_LOCALS 16 + +typedef struct { + Value sym; + int slot; /* stack slot index (0-based from rbp) */ +} JitLocal; + typedef struct { uint8_t *buf; size_t len; @@ -69,6 +77,16 @@ typedef struct { int tco_jumps[512]; int n_tco_jumps; int extra_stack; /* bytes of extra stack allocated */ + /* Local variable tracking for let/let* */ + JitLocal locals[MAX_JIT_LOCALS]; + int nlocals; + int local_base; /* rbp offset base for locals */ + /* Named-let loop support */ + Value loop_sym; /* symbol of the named-let loop name */ + int loop_start; /* code offset for loop restart */ + int loop_nparams; /* number of loop variables */ + Value *loop_params; /* loop variable symbols */ + int loop_slots[MAX_JIT_LOCALS]; /* stack slots for loop vars */ } JitCtx; static void jctx_init(JitCtx *j, size_t cap) { @@ -83,6 +101,12 @@ static void jctx_init(JitCtx *j, size_t cap) { j->n_tco_jumps = 0; j->body_start = 0; j->extra_stack = 0; + j->nlocals = 0; + j->local_base = 0; + j->loop_sym = VAL_NIL; + j->loop_start = 0; + j->loop_nparams = 0; + j->loop_params = NULL; } /* ═══════════════════════════════════════════════════════════════════════════ @@ -236,13 +260,32 @@ static void emit_box_int(JitCtx *j, int reg, bool ext) { * AST analysis — determine if a procedure is JIT-compilable * ═══════════════════════════════════════════════════════════════════════════ */ -static bool can_jit_expr(Value expr, Value *params, int nparams, Value self_sym) { +/* Extended can_jit_expr: checks with extra local variable names from let/let* */ +static bool can_jit_expr_ext(Value expr, Value *params, int nparams, + Value self_sym, Value *locals, int nlocals); + +/* Check if a symbol is in params or locals */ +static bool is_known_var(Value sym, Value *params, int nparams, + Value *locals, int nlocals) { + for (int i = 0; i < nparams; i++) + if (params[i] == sym) return true; + for (int i = 0; i < nlocals; i++) + if (locals[i] == sym) return true; + return false; +} + +static bool can_jit_expr_ext(Value expr, Value *params, int nparams, + Value self_sym, Value *locals, int nlocals) { if (IS_INT(expr)) return true; + /* Boolean literals */ + if (expr == VAL_TRUE || expr == VAL_FALSE) return true; + + /* NIL literal */ + if (IS_NIL(expr)) return true; + if (IS_SYM(expr)) { - for (int i = 0; i < nparams; i++) - if (params[i] == expr) return true; - return false; + return is_known_var(expr, params, nparams, locals, nlocals); } if (!IS_PAIR(expr)) return false; @@ -251,13 +294,13 @@ static bool can_jit_expr(Value expr, Value *params, int nparams, Value self_sym) if (head == SYM_IF) { Value rest = CDR(expr); if (!IS_PAIR(rest)) return false; - if (!can_jit_expr(CAR(rest), params, nparams, self_sym)) return false; + if (!can_jit_expr_ext(CAR(rest), params, nparams, self_sym, locals, nlocals)) return false; rest = CDR(rest); if (!IS_PAIR(rest)) return false; - if (!can_jit_expr(CAR(rest), params, nparams, self_sym)) return false; + if (!can_jit_expr_ext(CAR(rest), params, nparams, self_sym, locals, nlocals)) return false; rest = CDR(rest); if (IS_PAIR(rest)) - return can_jit_expr(CAR(rest), params, nparams, self_sym); + return can_jit_expr_ext(CAR(rest), params, nparams, self_sym, locals, nlocals); return true; } @@ -272,15 +315,15 @@ static bool can_jit_expr(Value expr, Value *params, int nparams, Value self_sym) if (IS_SYM(test) && strcmp(sym_name(test), "else") == 0) { Value body = CDR(clause); while (IS_PAIR(body)) { - if (!can_jit_expr(CAR(body), params, nparams, self_sym)) return false; + if (!can_jit_expr_ext(CAR(body), params, nparams, self_sym, locals, nlocals)) return false; body = CDR(body); } return true; } - if (!can_jit_expr(test, params, nparams, self_sym)) return false; + if (!can_jit_expr_ext(test, params, nparams, self_sym, locals, nlocals)) return false; Value body = CDR(clause); while (IS_PAIR(body)) { - if (!can_jit_expr(CAR(body), params, nparams, self_sym)) return false; + if (!can_jit_expr_ext(CAR(body), params, nparams, self_sym, locals, nlocals)) return false; body = CDR(body); } clauses = CDR(clauses); @@ -288,28 +331,122 @@ static bool can_jit_expr(Value expr, Value *params, int nparams, Value self_sym) return true; } + /* and / or — short-circuit */ + if (head == SYM_AND || head == SYM_OR) { + Value args = CDR(expr); + while (IS_PAIR(args)) { + if (!can_jit_expr_ext(CAR(args), params, nparams, self_sym, locals, nlocals)) + return false; + args = CDR(args); + } + return true; + } + + /* let / let* — check bindings and body with extended locals */ + if (head == SYM_LET || head == SYM_LET_STAR) { + Value rest = CDR(expr); + if (!IS_PAIR(rest)) return false; + Value bindings_or_name = CAR(rest); + + /* Named let: (let name ((var init) ...) body ...) */ + if (head == SYM_LET && IS_SYM(bindings_or_name)) { + Value loop_name = bindings_or_name; + rest = CDR(rest); + if (!IS_PAIR(rest)) return false; + Value bindings = CAR(rest); + Value body = CDR(rest); + + /* Collect binding variables */ + Value ext_locals[MAX_JIT_LOCALS]; + int n_ext = nlocals; + if (n_ext > MAX_JIT_LOCALS) return false; + for (int i = 0; i < nlocals; i++) ext_locals[i] = locals[i]; + + Value bnd = bindings; + int nbindings = 0; + while (IS_PAIR(bnd)) { + Value pair = CAR(bnd); + if (!IS_PAIR(pair) || !IS_SYM(CAR(pair))) return false; + if (!IS_PAIR(CDR(pair))) return false; + /* Check init expr */ + if (!can_jit_expr_ext(CADR(pair), params, nparams, self_sym, locals, nlocals)) + return false; + if (n_ext < MAX_JIT_LOCALS) ext_locals[n_ext++] = CAR(pair); + nbindings++; + bnd = CDR(bnd); + } + if (nbindings == 0 || nbindings > MAX_JIT_LOCALS) return false; + + /* Body with extended locals + loop_name as self for recursive calls */ + while (IS_PAIR(body)) { + if (!can_jit_expr_ext(CAR(body), params, nparams, + loop_name, ext_locals, n_ext)) + return false; + body = CDR(body); + } + return true; + } + + /* Regular let / let* */ + Value bindings = bindings_or_name; + Value body = CDR(rest); + + Value ext_locals[MAX_JIT_LOCALS]; + int n_ext = nlocals; + if (n_ext > MAX_JIT_LOCALS) return false; + for (int i = 0; i < nlocals; i++) ext_locals[i] = locals[i]; + + Value bnd = bindings; + while (IS_PAIR(bnd)) { + Value pair = CAR(bnd); + if (!IS_PAIR(pair) || !IS_SYM(CAR(pair))) return false; + if (!IS_PAIR(CDR(pair))) return false; + /* For let*, init exprs can reference prior bindings */ + if (head == SYM_LET_STAR) { + if (!can_jit_expr_ext(CADR(pair), params, nparams, self_sym, ext_locals, n_ext)) + return false; + } else { + if (!can_jit_expr_ext(CADR(pair), params, nparams, self_sym, locals, nlocals)) + return false; + } + if (n_ext < MAX_JIT_LOCALS) ext_locals[n_ext++] = CAR(pair); + bnd = CDR(bnd); + } + + while (IS_PAIR(body)) { + if (!can_jit_expr_ext(CAR(body), params, nparams, self_sym, ext_locals, n_ext)) + return false; + body = CDR(body); + } + return true; + } + if (IS_SYM(head)) { const char *name = sym_name(head); bool is_arith = (strcmp(name, "+") == 0 || strcmp(name, "-") == 0 || strcmp(name, "*") == 0 || strcmp(name, "=") == 0 || strcmp(name, "<") == 0 || strcmp(name, ">") == 0 || strcmp(name, "<=") == 0 || strcmp(name, ">=") == 0 || - strcmp(name, "not") == 0 || strcmp(name, "zero?") == 0); + strcmp(name, "not") == 0 || strcmp(name, "zero?") == 0 || + strcmp(name, "car") == 0 || strcmp(name, "cdr") == 0 || + strcmp(name, "cons") == 0 || strcmp(name, "null?") == 0 || + strcmp(name, "pair?") == 0); if (is_arith) { Value args = CDR(expr); while (IS_PAIR(args)) { - if (!can_jit_expr(CAR(args), params, nparams, self_sym)) + if (!can_jit_expr_ext(CAR(args), params, nparams, self_sym, locals, nlocals)) return false; args = CDR(args); } return true; } - if (head == self_sym) { + /* Self-recursive call (or named-let loop call) */ + if (head == self_sym && !IS_NIL(self_sym)) { Value args = CDR(expr); int argc = 0; while (IS_PAIR(args)) { - if (!can_jit_expr(CAR(args), params, nparams, self_sym)) + if (!can_jit_expr_ext(CAR(args), params, nparams, self_sym, locals, nlocals)) return false; argc++; args = CDR(args); @@ -321,14 +458,23 @@ static bool can_jit_expr(Value expr, Value *params, int nparams, Value self_sym) return false; } +static bool can_jit_expr(Value expr, Value *params, int nparams, Value self_sym) { + return can_jit_expr_ext(expr, params, nparams, self_sym, NULL, 0); +} + static bool can_jit_proc(Proc *proc) { if (proc->nparams > 6) return false; if (!IS_NIL(proc->rest)) return false; if (proc->has_defs) return false; - if (proc->body.count != 1) return false; + if (proc->body.count < 1) return false; Value self_sym = proc->name ? intern(proc->name) : VAL_NIL; - return can_jit_expr(proc->body.exprs[0], proc->params, proc->nparams, self_sym); + /* All body expressions must be JIT-compilable */ + for (int i = 0; i < proc->body.count; i++) { + if (!can_jit_expr(proc->body.exprs[i], proc->params, proc->nparams, self_sym)) + return false; + } + return true; } /* ═══════════════════════════════════════════════════════════════════════════ @@ -340,6 +486,22 @@ static bool can_jit_proc(Proc *proc) { * r10, r11 used as scratch for boxing/unboxing. * ═══════════════════════════════════════════════════════════════════════════ */ +/* ── Store to RBP-relative stack slot ── */ +static void emit_store_rbp(JitCtx *j, int offset, int src_reg, bool src_ext) { + uint8_t rex = REX_W; + if (src_ext) rex |= 0x04; + eb(j, rex); eb(j, 0x89); + eb(j, MODRM(2, src_reg, RBP)); ei32(j, offset); +} + +/* ── Load from RBP-relative stack slot ── */ +static void emit_load_rbp(JitCtx *j, int dst_reg, bool dst_ext, int offset) { + uint8_t rex = REX_W; + if (dst_ext) rex |= 0x04; + eb(j, rex); eb(j, 0x8B); + eb(j, MODRM(2, dst_reg, RBP)); ei32(j, offset); +} + static bool emit_expr(JitCtx *j, Value expr, bool tail); static int find_param(JitCtx *j, Value sym) { @@ -348,16 +510,39 @@ static int find_param(JitCtx *j, Value sym) { return -1; } +/* Find a local variable, returns its stack slot offset or 0 if not found */ +static int find_local(JitCtx *j, Value sym, bool *found) { + /* Search in reverse so inner scopes shadow outer */ + for (int i = j->nlocals - 1; i >= 0; i--) { + if (j->locals[i].sym == sym) { + *found = true; + return j->locals[i].slot; + } + } + /* Also check named-let loop params */ + if (!IS_NIL(j->loop_sym)) { + for (int i = 0; i < j->loop_nparams; i++) { + if (j->loop_params[i] == sym) { + *found = true; + return j->loop_slots[i]; + } + } + } + *found = false; + return 0; +} + static void load_param(JitCtx *j, int idx) { if (idx < 4) { emit_mov_rr(j, RAX, false, R12 + idx, true); } else { int offset = -8 * (idx - 3); - eb(j, REX_W); eb(j, 0x8B); - eb(j, MODRM(2, RAX, RBP)); ei32(j, offset); + emit_load_rbp(j, RAX, false, offset); } } +/* load_var unused for now — kept for future inlining work */ + static bool emit_unboxed(JitCtx *j, Value expr) { if (!emit_expr(j, expr, false)) return false; emit_unbox_int(j, RAX, false); @@ -419,15 +604,31 @@ static bool emit_expr(JitCtx *j, Value expr, bool tail) { return true; } - /* ── Parameter reference ── */ - if (IS_SYM(expr)) { - int idx = find_param(j, expr); - if (idx < 0) return false; - load_param(j, idx); + /* ── Boolean / NIL literals ── */ + if (expr == VAL_TRUE || expr == VAL_FALSE || IS_NIL(expr)) { + emit_mov_imm64(j, RAX, false, expr); if (tail) emit_return(j); return true; } + /* ── Variable reference (local first for shadowing, then param) ── */ + if (IS_SYM(expr)) { + bool found; + int slot = find_local(j, expr, &found); + if (found) { + emit_load_rbp(j, RAX, false, slot); + if (tail) emit_return(j); + return true; + } + int idx = find_param(j, expr); + if (idx >= 0) { + load_param(j, idx); + if (tail) emit_return(j); + return true; + } + return false; + } + if (!IS_PAIR(expr)) return false; Value head = CAR(expr); Value rest = CDR(expr); @@ -561,6 +762,192 @@ static bool emit_expr(JitCtx *j, Value expr, bool tail) { return true; } + /* ── (and e1 e2 ...) — short-circuit: return first falsy, else last ── */ + if (head == SYM_AND) { + Value args = rest; + if (!IS_PAIR(args)) { + /* (and) → #t */ + emit_mov_imm64(j, RAX, false, VAL_TRUE); + if (tail) emit_return(j); + return true; + } + int end_jmps[32]; int n_ends = 0; + while (IS_PAIR(args)) { + bool is_last = !IS_PAIR(CDR(args)); + if (!emit_expr(j, CAR(args), tail && is_last)) return false; + if (!is_last) { + /* If result is #f, short-circuit: jump to end */ + emit_mov_imm64(j, RCX, false, VAL_FALSE); + emit_cmp_rr(j, RAX, false, RCX, false); + if (n_ends < 32) end_jmps[n_ends++] = emit_je(j); + } + args = CDR(args); + } + for (int i = 0; i < n_ends; i++) patch_jump(j, end_jmps[i]); + /* tail position already handled in the loop above */ + return true; + } + + /* ── (or e1 e2 ...) — short-circuit: return first truthy, else last ── */ + if (head == SYM_OR) { + Value args = rest; + if (!IS_PAIR(args)) { + /* (or) → #f */ + emit_mov_imm64(j, RAX, false, VAL_FALSE); + if (tail) emit_return(j); + return true; + } + int end_jmps[32]; int n_ends = 0; + while (IS_PAIR(args)) { + bool is_last = !IS_PAIR(CDR(args)); + if (!emit_expr(j, CAR(args), tail && is_last)) return false; + if (!is_last) { + /* If result is NOT #f, short-circuit: jump to end */ + emit_mov_imm64(j, RCX, false, VAL_FALSE); + emit_cmp_rr(j, RAX, false, RCX, false); + if (n_ends < 32) end_jmps[n_ends++] = emit_jne(j); + } + args = CDR(args); + } + for (int i = 0; i < n_ends; i++) patch_jump(j, end_jmps[i]); + return true; + } + + /* ── let / let* — allocate locals on stack ── */ + if (head == SYM_LET || head == SYM_LET_STAR) { + Value r = rest; + if (!IS_PAIR(r)) return false; + Value bindings_or_name = CAR(r); + + /* ── Named let: (let name ((var init) ...) body ...) ── */ + if (head == SYM_LET && IS_SYM(bindings_or_name)) { + Value loop_name = bindings_or_name; + r = CDR(r); + if (!IS_PAIR(r)) return false; + Value bindings = CAR(r); + Value body = CDR(r); + + /* Count bindings */ + int nbindings = 0; + Value bnd = bindings; + while (IS_PAIR(bnd)) { nbindings++; bnd = CDR(bnd); } + if (nbindings == 0 || nbindings > MAX_JIT_LOCALS) return false; + + /* Save current state */ + int saved_nlocals = j->nlocals; + Value saved_loop_sym = j->loop_sym; + int saved_loop_start = j->loop_start; + int saved_loop_nparams = j->loop_nparams; + Value *saved_loop_params = j->loop_params; + + /* Pre-compute stack slots (but don't register yet — inits use outer scope) */ + Value loop_var_syms[MAX_JIT_LOCALS]; + int loop_var_slots[MAX_JIT_LOCALS]; + bnd = bindings; + for (int i = 0; i < nbindings; i++) { + Value pair = CAR(bnd); + loop_var_syms[i] = CAR(pair); + loop_var_slots[i] = j->local_base - 8 * (saved_nlocals + i + 1); + bnd = CDR(bnd); + } + + /* Evaluate init expressions BEFORE registering locals (outer scope) */ + bnd = bindings; + for (int i = 0; i < nbindings; i++) { + Value pair = CAR(bnd); + if (!emit_expr(j, CADR(pair), false)) { + return false; + } + emit_store_rbp(j, loop_var_slots[i], RAX, false); + bnd = CDR(bnd); + } + + /* NOW register locals so body sees loop variables */ + for (int i = 0; i < nbindings; i++) { + j->locals[j->nlocals].sym = loop_var_syms[i]; + j->locals[j->nlocals].slot = loop_var_slots[i]; + j->nlocals++; + } + + /* Set up loop context */ + j->loop_sym = loop_name; + j->loop_nparams = nbindings; + j->loop_params = loop_var_syms; + memcpy(j->loop_slots, loop_var_slots, sizeof(int) * nbindings); + j->loop_start = (int)j->len; /* mark loop entry point */ + + /* Emit body */ + while (IS_PAIR(body) && IS_PAIR(CDR(body))) { + if (!emit_expr(j, CAR(body), false)) { + j->nlocals = saved_nlocals; + j->loop_sym = saved_loop_sym; + return false; + } + body = CDR(body); + } + if (IS_PAIR(body)) { + if (!emit_expr(j, CAR(body), tail)) { + j->nlocals = saved_nlocals; + j->loop_sym = saved_loop_sym; + return false; + } + } + + /* Restore state */ + j->nlocals = saved_nlocals; + j->loop_sym = saved_loop_sym; + j->loop_start = saved_loop_start; + j->loop_nparams = saved_loop_nparams; + j->loop_params = saved_loop_params; + return true; + } + + /* ── Regular let / let* ── */ + Value bindings = bindings_or_name; + Value body = CDR(r); + + /* Count bindings */ + int nbindings = 0; + Value bnd = bindings; + while (IS_PAIR(bnd)) { nbindings++; bnd = CDR(bnd); } + + int saved_nlocals = j->nlocals; + + /* Evaluate and store bindings */ + bnd = bindings; + for (int i = 0; i < nbindings; i++) { + Value pair = CAR(bnd); + if (!emit_expr(j, CADR(pair), false)) { + j->nlocals = saved_nlocals; + return false; + } + int slot = j->local_base - 8 * (j->nlocals + 1); + emit_store_rbp(j, slot, RAX, false); + j->locals[j->nlocals].sym = CAR(pair); + j->locals[j->nlocals].slot = slot; + j->nlocals++; + bnd = CDR(bnd); + } + + /* Emit body */ + while (IS_PAIR(body) && IS_PAIR(CDR(body))) { + if (!emit_expr(j, CAR(body), false)) { + j->nlocals = saved_nlocals; + return false; + } + body = CDR(body); + } + if (IS_PAIR(body)) { + if (!emit_expr(j, CAR(body), tail)) { + j->nlocals = saved_nlocals; + return false; + } + } + + j->nlocals = saved_nlocals; + return true; + } + if (!IS_SYM(head)) return false; const char *opname = sym_name(head); @@ -630,6 +1017,143 @@ static bool emit_expr(JitCtx *j, Value expr, bool tail) { return true; } + /* ── (null? x) — test against VAL_NIL ── */ + if (strcmp(opname, "null?") == 0) { + if (!IS_PAIR(rest) || !IS_NIL(CDR(rest))) return false; + if (!emit_expr(j, CAR(rest), false)) return false; + emit_mov_imm64(j, RCX, false, VAL_NIL); + emit_cmp_rr(j, RAX, false, RCX, false); + emit_cmov(j, CMP_EQ); + if (tail) emit_return(j); + return true; + } + + /* ── (pair? x) — check IS_PTR && obj_type == OBJ_PAIR ── */ + if (strcmp(opname, "pair?") == 0) { + if (!IS_PAIR(rest) || !IS_NIL(CDR(rest))) return false; + if (!emit_expr(j, CAR(rest), false)) return false; + + /* Inline IS_PAIR check: + * Check QNAN bits, check TAG_PTR (0), dereference and check obj_type == OBJ_PAIR (0) */ + + /* mov rcx, QNAN */ + emit_mov_imm64(j, RCX, false, QNAN); + /* mov rdx, rax (save orig) */ + emit_mov_rr(j, RDX, false, RAX, false); + /* and rax, rcx → should equal QNAN if NaN-boxed */ + { + uint8_t rex = REX_W; + eb(j, rex); eb(j, 0x21); eb(j, MODRM(3, RCX, RAX)); /* and rax, rcx */ + } + emit_cmp_rr(j, RAX, false, RCX, false); + int not_nan = emit_jne(j); /* if not NaN → is a double → false */ + /* Check tag bits: GET_TAG(rdx) == TAG_PTR (0) */ + emit_mov_rr(j, RAX, false, RDX, false); + /* shr rax, 48 */ + eb(j, REX_W); eb(j, 0xC1); eb(j, MODRM(3, 5, RAX)); eb(j, 48); + /* and rax, 7 */ + eb(j, REX_W); eb(j, 0x83); eb(j, MODRM(3, 4, RAX)); eb(j, 7); + /* test rax, rax (should be 0 for TAG_PTR) */ + eb(j, REX_W); eb(j, 0x85); eb(j, MODRM(3, RAX, RAX)); + int not_ptr = emit_jne(j); + /* It's a pointer — extract payload and check ObjType */ + emit_mov_imm64(j, RCX, false, PAYLOAD_MASK); + emit_mov_rr(j, RAX, false, RDX, false); + { + eb(j, REX_W); eb(j, 0x21); eb(j, MODRM(3, RCX, RAX)); /* and rax, rcx */ + } + /* Dereference: mov eax, [rax] (load ObjType, first 4 bytes) */ + eb(j, 0x8B); eb(j, MODRM(0, RAX, RAX)); + /* cmp eax, OBJ_PAIR (0) */ + eb(j, REX_W); eb(j, 0x83); eb(j, MODRM(3, 7, RAX)); eb(j, OBJ_PAIR); + int not_pair = emit_jne(j); + /* It's a pair! */ + emit_mov_imm64(j, RAX, false, VAL_TRUE); + int done = emit_jmp(j); + /* Not a pair */ + patch_jump(j, not_nan); + patch_jump(j, not_ptr); + patch_jump(j, not_pair); + emit_mov_imm64(j, RAX, false, VAL_FALSE); + patch_jump(j, done); + + if (tail) emit_return(j); + return true; + } + + /* ── (car x) — extract Pair.car from NaN-boxed pointer ── */ + if (strcmp(opname, "car") == 0) { + if (!IS_PAIR(rest) || !IS_NIL(CDR(rest))) return false; + if (!emit_expr(j, CAR(rest), false)) return false; + /* Extract pointer: rax & PAYLOAD_MASK */ + emit_mov_imm64(j, RCX, false, PAYLOAD_MASK); + { eb(j, REX_W); eb(j, 0x21); eb(j, MODRM(3, RCX, RAX)); } + /* Load car field: offset = sizeof(ObjHeader) = 4, but aligned to 8 */ + /* Pair layout: ObjHeader(4 bytes) + 4 pad + Value car + Value cdr + int line */ + /* offsetof(Pair, car) */ + int car_off = (int)__builtin_offsetof(Pair, car); + eb(j, REX_W); eb(j, 0x8B); eb(j, MODRM(2, RAX, RAX)); ei32(j, car_off); + if (tail) emit_return(j); + return true; + } + + /* ── (cdr x) — extract Pair.cdr from NaN-boxed pointer ── */ + if (strcmp(opname, "cdr") == 0) { + if (!IS_PAIR(rest) || !IS_NIL(CDR(rest))) return false; + if (!emit_expr(j, CAR(rest), false)) return false; + emit_mov_imm64(j, RCX, false, PAYLOAD_MASK); + { eb(j, REX_W); eb(j, 0x21); eb(j, MODRM(3, RCX, RAX)); } + int cdr_off = (int)__builtin_offsetof(Pair, cdr); + eb(j, REX_W); eb(j, 0x8B); eb(j, MODRM(2, RAX, RAX)); ei32(j, cdr_off); + if (tail) emit_return(j); + return true; + } + + /* ── (cons a b) — call C cons() function ── */ + if (strcmp(opname, "cons") == 0) { + if (!IS_PAIR(rest) || !IS_PAIR(CDR(rest)) || !IS_NIL(CDDR(rest))) return false; + /* Evaluate first arg */ + if (!emit_expr(j, CAR(rest), false)) return false; + emit_push(j, RAX, false); + /* Evaluate second arg */ + if (!emit_expr(j, CADR(rest), false)) return false; + emit_mov_rr(j, RSI, false, RAX, false); /* arg1 = cdr */ + emit_pop(j, RDI, false); /* arg0 = car */ + /* Call cons() — must preserve callee-saved regs */ + emit_mov_imm64(j, RAX, false, (uint64_t)(uintptr_t)cons); + eb(j, 0xFF); eb(j, MODRM(3, 2, RAX)); /* call rax */ + /* Result in rax (NaN-boxed pointer) */ + if (tail) emit_return(j); + return true; + } + + /* ── Named-let loop call — compile to jmp back to loop start ── */ + if (!IS_NIL(j->loop_sym) && head == j->loop_sym) { + Value args = rest; + int argc = 0; + Value av[MAX_JIT_LOCALS]; + while (IS_PAIR(args) && argc < MAX_JIT_LOCALS) { + av[argc++] = CAR(args); + args = CDR(args); + } + if (argc != j->loop_nparams) return false; + + /* Evaluate all args to temp stack, then store into loop slots */ + for (int i = 0; i < argc; i++) { + if (!emit_expr(j, av[i], false)) return false; + emit_push(j, RAX, false); + } + for (int i = argc - 1; i >= 0; i--) { + emit_pop(j, RAX, false); + emit_store_rbp(j, j->loop_slots[i], RAX, false); + } + /* Jump back to loop start — true zero-overhead native loop */ + eb(j, 0xE9); + int32_t rel = (int32_t)(j->loop_start - ((int)j->len + 4)); + ei32(j, rel); + return true; + } + /* ── Self-recursive call ── */ if (head == j->self_sym && !IS_NIL(j->self_sym)) { Value args = rest; @@ -703,14 +1227,17 @@ JitBlock *jit_compile(Proc *proc) { emit_push(&j, R15, true); /* push r15 */ /* Stack alignment: 6 pushes = 48 bytes. Entry was 16-aligned - 8 (ret addr). - * After 6 pushes: offset = 56 = 16*3 + 8. Need sub 8 to align to 16. */ - j.extra_stack = 8; - if (proc->nparams > 4) { - j.extra_stack += 8 * (proc->nparams - 4); - j.extra_stack = (j.extra_stack + 15) & ~15; - } + * After 6 pushes: offset = 56 = 16*3 + 8. Need sub 8 to align to 16. + * Reserve extra space for local variables (let, named-let). */ + int param_slots = (proc->nparams > 4) ? 8 * (proc->nparams - 4) : 0; + int local_slots = MAX_JIT_LOCALS * 8; /* reserve space for locals */ + j.extra_stack = 8 + param_slots + local_slots; + j.extra_stack = (j.extra_stack + 15) & ~15; emit_sub_imm(&j, RSP, false, j.extra_stack); + /* Set local_base: locals start after param spill area */ + j.local_base = -(param_slots + 8); + /* Copy System V arg regs → callee-saved param regs */ for (int i = 0; i < proc->nparams && i < 4; i++) { emit_mov_rr(&j, R12 + i, true, ARG_REGS[i], ARG_REG_EXT[i]); @@ -727,9 +1254,12 @@ JitBlock *jit_compile(Proc *proc) { j.body_start = (int)j.len; /* ═══ Emit body ═══ */ - if (!emit_expr(&j, proc->body.exprs[0], true)) { - munmap(j.buf, j.cap); - return NULL; + for (int bi = 0; bi < proc->body.count; bi++) { + bool is_last = (bi == proc->body.count - 1); + if (!emit_expr(&j, proc->body.exprs[bi], is_last)) { + munmap(j.buf, j.cap); + return NULL; + } } /* ═══ Epilogue ═══ */ diff --git a/c/test.c b/c/test.c index e26ad6e..2dd5547 100644 --- a/c/test.c +++ b/c/test.c @@ -2,6 +2,7 @@ * test.c — Unit + integration tests for the C Scheme interpreter */ #include "uncommonlisp.h" +#include "jit.h" static int tests_run = 0; static int tests_passed = 0; @@ -512,6 +513,214 @@ TEST(vm_tail_call) { ul_free(exprs); } +/* ═══════════════════════════════════════════════════════════════════════════ + * JIT tests — verify JIT-compiled functions match interpreter results + * ═══════════════════════════════════════════════════════════════════════════ */ + +/* Helper: run code with JIT enabled */ +static Value run_jit(const char *src) { + Env *g = fresh_env(); + g_jit_enabled = true; + int count; + Value *exprs = read_all(src, &count, false); + Value result = VAL_VOID; + for (int i = 0; i < count; i++) result = leval(exprs[i], g); + g_jit_enabled = false; + ul_free(exprs); + return result; +} + +TEST(jit_arithmetic) { + /* Basic arithmetic under JIT */ + ASSERT_EQ_INT(as_int(run_jit( + "(define (add a b) (+ a b)) (add 20 22)")), 42); + ASSERT_EQ_INT(as_int(run_jit( + "(define (sub a b) (- a b)) (sub 50 8)")), 42); + ASSERT_EQ_INT(as_int(run_jit( + "(define (mul a b) (* a b)) (mul 6 7)")), 42); + ASSERT_EQ_INT(as_int(run_jit( + "(define (tri-add a b c) (+ a b c)) (tri-add 10 20 12)")), 42); +} + +TEST(jit_comparison) { + ASSERT(IS_TRUE(run_jit( + "(define (eq a b) (= a b)) (eq 5 5)")), "jit = true"); + ASSERT(IS_FALSE(run_jit( + "(define (eq a b) (= a b)) (eq 5 6)")), "jit = false"); + ASSERT(IS_TRUE(run_jit( + "(define (lt a b) (< a b)) (lt 3 5)")), "jit <"); + ASSERT(IS_TRUE(run_jit( + "(define (gt a b) (> a b)) (gt 5 3)")), "jit >"); + ASSERT(IS_TRUE(run_jit( + "(define (le a b) (<= a b)) (le 3 3)")), "jit <="); + ASSERT(IS_TRUE(run_jit( + "(define (ge a b) (>= a b)) (ge 5 5)")), "jit >="); +} + +TEST(jit_if_branching) { + ASSERT_EQ_INT(as_int(run_jit( + "(define (f x) (if (= x 0) 1 2)) (f 0)")), 1); + ASSERT_EQ_INT(as_int(run_jit( + "(define (f x) (if (= x 0) 1 2)) (f 5)")), 2); +} + +TEST(jit_cond) { + ASSERT_EQ_INT(as_int(run_jit( + "(define (f x) (cond ((= x 1) 10) ((= x 2) 20) (else 30))) (f 2)")), 20); +} + +TEST(jit_not_zero) { + ASSERT(IS_TRUE(run_jit( + "(define (f x) (not x)) (f #f)")), "jit not"); + ASSERT(IS_TRUE(run_jit( + "(define (f x) (zero? x)) (f 0)")), "jit zero?"); + ASSERT(IS_FALSE(run_jit( + "(define (f x) (zero? x)) (f 5)")), "jit zero? false"); +} + +TEST(jit_self_recursion) { + /* Factorial via JIT self-recursion */ + ASSERT_EQ_INT(as_int(run_jit( + "(define (fact n) (if (= n 0) 1 (* n (fact (- n 1)))))" + "(fact 10)" + )), 3628800); +} + +TEST(jit_tco) { + /* Tail-call optimization: count down from 100k without stack overflow */ + char *s = show(run_jit( + "(define (loop n) (if (= n 0) 42 (loop (- n 1))))" + "(loop 100000)"), false); + ASSERT_EQ_STR(s, "42"); + ul_free(s); +} + +TEST(jit_and_or) { + /* and: short-circuit, returns last truthy or first falsy */ + ASSERT_EQ_INT(as_int(run_jit( + "(define (f a b) (and a b)) (f 1 2)")), 2); + ASSERT(IS_FALSE(run_jit( + "(define (f a b) (and a b)) (f 1 #f)")), "and short-circuit"); + /* or: short-circuit, returns first truthy or last falsy */ + ASSERT_EQ_INT(as_int(run_jit( + "(define (f a b) (or a b)) (f #f 3)")), 3); + ASSERT_EQ_INT(as_int(run_jit( + "(define (f a b) (or a b)) (f 1 2)")), 1); + ASSERT(IS_FALSE(run_jit( + "(define (f a b) (or a b)) (f #f #f)")), "or all false"); +} + +TEST(jit_let) { + ASSERT_EQ_INT(as_int(run_jit( + "(define (f a b) (let ((x (+ a 1)) (y (+ b 2))) (+ x y)))" + "(f 10 20)")), 33); +} + +TEST(jit_let_star) { + ASSERT_EQ_INT(as_int(run_jit( + "(define (f a) (let* ((x (+ a 1)) (y (* x 2))) y))" + "(f 5)")), 12); +} + +TEST(jit_named_let) { + /* Named let compiles to native loop with jmp back */ + ASSERT_EQ_INT(as_int(run_jit( + "(define (fib n)" + " (let loop ((a 0) (b 1) (i 0))" + " (if (= i n) a (loop b (+ a b) (+ i 1)))))" + "(fib 30)")), 832040); +} + +TEST(jit_named_let_factorial) { + ASSERT_EQ_INT(as_int(run_jit( + "(define (fact n)" + " (let loop ((i n) (acc 1))" + " (if (= i 0) acc (loop (- i 1) (* acc i)))))" + "(fact 10)")), 3628800); +} + +TEST(jit_car_cdr) { + ASSERT_EQ_INT(as_int(run_jit( + "(define (my-car p) (car p)) (my-car '(42 2 3))")), 42); + char *s = show(run_jit( + "(define (my-cdr p) (cdr p)) (my-cdr '(1 2 3))"), false); + ASSERT_EQ_STR(s, "(2 3)"); + ul_free(s); +} + +TEST(jit_cons) { + char *s = show(run_jit( + "(define (f a b) (cons a b)) (f 1 2)"), false); + ASSERT_EQ_STR(s, "(1 . 2)"); + ul_free(s); +} + +TEST(jit_null_pair) { + ASSERT(IS_TRUE(run_jit( + "(define (f x) (null? x)) (f '())")), "jit null? true"); + ASSERT(IS_FALSE(run_jit( + "(define (f x) (null? x)) (f '(1))")), "jit null? false"); + ASSERT(IS_TRUE(run_jit( + "(define (f x) (pair? x)) (f '(1 2))")), "jit pair? true"); + ASSERT(IS_FALSE(run_jit( + "(define (f x) (pair? x)) (f 42)")), "jit pair? false"); +} + +TEST(jit_ackermann) { + /* ack(3,4) = 125 — the benchmark function */ + ASSERT_EQ_INT(as_int(run_jit( + "(define (ack m n)" + " (cond ((= m 0) (+ n 1))" + " ((= n 0) (ack (- m 1) 1))" + " (else (ack (- m 1) (ack m (- n 1))))))" + "(ack 3 4)")), 125); +} + +TEST(jit_list_sum) { + /* JIT compiled function that uses car/cdr/null? in a named-let loop */ + ASSERT_EQ_INT(as_int(run_jit( + "(define (list-sum lst)" + " (let loop ((l lst) (acc 0))" + " (if (null? l) acc (loop (cdr l) (+ acc (car l))))))" + "(list-sum '(1 2 3 4 5))")), 15); +} + +TEST(jit_functional_suite) { + /* Run the full shared functional test suite with JIT enabled */ + Env *g = fresh_env(); + g_jit_enabled = true; + int count; + Value *exprs = read_all(PRELUDE, &count, false); + for (int i = 0; i < count; i++) leval(exprs[i], g); + ul_free(exprs); + + /* Load and run the functional test file. + * Try both paths since CWD may be project root or c/ subdir. */ + TRY(ctx) { + FILE *fp = fopen("tests/functional.lsp", "r"); + if (fp) { + fclose(fp); + load_file("tests/functional.lsp", g); + } else { + load_file("../tests/functional.lsp", g); + } + } CATCH { + /* Functional tests may use features beyond JIT scope — that's OK, + * the interpreter handles what the JIT can't. */ + ASSERT(0, ctx.message); + } ENDTRY; + + /* Check results — *pass* should be 114, *fail* should be 0 */ + Value pass_sym = intern("*pass*"); + Value fail_sym = intern("*fail*"); + Value pass_val = env_lookup(g, pass_sym); + Value fail_val = env_lookup(g, fail_sym); + ASSERT_EQ_INT(as_int(pass_val), 114); + ASSERT_EQ_INT(as_int(fail_val), 0); + + g_jit_enabled = false; +} + /* ═══════════════════════════════════════════════════════════════════════════ * Main * ═══════════════════════════════════════════════════════════════════════════ */ @@ -588,6 +797,26 @@ int main(void) { run_test_vm_fibonacci(); run_test_vm_tail_call(); + printf("\n[jit]\n"); + run_test_jit_arithmetic(); + run_test_jit_comparison(); + run_test_jit_if_branching(); + run_test_jit_cond(); + run_test_jit_not_zero(); + run_test_jit_self_recursion(); + run_test_jit_tco(); + run_test_jit_and_or(); + run_test_jit_let(); + run_test_jit_let_star(); + run_test_jit_named_let(); + run_test_jit_named_let_factorial(); + run_test_jit_car_cdr(); + run_test_jit_cons(); + run_test_jit_null_pair(); + run_test_jit_ackermann(); + run_test_jit_list_sum(); + run_test_jit_functional_suite(); + printf("\n═══════════════════════════════════════════\n"); printf("Results: %d/%d passed", tests_passed, tests_run); if (tests_failed > 0) printf(" (%d failed)", tests_failed);