use shape_vm::bytecode::BytecodeProgram;
use shape_vm::tier::{CompilationBackend, CompilationRequest, CompilationResult, Tier};
use shape_vm::type_tracking::FrameDescriptor;
use crate::compiler::JITCompiler;
use crate::context::JITConfig;
use crate::loop_analysis;
use crate::osr_compiler;
pub struct JitCompilationBackend {
jit: JITCompiler,
}
impl JitCompilationBackend {
pub fn new() -> Result<Self, crate::error::JitError> {
Ok(Self {
jit: JITCompiler::new(JITConfig::default())?,
})
}
pub fn with_config(config: JITConfig) -> Result<Self, crate::error::JitError> {
Ok(Self {
jit: JITCompiler::new(config)?,
})
}
fn compile_osr(
&mut self,
request: &CompilationRequest,
program: &BytecodeProgram,
) -> CompilationResult {
let func_id = request.function_id;
let loop_header_ip = request.loop_header_ip;
let function = match program.functions.get(func_id as usize) {
Some(f) => f,
None => {
return CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some(format!("Function {} not found in program", func_id)),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip,
shape_guards: Vec::new(),
};
}
};
let entry = function.entry_point;
let end = find_function_end(program, func_id as usize);
if entry >= program.instructions.len() || end > program.instructions.len() {
return CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some(format!(
"Function {} instruction range [{}, {}) out of bounds",
func_id, entry, end
)),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip,
shape_guards: Vec::new(),
};
}
let func_instructions = &program.instructions[entry..end];
let sub_program = build_sub_program(program, entry, end);
let loop_infos = loop_analysis::analyze_loops(&sub_program);
let target_local_ip = match loop_header_ip {
Some(ip) => {
if ip < entry {
return CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some(format!(
"OSR loop header IP {} is before function entry {}",
ip, entry
)),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip: Some(ip),
shape_guards: Vec::new(),
};
}
ip - entry
}
None => {
return CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some("OSR request without loop_header_ip".to_string()),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip: None,
shape_guards: Vec::new(),
};
}
};
let loop_info = match loop_infos.get(&target_local_ip) {
Some(li) => li,
None => {
return CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some(format!(
"No loop found at local IP {} (global IP {:?})",
target_local_ip, loop_header_ip
)),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip,
shape_guards: Vec::new(),
};
}
};
let default_frame = FrameDescriptor::default();
let frame_descriptor = function.frame_descriptor.as_ref().unwrap_or(&default_frame);
match osr_compiler::compile_osr_loop(
&mut self.jit,
function,
func_instructions,
loop_info,
frame_descriptor,
) {
Ok(osr_result) => {
let mut entry_point = osr_result.entry_point;
entry_point.bytecode_ip += entry;
entry_point.exit_ip += entry;
CompilationResult {
function_id: func_id,
compiled_tier: Tier::BaselineJit,
native_code: Some(osr_result.native_code),
error: None,
osr_entry: Some(entry_point),
deopt_points: osr_result.deopt_points,
loop_header_ip,
shape_guards: Vec::new(),
}
}
Err(e) => CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some(e),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip,
shape_guards: Vec::new(),
},
}
}
}
unsafe impl Send for JitCompilationBackend {}
impl JitCompilationBackend {
fn compile_function(
&mut self,
request: &CompilationRequest,
program: &BytecodeProgram,
) -> CompilationResult {
let func_id = request.function_id;
if let Some(fv) = request.feedback.clone() {
return match self.jit.compile_optimizing_function(
program,
func_id as usize,
fv,
&request.callee_feedback,
) {
Ok((code_ptr, deopt_points, shape_guards)) => CompilationResult {
function_id: func_id,
compiled_tier: request.target_tier,
native_code: Some(code_ptr),
error: None,
osr_entry: None,
deopt_points,
loop_header_ip: None,
shape_guards,
},
Err(e) => CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some(e),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip: None,
shape_guards: Vec::new(),
},
};
}
match self
.jit
.compile_single_function(program, func_id as usize, None)
{
Ok((code_ptr, deopt_points, shape_guards)) => CompilationResult {
function_id: func_id,
compiled_tier: request.target_tier,
native_code: Some(code_ptr),
error: None,
osr_entry: None,
deopt_points,
loop_header_ip: None,
shape_guards,
},
Err(e) => CompilationResult {
function_id: func_id,
compiled_tier: Tier::Interpreted,
native_code: None,
error: Some(e),
osr_entry: None,
deopt_points: Vec::new(),
loop_header_ip: None,
shape_guards: Vec::new(),
},
}
}
}
impl CompilationBackend for JitCompilationBackend {
fn compile(
&mut self,
request: &CompilationRequest,
program: &BytecodeProgram,
) -> CompilationResult {
if request.osr {
self.compile_osr(request, program)
} else {
self.compile_function(request, program)
}
}
}
fn find_function_end(program: &BytecodeProgram, func_index: usize) -> usize {
let func = &program.functions[func_index];
func.entry_point + func.body_length
}
fn build_sub_program(program: &BytecodeProgram, start: usize, end: usize) -> BytecodeProgram {
BytecodeProgram {
instructions: program.instructions[start..end].to_vec(),
constants: program.constants.clone(),
strings: program.strings.clone(),
functions: vec![],
debug_info: Default::default(),
data_schema: None,
module_binding_names: vec![],
top_level_locals_count: 0,
top_level_local_storage_hints: vec![],
type_schema_registry: Default::default(),
module_binding_storage_hints: vec![],
function_local_storage_hints: vec![],
compiled_annotations: Default::default(),
trait_method_symbols: Default::default(),
expanded_function_defs: Default::default(),
string_index: Default::default(),
foreign_functions: Vec::new(),
native_struct_layouts: vec![],
content_addressed: None,
top_level_mir: None,
function_blob_hashes: vec![],
top_level_frame: None,
top_level_local_concrete_types: vec![],
function_local_concrete_types: vec![],
function_return_concrete_types: vec![],
monomorphized_method_call_sites: Default::default(),
value_call_return_concrete_types: Default::default(),
operator_trait_dispatch_sites: Default::default(),
monomorphization_keys: vec![],
closure_function_layouts: program.closure_function_layouts.clone(),
trait_vtables: program.trait_vtables.clone(),
has_imported_const_inline: program.has_imported_const_inline,
has_w17_marshal_residual: program.has_w17_marshal_residual,
}
}
#[cfg(test)]
mod tests {
use super::*;
use shape_vm::bytecode::*;
use shape_vm::type_tracking::{FrameDescriptor, NativeKind};
fn make_instr(opcode: OpCode, operand: Option<Operand>) -> Instruction {
Instruction { opcode, operand }
}
#[test]
#[ignore = "v2: Tier 1 whole-function JIT (compile_single_function) deprecated; tests dead path"]
fn test_backend_compiles_whole_function() {
let mut backend = JitCompilationBackend::new().unwrap();
let instrs = vec![
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), ];
let func = Function {
name: "add_two".to_string(),
arity: 2,
param_names: vec![],
locals_count: 2,
entry_point: 0,
body_length: 4,
is_closure: false,
captures_count: 0,
is_async: false,
ref_params: vec![],
mir_data: None,
ref_mutates: vec![],
mutable_captures: vec![],
frame_descriptor: Some(FrameDescriptor::from_slots(vec![
NativeKind::Int64, NativeKind::Int64, ])),
osr_entry_points: vec![],
};
let program = BytecodeProgram {
instructions: instrs,
constants: vec![],
strings: vec![],
functions: vec![func],
debug_info: Default::default(),
data_schema: None,
module_binding_names: vec![],
top_level_locals_count: 0,
top_level_local_storage_hints: vec![],
type_schema_registry: Default::default(),
module_binding_storage_hints: vec![],
function_local_storage_hints: vec![],
compiled_annotations: Default::default(),
trait_method_symbols: Default::default(),
expanded_function_defs: Default::default(),
string_index: Default::default(),
foreign_functions: Vec::new(),
native_struct_layouts: vec![],
content_addressed: None,
top_level_mir: None,
function_blob_hashes: vec![],
top_level_frame: None,
..Default::default()
};
let request = CompilationRequest {
function_id: 0,
target_tier: Tier::BaselineJit,
blob_hash: None,
osr: false,
loop_header_ip: None,
feedback: None,
callee_feedback: std::collections::HashMap::new(),
};
let result = backend.compile(&request, &program);
assert!(
result.error.is_none(),
"Expected successful whole-function compilation, got: {:?}",
result.error
);
assert!(result.native_code.is_some());
assert_eq!(result.compiled_tier, Tier::BaselineJit);
assert!(result.osr_entry.is_none()); }
#[test]
#[ignore = "v2: Tier 1 whole-function JIT deprecated; test asserts on error message that no longer matches"]
fn test_backend_whole_function_invalid_id() {
let mut backend = JitCompilationBackend::new().unwrap();
let program = BytecodeProgram {
instructions: vec![make_instr(OpCode::Halt, None)],
constants: vec![],
strings: vec![],
functions: vec![], debug_info: Default::default(),
data_schema: None,
module_binding_names: vec![],
top_level_locals_count: 0,
top_level_local_storage_hints: vec![],
type_schema_registry: Default::default(),
module_binding_storage_hints: vec![],
function_local_storage_hints: vec![],
compiled_annotations: Default::default(),
trait_method_symbols: Default::default(),
expanded_function_defs: Default::default(),
string_index: Default::default(),
foreign_functions: Vec::new(),
native_struct_layouts: vec![],
content_addressed: None,
top_level_mir: None,
function_blob_hashes: vec![],
top_level_frame: None,
..Default::default()
};
let request = CompilationRequest {
function_id: 99,
target_tier: Tier::BaselineJit,
blob_hash: None,
osr: false,
loop_header_ip: None,
feedback: None,
callee_feedback: std::collections::HashMap::new(),
};
let result = backend.compile(&request, &program);
assert!(result.error.is_some());
assert!(result.error.unwrap().contains("not found"));
}
#[test]
fn test_backend_osr_compiles_simple_loop() {
let mut backend = JitCompilationBackend::new().unwrap();
let instrs = vec![
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), ];
let func = Function {
name: "test_loop".to_string(),
arity: 0,
param_names: vec![],
locals_count: 3,
entry_point: 0,
body_length: 15,
is_closure: false,
captures_count: 0,
is_async: false,
ref_params: vec![],
mir_data: None,
ref_mutates: vec![],
mutable_captures: vec![],
frame_descriptor: Some(FrameDescriptor::from_slots(vec![
NativeKind::Int64, NativeKind::Int64, NativeKind::Int64, ])),
osr_entry_points: vec![],
};
let program = BytecodeProgram {
instructions: instrs,
constants: vec![Constant::Int(1)],
strings: vec![],
functions: vec![func],
debug_info: Default::default(),
data_schema: None,
module_binding_names: vec![],
top_level_locals_count: 0,
top_level_local_storage_hints: vec![],
type_schema_registry: Default::default(),
module_binding_storage_hints: vec![],
function_local_storage_hints: vec![],
compiled_annotations: Default::default(),
trait_method_symbols: Default::default(),
expanded_function_defs: Default::default(),
string_index: Default::default(),
foreign_functions: Vec::new(),
native_struct_layouts: vec![],
content_addressed: None,
top_level_mir: None,
function_blob_hashes: vec![],
top_level_frame: None,
..Default::default()
};
let request = CompilationRequest {
function_id: 0,
target_tier: Tier::BaselineJit,
blob_hash: None,
osr: true,
loop_header_ip: Some(0), feedback: None,
callee_feedback: std::collections::HashMap::new(),
};
let result = backend.compile(&request, &program);
assert!(
result.error.is_none(),
"Expected successful compilation, got: {:?}",
result.error
);
assert!(result.native_code.is_some());
assert!(result.osr_entry.is_some());
assert_eq!(result.compiled_tier, Tier::BaselineJit);
let entry = result.osr_entry.unwrap();
assert_eq!(entry.bytecode_ip, 0);
assert!(entry.live_locals.contains(&0)); assert!(entry.live_locals.contains(&1)); assert!(entry.live_locals.contains(&2)); }
#[test]
fn test_backend_osr_blacklists_unsupported_loop() {
let mut backend = JitCompilationBackend::new().unwrap();
let instrs = vec![
make_instr(OpCode::LoopStart, None),
make_instr(OpCode::LoadLocal, Some(Operand::Local(0))),
make_instr(OpCode::CallMethod, None), make_instr(OpCode::Pop, None),
make_instr(OpCode::LoopEnd, None),
make_instr(OpCode::Halt, None),
];
let func = Function {
name: "unsupported_loop".to_string(),
arity: 0,
param_names: vec![],
locals_count: 1,
entry_point: 0,
body_length: 6,
is_closure: false,
captures_count: 0,
is_async: false,
ref_params: vec![],
mir_data: None,
ref_mutates: vec![],
mutable_captures: vec![],
frame_descriptor: Some(FrameDescriptor::from_slots(vec![NativeKind::Bool])),
osr_entry_points: vec![],
};
let program = BytecodeProgram {
instructions: instrs,
constants: vec![],
strings: vec![],
functions: vec![func],
debug_info: Default::default(),
data_schema: None,
module_binding_names: vec![],
top_level_locals_count: 0,
top_level_local_storage_hints: vec![],
type_schema_registry: Default::default(),
module_binding_storage_hints: vec![],
function_local_storage_hints: vec![],
compiled_annotations: Default::default(),
trait_method_symbols: Default::default(),
expanded_function_defs: Default::default(),
string_index: Default::default(),
foreign_functions: Vec::new(),
native_struct_layouts: vec![],
content_addressed: None,
top_level_mir: None,
function_blob_hashes: vec![],
top_level_frame: None,
..Default::default()
};
let request = CompilationRequest {
function_id: 0,
target_tier: Tier::BaselineJit,
blob_hash: None,
osr: true,
loop_header_ip: Some(0),
feedback: None,
callee_feedback: std::collections::HashMap::new(),
};
let result = backend.compile(&request, &program);
assert!(result.error.is_some());
assert!(result.error.unwrap().contains("unsupported opcode"));
assert_eq!(result.loop_header_ip, Some(0)); }
}