Skip to main content

shape_jit/
worker.rs

1//! JIT Compilation Backend
2//!
3//! Implements the `CompilationBackend` trait from `shape-vm` so the TierManager
4//! can drive JIT compilation on a background worker thread.
5
6use shape_vm::bytecode::BytecodeProgram;
7use shape_vm::tier::{CompilationBackend, CompilationRequest, CompilationResult, Tier};
8use shape_vm::type_tracking::FrameDescriptor;
9
10use crate::compiler::JITCompiler;
11use crate::context::JITConfig;
12use crate::loop_analysis;
13use crate::osr_compiler;
14
15/// JIT compilation backend that compiles hot loops to native code via Cranelift.
16///
17/// Owns a `JITCompiler` instance and implements the `CompilationBackend` trait.
18/// The `TierManager::set_backend()` spawns a worker thread that drives this.
19pub struct JitCompilationBackend {
20    jit: JITCompiler,
21}
22
23impl JitCompilationBackend {
24    /// Create a new JIT compilation backend with default configuration.
25    pub fn new() -> Result<Self, crate::error::JitError> {
26        Ok(Self {
27            jit: JITCompiler::new(JITConfig::default())?,
28        })
29    }
30
31    /// Create a new JIT compilation backend with custom configuration.
32    pub fn with_config(config: JITConfig) -> Result<Self, crate::error::JitError> {
33        Ok(Self {
34            jit: JITCompiler::new(config)?,
35        })
36    }
37
38    /// Compile an OSR loop from a compilation request.
39    fn compile_osr(
40        &mut self,
41        request: &CompilationRequest,
42        program: &BytecodeProgram,
43    ) -> CompilationResult {
44        let func_id = request.function_id;
45        let loop_header_ip = request.loop_header_ip;
46
47        // Get the target function
48        let function = match program.functions.get(func_id as usize) {
49            Some(f) => f,
50            None => {
51                return CompilationResult {
52                    function_id: func_id,
53                    compiled_tier: Tier::Interpreted,
54                    native_code: None,
55                    error: Some(format!("Function {} not found in program", func_id)),
56                    osr_entry: None,
57                    deopt_points: Vec::new(),
58                    loop_header_ip,
59                    shape_guards: Vec::new(),
60                };
61            }
62        };
63
64        // Extract the function's instruction range
65        let entry = function.entry_point;
66        let end = find_function_end(program, func_id as usize);
67        if entry >= program.instructions.len() || end > program.instructions.len() {
68            return CompilationResult {
69                function_id: func_id,
70                compiled_tier: Tier::Interpreted,
71                native_code: None,
72                error: Some(format!(
73                    "Function {} instruction range [{}, {}) out of bounds",
74                    func_id, entry, end
75                )),
76                osr_entry: None,
77                deopt_points: Vec::new(),
78                loop_header_ip,
79                shape_guards: Vec::new(),
80            };
81        }
82        let func_instructions = &program.instructions[entry..end];
83
84        // Run loop analysis on a sub-program containing just this function's instructions
85        let sub_program = build_sub_program(program, entry, end);
86        let loop_infos = loop_analysis::analyze_loops(&sub_program);
87
88        // Find the target loop. The loop_header_ip from the request is in
89        // global instruction coordinates; convert to function-local offset.
90        let target_local_ip = match loop_header_ip {
91            Some(ip) => {
92                if ip < entry {
93                    return CompilationResult {
94                        function_id: func_id,
95                        compiled_tier: Tier::Interpreted,
96                        native_code: None,
97                        error: Some(format!(
98                            "OSR loop header IP {} is before function entry {}",
99                            ip, entry
100                        )),
101                        osr_entry: None,
102                        deopt_points: Vec::new(),
103                        loop_header_ip: Some(ip),
104                        shape_guards: Vec::new(),
105                    };
106                }
107                ip - entry
108            }
109            None => {
110                return CompilationResult {
111                    function_id: func_id,
112                    compiled_tier: Tier::Interpreted,
113                    native_code: None,
114                    error: Some("OSR request without loop_header_ip".to_string()),
115                    osr_entry: None,
116                    deopt_points: Vec::new(),
117                    loop_header_ip: None,
118                    shape_guards: Vec::new(),
119                };
120            }
121        };
122
123        let loop_info = match loop_infos.get(&target_local_ip) {
124            Some(li) => li,
125            None => {
126                return CompilationResult {
127                    function_id: func_id,
128                    compiled_tier: Tier::Interpreted,
129                    native_code: None,
130                    error: Some(format!(
131                        "No loop found at local IP {} (global IP {:?})",
132                        target_local_ip, loop_header_ip
133                    )),
134                    osr_entry: None,
135                    deopt_points: Vec::new(),
136                    loop_header_ip,
137                    shape_guards: Vec::new(),
138                };
139            }
140        };
141
142        // Build frame descriptor (use function's if available, else default)
143        let default_frame = FrameDescriptor::default();
144        let frame_descriptor = function.frame_descriptor.as_ref().unwrap_or(&default_frame);
145
146        // Compile the loop
147        match osr_compiler::compile_osr_loop(
148            &mut self.jit,
149            function,
150            func_instructions,
151            loop_info,
152            frame_descriptor,
153        ) {
154            Ok(osr_result) => {
155                // Adjust entry point bytecode_ip back to global coordinates
156                let mut entry_point = osr_result.entry_point;
157                entry_point.bytecode_ip += entry;
158                entry_point.exit_ip += entry;
159
160                CompilationResult {
161                    function_id: func_id,
162                    compiled_tier: Tier::BaselineJit,
163                    native_code: Some(osr_result.native_code),
164                    error: None,
165                    osr_entry: Some(entry_point),
166                    deopt_points: osr_result.deopt_points,
167                    loop_header_ip,
168                    shape_guards: Vec::new(),
169                }
170            }
171            Err(e) => CompilationResult {
172                function_id: func_id,
173                compiled_tier: Tier::Interpreted,
174                native_code: None,
175                error: Some(e),
176                osr_entry: None,
177                deopt_points: Vec::new(),
178                loop_header_ip,
179                shape_guards: Vec::new(),
180            },
181        }
182    }
183}
184
185// SAFETY: JitCompilationBackend is used exclusively on its own worker thread.
186// The raw pointers in JITCompiler (compiled_functions, function_table) point
187// to JIT code that is immutable after compilation and valid for the module's
188// lifetime. Access is single-threaded (the worker thread).
189unsafe impl Send for JitCompilationBackend {}
190
191impl JitCompilationBackend {
192    /// Compile a whole function for Tier 1/2 promotion.
193    ///
194    /// Tier 1 (BaselineJit, no feedback): uses `compile_single_function` with
195    /// empty user_funcs — cross-function calls deopt to interpreter.
196    ///
197    /// Tier 2 (OptimizingJit, with feedback): uses `compile_optimizing_function`
198    /// which enables speculative calls based on monomorphic call feedback.
199    /// Self-recursive calls get direct-call FuncRefs. Cross-function monomorphic
200    /// calls get callee identity guard + FFI fallthrough (guard deopt on mismatch).
201    fn compile_function(
202        &mut self,
203        request: &CompilationRequest,
204        program: &BytecodeProgram,
205    ) -> CompilationResult {
206        let func_id = request.function_id;
207
208        // Tier 2: feedback-guided optimizing compilation with populated user_funcs
209        if let Some(fv) = request.feedback.clone() {
210            return match self.jit.compile_optimizing_function(
211                program,
212                func_id as usize,
213                fv,
214                &request.callee_feedback,
215            ) {
216                Ok((code_ptr, deopt_points, shape_guards)) => CompilationResult {
217                    function_id: func_id,
218                    compiled_tier: request.target_tier,
219                    native_code: Some(code_ptr),
220                    error: None,
221                    osr_entry: None,
222                    deopt_points,
223                    loop_header_ip: None,
224                    shape_guards,
225                },
226                Err(e) => CompilationResult {
227                    function_id: func_id,
228                    compiled_tier: Tier::Interpreted,
229                    native_code: None,
230                    error: Some(e),
231                    osr_entry: None,
232                    deopt_points: Vec::new(),
233                    loop_header_ip: None,
234                    shape_guards: Vec::new(),
235                },
236            };
237        }
238
239        // Tier 1: baseline compilation without cross-function speculation
240        match self
241            .jit
242            .compile_single_function(program, func_id as usize, None)
243        {
244            Ok((code_ptr, deopt_points, shape_guards)) => CompilationResult {
245                function_id: func_id,
246                compiled_tier: request.target_tier,
247                native_code: Some(code_ptr),
248                error: None,
249                osr_entry: None,
250                deopt_points,
251                loop_header_ip: None,
252                shape_guards,
253            },
254            Err(e) => CompilationResult {
255                function_id: func_id,
256                compiled_tier: Tier::Interpreted,
257                native_code: None,
258                error: Some(e),
259                osr_entry: None,
260                deopt_points: Vec::new(),
261                loop_header_ip: None,
262                shape_guards: Vec::new(),
263            },
264        }
265    }
266}
267
268impl CompilationBackend for JitCompilationBackend {
269    fn compile(
270        &mut self,
271        request: &CompilationRequest,
272        program: &BytecodeProgram,
273    ) -> CompilationResult {
274        if request.osr {
275            self.compile_osr(request, program)
276        } else {
277            self.compile_function(request, program)
278        }
279    }
280}
281
282/// Find the end of a function's instruction range.
283///
284/// For the last function, this is the end of the instruction stream.
285/// For other functions, this is the entry point of the next function.
286fn find_function_end(program: &BytecodeProgram, func_index: usize) -> usize {
287    let func = &program.functions[func_index];
288    func.entry_point + func.body_length
289}
290
291/// Build a minimal sub-program containing only the instructions in [start, end).
292///
293/// The sub-program's instructions are indexed from 0, making it compatible
294/// with `analyze_loops()` which expects a contiguous instruction stream.
295fn build_sub_program(program: &BytecodeProgram, start: usize, end: usize) -> BytecodeProgram {
296    BytecodeProgram {
297        instructions: program.instructions[start..end].to_vec(),
298        constants: program.constants.clone(),
299        strings: program.strings.clone(),
300        functions: vec![],
301        debug_info: Default::default(),
302        data_schema: None,
303        module_binding_names: vec![],
304        top_level_locals_count: 0,
305        top_level_local_storage_hints: vec![],
306        type_schema_registry: Default::default(),
307        module_binding_storage_hints: vec![],
308        function_local_storage_hints: vec![],
309        compiled_annotations: Default::default(),
310        trait_method_symbols: Default::default(),
311        expanded_function_defs: Default::default(),
312        string_index: Default::default(),
313        foreign_functions: Vec::new(),
314        native_struct_layouts: vec![],
315        content_addressed: None,
316            top_level_mir: None,
317        function_blob_hashes: vec![],
318        top_level_frame: None,
319        top_level_local_concrete_types: vec![],
320        function_local_concrete_types: vec![],
321        function_return_concrete_types: vec![],
322        monomorphized_method_call_sites: Default::default(),
323        value_call_return_concrete_types: Default::default(),
324        operator_trait_dispatch_sites: Default::default(),
325        monomorphization_keys: vec![],
326        closure_function_layouts: program.closure_function_layouts.clone(),
327        trait_vtables: program.trait_vtables.clone(),
328        has_imported_const_inline: program.has_imported_const_inline,
329        has_w17_marshal_residual: program.has_w17_marshal_residual,
330    }
331}
332
333#[cfg(test)]
334mod tests {
335    use super::*;
336    use shape_vm::bytecode::*;
337    use shape_vm::type_tracking::{FrameDescriptor, NativeKind};
338
339    fn make_instr(opcode: OpCode, operand: Option<Operand>) -> Instruction {
340        Instruction { opcode, operand }
341    }
342
343    #[test]
344    #[ignore = "v2: Tier 1 whole-function JIT (compile_single_function) deprecated; tests dead path"]
345    fn test_backend_compiles_whole_function() {
346        let mut backend = JitCompilationBackend::new().unwrap();
347
348        // Simple function: return local 0 + local 1
349        let instrs = vec![
350            // Function body at entry_point=0
351            make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), // 0
352            make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), // 1
353            make_instr(OpCode::AddInt, None),                       // 2
354            make_instr(OpCode::ReturnValue, None),                  // 3
355            // Main code (trampoline target)
356            make_instr(OpCode::Halt, None), // 4
357        ];
358
359        let func = Function {
360            name: "add_two".to_string(),
361            arity: 2,
362            param_names: vec![],
363            locals_count: 2,
364            entry_point: 0,
365            body_length: 4,
366            is_closure: false,
367            captures_count: 0,
368            is_async: false,
369            ref_params: vec![],
370                    mir_data: None,
371            ref_mutates: vec![],
372            mutable_captures: vec![],
373            frame_descriptor: Some(FrameDescriptor::from_slots(vec![
374                NativeKind::Int64, // arg0
375                NativeKind::Int64, // arg1
376            ])),
377            osr_entry_points: vec![],
378        };
379
380        let program = BytecodeProgram {
381            instructions: instrs,
382            constants: vec![],
383            strings: vec![],
384            functions: vec![func],
385            debug_info: Default::default(),
386            data_schema: None,
387            module_binding_names: vec![],
388            top_level_locals_count: 0,
389            top_level_local_storage_hints: vec![],
390            type_schema_registry: Default::default(),
391            module_binding_storage_hints: vec![],
392            function_local_storage_hints: vec![],
393            compiled_annotations: Default::default(),
394            trait_method_symbols: Default::default(),
395            expanded_function_defs: Default::default(),
396            string_index: Default::default(),
397            foreign_functions: Vec::new(),
398            native_struct_layouts: vec![],
399            content_addressed: None,
400            top_level_mir: None,
401            function_blob_hashes: vec![],
402            top_level_frame: None,
403            ..Default::default()
404        };
405
406        let request = CompilationRequest {
407            function_id: 0,
408            target_tier: Tier::BaselineJit,
409            blob_hash: None,
410            osr: false,
411            loop_header_ip: None,
412            feedback: None,
413            callee_feedback: std::collections::HashMap::new(),
414        };
415
416        let result = backend.compile(&request, &program);
417        assert!(
418            result.error.is_none(),
419            "Expected successful whole-function compilation, got: {:?}",
420            result.error
421        );
422        assert!(result.native_code.is_some());
423        assert_eq!(result.compiled_tier, Tier::BaselineJit);
424        assert!(result.osr_entry.is_none()); // Not an OSR result
425    }
426
427    #[test]
428    #[ignore = "v2: Tier 1 whole-function JIT deprecated; test asserts on error message that no longer matches"]
429    fn test_backend_whole_function_invalid_id() {
430        let mut backend = JitCompilationBackend::new().unwrap();
431        let program = BytecodeProgram {
432            instructions: vec![make_instr(OpCode::Halt, None)],
433            constants: vec![],
434            strings: vec![],
435            functions: vec![], // No functions
436            debug_info: Default::default(),
437            data_schema: None,
438            module_binding_names: vec![],
439            top_level_locals_count: 0,
440            top_level_local_storage_hints: vec![],
441            type_schema_registry: Default::default(),
442            module_binding_storage_hints: vec![],
443            function_local_storage_hints: vec![],
444            compiled_annotations: Default::default(),
445            trait_method_symbols: Default::default(),
446            expanded_function_defs: Default::default(),
447            string_index: Default::default(),
448            foreign_functions: Vec::new(),
449            native_struct_layouts: vec![],
450            content_addressed: None,
451            top_level_mir: None,
452            function_blob_hashes: vec![],
453            top_level_frame: None,
454            ..Default::default()
455        };
456        let request = CompilationRequest {
457            function_id: 99,
458            target_tier: Tier::BaselineJit,
459            blob_hash: None,
460            osr: false,
461            loop_header_ip: None,
462            feedback: None,
463            callee_feedback: std::collections::HashMap::new(),
464        };
465        let result = backend.compile(&request, &program);
466        assert!(result.error.is_some());
467        assert!(result.error.unwrap().contains("not found"));
468    }
469
470    #[test]
471    fn test_backend_osr_compiles_simple_loop() {
472        let mut backend = JitCompilationBackend::new().unwrap();
473
474        // Function at entry_point=0: for (i=0; i<n; i++) { sum += i }
475        let instrs = vec![
476            make_instr(OpCode::LoopStart, None),                       // 0
477            make_instr(OpCode::LoadLocal, Some(Operand::Local(0))),    // 1: i
478            make_instr(OpCode::LoadLocal, Some(Operand::Local(1))),    // 2: n
479            make_instr(OpCode::LtInt, None),                           // 3
480            make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(7))), // 4
481            make_instr(OpCode::LoadLocal, Some(Operand::Local(2))),    // 5: sum
482            make_instr(OpCode::LoadLocal, Some(Operand::Local(0))),    // 6: i
483            make_instr(OpCode::AddInt, None),                          // 7
484            make_instr(OpCode::StoreLocal, Some(Operand::Local(2))),   // 8
485            make_instr(OpCode::LoadLocal, Some(Operand::Local(0))),    // 9: i
486            make_instr(OpCode::PushConst, Some(Operand::Const(0))),    // 10: 1
487            make_instr(OpCode::AddInt, None),                          // 11
488            make_instr(OpCode::StoreLocal, Some(Operand::Local(0))),   // 12
489            make_instr(OpCode::LoopEnd, None),                         // 13
490            make_instr(OpCode::ReturnValue, None),                     // 14
491        ];
492
493        let func = Function {
494            name: "test_loop".to_string(),
495            arity: 0,
496            param_names: vec![],
497            locals_count: 3,
498            entry_point: 0,
499            body_length: 15,
500            is_closure: false,
501            captures_count: 0,
502            is_async: false,
503            ref_params: vec![],
504                    mir_data: None,
505            ref_mutates: vec![],
506            mutable_captures: vec![],
507            frame_descriptor: Some(FrameDescriptor::from_slots(vec![
508                NativeKind::Int64, // i
509                NativeKind::Int64, // n
510                NativeKind::Int64, // sum
511            ])),
512            osr_entry_points: vec![],
513        };
514
515        let program = BytecodeProgram {
516            instructions: instrs,
517            constants: vec![Constant::Int(1)],
518            strings: vec![],
519            functions: vec![func],
520            debug_info: Default::default(),
521            data_schema: None,
522            module_binding_names: vec![],
523            top_level_locals_count: 0,
524            top_level_local_storage_hints: vec![],
525            type_schema_registry: Default::default(),
526            module_binding_storage_hints: vec![],
527            function_local_storage_hints: vec![],
528            compiled_annotations: Default::default(),
529            trait_method_symbols: Default::default(),
530            expanded_function_defs: Default::default(),
531            string_index: Default::default(),
532            foreign_functions: Vec::new(),
533            native_struct_layouts: vec![],
534            content_addressed: None,
535            top_level_mir: None,
536            function_blob_hashes: vec![],
537            top_level_frame: None,
538            ..Default::default()
539        };
540
541        let request = CompilationRequest {
542            function_id: 0,
543            target_tier: Tier::BaselineJit,
544            blob_hash: None,
545            osr: true,
546            loop_header_ip: Some(0), // Global IP of LoopStart
547            feedback: None,
548            callee_feedback: std::collections::HashMap::new(),
549        };
550
551        let result = backend.compile(&request, &program);
552        assert!(
553            result.error.is_none(),
554            "Expected successful compilation, got: {:?}",
555            result.error
556        );
557        assert!(result.native_code.is_some());
558        assert!(result.osr_entry.is_some());
559        assert_eq!(result.compiled_tier, Tier::BaselineJit);
560
561        let entry = result.osr_entry.unwrap();
562        assert_eq!(entry.bytecode_ip, 0);
563        assert!(entry.live_locals.contains(&0)); // i
564        assert!(entry.live_locals.contains(&1)); // n
565        assert!(entry.live_locals.contains(&2)); // sum
566    }
567
568    #[test]
569    fn test_backend_osr_blacklists_unsupported_loop() {
570        let mut backend = JitCompilationBackend::new().unwrap();
571
572        // Function with a loop containing CallMethod (unsupported)
573        let instrs = vec![
574            make_instr(OpCode::LoopStart, None),
575            make_instr(OpCode::LoadLocal, Some(Operand::Local(0))),
576            make_instr(OpCode::CallMethod, None), // Unsupported!
577            make_instr(OpCode::Pop, None),
578            make_instr(OpCode::LoopEnd, None),
579            make_instr(OpCode::Halt, None),
580        ];
581
582        let func = Function {
583            name: "unsupported_loop".to_string(),
584            arity: 0,
585            param_names: vec![],
586            locals_count: 1,
587            entry_point: 0,
588            body_length: 6,
589            is_closure: false,
590            captures_count: 0,
591            is_async: false,
592            ref_params: vec![],
593                    mir_data: None,
594            ref_mutates: vec![],
595            mutable_captures: vec![],
596            // W11: `NativeKind::Unknown` deleted; `Bool` is a benign stand-in
597            // for this single-slot test descriptor (the slot is never read).
598            frame_descriptor: Some(FrameDescriptor::from_slots(vec![NativeKind::Bool])),
599            osr_entry_points: vec![],
600        };
601
602        let program = BytecodeProgram {
603            instructions: instrs,
604            constants: vec![],
605            strings: vec![],
606            functions: vec![func],
607            debug_info: Default::default(),
608            data_schema: None,
609            module_binding_names: vec![],
610            top_level_locals_count: 0,
611            top_level_local_storage_hints: vec![],
612            type_schema_registry: Default::default(),
613            module_binding_storage_hints: vec![],
614            function_local_storage_hints: vec![],
615            compiled_annotations: Default::default(),
616            trait_method_symbols: Default::default(),
617            expanded_function_defs: Default::default(),
618            string_index: Default::default(),
619            foreign_functions: Vec::new(),
620            native_struct_layouts: vec![],
621            content_addressed: None,
622            top_level_mir: None,
623            function_blob_hashes: vec![],
624            top_level_frame: None,
625            ..Default::default()
626        };
627
628        let request = CompilationRequest {
629            function_id: 0,
630            target_tier: Tier::BaselineJit,
631            blob_hash: None,
632            osr: true,
633            loop_header_ip: Some(0),
634            feedback: None,
635            callee_feedback: std::collections::HashMap::new(),
636        };
637
638        let result = backend.compile(&request, &program);
639        assert!(result.error.is_some());
640        assert!(result.error.unwrap().contains("unsupported opcode"));
641        assert_eq!(result.loop_header_ip, Some(0)); // For blacklisting
642    }
643}