Skip to main content

byteflow/jit/
compiler.rs

1#![allow(unsafe_code)]
2
3use std::collections::{HashMap, HashSet, VecDeque};
4
5use cranelift_codegen::ir::types;
6use cranelift_codegen::ir::{AbiParam, Block, Function, InstBuilder, MachMemFlags, Signature, Value as ClifValue};
7use cranelift_frontend::{FunctionBuilder, FunctionBuilderContext};
8use cranelift_module::{Linkage, Module};
9
10use crate::{Chunk, Instruction, Opcode, Value};
11
12use super::error::CompileError;
13use super::exit::{JIT_BUDGET, JIT_CONTINUE, JIT_RETURN, JIT_TRAP};
14use super::frame::{
15    JitEntry, OFF_BUDGET, OFF_CALL_DEPTH, OFF_CALL_STACK, OFF_EXIT_KIND, OFF_FUNCTION,
16    OFF_PC, OFF_REGISTER_COUNT, OFF_RETURN_REG, MAX_JIT_CALL_DEPTH,
17};
18use super::trace::{CompiledTrace, TraceKey, TraceSpan, MAX_TRACE_LENGTH};
19
20/// Maps bytecode register indices to Cranelift SSA values for the current trace.
21pub struct RegisterMap {
22    values: Vec<Option<ClifValue>>,
23}
24
25impl RegisterMap {
26    pub fn new(registers: usize) -> Self {
27        Self {
28            values: vec![None; registers],
29        }
30    }
31
32    pub fn get(&self, reg: usize) -> Option<ClifValue> {
33        self.values.get(reg).copied().flatten()
34    }
35
36    pub fn set(&mut self, reg: usize, value: ClifValue) {
37        if let Some(slot) = self.values.get_mut(reg) {
38            *slot = Some(value);
39        }
40    }
41
42    pub fn clear(&mut self) {
43        for slot in &mut self.values {
44            *slot = None;
45        }
46    }
47}
48
49pub struct TraceCompiler<'a> {
50    module: &'a mut cranelift_jit::JITModule,
51}
52
53impl<'a> TraceCompiler<'a> {
54    pub fn new(module: &'a mut cranelift_jit::JITModule) -> Self {
55        Self { module }
56    }
57
58    pub fn compile_trace(
59        &mut self,
60        chunk: &Chunk,
61        key: TraceKey,
62    ) -> Result<CompiledTrace, CompileError> {
63        let def = chunk.function(key.function).ok_or(CompileError::EmptyTrace {
64            function: key.function,
65            pc: key.entry_pc,
66        })?;
67        let region = collect_trace_region(chunk, key)?;
68        if region.blocks.is_empty() {
69            return Err(CompileError::EmptyTrace {
70                function: key.function,
71                pc: key.entry_pc,
72            });
73        }
74
75        let end_pc = region_max_pc(&region, key)?;
76        let span = TraceSpan {
77            function: key.function,
78            range: key.entry_pc..end_pc.saturating_add(1),
79        };
80
81        let target = self.module.target_config();
82        let pointer_type = target.pointer_type();
83        let mut sig = Signature::new(target.default_call_conv);
84        sig.params.push(AbiParam::new(pointer_type));
85
86        let func_id = self
87            .module
88            .declare_function(
89                &format!("trace_{}_{}", key.function, key.entry_pc),
90                Linkage::Local,
91                &sig,
92            )
93            .map_err(|e| CompileError::Module(e.to_string()))?;
94
95        let mut ctx = cranelift_codegen::Context::new();
96        ctx.func = Function::with_name_signature(
97            cranelift_codegen::ir::UserFuncName::user(0, func_id.as_u32()),
98            sig,
99        );
100
101        let mut func_ctx = FunctionBuilderContext::new();
102        let mut builder = FunctionBuilder::new(&mut ctx.func, &mut func_ctx);
103        let entry = builder.create_block();
104        builder.append_block_params_for_function_params(entry);
105        builder.switch_to_block(entry);
106
107        let frame_ptr = builder.block_params(entry)[0];
108        let slots_ptr = builder.ins().load(pointer_type, MachMemFlags::new(), frame_ptr, 0);
109        let budget_ptr = field_ptr(&mut builder, frame_ptr, pointer_type, OFF_BUDGET);
110        let call_depth_ptr = field_ptr(&mut builder, frame_ptr, pointer_type, OFF_CALL_DEPTH);
111        let call_stack_ptr = {
112            let addr = field_ptr(&mut builder, frame_ptr, pointer_type, OFF_CALL_STACK);
113            builder
114                .ins()
115                .load(pointer_type, MachMemFlags::new(), addr, 0)
116        };
117
118        let supports_calls = region.blocks.iter().any(|block| {
119            block
120                .pcs
121                .iter()
122                .any(|pc| chunk.code[*pc as usize].op == Opcode::Call)
123        });
124
125        let trap_block = builder.create_block();
126        let return_block = builder.create_block();
127
128        let mut pc_to_block = HashMap::new();
129        for block in &region.blocks {
130            pc_to_block.insert(block.start_pc, builder.create_block());
131        }
132
133        let entry_clif = *pc_to_block
134            .get(&key.entry_pc)
135            .ok_or(CompileError::EmptyTrace {
136                function: key.function,
137                pc: key.entry_pc,
138            })?;
139        builder.ins().jump(entry_clif, &[]);
140
141        for block in &region.blocks {
142            let clif_block = pc_to_block[&block.start_pc];
143            builder.switch_to_block(clif_block);
144            let mut registers = RegisterMap::new(def.num_registers as usize);
145
146            for pc in &block.pcs {
147                emit_budget_tick(
148                    &mut builder,
149                    budget_ptr,
150                    frame_ptr,
151                    pointer_type,
152                    *pc,
153                );
154                let instr = chunk.code[*pc as usize];
155                let mut emit_ctx = EmitContext {
156                    builder: &mut builder,
157                    chunk,
158                    registers: &mut registers,
159                    slots_ptr,
160                    frame_ptr,
161                    pointer_type,
162                    trap_block,
163                    return_block,
164                    pc_to_block: &pc_to_block,
165                    call_depth_ptr,
166                    call_stack_ptr,
167                    supports_calls,
168                };
169                emit_instruction(&mut emit_ctx, *pc, instr)?;
170            }
171        }
172
173        builder.switch_to_block(trap_block);
174        builder.seal_block(trap_block);
175        let trap_pc = region_max_pc(&region, key)?;
176        write_exit(
177            &mut builder,
178            frame_ptr,
179            pointer_type,
180            JIT_TRAP,
181            trap_pc,
182            0,
183        );
184
185        builder.switch_to_block(return_block);
186        builder.seal_block(return_block);
187        let (last_pc, return_reg) = match last_return_in_region(chunk, &region) {
188            Some(pair) => pair,
189            // Side-exit-only traces never reach this block; metadata is unused.
190            None => (key.entry_pc, 0),
191        };
192        write_exit(
193            &mut builder,
194            frame_ptr,
195            pointer_type,
196            JIT_RETURN,
197            last_pc,
198            return_reg,
199        );
200
201        builder.seal_all_blocks();
202        builder.finalize(self.module.target_config());
203
204        self.module
205            .define_function(func_id, &mut ctx)
206            .map_err(|e| CompileError::Module(e.to_string()))?;
207        self.module.clear_context(&mut ctx);
208        self.module
209            .finalize_definitions()
210            .map_err(|e| CompileError::Module(e.to_string()))?;
211
212        let code = self.module.get_finalized_function(func_id);
213        let entry: JitEntry = unsafe { std::mem::transmute(code) };
214
215        Ok(CompiledTrace { span, entry })
216    }
217}
218
219struct TraceBlock {
220    start_pc: u32,
221    pcs: Vec<u32>,
222}
223
224struct TraceRegion {
225    blocks: Vec<TraceBlock>,
226}
227
228fn region_max_pc(region: &TraceRegion, key: TraceKey) -> Result<u32, CompileError> {
229    region
230        .blocks
231        .iter()
232        .flat_map(|b| b.pcs.iter())
233        .copied()
234        .max()
235        .ok_or(CompileError::EmptyTrace {
236            function: key.function,
237            pc: key.entry_pc,
238        })
239}
240
241fn last_return_in_region(chunk: &Chunk, region: &TraceRegion) -> Option<(u32, u32)> {
242    region
243        .blocks
244        .iter()
245        .flat_map(|b| b.pcs.iter().map(|pc| (*pc, chunk.code[*pc as usize])))
246        .rev()
247        .find(|(_, instr)| instr.op == Opcode::Return)
248        .map(|(pc, instr)| (pc, u32::from(instr.a)))
249}
250
251fn jump_target(jpc: u32, imm: i32) -> u32 {
252    (i64::from(jpc) + 1 + i64::from(imm)) as u32
253}
254
255fn branch_target(bpc: u32, imm: i32) -> u32 {
256    jump_target(bpc, imm)
257}
258
259struct EmitContext<'a, 'b> {
260    builder: &'a mut FunctionBuilder<'b>,
261    chunk: &'a Chunk,
262    registers: &'a mut RegisterMap,
263    slots_ptr: ClifValue,
264    frame_ptr: ClifValue,
265    pointer_type: types::Type,
266    trap_block: Block,
267    return_block: Block,
268    pc_to_block: &'a HashMap<u32, Block>,
269    call_depth_ptr: ClifValue,
270    call_stack_ptr: ClifValue,
271    supports_calls: bool,
272}
273
274fn emit_instruction(
275    ctx: &mut EmitContext<'_, '_>,
276    pc: u32,
277    instr: Instruction,
278) -> Result<(), CompileError> {
279    let builder = &mut *ctx.builder;
280    let chunk = ctx.chunk;
281    let registers = &mut *ctx.registers;
282    let slots_ptr = ctx.slots_ptr;
283    let frame_ptr = ctx.frame_ptr;
284    let pointer_type = ctx.pointer_type;
285    let trap_block = ctx.trap_block;
286    let return_block = ctx.return_block;
287    let supports_calls = ctx.supports_calls;
288    let pc_to_block = ctx.pc_to_block;
289    let call_depth_ptr = ctx.call_depth_ptr;
290    let call_stack_ptr = ctx.call_stack_ptr;
291    match instr.op {
292        Opcode::LoadImm => {
293            let v = builder.ins().iconst(types::I64, i64::from(instr.imm));
294            registers.set(instr.a as usize, v);
295        }
296        Opcode::LoadConst => {
297            let idx = instr.imm as u32;
298            let konst = chunk.constant(idx).ok_or(CompileError::UnsupportedOpcode {
299                opcode: Opcode::LoadConst,
300                pc,
301            })?;
302            let v = match konst {
303                Value::Int(i) => builder.ins().iconst(types::I64, *i),
304                Value::Bool(b) => builder.ins().iconst(types::I64, i64::from(*b)),
305                _ => {
306                    return Err(CompileError::UnsupportedOpcode {
307                        opcode: Opcode::LoadConst,
308                        pc,
309                    });
310                }
311            };
312            registers.set(instr.a as usize, v);
313        }
314        Opcode::Move => {
315            let v = load_reg(
316                builder,
317                registers,
318                slots_ptr,
319                pointer_type,
320                instr.b,
321                pc,
322            )?;
323            registers.set(instr.a as usize, v);
324        }
325        Opcode::Add | Opcode::Sub | Opcode::Mul | Opcode::Eq | Opcode::Lt | Opcode::Le => {
326            let lhs = load_reg(
327                builder,
328                registers,
329                slots_ptr,
330                pointer_type,
331                instr.b,
332                pc,
333            )?;
334            let rhs = load_reg(
335                builder,
336                registers,
337                slots_ptr,
338                pointer_type,
339                instr.c,
340                pc,
341            )?;
342            let result = match instr.op {
343                Opcode::Add => builder.ins().iadd(lhs, rhs),
344                Opcode::Sub => builder.ins().isub(lhs, rhs),
345                Opcode::Mul => builder.ins().imul(lhs, rhs),
346                Opcode::Eq => {
347                    let cmp = builder.ins().icmp(
348                        cranelift_codegen::ir::condcodes::IntCC::Equal,
349                        lhs,
350                        rhs,
351                    );
352                    bool_as_i64(builder, cmp)
353                }
354                Opcode::Lt => {
355                    let cmp = builder.ins().icmp(
356                        cranelift_codegen::ir::condcodes::IntCC::SignedLessThan,
357                        lhs,
358                        rhs,
359                    );
360                    bool_as_i64(builder, cmp)
361                }
362                Opcode::Le => {
363                    let cmp = builder.ins().icmp(
364                        cranelift_codegen::ir::condcodes::IntCC::SignedLessThanOrEqual,
365                        lhs,
366                        rhs,
367                    );
368                    bool_as_i64(builder, cmp)
369                }
370                _ => unreachable!(),
371            };
372            registers.set(instr.a as usize, result);
373        }
374        Opcode::Div | Opcode::Mod => {
375            let lhs = load_reg(
376                builder,
377                registers,
378                slots_ptr,
379                pointer_type,
380                instr.b,
381                pc,
382            )?;
383            let rhs = load_reg(
384                builder,
385                registers,
386                slots_ptr,
387                pointer_type,
388                instr.c,
389                pc,
390            )?;
391            flush_registers(builder, registers, slots_ptr, pointer_type);
392            let zero = builder.ins().iconst(types::I64, 0);
393            let is_zero = builder.ins().icmp(
394                cranelift_codegen::ir::condcodes::IntCC::Equal,
395                rhs,
396                zero,
397            );
398            let ok = builder.create_block();
399            builder.ins().brif(is_zero, trap_block, &[], ok, &[]);
400            builder.switch_to_block(ok);
401            builder.seal_block(ok);
402            let result = if instr.op == Opcode::Div {
403                builder.ins().sdiv(lhs, rhs)
404            } else {
405                builder.ins().srem(lhs, rhs)
406            };
407            registers.set(instr.a as usize, result);
408        }
409        Opcode::Neg => {
410            let v = load_reg(
411                builder,
412                registers,
413                slots_ptr,
414                pointer_type,
415                instr.b,
416                pc,
417            )?;
418            let zero = builder.ins().iconst(types::I64, 0);
419            registers.set(instr.a as usize, builder.ins().isub(zero, v));
420        }
421        Opcode::Jump => {
422            flush_registers(builder, registers, slots_ptr, pointer_type);
423            let target_pc = jump_target(pc, instr.imm);
424            branch_to_pc(
425                builder,
426                frame_ptr,
427                pointer_type,
428                pc_to_block,
429                target_pc,
430            );
431        }
432        Opcode::Branch => {
433            let cond = load_reg(
434                builder,
435                registers,
436                slots_ptr,
437                pointer_type,
438                instr.a,
439                pc,
440            )?;
441            flush_registers(builder, registers, slots_ptr, pointer_type);
442            let zero = builder.ins().iconst(types::I64, 0);
443            let is_falsy = builder.ins().icmp(
444                cranelift_codegen::ir::condcodes::IntCC::Equal,
445                cond,
446                zero,
447            );
448            let fall_pc = pc + 1;
449            let falsy_pc = branch_target(pc, instr.imm);
450            match (pc_to_block.get(&fall_pc), pc_to_block.get(&falsy_pc)) {
451                (Some(fall), Some(falsy)) => {
452                    builder.ins().brif(is_falsy, *falsy, &[], *fall, &[]);
453                }
454                (Some(fall), None) => {
455                    let side = builder.create_block();
456                    builder.ins().brif(is_falsy, side, &[], *fall, &[]);
457                    builder.switch_to_block(side);
458                    builder.seal_block(side);
459                    write_exit(
460                        builder,
461                        frame_ptr,
462                        pointer_type,
463                        JIT_CONTINUE,
464                        falsy_pc,
465                        0,
466                    );
467                }
468                (None, Some(falsy)) => {
469                    let side = builder.create_block();
470                    builder.ins().brif(is_falsy, *falsy, &[], side, &[]);
471                    builder.switch_to_block(side);
472                    builder.seal_block(side);
473                    write_exit(
474                        builder,
475                        frame_ptr,
476                        pointer_type,
477                        JIT_CONTINUE,
478                        fall_pc,
479                        0,
480                    );
481                }
482                (None, None) => {
483                    write_exit(
484                        builder,
485                        frame_ptr,
486                        pointer_type,
487                        JIT_CONTINUE,
488                        falsy_pc,
489                        0,
490                    );
491                }
492            }
493        }
494        Opcode::Return => {
495            let v = load_reg(
496                builder,
497                registers,
498                slots_ptr,
499                pointer_type,
500                instr.a,
501                pc,
502            )?;
503            store_slot(builder, slots_ptr, pointer_type, instr.a, v);
504            if supports_calls {
505                emit_return(
506                    builder,
507                    frame_ptr,
508                    pointer_type,
509                    call_depth_ptr,
510                    call_stack_ptr,
511                    slots_ptr,
512                    registers,
513                    instr.a,
514                    pc,
515                );
516            } else {
517                write_u32_field(builder, frame_ptr, pointer_type, OFF_PC, pc);
518                builder.ins().jump(return_block, &[]);
519            }
520        }
521        Opcode::Nop => {}
522        Opcode::Call => {
523            emit_call(
524                builder,
525                chunk,
526                registers,
527                slots_ptr,
528                frame_ptr,
529                pointer_type,
530                pc_to_block,
531                call_depth_ptr,
532                call_stack_ptr,
533                pc,
534                instr,
535            )?;
536        }
537        other => {
538            return Err(CompileError::UnsupportedOpcode { opcode: other, pc });
539        }
540    }
541    Ok(())
542}
543
544fn branch_to_pc(
545    builder: &mut FunctionBuilder<'_>,
546    frame_ptr: ClifValue,
547    pointer_type: types::Type,
548    pc_to_block: &HashMap<u32, Block>,
549    target_pc: u32,
550) {
551    if let Some(block) = pc_to_block.get(&target_pc) {
552        builder.ins().jump(*block, &[]);
553    } else {
554        write_exit(
555            builder,
556            frame_ptr,
557            pointer_type,
558            JIT_CONTINUE,
559            target_pc,
560            0,
561        );
562    }
563}
564
565fn flush_registers(
566    builder: &mut FunctionBuilder<'_>,
567    registers: &RegisterMap,
568    slots_ptr: ClifValue,
569    pointer_type: types::Type,
570) {
571    for (reg, value) in registers.values.iter().enumerate() {
572        if let Some(v) = *value {
573            store_slot(builder, slots_ptr, pointer_type, reg as u8, v);
574        }
575    }
576}
577
578fn load_reg(
579    builder: &mut FunctionBuilder<'_>,
580    registers: &mut RegisterMap,
581    slots_ptr: ClifValue,
582    pointer_type: types::Type,
583    reg: u8,
584    _pc: u32,
585) -> Result<ClifValue, CompileError> {
586    if let Some(v) = registers.get(reg as usize) {
587        return Ok(v);
588    }
589    let offset = builder.ins().iconst(pointer_type, i64::from(u32::from(reg) * 8));
590    let addr = builder.ins().iadd(slots_ptr, offset);
591    let v = builder
592        .ins()
593        .load(types::I64, MachMemFlags::new(), addr, 0);
594    registers.set(reg as usize, v);
595    Ok(v)
596}
597
598fn collect_trace_region(chunk: &Chunk, key: TraceKey) -> Result<TraceRegion, CompileError> {
599    let mut queue = VecDeque::from([key.entry_pc]);
600    let mut seen_starts = HashSet::new();
601    let mut blocks = Vec::new();
602    let mut total = 0usize;
603
604    while let Some(start) = queue.pop_front() {
605        if !seen_starts.insert(start) {
606            continue;
607        }
608        if total >= MAX_TRACE_LENGTH {
609            break;
610        }
611
612        let mut pcs = Vec::new();
613        let mut pc = start;
614        while let Some(instr) = chunk.code.get(pc as usize).copied() {
615            if is_effect_opcode(instr.op) {
616                break;
617            }
618            if instr.op == Opcode::LoadConst {
619                let idx = instr.imm as u32;
620                match chunk.constant(idx) {
621                    Some(Value::Int(_) | Value::Bool(_)) => {}
622                    _ => break,
623                }
624            }
625            pcs.push(pc);
626            total += 1;
627            if total >= MAX_TRACE_LENGTH {
628                break;
629            }
630
631            match instr.op {
632                Opcode::Return => break,
633                Opcode::Jump => {
634                    queue.push_back(jump_target(pc, instr.imm));
635                    break;
636                }
637                Opcode::Branch => {
638                    queue.push_back(pc + 1);
639                    queue.push_back(branch_target(pc, instr.imm));
640                    break;
641                }
642                Opcode::Call => {
643                    let callee = instr.imm as u32;
644                    if let Some(def) = chunk.function(callee) {
645                        queue.push_back(def.entry);
646                    }
647                    queue.push_back(pc + 1);
648                    break;
649                }
650                _ => pc += 1,
651            }
652        }
653
654        if !pcs.is_empty() {
655            blocks.push(TraceBlock { start_pc: start, pcs });
656        }
657    }
658
659    Ok(TraceRegion { blocks })
660}
661
662fn is_effect_opcode(op: Opcode) -> bool {
663    matches!(
664        op,
665        Opcode::CallNative
666            | Opcode::Spawn
667            | Opcode::Yield
668            | Opcode::Sleep
669            | Opcode::Exit
670            | Opcode::SelfPid
671            | Opcode::Send
672            | Opcode::Receive
673            | Opcode::ReceiveTimeout
674            | Opcode::ReceiveMatch
675            | Opcode::ReceiveMatchImm
676            | Opcode::Ask
677            | Opcode::AskTimeout
678            | Opcode::Monitor
679            | Opcode::Demonitor
680            | Opcode::Link
681            | Opcode::Unlink
682            | Opcode::Delegate
683            | Opcode::FreshRequestId
684            | Opcode::ReceiveMatchCorr
685            | Opcode::ReceiveMatchCorrImm
686            | Opcode::RegisterName
687            | Opcode::Whereis
688            | Opcode::Trap
689            | Opcode::Halt
690    )
691}
692
693fn field_ptr(
694    builder: &mut FunctionBuilder<'_>,
695    base: ClifValue,
696    pointer_type: types::Type,
697    offset: i32,
698) -> ClifValue {
699    let off = builder.ins().iconst(pointer_type, i64::from(offset));
700    builder.ins().iadd(base, off)
701}
702
703fn store_slot(
704    builder: &mut FunctionBuilder<'_>,
705    slots: ClifValue,
706    pointer_type: types::Type,
707    reg: u8,
708    value: ClifValue,
709) {
710    let offset = builder.ins().iconst(pointer_type, i64::from(u32::from(reg) * 8));
711    let addr = builder.ins().iadd(slots, offset);
712    builder.ins().store(MachMemFlags::new(), value, addr, 0);
713}
714
715fn bool_as_i64(builder: &mut FunctionBuilder<'_>, cond: ClifValue) -> ClifValue {
716    let one = builder.ins().iconst(types::I64, 1);
717    let zero = builder.ins().iconst(types::I64, 0);
718    builder.ins().select(cond, one, zero)
719}
720
721fn emit_budget_tick(
722    builder: &mut FunctionBuilder<'_>,
723    budget_ptr: ClifValue,
724    frame_ptr: ClifValue,
725    pointer_type: types::Type,
726    pc: u32,
727) {
728    let budget = builder.ins().load(types::I32, MachMemFlags::new(), budget_ptr, 0);
729    let zero = builder.ins().iconst(types::I32, 0);
730    let exhausted = builder.ins().icmp(
731        cranelift_codegen::ir::condcodes::IntCC::Equal,
732        budget,
733        zero,
734    );
735    let continue_insn = builder.create_block();
736    let budget_exit = builder.create_block();
737    builder.ins().brif(exhausted, budget_exit, &[], continue_insn, &[]);
738    builder.switch_to_block(budget_exit);
739    builder.seal_block(budget_exit);
740    write_exit(builder, frame_ptr, pointer_type, JIT_BUDGET, pc, 0);
741    builder.switch_to_block(continue_insn);
742    builder.seal_block(continue_insn);
743    let one = builder.ins().iconst(types::I32, 1);
744    let new_budget = builder.ins().isub(budget, one);
745    builder.ins().store(MachMemFlags::new(), new_budget, budget_ptr, 0);
746}
747
748fn write_u32_field(
749    builder: &mut FunctionBuilder<'_>,
750    frame_ptr: ClifValue,
751    pointer_type: types::Type,
752    offset: i32,
753    value: u32,
754) {
755    let addr = field_ptr(builder, frame_ptr, pointer_type, offset);
756    let val = builder.ins().iconst(types::I32, i64::from(value));
757    builder.ins().store(MachMemFlags::new(), val, addr, 0);
758}
759
760fn write_exit(
761    builder: &mut FunctionBuilder<'_>,
762    frame_ptr: ClifValue,
763    pointer_type: types::Type,
764    kind: u32,
765    pc: u32,
766    return_reg: u32,
767) {
768    write_u32_field(builder, frame_ptr, pointer_type, OFF_EXIT_KIND, kind);
769    write_u32_field(builder, frame_ptr, pointer_type, OFF_PC, pc);
770    write_u32_field(builder, frame_ptr, pointer_type, OFF_RETURN_REG, return_reg);
771    builder.ins().return_(&[]);
772}
773
774#[allow(clippy::too_many_arguments)]
775fn emit_call(
776    builder: &mut FunctionBuilder<'_>,
777    chunk: &Chunk,
778    registers: &mut RegisterMap,
779    slots_ptr: ClifValue,
780    frame_ptr: ClifValue,
781    pointer_type: types::Type,
782    pc_to_block: &HashMap<u32, Block>,
783    call_depth_ptr: ClifValue,
784    call_stack_ptr: ClifValue,
785    pc: u32,
786    instr: Instruction,
787) -> Result<(), CompileError> {
788    let callee = instr.imm as u32;
789    let argc = instr.b;
790    let dst = instr.a;
791    let def = chunk.function(callee).ok_or(CompileError::UnsupportedOpcode {
792        opcode: Opcode::Call,
793        pc,
794    })?;
795    let callee_entry = def.entry;
796    let Some(callee_block) = pc_to_block.get(&callee_entry) else {
797        flush_registers(builder, registers, slots_ptr, pointer_type);
798        write_exit(
799            builder,
800            frame_ptr,
801            pointer_type,
802            JIT_CONTINUE,
803            pc,
804            0,
805        );
806        return Ok(());
807    };
808
809    flush_registers(builder, registers, slots_ptr, pointer_type);
810
811    let depth = builder
812        .ins()
813        .load(types::I32, MachMemFlags::new(), call_depth_ptr, 0);
814    let max_depth = builder
815        .ins()
816        .iconst(types::I32, i64::from(MAX_JIT_CALL_DEPTH as u32));
817    let too_deep = builder.ins().icmp(
818        cranelift_codegen::ir::condcodes::IntCC::SignedGreaterThanOrEqual,
819        depth,
820        max_depth,
821    );
822    let effect_exit = builder.create_block();
823    let call_body = builder.create_block();
824    builder
825        .ins()
826        .brif(too_deep, effect_exit, &[], call_body, &[]);
827    builder.switch_to_block(effect_exit);
828    builder.seal_block(effect_exit);
829    write_exit(
830        builder,
831        frame_ptr,
832        pointer_type,
833        super::exit::JIT_EFFECT,
834        pc,
835        0,
836    );
837    builder.switch_to_block(call_body);
838    builder.seal_block(call_body);
839
840    let record_bytes = builder.ins().iconst(
841        types::I32,
842        i64::from(std::mem::size_of::<super::frame::JitCallRecord>() as u32),
843    );
844    let record_off = builder.ins().imul(depth, record_bytes);
845    let record_off_ptr = builder.ins().uextend(pointer_type, record_off);
846    let record_addr = builder.ins().iadd(call_stack_ptr, record_off_ptr);
847
848    let return_pc = builder.ins().iconst(types::I32, i64::from(pc + 1));
849    builder
850        .ins()
851        .store(MachMemFlags::new(), return_pc, record_addr, 0);
852
853    let func_ptr = field_ptr(builder, frame_ptr, pointer_type, OFF_FUNCTION);
854    let caller_fn = builder
855        .ins()
856        .load(types::I32, MachMemFlags::new(), func_ptr, 0);
857    let caller_fn_addr = {
858        let off = builder.ins().iconst(pointer_type, 4);
859        builder.ins().iadd(record_addr, off)
860    };
861    builder
862        .ins()
863        .store(MachMemFlags::new(), caller_fn, caller_fn_addr, 0);
864
865    let dest = builder.ins().iconst(types::I32, i64::from(u32::from(dst)));
866    let dest_addr = {
867        let off = builder.ins().iconst(pointer_type, 8);
868        builder.ins().iadd(record_addr, off)
869    };
870    builder
871        .ins()
872        .store(MachMemFlags::new(), dest, dest_addr, 0);
873
874    let reg_ptr = field_ptr(builder, frame_ptr, pointer_type, OFF_REGISTER_COUNT);
875    let caller_regs = builder
876        .ins()
877        .load(types::I32, MachMemFlags::new(), reg_ptr, 0);
878    let caller_regs_addr = {
879        let off = builder.ins().iconst(pointer_type, 12);
880        builder.ins().iadd(record_addr, off)
881    };
882    builder
883        .ins()
884        .store(MachMemFlags::new(), caller_regs, caller_regs_addr, 0);
885
886    let one = builder.ins().iconst(types::I32, 1);
887    let new_depth = builder.ins().iadd(depth, one);
888    builder
889        .ins()
890        .store(MachMemFlags::new(), new_depth, call_depth_ptr, 0);
891
892    for i in 0..argc {
893        let src_reg = u32::from(dst) + u32::from(i);
894        let src_offset = builder.ins().iconst(pointer_type, i64::from(src_reg * 8));
895        let src_addr = builder.ins().iadd(slots_ptr, src_offset);
896        let value = builder
897            .ins()
898            .load(types::I64, MachMemFlags::new(), src_addr, 0);
899        let dst_offset = builder.ins().iconst(pointer_type, i64::from(u32::from(i) * 8));
900        let dst_addr = builder.ins().iadd(slots_ptr, dst_offset);
901        builder
902            .ins()
903            .store(MachMemFlags::new(), value, dst_addr, 0);
904    }
905
906    let callee_fn = builder.ins().iconst(types::I32, i64::from(callee));
907    builder
908        .ins()
909        .store(MachMemFlags::new(), callee_fn, func_ptr, 0);
910    let callee_regs = builder.ins().iconst(types::I32, i64::from(def.num_registers));
911    builder
912        .ins()
913        .store(MachMemFlags::new(), callee_regs, reg_ptr, 0);
914
915    registers.clear();
916    builder.ins().jump(*callee_block, &[]);
917    Ok(())
918}
919
920#[allow(clippy::too_many_arguments)]
921fn emit_return(
922    builder: &mut FunctionBuilder<'_>,
923    frame_ptr: ClifValue,
924    pointer_type: types::Type,
925    call_depth_ptr: ClifValue,
926    call_stack_ptr: ClifValue,
927    slots_ptr: ClifValue,
928    registers: &mut RegisterMap,
929    return_reg: u8,
930    pc: u32,
931) {
932    let depth = builder
933        .ins()
934        .load(types::I32, MachMemFlags::new(), call_depth_ptr, 0);
935    let zero = builder.ins().iconst(types::I32, 0);
936    let is_outer = builder.ins().icmp(
937        cranelift_codegen::ir::condcodes::IntCC::Equal,
938        depth,
939        zero,
940    );
941    let outer = builder.create_block();
942    let inner = builder.create_block();
943    builder
944        .ins()
945        .brif(is_outer, outer, &[], inner, &[]);
946    builder.switch_to_block(outer);
947    builder.seal_block(outer);
948    write_exit(
949        builder,
950        frame_ptr,
951        pointer_type,
952        JIT_RETURN,
953        pc,
954        u32::from(return_reg),
955    );
956    builder.switch_to_block(inner);
957    builder.seal_block(inner);
958
959    let one = builder.ins().iconst(types::I32, 1);
960    let new_depth = builder.ins().isub(depth, one);
961    builder
962        .ins()
963        .store(MachMemFlags::new(), new_depth, call_depth_ptr, 0);
964
965    let record_bytes = builder.ins().iconst(
966        types::I32,
967        i64::from(std::mem::size_of::<super::frame::JitCallRecord>() as u32),
968    );
969    let record_off = builder.ins().imul(new_depth, record_bytes);
970    let record_off_ptr = builder.ins().uextend(pointer_type, record_off);
971    let record_addr = builder.ins().iadd(call_stack_ptr, record_off_ptr);
972
973    let return_pc = builder
974        .ins()
975        .load(types::I32, MachMemFlags::new(), record_addr, 0);
976    let return_fn_addr = {
977        let off = builder.ins().iconst(pointer_type, 4);
978        builder.ins().iadd(record_addr, off)
979    };
980    let return_fn = builder
981        .ins()
982        .load(types::I32, MachMemFlags::new(), return_fn_addr, 0);
983    let dest_addr = {
984        let off = builder.ins().iconst(pointer_type, 8);
985        builder.ins().iadd(record_addr, off)
986    };
987    let dest_reg = builder
988        .ins()
989        .load(types::I32, MachMemFlags::new(), dest_addr, 0);
990    let caller_regs_addr = {
991        let off = builder.ins().iconst(pointer_type, 12);
992        builder.ins().iadd(record_addr, off)
993    };
994    let caller_regs = builder
995        .ins()
996        .load(types::I32, MachMemFlags::new(), caller_regs_addr, 0);
997
998    let ret_offset = builder.ins().iconst(pointer_type, i64::from(u32::from(return_reg) * 8));
999    let ret_addr = builder.ins().iadd(slots_ptr, ret_offset);
1000    let ret_val = builder
1001        .ins()
1002        .load(types::I64, MachMemFlags::new(), ret_addr, 0);
1003    let eight = builder.ins().iconst(types::I32, 8);
1004    let dest_offset = builder.ins().imul(dest_reg, eight);
1005    let dest_offset_ptr = builder.ins().uextend(pointer_type, dest_offset);
1006    let dest_slot = builder.ins().iadd(slots_ptr, dest_offset_ptr);
1007    builder
1008        .ins()
1009        .store(MachMemFlags::new(), ret_val, dest_slot, 0);
1010
1011    let func_ptr = field_ptr(builder, frame_ptr, pointer_type, OFF_FUNCTION);
1012    builder
1013        .ins()
1014        .store(MachMemFlags::new(), return_fn, func_ptr, 0);
1015    let reg_ptr = field_ptr(builder, frame_ptr, pointer_type, OFF_REGISTER_COUNT);
1016    builder
1017        .ins()
1018        .store(MachMemFlags::new(), caller_regs, reg_ptr, 0);
1019
1020    registers.clear();
1021    write_dynamic_exit(
1022        builder,
1023        frame_ptr,
1024        pointer_type,
1025        JIT_CONTINUE,
1026        return_pc,
1027        0,
1028    );
1029    let _ = pc;
1030}
1031
1032fn write_dynamic_exit(
1033    builder: &mut FunctionBuilder<'_>,
1034    frame_ptr: ClifValue,
1035    pointer_type: types::Type,
1036    kind: u32,
1037    pc: ClifValue,
1038    return_reg: u32,
1039) {
1040    write_u32_field(builder, frame_ptr, pointer_type, OFF_EXIT_KIND, kind);
1041    let pc_ptr = field_ptr(builder, frame_ptr, pointer_type, OFF_PC);
1042    builder.ins().store(MachMemFlags::new(), pc, pc_ptr, 0);
1043    write_u32_field(builder, frame_ptr, pointer_type, OFF_RETURN_REG, return_reg);
1044    builder.ins().return_(&[]);
1045}
1046
1047#[cfg(test)]
1048mod tests {
1049    use std::sync::Arc;
1050
1051    use crate::{bytecode::builder::ChunkBuilder, NativeTable, Opcode, Vm};
1052
1053    use super::*;
1054    use crate::jit::dispatch::force_compile;
1055    use crate::jit::trace::{JitContext, TraceKey};
1056
1057    type TestResult = Result<(), Box<dyn std::error::Error>>;
1058
1059    #[test]
1060    fn compiles_load_imm_add_return() -> TestResult {
1061        let mut b = ChunkBuilder::new("jit");
1062        b.begin_function("main", 0, 3);
1063        b.emit_load_imm(0, 41);
1064        b.emit_load_imm(1, 1);
1065        b.emit_binop(Opcode::Add, 2, 0, 1);
1066        b.emit_return(2);
1067        let chunk = Arc::new(b.finish());
1068
1069        let mut ctx = JitContext::new(chunk.clone())?;
1070        let key = TraceKey {
1071            function: 0,
1072            entry_pc: 0,
1073        };
1074        let mut compiler = TraceCompiler::new(ctx.module_mut());
1075        let trace = compiler.compile_trace(&chunk, key)?;
1076        assert!(trace.span.range.contains(&0));
1077        Ok(())
1078    }
1079
1080    #[test]
1081    fn compiles_branch_falsy_taken() -> TestResult {
1082        let mut b = ChunkBuilder::new("jit-branch");
1083        b.begin_function("main", 0, 1);
1084        let done = b.new_label();
1085        let ret = b.new_label();
1086        b.emit_load_imm(0, 0);
1087        b.emit_branch(0, done);
1088        b.emit_load_imm(0, 99);
1089        b.emit_jump(ret);
1090        b.bind_label(done);
1091        b.emit_load_imm(0, 7);
1092        b.bind_label(ret);
1093        b.emit_return(0);
1094        let chunk = Arc::new(b.finish());
1095
1096        let mut ctx = JitContext::new(chunk.clone())?;
1097        let key = TraceKey {
1098            function: 0,
1099            entry_pc: 0,
1100        };
1101        let mut compiler = TraceCompiler::new(ctx.module_mut());
1102        let trace = compiler.compile_trace(&chunk, key)?;
1103        let mut slots = vec![0i64; 1];
1104        let ret = super::super::dispatch::run_compiled_trace_ref(
1105            &trace,
1106            &mut slots,
1107            0,
1108            10_000,
1109            0,
1110            1,
1111        );
1112        assert_eq!(slots[0], 7);
1113        assert!(matches!(
1114            ret.into_reason(),
1115            super::super::exit::ExitReason::Return { .. }
1116        ));
1117        Ok(())
1118    }
1119
1120    #[test]
1121    fn compiles_unconditional_jump() -> TestResult {
1122        let mut b = ChunkBuilder::new("jit-jump");
1123        b.begin_function("main", 0, 1);
1124        let skip = b.new_label();
1125        let ret = b.new_label();
1126        b.emit_load_imm(0, 1);
1127        b.emit_jump(skip);
1128        b.emit_load_imm(0, 99);
1129        b.bind_label(skip);
1130        b.emit_load_imm(0, 41);
1131        b.bind_label(ret);
1132        b.emit_return(0);
1133        let chunk = Arc::new(b.finish());
1134
1135        let mut ctx = JitContext::new(chunk.clone())?;
1136        let key = TraceKey {
1137            function: 0,
1138            entry_pc: 0,
1139        };
1140        let mut compiler = TraceCompiler::new(ctx.module_mut());
1141        let trace = compiler.compile_trace(&chunk, key)?;
1142        let mut slots = vec![0i64; 1];
1143        let ret = super::super::dispatch::run_compiled_trace_ref(
1144            &trace,
1145            &mut slots,
1146            0,
1147            10_000,
1148            0,
1149            1,
1150        );
1151        assert!(matches!(
1152            ret.into_reason(),
1153            super::super::exit::ExitReason::Return { return_reg: 0 }
1154        ));
1155        assert_eq!(slots[0], 41);
1156        Ok(())
1157    }
1158
1159    #[test]
1160    fn compiles_intra_chunk_call() -> TestResult {
1161        let mut b = ChunkBuilder::new("jit-call");
1162        b.begin_function("double", 0, 1);
1163        b.emit_load_imm(0, 21);
1164        b.emit_return(0);
1165        b.begin_function("main", 1, 1);
1166        b.emit_load_imm(0, 0);
1167        b.emit_call(0, 0, 1);
1168        b.emit_return(0);
1169        let chunk = Arc::new(b.finish());
1170
1171        let mut ctx = JitContext::new(chunk.clone())?;
1172        let key = TraceKey {
1173            function: 1,
1174            entry_pc: chunk.functions[1].entry,
1175        };
1176        force_compile(&mut ctx, &chunk, key)?;
1177        let mut vm = Vm::new(chunk, NativeTable::empty(), key.function, &[])?;
1178        let result = super::super::dispatch::run_vm_with_jit(&mut vm, 10_000, &mut ctx);
1179        assert!(matches!(
1180            result,
1181            crate::VmResult::Complete(crate::Value::Int(21))
1182        ));
1183        Ok(())
1184    }
1185}