use std::collections::HashMap;
use rayon::prelude::*;
use crate::bytecode::{
BytecodeProgram, Constant, DebugInfo, Function, FunctionBlob, FunctionHash, Instruction,
LinkedFunction, LinkedProgram, Operand, Program, SourceMap,
};
use shape_abi_v1::PermissionSet;
use shape_value::{FunctionId, StringId};
#[derive(Debug, thiserror::Error)]
pub enum LinkError {
#[error("Missing function blob: {0}")]
MissingBlob(FunctionHash),
#[error("Circular dependency detected")]
CircularDependency,
#[error("Constant pool overflow: {0} constants exceeds u16 max")]
ConstantPoolOverflow(usize),
#[error("String pool overflow: {0} strings exceeds u32 max")]
StringPoolOverflow(usize),
}
fn topo_sort(program: &Program) -> Result<Vec<FunctionHash>, LinkError> {
let mut state: HashMap<FunctionHash, u8> = HashMap::new();
let mut order: Vec<FunctionHash> = Vec::with_capacity(program.function_store.len());
fn visit(
hash: FunctionHash,
program: &Program,
state: &mut HashMap<FunctionHash, u8>,
order: &mut Vec<FunctionHash>,
) -> Result<(), LinkError> {
match state.get(&hash).copied().unwrap_or(0) {
2 => return Ok(()), 1 => return Err(LinkError::CircularDependency),
_ => {}
}
state.insert(hash, 1);
let blob = program
.function_store
.get(&hash)
.ok_or(LinkError::MissingBlob(hash))?;
for dep in &blob.dependencies {
if *dep == FunctionHash::ZERO {
continue;
}
visit(*dep, program, state, order)?;
}
state.insert(hash, 2); order.push(hash);
Ok(())
}
visit(program.entry, program, &mut state, &mut order)?;
let remaining: Vec<FunctionHash> = program
.function_store
.keys()
.copied()
.filter(|h| state.get(h).copied().unwrap_or(0) != 2)
.collect();
for hash in remaining {
visit(hash, program, &mut state, &mut order)?;
}
Ok(order)
}
fn remap_fid(
dep_idx: u16,
blob: &FunctionBlob,
current_function_id: usize,
hash_to_id: &HashMap<FunctionHash, usize>,
name_to_id: &HashMap<&str, usize>,
) -> u16 {
if let Some(dep_hash) = blob.dependencies.get(dep_idx as usize) {
if *dep_hash == FunctionHash::ZERO {
if let Some(callee_name) = blob.callee_names.get(dep_idx as usize) {
if callee_name != &blob.name {
if let Some(target_id) = name_to_id.get(callee_name.as_str()) {
*target_id as u16
} else {
current_function_id as u16
}
} else {
current_function_id as u16
}
} else {
current_function_id as u16
}
} else {
hash_to_id[dep_hash] as u16
}
} else {
dep_idx
}
}
fn remap_operand(
operand: Operand,
const_base: usize,
string_base: usize,
blob: &FunctionBlob,
current_function_id: usize,
hash_to_id: &HashMap<FunctionHash, usize>,
name_to_id: &HashMap<&str, usize>,
) -> Operand {
match operand {
Operand::Const(i) => Operand::Const((const_base + i as usize) as u16),
Operand::Property(i) => Operand::Property((string_base + i as usize) as u16),
Operand::Name(StringId(i)) => Operand::Name(StringId((string_base + i as usize) as u32)),
Operand::Function(FunctionId(dep_idx)) => {
Operand::Function(FunctionId(remap_fid(
dep_idx,
blob,
current_function_id,
hash_to_id,
name_to_id,
)))
}
Operand::ClosureAlloc { fid: FunctionId(dep_idx), escapes } => {
Operand::ClosureAlloc {
fid: FunctionId(remap_fid(
dep_idx,
blob,
current_function_id,
hash_to_id,
name_to_id,
)),
escapes,
}
}
Operand::TypedMethodCall {
method_id,
arg_count,
string_id,
receiver_type_tag,
} => Operand::TypedMethodCall {
method_id,
arg_count,
string_id: (string_base + string_id as usize) as u16,
receiver_type_tag,
},
Operand::Offset(_)
| Operand::Local(_)
| Operand::ModuleBinding(_)
| Operand::Builtin(_)
| Operand::Count(_)
| Operand::ColumnIndex(_)
| Operand::TypedField { .. }
| Operand::TypedObjectAlloc { .. }
| Operand::TypedMerge { .. }
| Operand::ColumnAccess { .. }
| Operand::ForeignFunction(_)
| Operand::MatrixDims { .. }
| Operand::Width(_)
| Operand::TypedLocal(_, _)
| Operand::TypedModuleBinding(_, _)
| Operand::FieldOffset(_) => operand,
}
}
fn remap_constant(
constant: &Constant,
blob: &FunctionBlob,
current_function_id: usize,
hash_to_id: &HashMap<FunctionHash, usize>,
name_to_id: &HashMap<&str, usize>,
) -> Constant {
match constant {
Constant::Function(dep_idx) => {
let dep_idx = *dep_idx as usize;
if dep_idx < blob.dependencies.len() {
let dep_hash = blob.dependencies[dep_idx];
if dep_hash == FunctionHash::ZERO {
if let Some(callee_name) = blob.callee_names.get(dep_idx) {
if callee_name != &blob.name {
if let Some(target_id) = name_to_id.get(callee_name.as_str()) {
Constant::Function(*target_id as u16)
} else {
Constant::Function(current_function_id as u16)
}
} else {
Constant::Function(current_function_id as u16)
}
} else {
Constant::Function(current_function_id as u16)
}
} else {
let linked_id = hash_to_id[&dep_hash];
Constant::Function(linked_id as u16)
}
} else {
constant.clone()
}
}
other => other.clone(),
}
}
const PARALLEL_THRESHOLD: usize = 50;
struct BlobOffsets {
instruction_base: usize,
const_base: usize,
string_base: usize,
}
pub fn link(program: &Program) -> Result<LinkedProgram, LinkError> {
let sorted = topo_sort(program)?;
let blobs: Vec<&FunctionBlob> = sorted
.iter()
.map(|h| {
program
.function_store
.get(h)
.ok_or(LinkError::MissingBlob(*h))
})
.collect::<Result<Vec<_>, _>>()?;
let mut offsets: Vec<BlobOffsets> = Vec::with_capacity(blobs.len());
let mut hash_to_id: HashMap<FunctionHash, usize> = HashMap::with_capacity(blobs.len());
let mut name_to_id: HashMap<&str, usize> = HashMap::with_capacity(blobs.len());
let mut total_instructions: usize = 0;
let mut total_constants: usize = 0;
let mut total_strings: usize = 0;
for (i, blob) in blobs.iter().enumerate() {
offsets.push(BlobOffsets {
instruction_base: total_instructions,
const_base: total_constants,
string_base: total_strings,
});
hash_to_id.insert(blob.content_hash, i);
name_to_id.insert(&blob.name, i);
total_instructions += blob.instructions.len();
total_constants += blob.constants.len();
total_strings += blob.strings.len();
}
if total_constants > u16::MAX as usize + 1 {
return Err(LinkError::ConstantPoolOverflow(total_constants));
}
if total_strings > u32::MAX as usize + 1 {
return Err(LinkError::StringPoolOverflow(total_strings));
}
let total_required_permissions = blobs.iter().fold(PermissionSet::pure(), |acc, blob| {
acc.union(&blob.required_permissions)
});
let use_parallel = blobs.len() > PARALLEL_THRESHOLD;
let mut instructions: Vec<Instruction> = Vec::with_capacity(total_instructions);
let mut constants: Vec<Constant> = Vec::with_capacity(total_constants);
let mut strings: Vec<String> = Vec::with_capacity(total_strings);
if use_parallel {
struct BlobResult {
instructions: Vec<Instruction>,
constants: Vec<Constant>,
strings: Vec<String>,
source_map: Vec<(usize, u16, u32)>,
}
let results: Vec<BlobResult> = blobs
.par_iter()
.zip(offsets.par_iter())
.enumerate()
.map(|(function_id, (blob, off))| {
let remapped_instrs: Vec<Instruction> = blob
.instructions
.iter()
.map(|instr| {
let remapped_operand = instr.operand.map(|op| {
remap_operand(
op,
off.const_base,
off.string_base,
blob,
function_id,
&hash_to_id,
&name_to_id,
)
});
Instruction {
opcode: instr.opcode,
operand: remapped_operand,
}
})
.collect();
let remapped_consts: Vec<Constant> = blob
.constants
.iter()
.map(|c| remap_constant(c, blob, function_id, &hash_to_id, &name_to_id))
.collect();
let cloned_strings: Vec<String> = blob.strings.clone();
let source_entries: Vec<(usize, u16, u32)> = blob
.source_map
.iter()
.map(|&(local_offset, file_id, line)| {
(off.instruction_base + local_offset, file_id as u16, line)
})
.collect();
BlobResult {
instructions: remapped_instrs,
constants: remapped_consts,
strings: cloned_strings,
source_map: source_entries,
}
})
.collect();
let mut merged_line_numbers: Vec<(usize, u16, u32)> = Vec::new();
for result in results {
instructions.extend(result.instructions);
constants.extend(result.constants);
strings.extend(result.strings);
merged_line_numbers.extend(result.source_map);
}
merged_line_numbers.sort_by_key(|&(offset, _, _)| offset);
let functions: Vec<LinkedFunction> = blobs
.iter()
.zip(offsets.iter())
.map(|(blob, off)| LinkedFunction {
blob_hash: blob.content_hash,
entry_point: off.instruction_base,
body_length: blob.instructions.len(),
name: blob.name.clone(),
arity: blob.arity,
param_names: blob.param_names.clone(),
locals_count: blob.locals_count,
is_closure: blob.is_closure,
captures_count: blob.captures_count,
is_async: blob.is_async,
ref_params: blob.ref_params.clone(),
ref_mutates: blob.ref_mutates.clone(),
mutable_captures: blob.mutable_captures.clone(),
frame_descriptor: blob.frame_descriptor.clone(),
})
.collect();
let debug_info = DebugInfo {
source_map: SourceMap {
files: program.debug_info.source_map.files.clone(),
source_texts: program.debug_info.source_map.source_texts.clone(),
},
line_numbers: merged_line_numbers,
variable_names: program.debug_info.variable_names.clone(),
source_text: String::new(),
};
return Ok(LinkedProgram {
entry: program.entry,
instructions,
constants,
strings,
functions,
hash_to_id,
debug_info,
data_schema: program.data_schema.clone(),
module_binding_names: program.module_binding_names.clone(),
top_level_locals_count: program.top_level_locals_count,
top_level_local_storage_hints: program.top_level_local_storage_hints.clone(),
type_schema_registry: program.type_schema_registry.clone(),
module_binding_storage_hints: program.module_binding_storage_hints.clone(),
function_local_storage_hints: program.function_local_storage_hints.clone(),
top_level_frame: program.top_level_frame.clone(),
top_level_local_concrete_types: program.top_level_local_concrete_types.clone(),
function_local_concrete_types: program.function_local_concrete_types.clone(),
function_return_concrete_types: program.function_return_concrete_types.clone(),
monomorphized_method_call_sites:
program.monomorphized_method_call_sites.clone(),
value_call_return_concrete_types:
program.value_call_return_concrete_types.clone(),
operator_trait_dispatch_sites:
program.operator_trait_dispatch_sites.clone(),
trait_method_symbols: program.trait_method_symbols.clone(),
foreign_functions: program.foreign_functions.clone(),
native_struct_layouts: program.native_struct_layouts.clone(),
total_required_permissions: total_required_permissions.clone(),
closure_function_layouts: remap_closure_function_layouts(
program,
&blobs,
),
trait_vtables: program.trait_vtables.clone(),
has_imported_const_inline: program.has_imported_const_inline,
has_w17_marshal_residual: program.has_w17_marshal_residual,
});
}
let mut merged_line_numbers: Vec<(usize, u16, u32)> = Vec::new();
for (function_id, (blob, off)) in blobs.iter().zip(offsets.iter()).enumerate() {
for instr in &blob.instructions {
let remapped_operand = instr.operand.map(|op| {
remap_operand(
op,
off.const_base,
off.string_base,
blob,
function_id,
&hash_to_id,
&name_to_id,
)
});
instructions.push(Instruction {
opcode: instr.opcode,
operand: remapped_operand,
});
}
for c in &blob.constants {
constants.push(remap_constant(
c,
blob,
function_id,
&hash_to_id,
&name_to_id,
));
}
strings.extend(blob.strings.iter().cloned());
for &(local_offset, file_id, line) in &blob.source_map {
let global_offset = off.instruction_base + local_offset;
merged_line_numbers.push((global_offset, file_id as u16, line));
}
}
merged_line_numbers.sort_by_key(|&(offset, _, _)| offset);
let functions: Vec<LinkedFunction> = blobs
.iter()
.zip(offsets.iter())
.map(|(blob, off)| LinkedFunction {
blob_hash: blob.content_hash,
entry_point: off.instruction_base,
body_length: blob.instructions.len(),
name: blob.name.clone(),
arity: blob.arity,
param_names: blob.param_names.clone(),
locals_count: blob.locals_count,
is_closure: blob.is_closure,
captures_count: blob.captures_count,
is_async: blob.is_async,
ref_params: blob.ref_params.clone(),
ref_mutates: blob.ref_mutates.clone(),
mutable_captures: blob.mutable_captures.clone(),
frame_descriptor: blob.frame_descriptor.clone(),
})
.collect();
let debug_info = DebugInfo {
source_map: SourceMap {
files: program.debug_info.source_map.files.clone(),
source_texts: program.debug_info.source_map.source_texts.clone(),
},
line_numbers: merged_line_numbers,
variable_names: program.debug_info.variable_names.clone(),
source_text: String::new(),
};
Ok(LinkedProgram {
entry: program.entry,
instructions,
constants,
strings,
functions,
hash_to_id,
debug_info,
data_schema: program.data_schema.clone(),
module_binding_names: program.module_binding_names.clone(),
top_level_locals_count: program.top_level_locals_count,
top_level_local_storage_hints: program.top_level_local_storage_hints.clone(),
type_schema_registry: program.type_schema_registry.clone(),
module_binding_storage_hints: program.module_binding_storage_hints.clone(),
function_local_storage_hints: program.function_local_storage_hints.clone(),
top_level_frame: program.top_level_frame.clone(),
top_level_local_concrete_types: program.top_level_local_concrete_types.clone(),
function_local_concrete_types: program.function_local_concrete_types.clone(),
function_return_concrete_types: program.function_return_concrete_types.clone(),
monomorphized_method_call_sites:
program.monomorphized_method_call_sites.clone(),
value_call_return_concrete_types:
program.value_call_return_concrete_types.clone(),
operator_trait_dispatch_sites:
program.operator_trait_dispatch_sites.clone(),
trait_method_symbols: program.trait_method_symbols.clone(),
foreign_functions: program.foreign_functions.clone(),
native_struct_layouts: program.native_struct_layouts.clone(),
total_required_permissions,
closure_function_layouts: remap_closure_function_layouts(program, &blobs),
trait_vtables: program.trait_vtables.clone(),
has_imported_const_inline: program.has_imported_const_inline,
has_w17_marshal_residual: program.has_w17_marshal_residual,
})
}
fn remap_closure_function_layouts(
program: &Program,
blobs: &[&FunctionBlob],
) -> Vec<Option<std::sync::Arc<shape_value::v2::closure_layout::ClosureLayout>>> {
if program.closure_function_layouts_by_name.is_empty() {
return Vec::new();
}
blobs
.iter()
.map(|blob| {
program
.closure_function_layouts_by_name
.get(&blob.name)
.cloned()
})
.collect()
}
pub fn linked_to_bytecode_program(linked: &LinkedProgram) -> BytecodeProgram {
let functions: Vec<Function> = linked
.functions
.iter()
.map(|lf| Function {
name: lf.name.clone(),
arity: lf.arity,
param_names: lf.param_names.clone(),
locals_count: lf.locals_count,
entry_point: lf.entry_point,
body_length: lf.body_length,
is_closure: lf.is_closure,
captures_count: lf.captures_count,
is_async: lf.is_async,
ref_params: lf.ref_params.clone(),
ref_mutates: lf.ref_mutates.clone(),
mutable_captures: lf.mutable_captures.clone(),
frame_descriptor: lf.frame_descriptor.clone(),
osr_entry_points: Vec::new(),
mir_data: None,
})
.collect();
BytecodeProgram {
instructions: linked.instructions.clone(),
constants: linked.constants.clone(),
strings: linked.strings.clone(),
functions,
has_imported_const_inline: linked.has_imported_const_inline,
has_w17_marshal_residual: linked.has_w17_marshal_residual,
debug_info: linked.debug_info.clone(),
data_schema: linked.data_schema.clone(),
module_binding_names: linked.module_binding_names.clone(),
top_level_locals_count: linked.top_level_locals_count,
top_level_local_storage_hints: linked.top_level_local_storage_hints.clone(),
type_schema_registry: linked.type_schema_registry.clone(),
module_binding_storage_hints: linked.module_binding_storage_hints.clone(),
function_local_storage_hints: linked.function_local_storage_hints.clone(),
top_level_frame: linked.top_level_frame.clone(),
top_level_local_concrete_types: linked.top_level_local_concrete_types.clone(),
function_local_concrete_types: linked.function_local_concrete_types.clone(),
function_return_concrete_types: linked.function_return_concrete_types.clone(),
monomorphized_method_call_sites:
linked.monomorphized_method_call_sites.clone(),
value_call_return_concrete_types:
linked.value_call_return_concrete_types.clone(),
operator_trait_dispatch_sites:
linked.operator_trait_dispatch_sites.clone(),
top_level_mir: None,
compiled_annotations: HashMap::new(),
trait_method_symbols: linked.trait_method_symbols.clone(),
expanded_function_defs: HashMap::new(),
string_index: HashMap::new(),
foreign_functions: linked.foreign_functions.clone(),
native_struct_layouts: linked.native_struct_layouts.clone(),
content_addressed: None,
function_blob_hashes: linked
.functions
.iter()
.map(|lf| {
if lf.blob_hash == FunctionHash::ZERO {
None
} else {
Some(lf.blob_hash)
}
})
.collect(),
monomorphization_keys: Vec::new(),
closure_function_layouts: linked.closure_function_layouts.clone(),
trait_vtables: linked.trait_vtables.clone(),
}
}
#[cfg(test)]
#[path = "linker_tests.rs"]
mod tests;