1use 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
15pub struct JitCompilationBackend {
20 jit: JITCompiler,
21}
22
23impl JitCompilationBackend {
24 pub fn new() -> Result<Self, crate::error::JitError> {
26 Ok(Self {
27 jit: JITCompiler::new(JITConfig::default())?,
28 })
29 }
30
31 pub fn with_config(config: JITConfig) -> Result<Self, crate::error::JitError> {
33 Ok(Self {
34 jit: JITCompiler::new(config)?,
35 })
36 }
37
38 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 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 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 let sub_program = build_sub_program(program, entry, end);
86 let loop_infos = loop_analysis::analyze_loops(&sub_program);
87
88 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 let default_frame = FrameDescriptor::default();
144 let frame_descriptor = function.frame_descriptor.as_ref().unwrap_or(&default_frame);
145
146 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 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
185unsafe impl Send for JitCompilationBackend {}
190
191impl JitCompilationBackend {
192 fn compile_function(
202 &mut self,
203 request: &CompilationRequest,
204 program: &BytecodeProgram,
205 ) -> CompilationResult {
206 let func_id = request.function_id;
207
208 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 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
282fn 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
291fn 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 let instrs = vec![
350 make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::AddInt, None), make_instr(OpCode::ReturnValue, None), make_instr(OpCode::Halt, None), ];
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, NativeKind::Int64, ])),
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()); }
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![], 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 let instrs = vec![
476 make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::LtInt, None), make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(7))), make_instr(OpCode::LoadLocal, Some(Operand::Local(2))), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(2))), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(0))), make_instr(OpCode::LoopEnd, None), make_instr(OpCode::ReturnValue, None), ];
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, NativeKind::Int64, NativeKind::Int64, ])),
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), 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)); assert!(entry.live_locals.contains(&1)); assert!(entry.live_locals.contains(&2)); }
567
568 #[test]
569 fn test_backend_osr_blacklists_unsupported_loop() {
570 let mut backend = JitCompilationBackend::new().unwrap();
571
572 let instrs = vec![
574 make_instr(OpCode::LoopStart, None),
575 make_instr(OpCode::LoadLocal, Some(Operand::Local(0))),
576 make_instr(OpCode::CallMethod, None), 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 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)); }
643}