use anyhow::Context;
use itertools::Itertools;
use object::{
Endianness, Object, ObjectSection, ObjectSymbol, SymbolFlags, SymbolKind, SymbolScope,
macho::{self},
read::File,
write::{MachOBuildVersion, SectionId, StandardSection, Symbol, SymbolId, SymbolSection},
};
use rayon::prelude::{IntoParallelRefIterator, ParallelIterator};
use std::{
collections::{BTreeMap, HashMap, HashSet},
io::Read,
ops::Range,
path::Path,
path::PathBuf,
sync::{Arc, RwLock},
};
use subsecond_types::{AddressMap, JumpTable};
use target_lexicon::{Architecture, OperatingSystem, PointerWidth, Triple};
use thiserror::Error;
use walrus::{
ConstExpr, DataKind, ElementItems, ElementKind, FunctionBuilder, FunctionId, FunctionKind,
ImportKind, Module, ModuleConfig, TableId,
};
use wasmparser::{
BinaryReader, BinaryReaderError, Linking, LinkingSectionReader, Payload, SymbolInfo,
};
type Result<T, E = PatchError> = std::result::Result<T, E>;
#[derive(Debug, Error)]
pub enum PatchError {
#[error("Failed to read file: {0}")]
ReadFs(#[from] std::io::Error),
#[error(
"No debug symbols in the patch output. Check your profile's `opt-level` and debug symbols config."
)]
MissingSymbols,
#[error("Failed to parse wasm section: {0}")]
ParseSection(#[from] wasmparser::BinaryReaderError),
#[error("Failed to parse object file, {0}")]
ParseObjectFile(#[from] object::read::Error),
#[error("Failed to write object file: {0}")]
WriteObjectFIle(#[from] object::write::Error),
#[error("Failed to emit module: {0}")]
RuntimeError(#[from] anyhow::Error),
#[error("Failed to read module's PDB file: {0}")]
PdbLoadError(#[from] pdb::Error),
#[error("{0}")]
InvalidModule(String),
#[error("Unsupported platform: {0}")]
UnsupportedPlatform(String),
}
#[derive(Default)]
pub struct HotpatchModuleCache {
pub path: PathBuf,
pub symbol_ifunc_map: HashMap<String, i32>,
pub old_wasm: Module,
pub old_bytes: Vec<u8>,
pub old_exports: HashSet<String>,
pub old_imports: HashSet<String>,
pub symbol_table: HashMap<String, CachedSymbol>,
pub tls_init_data: Vec<u8>,
pub tls_init_sizes: HashMap<String, (u64, u64)>,
}
pub struct CachedSymbol {
pub address: u64,
pub kind: SymbolKind,
pub is_undefined: bool,
pub is_weak: bool,
pub size: u64,
pub flags: SymbolFlags<SectionId, SymbolId>,
}
impl PartialEq for HotpatchModuleCache {
fn eq(&self, other: &Self) -> bool {
self.path == other.path
}
}
impl std::fmt::Debug for HotpatchModuleCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HotpatchModuleCache")
.field("_path", &self.path)
.finish()
}
}
impl HotpatchModuleCache {
pub fn new(original: &Path, triple: &Triple) -> Result<Self> {
let cache = match triple.operating_system {
OperatingSystem::Windows => {
use pdb::FallibleIterator;
let old_pdb_file = original.with_extension("pdb");
let old_pdb_file_handle = std::fs::File::open(old_pdb_file)?;
let mut pdb_file = pdb::PDB::open(old_pdb_file_handle)?;
let global_symbols = pdb_file.global_symbols()?;
let address_map = pdb_file.address_map()?;
let mut symbol_table = HashMap::new();
let mut symbols = global_symbols.iter();
while let Ok(Some(symbol)) = symbols.next() {
match symbol.parse() {
Ok(pdb::SymbolData::Public(data)) => {
let rva = data.offset.to_rva(&address_map);
let is_undefined = rva.is_none();
let rva = rva.unwrap_or_default();
symbol_table.insert(
data.name.to_string().to_string(),
CachedSymbol {
address: rva.0 as u64,
kind: if data.function {
SymbolKind::Text
} else {
SymbolKind::Data
},
is_undefined,
is_weak: false,
size: 0,
flags: SymbolFlags::None,
},
);
}
Ok(pdb::SymbolData::Data(data)) => {
let rva = data.offset.to_rva(&address_map);
let is_undefined = rva.is_none();
let rva = rva.unwrap_or_default();
symbol_table.insert(
data.name.to_string().to_string(),
CachedSymbol {
address: rva.0 as u64,
kind: SymbolKind::Data,
is_undefined,
is_weak: false,
size: 0,
flags: SymbolFlags::None,
},
);
}
_ => {}
}
}
HotpatchModuleCache {
symbol_table,
path: original.to_path_buf(),
..Default::default()
}
}
_ if triple.architecture == Architecture::Wasm32 => {
let bytes = std::fs::read(original)?;
let ParsedModule {
module, symbols, ..
} = parse_module_with_ids(&bytes)?;
if symbols.symbols.is_empty() {
return Err(PatchError::MissingSymbols);
}
let direct_name_to_ifunc = collect_func_ifuncs(&module);
let func_to_index = module
.funcs
.par_iter()
.filter_map(|f| {
let name = f.name.as_deref()?;
Some((*symbols.code_symbol_map.get(name)?, name))
})
.collect::<HashMap<usize, &str>>();
let mut symbol_ifunc_map: HashMap<String, i32> = symbols
.code_symbol_map
.par_iter()
.filter_map(|(name, idx)| {
let new_modules_unified_function = func_to_index.get(idx)?;
let offset = direct_name_to_ifunc.get(new_modules_unified_function)?;
Some((name.to_string(), *offset))
})
.collect();
for (name, offset) in &direct_name_to_ifunc {
symbol_ifunc_map
.entry((*name).to_string())
.or_insert(*offset);
}
let old_exports = module
.exports
.iter()
.map(|e| e.name.to_string())
.collect::<HashSet<_>>();
let old_imports = module
.imports
.iter()
.map(|i| i.name.to_string())
.collect::<HashSet<_>>();
HotpatchModuleCache {
path: original.to_path_buf(),
old_bytes: bytes,
symbol_ifunc_map,
old_exports,
old_imports,
old_wasm: module,
..Default::default()
}
}
_ => {
let old_bytes = std::fs::read(original)?;
let obj = File::parse(&old_bytes as &[u8])?;
let symbol_table = obj
.symbols()
.filter_map(|s| {
let flags = match s.flags() {
SymbolFlags::None => SymbolFlags::None,
SymbolFlags::Elf { st_info, st_other } => {
SymbolFlags::Elf { st_info, st_other }
}
SymbolFlags::MachO { n_desc } => SymbolFlags::MachO { n_desc },
_ => SymbolFlags::None,
};
Some((
s.name().ok()?.to_string(),
CachedSymbol {
address: s.address(),
is_undefined: s.is_undefined(),
is_weak: s.is_weak(),
kind: s.kind(),
size: s.size(),
flags,
},
))
})
.collect::<HashMap<_, _>>();
let tls_section = obj
.sections()
.find(|s| matches!(s.name(), Ok(".tdata" | "__thread_data")));
let tls_init_data = tls_section
.as_ref()
.and_then(|s| s.data().ok())
.unwrap_or(&[])
.to_vec();
let tls_data_addr = tls_section.as_ref().map(|s| s.address()).unwrap_or(0);
let tls_data_size = tls_section.as_ref().map(|s| s.size()).unwrap_or(0);
let tls_section_index = tls_section.as_ref().map(|s| s.index());
let mut tls_init_syms: Vec<(u64, String)> = Vec::new();
for sym in obj.symbols() {
if let (Some(section_idx), Ok(sname)) = (sym.section_index(), sym.name()) {
if Some(section_idx) == tls_section_index {
let offset = sym.address().saturating_sub(tls_data_addr);
tls_init_syms.push((offset, sname.to_string()));
}
}
}
tls_init_syms.sort_by_key(|(addr, _)| *addr);
tls_init_syms.dedup_by_key(|(addr, _)| *addr);
let mut tls_init_sizes: HashMap<String, (u64, u64)> = HashMap::new();
for (i, (offset, sname)) in tls_init_syms.iter().enumerate() {
let size = if i + 1 < tls_init_syms.len() {
tls_init_syms[i + 1].0 - offset
} else {
tls_data_size.saturating_sub(*offset)
};
tls_init_sizes.insert(sname.clone(), (*offset, size));
}
HotpatchModuleCache {
symbol_table,
path: original.to_path_buf(),
old_bytes,
tls_init_data,
tls_init_sizes,
..Default::default()
}
}
};
Ok(cache)
}
}
pub fn create_windows_jump_table(patch: &Path, cache: &HotpatchModuleCache) -> Result<JumpTable> {
use pdb::FallibleIterator;
let old_name_to_addr = &cache.symbol_table;
let mut new_name_to_addr = HashMap::new();
let new_pdb_file_handle = std::fs::File::open(patch.with_extension("pdb"))?;
let mut pdb_file = pdb::PDB::open(new_pdb_file_handle)?;
let symbol_table = pdb_file.global_symbols()?;
let address_map = pdb_file.address_map()?;
let mut symbol_iter = symbol_table.iter();
while let Ok(Some(symbol)) = symbol_iter.next() {
if let Ok(pdb::SymbolData::Public(data)) = symbol.parse() {
let rva = data.offset.to_rva(&address_map);
if let Some(rva) = rva {
new_name_to_addr.insert(data.name.to_string(), rva.0 as u64);
}
}
}
let mut map = AddressMap::default();
for (new_name, new_addr) in new_name_to_addr.iter() {
if let Some(old_addr) = old_name_to_addr.get(new_name.as_ref()) {
map.insert(old_addr.address, *new_addr);
}
}
let new_base_address = new_name_to_addr
.get("main")
.cloned()
.context("failed to find 'main' symbol in patch")?;
let aslr_reference = old_name_to_addr
.get("main")
.map(|s| s.address)
.context("failed to find '_main' symbol in original module")?;
Ok(JumpTable {
lib: patch.to_path_buf(),
map,
new_base_address,
aslr_reference,
ifunc_count: 0,
})
}
pub fn create_native_jump_table(
patch: &Path,
triple: &Triple,
cache: &HotpatchModuleCache,
) -> Result<JumpTable> {
let old_name_to_addr = &cache.symbol_table;
let obj2_bytes = std::fs::read(patch)?;
let obj2 = File::parse(&obj2_bytes as &[u8])?;
let mut map = AddressMap::default();
let new_syms = obj2.symbol_map();
let new_name_to_addr = new_syms
.symbols()
.par_iter()
.map(|s| (s.name(), s.address()))
.collect::<HashMap<_, _>>();
for (new_name, new_addr) in new_name_to_addr.iter() {
if let Some(old_addr) = old_name_to_addr.get(*new_name) {
map.insert(old_addr.address, *new_addr);
}
}
let sentinel = main_sentinel(triple);
let new_base_address = new_name_to_addr
.get(sentinel)
.cloned()
.context("failed to find 'main' symbol in base - are deubg symbols enabled?")?;
let aslr_reference = old_name_to_addr
.get(sentinel)
.map(|s| s.address)
.context("failed to find 'main' symbol in original module - are debug symbols enabled?")?;
Ok(JumpTable {
lib: patch.to_path_buf(),
map,
new_base_address,
aslr_reference,
ifunc_count: 0,
})
}
pub fn create_wasm_jump_table(patch: &Path, cache: &HotpatchModuleCache) -> Result<JumpTable> {
let name_to_ifunc_old = &cache.symbol_ifunc_map;
let old = &cache.old_wasm;
let old_symbols =
parse_bytes_to_data_segment(&cache.old_bytes).context("Failed to parse data segment")?;
let new_bytes = std::fs::read(patch).context("Could not read patch file")?;
let mut new = Module::from_buffer(&new_bytes)?;
let mut got_mems = vec![];
let mut got_funcs = vec![];
let mut wbg_funcs = vec![];
let mut env_funcs = vec![];
for import in new.imports.iter() {
match import.module.as_str() {
"GOT.func" => {
let Some(entry) = name_to_ifunc_old.get(import.name.as_str()).cloned() else {
return Err(PatchError::InvalidModule(format!(
"Expected to find GOT.func entry in ifunc table: {}",
import.name.as_str()
)));
};
got_funcs.push((import.id(), entry));
}
"GOT.mem" => got_mems.push(import.id()),
"env" => env_funcs.push(import.id()),
"__wbindgen_placeholder__" => wbg_funcs.push(import.id()),
m => tracing::trace!("Unknown import: {m}:{}", import.name),
}
}
for (import_id, ifunc_index) in got_funcs {
let import = new.imports.get(import_id);
let ImportKind::Global(id) = import.kind else {
return Err(PatchError::InvalidModule(format!(
"Expected GOT.func import to be a global: {}",
import.name
)));
};
new.imports.delete(import_id);
new.globals.get_mut(id).kind =
walrus::GlobalKind::Local(ConstExpr::Value(walrus::ir::Value::I32(ifunc_index)));
}
for mem in got_mems {
let import = new.imports.get(mem);
let data_symbol_idx = *old_symbols
.data_symbol_map
.get(import.name.as_str())
.with_context(|| {
format!("Failed to find GOT.mem import by its name: {}", import.name)
})?;
let data_symbol = old_symbols
.data_symbols
.get(&data_symbol_idx)
.context("Failed to find data symbol by its index")?;
let data = old
.data
.iter()
.nth(data_symbol.which_data_segment)
.context("Missing data segment in the main module")?;
let offset = match data.kind {
DataKind::Active {
offset: ConstExpr::Value(walrus::ir::Value::I32(idx)),
..
} => idx,
DataKind::Active {
offset: ConstExpr::Value(walrus::ir::Value::I64(idx)),
..
} => idx as i32,
_ => {
return Err(PatchError::InvalidModule(format!(
"Data segment of invalid table: {:?}",
data.kind
)));
}
};
let ImportKind::Global(global_id) = import.kind else {
return Err(PatchError::InvalidModule(
"Expected GOT.mem import to be a global".to_string(),
));
};
new.imports.delete(mem);
new.globals.get_mut(global_id).kind = walrus::GlobalKind::Local(ConstExpr::Value(
walrus::ir::Value::I32(offset + data_symbol.segment_offset as i32),
));
}
let ifunc_table_initializer = new
.elements
.iter()
.find_map(|e| match e.kind {
ElementKind::Active { table, .. } => Some(table),
_ => None,
})
.context("Missing ifunc table")?;
for env_func_import in env_funcs {
let import = new.imports.get(env_func_import);
let ImportKind::Function(func_id) = import.kind else {
continue;
};
if cache.old_exports.contains(import.name.as_str())
|| cache.old_imports.contains(import.name.as_str())
{
continue;
}
let name = import.name.as_str().to_string();
if let Some(table_idx) = name_to_ifunc_old.get(import.name.as_str()) {
new.imports.delete(env_func_import);
convert_func_to_ifunc_call(
&mut new,
ifunc_table_initializer,
func_id,
*table_idx,
name.clone(),
);
continue;
}
if name_is_bindgen_symbol(&name) {
new.imports.delete(env_func_import);
convert_func_to_ifunc_call(&mut new, ifunc_table_initializer, func_id, 0, name);
continue;
}
tracing::warn!("[hotpatching]: Symbol slipped through the cracks: {}", name);
}
for import_id in wbg_funcs {
let import = new.imports.get_mut(import_id);
let ImportKind::Function(func_id) = import.kind else {
continue;
};
import.module = "env".into();
import.name = format!("__saved_wbg_{}", import.name);
if name_is_bindgen_symbol(&import.name) {
let name = import.name.as_str().to_string();
new.imports.delete(import_id);
convert_func_to_ifunc_call(&mut new, ifunc_table_initializer, func_id, 0, name);
}
}
let new_func_ids = new.funcs.iter().map(|f| f.id()).collect::<Vec<_>>();
for func_id in new_func_ids {
let Some(name) = new.funcs.get(func_id).name.as_deref() else {
continue;
};
if name_is_wbg_cast_symbol(name) {
let name = name.to_string();
let old_idx = name_to_ifunc_old
.get(&name)
.copied()
.ok_or_else(|| anyhow::anyhow!("Could not find matching wbg_cast function for [{name}] - must generate new JS bindings."))?;
convert_func_to_ifunc_call(&mut new, ifunc_table_initializer, func_id, old_idx, name);
}
}
let customs = new.customs.iter().map(|f| f.0).collect::<Vec<_>>();
for custom_id in customs {
if let Some(custom) = new.customs.get_mut(custom_id) {
if custom.name().contains("manganis") || custom.name().contains("__wasm_bindgen") {
new.customs.delete(custom_id);
}
}
}
new.start = None;
const APPLY_RELOCS: &str = "__wasm_apply_global_relocs";
if let Some(func) = new
.funcs
.iter()
.find(|f| f.name.as_deref() == Some(APPLY_RELOCS))
{
new.exports.add(APPLY_RELOCS, func.id());
}
let lib = patch.to_path_buf();
std::fs::write(&lib, new.emit_wasm())?;
let name_to_ifunc_new = collect_func_ifuncs(&new);
let ifunc_count = name_to_ifunc_new.len() as u64;
let mut map = AddressMap::default();
for (name, idx) in name_to_ifunc_new.iter() {
if let Some(old_idx) = name_to_ifunc_old.get(*name) {
map.insert(*old_idx as u64, *idx as u64);
continue;
}
}
Ok(JumpTable {
map,
lib,
ifunc_count,
aslr_reference: 0,
new_base_address: 0,
})
}
fn convert_func_to_ifunc_call(
new: &mut Module,
ifunc_table_initializer: TableId,
func_id: FunctionId,
table_idx: i32,
name: String,
) {
use walrus::ir;
let func = new.funcs.get_mut(func_id);
let ty_id = func.ty();
let ty = new.types.get(ty_id);
let params = ty.params().to_vec();
let results = ty.results().to_vec();
let locals: Vec<_> = params.iter().map(|ty| new.locals.add(*ty)).collect();
let mut builder = FunctionBuilder::new(&mut new.types, ¶ms, &results);
let mut body = builder.name(name).func_body();
for arg in locals.iter() {
body.local_get(*arg);
}
body.instr(ir::Instr::Const(ir::Const {
value: ir::Value::I32(table_idx),
}));
body.instr(ir::Instr::CallIndirect(ir::CallIndirect {
ty: ty_id,
table: ifunc_table_initializer,
}));
new.funcs.get_mut(func_id).kind = FunctionKind::Local(builder.local_func(locals));
}
fn collect_func_ifuncs(m: &Module) -> HashMap<&str, i32> {
let mut func_to_offset = HashMap::new();
for el in m.elements.iter() {
let ElementKind::Active { offset, .. } = &el.kind else {
continue;
};
let offset = match offset {
ConstExpr::Value(value) => match value {
walrus::ir::Value::I32(idx) => *idx,
walrus::ir::Value::I64(idx) => *idx as i32,
_ => continue,
},
ConstExpr::Global(_) => 0,
_ => continue,
};
match &el.items {
ElementItems::Functions(ids) => {
for (idx, id) in ids.iter().enumerate() {
if let Some(name) = m.funcs.get(*id).name.as_deref() {
func_to_offset.insert(name, offset + idx as i32);
}
}
}
ElementItems::Expressions(_ref_type, _const_exprs) => {}
}
}
func_to_offset
}
pub fn create_undefined_symbol_stub(
cache: &HotpatchModuleCache,
incrementals: &[PathBuf],
triple: &Triple,
aslr_reference: u64,
) -> Result<Vec<u8>> {
let sorted: Vec<_> = incrementals.iter().sorted().collect();
let mut undefined_symbols = HashSet::new();
let mut defined_symbols = HashSet::new();
for path in sorted {
collect_stub_symbols_from_path(path, &mut undefined_symbols, &mut defined_symbols)?;
}
let undefined_symbols: Vec<_> = undefined_symbols
.difference(&defined_symbols)
.cloned()
.collect();
tracing::trace!("Undefined symbols: {:#?}", undefined_symbols);
let mut obj = object::write::Object::new(
match triple.binary_format {
target_lexicon::BinaryFormat::Elf => object::BinaryFormat::Elf,
target_lexicon::BinaryFormat::Macho => object::BinaryFormat::MachO,
target_lexicon::BinaryFormat::Coff => object::BinaryFormat::Coff,
target_lexicon::BinaryFormat::Wasm => object::BinaryFormat::Wasm,
target_lexicon::BinaryFormat::Xcoff => object::BinaryFormat::Xcoff,
_ => return Err(PatchError::UnsupportedPlatform(triple.to_string())),
},
match triple.architecture {
Architecture::Aarch64(_) => object::Architecture::Aarch64,
Architecture::Wasm32 => object::Architecture::Wasm32,
Architecture::X86_64 => object::Architecture::X86_64,
_ => return Err(PatchError::UnsupportedPlatform(triple.to_string())),
},
match triple.endianness() {
Ok(target_lexicon::Endianness::Little) => Endianness::Little,
Ok(target_lexicon::Endianness::Big) => Endianness::Big,
_ => Endianness::Little,
},
);
#[allow(clippy::identity_op)]
match triple.operating_system {
OperatingSystem::Darwin(_) => {
obj.set_macho_build_version({
let mut build_version = MachOBuildVersion::default();
build_version.platform = macho::PLATFORM_MACOS;
build_version.minos = (11 << 16) | (0 << 8) | 0; build_version.sdk = (11 << 16) | (0 << 8) | 0; build_version
});
}
OperatingSystem::IOS(_) => {
obj.set_macho_build_version({
let mut build_version = MachOBuildVersion::default();
build_version.platform = match triple.environment {
target_lexicon::Environment::Sim => macho::PLATFORM_IOSSIMULATOR,
_ => macho::PLATFORM_IOS,
};
build_version.minos = (14 << 16) | (0 << 8) | 0; build_version.sdk = (14 << 16) | (0 << 8) | 0; build_version
});
}
_ => {}
}
let aslr_ref_address = cache
.symbol_table
.get(main_sentinel(triple))
.context("failed to find '_main' symbol in patch")?
.address;
if aslr_reference < aslr_ref_address {
return Err(PatchError::InvalidModule(format!(
"ASLR reference is less than the main module's address - is there a `main`?. {aslr_reference:x} < {aslr_ref_address:x}"
)));
}
let aslr_offset = aslr_reference - aslr_ref_address;
let text_section = obj.section_id(StandardSection::Text);
for name in undefined_symbols {
let Some(sym) = cache
.symbol_table
.get(name.as_str().trim_start_matches("__imp_"))
else {
tracing::debug!("Symbol not found: {}", name);
continue;
};
if sym.is_undefined {
continue;
}
let name_offset = match triple.operating_system {
OperatingSystem::MacOSX(_) | OperatingSystem::Darwin(_) | OperatingSystem::IOS(_) => 1,
_ => 0,
};
let abs_addr = sym.address + aslr_offset;
match sym.kind {
_ if name.starts_with("__imp_") => {
let data_section = obj.section_id(StandardSection::Data);
let offset = obj.append_section_data(
data_section,
&abs_addr.to_le_bytes(),
8, );
obj.add_symbol(Symbol {
name: name.as_bytes().to_vec(),
value: offset, size: 8, scope: SymbolScope::Linkage,
kind: SymbolKind::Data, weak: false,
section: SymbolSection::Section(data_section),
flags: SymbolFlags::None,
});
}
SymbolKind::Text => {
let jump_asm = match triple.operating_system {
OperatingSystem::Windows => match triple.architecture {
Architecture::X86_64 => {
let mut code = vec![
0x48, 0xB8, ];
code.extend_from_slice(&abs_addr.to_le_bytes());
code.extend_from_slice(&[0xFF, 0xE0]);
code
}
Architecture::X86_32(_) => {
let mut code = vec![
0xB8, ];
code.extend_from_slice(&(abs_addr as u32).to_le_bytes());
code.extend_from_slice(&[0xFF, 0xE0]);
code
}
Architecture::Aarch64(_) => {
let mut code = Vec::new();
let imm16_0 = (abs_addr & 0xFFFF) as u16;
let movz = 0xD2800010u32 | ((imm16_0 as u32) << 5);
code.extend_from_slice(&movz.to_le_bytes());
let imm16_1 = ((abs_addr >> 16) & 0xFFFF) as u16;
let movk1 = 0xF2A00010u32 | ((imm16_1 as u32) << 5);
code.extend_from_slice(&movk1.to_le_bytes());
let imm16_2 = ((abs_addr >> 32) & 0xFFFF) as u16;
let movk2 = 0xF2C00010u32 | ((imm16_2 as u32) << 5);
code.extend_from_slice(&movk2.to_le_bytes());
let imm16_3 = ((abs_addr >> 48) & 0xFFFF) as u16;
let movk3 = 0xF2E00010u32 | ((imm16_3 as u32) << 5);
code.extend_from_slice(&movk3.to_le_bytes());
code.extend_from_slice(&[0x00, 0x02, 0x1F, 0xD6]);
code
}
Architecture::Arm(_) => {
let mut code = Vec::new();
code.extend_from_slice(&[0x08, 0xC0, 0x9F, 0xE5]);
code.extend_from_slice(&[0x1C, 0xFF, 0x2F, 0xE1]);
code.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]);
code.extend_from_slice(&(abs_addr as u32).to_le_bytes());
code
}
_ => return Err(PatchError::UnsupportedPlatform(triple.to_string())),
},
_ => match triple.architecture {
Architecture::X86_64 => {
let mut code = vec![0xFF, 0x25, 0x00, 0x00, 0x00, 0x00]; code.extend_from_slice(&abs_addr.to_le_bytes());
code
}
Architecture::X86_32(_) => {
let mut code = vec![0xE9]; let rel_addr = abs_addr as i32 - 5; code.extend_from_slice(&rel_addr.to_le_bytes());
code
}
Architecture::Aarch64(_) => {
let mut code = Vec::new();
code.extend_from_slice(&[0x50, 0x00, 0x00, 0x58]);
code.extend_from_slice(&[0x00, 0x02, 0x1F, 0xD6]);
code.extend_from_slice(&abs_addr.to_le_bytes());
code
}
Architecture::Arm(_) => {
let mut code = Vec::new();
code.extend_from_slice(&[0x04, 0xF0, 0x1F, 0xE5]);
code.extend_from_slice(&(abs_addr as u32).to_le_bytes());
code
}
_ => return Err(PatchError::UnsupportedPlatform(triple.to_string())),
},
};
let offset = obj.append_section_data(text_section, &jump_asm, 8);
obj.add_symbol(Symbol {
name: name.as_bytes()[name_offset..].to_vec(),
value: offset,
size: jump_asm.len() as u64,
scope: SymbolScope::Linkage,
kind: SymbolKind::Text,
weak: false,
section: SymbolSection::Section(text_section),
flags: SymbolFlags::None, });
}
SymbolKind::Tls => {
let tls_section = obj.section_id(StandardSection::Tls);
let pointer_width = match triple.pointer_width().unwrap() {
PointerWidth::U16 => 2,
PointerWidth::U32 => 4,
PointerWidth::U64 => 8,
};
let init_key = format!("{}$tlv$init", name);
let (tls_offset, size) =
if let Some(&(offset, size)) = cache.tls_init_sizes.get(&init_key) {
(offset, size)
} else if sym.size > 0 {
(sym.address, sym.size)
} else if !cache.tls_init_sizes.is_empty() {
(0, cache.tls_init_data.len() as u64)
} else {
(sym.address, pointer_width)
};
let align = size.min(pointer_width).next_power_of_two();
let start = tls_offset as usize;
let end = start + size as usize;
let init = if end <= cache.tls_init_data.len() {
cache.tls_init_data[start..end].to_vec()
} else {
vec![0u8; size as usize]
};
let sym_id = obj.add_symbol(Symbol {
name: name.as_bytes()[name_offset..].to_vec(),
value: 0,
size: 0,
scope: SymbolScope::Linkage,
kind: SymbolKind::Tls,
weak: false,
section: SymbolSection::Undefined,
flags: SymbolFlags::None,
});
obj.add_symbol_data(sym_id, tls_section, &init, align);
}
_ => {
let kind = match sym.kind {
SymbolKind::Unknown => SymbolKind::Data,
k => k,
};
let flags = match triple.environment {
target_lexicon::Environment::Android => SymbolFlags::None,
_ => sym.flags,
};
obj.add_symbol(Symbol {
name: name.as_bytes()[name_offset..].to_vec(),
value: abs_addr,
size: 0,
scope: SymbolScope::Linkage,
kind,
weak: sym.is_weak,
section: SymbolSection::Absolute,
flags,
});
}
}
}
Ok(obj.write()?)
}
fn collect_stub_symbols_from_path(
path: &Path,
undefined_symbols: &mut HashSet<String>,
defined_symbols: &mut HashSet<String>,
) -> Result<()> {
let bytes = std::fs::read(path).with_context(|| format!("failed to read {path:?}"))?;
if path
.extension()
.is_some_and(|ext| matches!(ext.to_str(), Some("rlib" | "a")))
{
let mut archive = ar::Archive::new(std::io::Cursor::new(bytes));
while let Some(entry) = archive.next_entry() {
let mut entry = entry?;
let name = std::str::from_utf8(entry.header().identifier()).unwrap_or_default();
if name.ends_with(".rmeta") || !(name.ends_with(".o") || name.ends_with(".obj")) {
continue;
}
let mut entry_bytes = Vec::with_capacity(entry.header().size() as usize);
entry.read_to_end(&mut entry_bytes)?;
collect_stub_symbols_from_bytes(&entry_bytes, undefined_symbols, defined_symbols)?;
}
return Ok(());
}
collect_stub_symbols_from_bytes(&bytes, undefined_symbols, defined_symbols)
}
fn collect_stub_symbols_from_bytes(
bytes: &[u8],
undefined_symbols: &mut HashSet<String>,
defined_symbols: &mut HashSet<String>,
) -> Result<()> {
let file = File::parse(bytes)?;
for symbol in file.symbols() {
if symbol.is_undefined() {
undefined_symbols.insert(symbol.name()?.to_string());
} else if symbol.is_global() {
defined_symbols.insert(symbol.name()?.to_string());
}
}
Ok(())
}
pub fn prepare_wasm_base_module(bytes: &[u8]) -> Result<Vec<u8>> {
let ParsedModule {
mut module,
ids,
symbols,
..
} = parse_module_with_ids(bytes)?;
let ifunc_map = collect_func_ifuncs(&module);
let ifuncs = module
.funcs
.par_iter()
.filter_map(|f| ifunc_map.get(f.name.as_deref()?).map(|_| f.id()))
.collect::<HashSet<_>>();
let imported_funcs = module
.imports
.iter()
.filter_map(|i| match i.kind {
ImportKind::Function(id) => Some((id, i.id())),
_ => None,
})
.collect::<HashMap<_, _>>();
let mut exported = HashSet::new();
let mut make_indirect = vec![];
for (imported_func, importid) in imported_funcs {
let (import_module, import_name) = {
let import = module.imports.get(importid);
(import.module.to_string(), import.name.to_string())
};
let name_is_wbg =
import_name.starts_with("__wbindgen") || import_name.starts_with("__wbg_");
if name_is_wbg && !name_is_bindgen_symbol(&import_name) {
let func = module.funcs.get(imported_func);
let ty = module.types.get(func.ty());
let params = ty.params().to_vec();
let results = ty.results().to_vec();
let mut builder = FunctionBuilder::new(&mut module.types, ¶ms, &results);
let mut body = builder
.name(format!("__saved_wbg_{}", import_name))
.func_body();
let locals = params
.iter()
.map(|ty| module.locals.add(*ty))
.collect::<Vec<_>>();
for l in locals.iter() {
body.local_get(*l);
}
body.call(imported_func);
let new_func_id = module.funcs.add_local(builder.local_func(locals));
let saved_name = format!("__saved_wbg_{}", import_name);
if exported.insert(saved_name.clone()) {
module.exports.add(&saved_name, new_func_id);
}
make_indirect.push(new_func_id);
} else if import_module == "env" && !name_is_wbg {
let original_name = module.funcs.get(imported_func).name.clone();
let new_fid = module
.replace_imported_func(imported_func, |(body, _args)| {
body.unreachable();
})
.map_err(|e| {
PatchError::InvalidModule(format!(
"Failed to stub env import {import_name}: {e}"
))
})?;
module.funcs.get_mut(new_fid).name = original_name;
make_indirect.push(new_fid);
}
}
for (name, index) in symbols.code_symbol_map.iter() {
if name_is_bindgen_symbol(name) {
continue;
}
let func = module.funcs.get(ids[*index]);
if name.starts_with("__wbindgen") {
let saved_name = format!("__saved_wbg_{}", name);
if exported.insert(saved_name.clone()) {
module.exports.add(&saved_name, func.id());
}
}
if let FunctionKind::Local(_) = &func.kind {
if !ifuncs.contains(&func.id()) {
make_indirect.push(func.id());
}
}
}
let segment = module
.elements
.iter_mut()
.last()
.context("Missing ifunc table")?;
let make_indirect_count = make_indirect.len() as u64;
let ElementItems::Functions(segment_ids) = &mut segment.items else {
return Err(PatchError::InvalidModule(
"Expected ifunc table to be a function table".into(),
));
};
for func in make_indirect {
segment_ids.push(func);
}
if let ElementKind::Active { table, .. } = segment.kind {
let table = module.tables.get_mut(table);
table.initial += make_indirect_count;
if let Some(max) = table.maximum {
table.maximum = Some(max + make_indirect_count);
}
}
Ok(module.emit_wasm())
}
fn name_is_wbg_cast_symbol(name: &str) -> bool {
(name.contains("wasm_bindgen4__rt8wbg_cast") || name.contains("wasm_bindgen4___rt8wbg_cast"))
&& !name.contains("breaks_if_inline")
}
fn name_is_bindgen_symbol(name: &str) -> bool {
name.contains("__wbindgen_describe")
|| name.contains("__wbindgen_externref")
|| name.contains("wasm_bindgen8describe6inform")
|| name.contains("wasm_bindgen..describe..WasmDescribe")
|| name.contains("12WasmDescribe8describe")
|| name.contains("18WasmDescribeVector15describe_vector")
|| (name.contains("wasm_bindgen..closure..WasmClosure") && name.contains("describe"))
|| (name.contains("11WasmClosure") && name.contains("describe"))
|| (name.contains("wasm_bindgen7closure16Closure") && name.contains("describe"))
|| (name.contains("7closure7Closure") && name.contains("describe"))
|| (name.contains("wasm_bindgen7convert8closures") && name.contains("describe_invoke"))
}
#[test]
fn bindgen_symbol_catch() {
let symbol = "_ZN12wasm_bindgen7convert8closures1_142_$LT$impl$u20$wasm_bindgen..closure..WasmClosure$u20$for$u20$dyn$u20$core..ops..function..FnMut$LT$$LP$$RP$$GT$$u2b$Output$u20$$u3d$$u20$R$GT$15describe_invoke17h4373f8b6570333dcE";
assert!(name_is_bindgen_symbol(symbol));
let symbol = "_ZN12wasm_bindgen7closure16Closure$LT$T$GT$4wrap8describe17h1234567890abcdefE";
assert!(name_is_bindgen_symbol(symbol));
let symbol = "_RNvXNvNtNtCs9jB4f2OZCsR_7web_sys8features36gen_TransformStreamDefaultController1__NtB4_32TransformStreamDefaultControllerNtNtCs9tRDgkfeYnK_12wasm_bindgen8describe12WasmDescribe8describe";
assert!(name_is_bindgen_symbol(symbol));
let symbol = "_RNvXNvNtNtCs9jB4f2OZCsR_7web_sys8features8gen_Node4NodeENtB5_12WasmDescribe8describeCs3X5Dvr2wWzv_21dioxus_interpreter_js";
assert!(name_is_bindgen_symbol(symbol));
let symbol = "_RNvXs4_NtNtCs9tRDgkfeYnK_12wasm_bindgen7convert6slicesNtNtCscHiZRFGp0KF_5alloc6string6StringNtNtB9_8describe18WasmDescribeVector15describe_vector";
assert!(name_is_bindgen_symbol(symbol));
let symbol = "_RNvXsf_NtCs4ofacjxbDm2_10dioxus_web8documentNtB5_7JSOwnerNtNtCs9tRDgkfeYnK_12wasm_bindgen8describe18WasmDescribeVector15describe_vector";
assert!(name_is_bindgen_symbol(symbol));
let symbol = "_RINvXs1_NvNtNtCs9tRDgkfeYnK_12wasm_bindgen7convert8closuress8_1__DINtNtNtCs9WN6KVdqFxk_4core3ops8function5FnMutTNtBc_7JsValueB1M_mNtCsezy3jvZZ1sp_6js_sys5ArrayEEp6OutputB1M_EL_NtNtBc_7closure11WasmClosure15describe_invokeKb1_EB26_";
assert!(name_is_bindgen_symbol(symbol));
let symbol = "_RNvCs9tRDgkfeYnK_12wasm_bindgen19___wbindgen_describe";
assert!(name_is_bindgen_symbol(symbol));
assert!(!name_is_bindgen_symbol("__wbindgen_malloc"));
assert!(!name_is_bindgen_symbol("__wbindgen_realloc"));
assert!(!name_is_bindgen_symbol("__wbindgen_free"));
assert!(!name_is_bindgen_symbol(
"_ZN5alloc7raw_vec19RawVec$LT$T$C$A$GT$8grow_one17h1234567890abcdefE"
));
assert!(!name_is_bindgen_symbol(
"_RNvXs5_NtCs9tRDgkfeYnK_12wasm_bindgen5__rt5LazyINtB5_4LazyNtNtCsezy3jvZZ1sp_6js_sys6ObjectE5force"
));
}
#[test]
fn wbg_cast_symbol_catch() {
assert!(name_is_wbg_cast_symbol(
"_ZN12wasm_bindgen4__rt8wbg_cast17h1234567890abcdefE"
));
assert!(!name_is_wbg_cast_symbol(
"_ZN12wasm_bindgen4__rt8wbg_cast17breaks_if_inlined17h1234567890abcdefE"
));
assert!(name_is_wbg_cast_symbol(
"_RINvNtCsa7akE1TfegA_12wasm_bindgen4___rt8wbg_castINtB4_7closure12OwnedClosureDINtNtNtCs9WN6KVdqFxk_4core3ops8function5FnMutTNtNtNtCs5qlPUvWaqlJ_7web_sys8features14gen_MouseEvent10MouseEventEEp6OutputuEL_Kb1_ENtBO_9JsClosureECsjFep1nV9Dzo_32dioxus_playwright_web_patch_test"
));
assert!(!name_is_wbg_cast_symbol(
"_RINvNvNtCsa7akE1TfegA_12wasm_bindgen4___rt8wbg_cast17breaks_if_inlinedINtNtB6_7closure12OwnedClosureDINtNtNtCs9WN6KVdqFxk_4core3ops8function5FnMutTNtNtNtCs5qlPUvWaqlJ_7web_sys8features14gen_MouseEvent10MouseEventEEp6OutputuEL_Kb1_ENtB19_9JsClosureECsjFep1nV9Dzo_32dioxus_playwright_web_patch_test"
));
}
fn parse_bytes_to_data_segment(bytes: &[u8]) -> Result<RawDataSection<'_>> {
let parser = wasmparser::Parser::new(0);
let mut parser = parser.parse_all(bytes);
let mut segments = vec![];
let mut data_range = 0..0;
let mut symbols = vec![];
while let Some(Ok(payload)) = parser.next() {
match payload {
Payload::DataSection(section) => {
data_range = section.range();
segments = section
.into_iter()
.collect::<Result<Vec<_>, BinaryReaderError>>()?
}
Payload::CustomSection(section) if section.name() == "linking" => {
let reader = BinaryReader::new(section.data(), 0);
let reader = LinkingSectionReader::new(reader)?;
for subsection in reader.subsections() {
if let Linking::SymbolTable(map) = subsection? {
symbols = map.into_iter().collect::<Result<Vec<_>, _>>()?;
}
}
}
Payload::CustomSection(section) => {
tracing::trace!("Skipping Custom section: {:?}", section.name());
}
_ => {}
}
}
let mut data_symbols = BTreeMap::new();
let mut data_symbol_map = HashMap::new();
let mut code_symbol_map = BTreeMap::new();
for (index, symbol) in symbols.iter().enumerate() {
if let SymbolInfo::Func { name, index, .. } = symbol {
if let Some(name) = name {
code_symbol_map.insert(*name, *index as usize);
}
continue;
}
let SymbolInfo::Data {
symbol: Some(symbol),
name,
..
} = symbol
else {
continue;
};
data_symbol_map.insert(*name, index);
let data_segment = segments
.get(symbol.index as usize)
.context("Failed to find data segment")?;
let offset: usize =
data_segment.range.end - data_segment.data.len() + (symbol.offset as usize);
let range = offset..(offset + symbol.size as usize);
data_symbols.insert(
index,
DataSymbol {
_index: index,
_range: range,
segment_offset: symbol.offset as usize,
_symbol_size: symbol.size as usize,
which_data_segment: symbol.index as usize,
},
);
}
Ok(RawDataSection {
_data_range: data_range,
symbols,
data_symbols,
data_symbol_map,
code_symbol_map,
})
}
struct RawDataSection<'a> {
_data_range: Range<usize>,
symbols: Vec<SymbolInfo<'a>>,
code_symbol_map: BTreeMap<&'a str, usize>,
data_symbols: BTreeMap<usize, DataSymbol>,
data_symbol_map: HashMap<&'a str, usize>,
}
#[derive(Debug)]
struct DataSymbol {
_index: usize,
_range: Range<usize>,
segment_offset: usize,
_symbol_size: usize,
which_data_segment: usize,
}
struct ParsedModule<'a> {
module: Module,
ids: Vec<FunctionId>,
symbols: RawDataSection<'a>,
}
fn parse_module_with_ids(bindgened: &[u8]) -> Result<ParsedModule<'_>> {
let ids = Arc::new(RwLock::new(Vec::new()));
let ids_ = ids.clone();
let module = Module::from_buffer_with_config(
bindgened,
ModuleConfig::new().on_parse(move |_m, our_ids| {
let mut ids = ids_.write().expect("No shared writers");
let mut idx = 0;
while let Ok(entry) = our_ids.get_func(idx) {
ids.push(entry);
idx += 1;
}
Ok(())
}),
)?;
let mut ids_ = ids.write().expect("No shared writers");
let mut ids = vec![];
std::mem::swap(&mut ids, &mut *ids_);
let symbols = parse_bytes_to_data_segment(bindgened).context("Failed to parse data segment")?;
Ok(ParsedModule {
module,
ids,
symbols,
})
}
fn main_sentinel(triple: &Triple) -> &'static str {
match triple.operating_system {
OperatingSystem::MacOSX(_) | OperatingSystem::Darwin(_) | OperatingSystem::IOS(_) => {
"_main"
}
_ => "main",
}
}