1use crate::{
4 compiled_blob::CompiledBlob,
5 memory::{BranchProtection, JITMemoryKind, JITMemoryProvider, SystemMemoryProvider},
6};
7use cranelift_codegen::binemit::Reloc;
8use cranelift_codegen::isa::{OwnedTargetIsa, TargetIsa};
9use cranelift_codegen::settings::Configurable;
10use cranelift_codegen::{ir, settings};
11use cranelift_control::ControlPlane;
12use cranelift_entity::SecondaryMap;
13use cranelift_module::{
14 DataDescription, DataId, FuncId, Init, Linkage, Module, ModuleDeclarations, ModuleError,
15 ModuleReloc, ModuleRelocTarget, ModuleResult,
16};
17use log::info;
18use std::cell::RefCell;
19use std::collections::BTreeMap;
20use std::collections::HashMap;
21use std::ffi::CString;
22use std::io::Write;
23use target_lexicon::{Architecture, PointerWidth};
24
25const WRITABLE_DATA_ALIGNMENT: u64 = 0x8;
26const READONLY_DATA_ALIGNMENT: u64 = 0x1;
27
28pub struct JITBuilder {
30 isa: OwnedTargetIsa,
31 symbols: HashMap<String, SendWrapper<*const u8>>,
32 lookup_symbols: Vec<Box<dyn Fn(&str) -> Option<*const u8> + Send>>,
33 libcall_names: Box<dyn Fn(ir::LibCall) -> String + Send + Sync>,
34 memory: Option<Box<dyn JITMemoryProvider + Send>>,
35}
36
37impl JITBuilder {
38 pub fn new(
45 libcall_names: Box<dyn Fn(ir::LibCall) -> String + Send + Sync>,
46 ) -> ModuleResult<Self> {
47 Self::with_flags(&[], libcall_names)
48 }
49
50 pub fn with_flags(
57 flags: &[(&str, &str)],
58 libcall_names: Box<dyn Fn(ir::LibCall) -> String + Send + Sync>,
59 ) -> ModuleResult<Self> {
60 let mut flag_builder = settings::builder();
61 for (name, value) in flags {
62 flag_builder.set(name, value)?;
63 }
64
65 flag_builder.set("use_colocated_libcalls", "false").unwrap();
69 let is_pic = if cfg!(target_arch = "x86_64") {
76 "true"
77 } else {
78 "false"
79 };
80 flag_builder.set("is_pic", is_pic).unwrap();
81 let isa_builder = cranelift_native::builder().unwrap_or_else(|msg| {
82 panic!("host machine is not supported: {msg}");
83 });
84 let isa = isa_builder.finish(settings::Flags::new(flag_builder))?;
85 Ok(Self::with_isa(isa, libcall_names))
86 }
87
88 pub fn with_isa(
99 isa: OwnedTargetIsa,
100 libcall_names: Box<dyn Fn(ir::LibCall) -> String + Send + Sync>,
101 ) -> Self {
102 let symbols = HashMap::new();
103 let lookup_symbols = vec![Box::new(lookup_with_dlsym) as Box<_>];
104 Self {
105 isa,
106 symbols,
107 lookup_symbols,
108 libcall_names,
109 memory: None,
110 }
111 }
112
113 pub fn symbol<K>(&mut self, name: K, ptr: *const u8) -> &mut Self
128 where
129 K: Into<String>,
130 {
131 self.symbols.insert(name.into(), SendWrapper(ptr));
132 self
133 }
134
135 pub fn symbols<It, K>(&mut self, symbols: It) -> &mut Self
139 where
140 It: IntoIterator<Item = (K, *const u8)>,
141 K: Into<String>,
142 {
143 for (name, ptr) in symbols {
144 self.symbols.insert(name.into(), SendWrapper(ptr));
145 }
146 self
147 }
148
149 pub fn symbol_lookup_fn(
154 &mut self,
155 symbol_lookup_fn: Box<dyn Fn(&str) -> Option<*const u8> + Send>,
156 ) -> &mut Self {
157 self.lookup_symbols.push(symbol_lookup_fn);
158 self
159 }
160
161 pub fn memory_provider(&mut self, provider: Box<dyn JITMemoryProvider + Send>) -> &mut Self {
165 self.memory = Some(provider);
166 self
167 }
168}
169
170#[derive(Copy, Clone)]
174struct SendWrapper<T>(T);
175unsafe impl<T> Send for SendWrapper<T> {}
176
177pub struct JITModule {
182 isa: OwnedTargetIsa,
183 symbols: RefCell<HashMap<String, SendWrapper<*const u8>>>,
184 lookup_symbols: Vec<Box<dyn Fn(&str) -> Option<*const u8> + Send>>,
185 libcall_names: Box<dyn Fn(ir::LibCall) -> String + Send + Sync>,
186 memory: Box<dyn JITMemoryProvider + Send>,
187 declarations: ModuleDeclarations,
188 compiled_functions: SecondaryMap<FuncId, Option<CompiledBlob>>,
189 compiled_data_objects: SecondaryMap<DataId, Option<CompiledBlob>>,
190 code_ranges: BTreeMap<usize, (usize, FuncId)>,
193 functions_to_finalize: Vec<FuncId>,
194 data_objects_to_finalize: Vec<DataId>,
195}
196
197impl JITModule {
198 pub unsafe fn free_memory(mut self) {
207 self.memory.free_memory();
208 }
209
210 fn lookup_symbol(&self, name: &str) -> Option<*const u8> {
211 match self.symbols.borrow_mut().entry(name.to_owned()) {
212 std::collections::hash_map::Entry::Occupied(occ) => Some(occ.get().0),
213 std::collections::hash_map::Entry::Vacant(vac) => {
214 let ptr = self
215 .lookup_symbols
216 .iter()
217 .rev() .find_map(|lookup| lookup(name));
219 if let Some(ptr) = ptr {
220 vac.insert(SendWrapper(ptr));
221 }
222 ptr
223 }
224 }
225 }
226
227 pub fn get_address(&self, name: &ModuleRelocTarget) -> *const u8 {
234 match name {
235 ModuleRelocTarget::User { .. } => {
236 let (name, linkage) = if ModuleDeclarations::is_function(name) {
237 let func_id = FuncId::from_name(name);
238 match &self.compiled_functions[func_id] {
239 Some(compiled) => return compiled.ptr(),
240 None => {
241 let decl = self.declarations.get_function_decl(func_id);
242 (&decl.name, decl.linkage)
243 }
244 }
245 } else {
246 let data_id = DataId::from_name(name);
247 match &self.compiled_data_objects[data_id] {
248 Some(compiled) => return compiled.ptr(),
249 None => {
250 let decl = self.declarations.get_data_decl(data_id);
251 (&decl.name, decl.linkage)
252 }
253 }
254 };
255 let name = name
256 .as_ref()
257 .expect("anonymous symbol must be defined locally");
258 if let Some(ptr) = self.lookup_symbol(name) {
259 ptr
260 } else if linkage == Linkage::Preemptible {
261 0 as *const u8
262 } else {
263 panic!("can't resolve symbol {name}");
264 }
265 }
266 ModuleRelocTarget::LibCall(libcall) => {
267 let sym = (self.libcall_names)(*libcall);
268 self.lookup_symbol(&sym)
269 .unwrap_or_else(|| panic!("can't resolve libcall {sym}"))
270 }
271 ModuleRelocTarget::FunctionOffset(func_id, offset) => {
272 match &self.compiled_functions[*func_id] {
273 Some(compiled) => return compiled.ptr().wrapping_add(*offset as usize),
274 None => todo!(),
275 }
276 }
277 name => panic!("invalid name {name:?}"),
278 }
279 }
280
281 pub fn get_finalized_function(&self, func_id: FuncId) -> *const u8 {
286 let info = &self.compiled_functions[func_id];
287 assert!(
288 !self.functions_to_finalize.iter().any(|x| *x == func_id),
289 "function not yet finalized"
290 );
291 info.as_ref()
292 .expect("function must be compiled before it can be finalized")
293 .ptr()
294 }
295
296 pub fn get_finalized_data(&self, data_id: DataId) -> (*const u8, usize) {
301 let info = &self.compiled_data_objects[data_id];
302 assert!(
303 !self.data_objects_to_finalize.iter().any(|x| *x == data_id),
304 "data object not yet finalized"
305 );
306 let compiled = info
307 .as_ref()
308 .expect("data object must be compiled before it can be finalized");
309
310 (compiled.ptr(), compiled.size())
311 }
312
313 fn record_function_for_perf(&self, ptr: *const u8, size: usize, name: &str) {
314 if cfg!(unix) && ::std::env::var_os("PERF_BUILDID_DIR").is_some() {
320 let mut map_file = ::std::fs::OpenOptions::new()
321 .create(true)
322 .append(true)
323 .open(format!("/tmp/perf-{}.map", ::std::process::id()))
324 .unwrap();
325
326 let _ = writeln!(map_file, "{:x} {:x} {}", ptr as usize, size, name);
327 }
328 }
329
330 pub fn finalize_definitions(&mut self) -> ModuleResult<()> {
339 for func in std::mem::take(&mut self.functions_to_finalize) {
340 let decl = self.declarations.get_function_decl(func);
341 assert!(decl.linkage.is_definable());
342 let func = self.compiled_functions[func]
343 .as_ref()
344 .expect("function must be compiled before it can be finalized");
345 func.perform_relocations(|name| self.get_address(name));
346 }
347
348 for data in std::mem::take(&mut self.data_objects_to_finalize) {
349 let decl = self.declarations.get_data_decl(data);
350 assert!(decl.linkage.is_definable());
351 let data = self.compiled_data_objects[data]
352 .as_ref()
353 .expect("data object must be compiled before it can be finalized");
354 data.perform_relocations(|name| self.get_address(name));
355 }
356
357 let branch_protection = if cfg!(target_arch = "aarch64") && use_bti(&self.isa.isa_flags()) {
359 BranchProtection::BTI
360 } else {
361 BranchProtection::None
362 };
363 self.memory.finalize(branch_protection)?;
364
365 Ok(())
366 }
367
368 pub fn new(builder: JITBuilder) -> Self {
370 assert!(
371 !builder.isa.flags().is_pic()
372 || builder.isa.triple().architecture == Architecture::X86_64,
373 "cranelift-jit only supports is_pic=true on x86_64"
374 );
375
376 let memory = builder
377 .memory
378 .unwrap_or_else(|| Box::new(SystemMemoryProvider::new()));
379 Self {
380 isa: builder.isa,
381 symbols: RefCell::new(builder.symbols),
382 lookup_symbols: builder.lookup_symbols,
383 libcall_names: builder.libcall_names,
384 memory,
385 declarations: ModuleDeclarations::default(),
386 compiled_functions: SecondaryMap::new(),
387 compiled_data_objects: SecondaryMap::new(),
388 code_ranges: BTreeMap::new(),
389 functions_to_finalize: Vec::new(),
390 data_objects_to_finalize: Vec::new(),
391 }
392 }
393
394 #[cfg(feature = "wasmtime-unwinder")]
398 pub fn lookup_wasmtime_exception_data<'a>(
399 &'a self,
400 pc: usize,
401 ) -> Option<(usize, wasmtime_unwinder::ExceptionTable<'a>)> {
402 let (&start, &(end, func)) = self.code_ranges.range(..=pc).next_back()?;
403 if pc >= end {
404 return None;
405 }
406
407 let data = self.compiled_functions[func]
411 .as_ref()
412 .unwrap()
413 .wasmtime_exception_data()?;
414 let exception_table = wasmtime_unwinder::ExceptionTable::parse(data).ok()?;
415 Some((start, exception_table))
416 }
417}
418
419impl Module for JITModule {
420 fn isa(&self) -> &dyn TargetIsa {
421 &*self.isa
422 }
423
424 fn declarations(&self) -> &ModuleDeclarations {
425 &self.declarations
426 }
427
428 fn declare_function(
429 &mut self,
430 name: &str,
431 linkage: Linkage,
432 signature: &ir::Signature,
433 ) -> ModuleResult<FuncId> {
434 let (id, _linkage) = self
435 .declarations
436 .declare_function(name, linkage, signature)?;
437 Ok(id)
438 }
439
440 fn declare_anonymous_function(&mut self, signature: &ir::Signature) -> ModuleResult<FuncId> {
441 let id = self.declarations.declare_anonymous_function(signature)?;
442 Ok(id)
443 }
444
445 fn declare_data(
446 &mut self,
447 name: &str,
448 linkage: Linkage,
449 writable: bool,
450 tls: bool,
451 ) -> ModuleResult<DataId> {
452 assert!(!tls, "JIT doesn't yet support TLS");
453 let (id, _linkage) = self
454 .declarations
455 .declare_data(name, linkage, writable, tls)?;
456 Ok(id)
457 }
458
459 fn declare_anonymous_data(&mut self, writable: bool, tls: bool) -> ModuleResult<DataId> {
460 assert!(!tls, "JIT doesn't yet support TLS");
461 let id = self.declarations.declare_anonymous_data(writable, tls)?;
462 Ok(id)
463 }
464
465 fn define_function_with_control_plane(
466 &mut self,
467 id: FuncId,
468 ctx: &mut cranelift_codegen::Context,
469 ctrl_plane: &mut ControlPlane,
470 ) -> ModuleResult<()> {
471 info!("defining function {}: {}", id, ctx.func.display());
472 let decl = self.declarations.get_function_decl(id);
473 if !decl.linkage.is_definable() {
474 return Err(ModuleError::InvalidImportDefinition(
475 decl.linkage_name(id).into_owned(),
476 ));
477 }
478
479 if !self.compiled_functions[id].is_none() {
480 return Err(ModuleError::DuplicateDefinition(
481 decl.linkage_name(id).into_owned(),
482 ));
483 }
484
485 let res = ctx.compile(self.isa(), ctrl_plane)?;
487 let alignment = res.buffer.min_alignment as u64;
488 let compiled_code = ctx.compiled_code().unwrap();
489
490 let align = alignment
491 .max(self.isa.function_alignment().minimum as u64)
492 .max(self.isa.symbol_alignment());
493
494 let relocs = compiled_code
495 .buffer
496 .relocs()
497 .iter()
498 .map(|reloc| ModuleReloc::from_mach_reloc(reloc, &ctx.func, id))
499 .collect();
500
501 #[cfg(feature = "wasmtime-unwinder")]
502 let wasmtime_exception_data = {
503 let mut exception_builder = wasmtime_unwinder::ExceptionTableBuilder::default();
504 exception_builder
505 .add_func(0, compiled_code.buffer.call_sites())
506 .map_err(|_| {
507 ModuleError::Compilation(cranelift_codegen::CodegenError::Unsupported(
508 "Invalid exception data".into(),
509 ))
510 })?;
511 Some(exception_builder.to_vec())
512 };
513
514 let blob = self.compiled_functions[id].insert(CompiledBlob::new(
515 &mut *self.memory,
516 compiled_code.code_buffer(),
517 align,
518 relocs,
519 #[cfg(feature = "wasmtime-unwinder")]
520 wasmtime_exception_data,
521 JITMemoryKind::Executable,
522 )?);
523 let (ptr, size) = (blob.ptr(), blob.size());
524 self.record_function_for_perf(ptr, size, &decl.linkage_name(id));
525
526 let range_start = ptr.addr();
527 let range_end = range_start + size;
528 self.code_ranges.insert(range_start, (range_end, id));
529
530 self.functions_to_finalize.push(id);
531
532 Ok(())
533 }
534
535 fn define_function_bytes(
536 &mut self,
537 id: FuncId,
538 alignment: u64,
539 bytes: &[u8],
540 relocs: &[ModuleReloc],
541 ) -> ModuleResult<()> {
542 info!("defining function {id} with bytes");
543 let decl = self.declarations.get_function_decl(id);
544 if !decl.linkage.is_definable() {
545 return Err(ModuleError::InvalidImportDefinition(
546 decl.linkage_name(id).into_owned(),
547 ));
548 }
549
550 if !self.compiled_functions[id].is_none() {
551 return Err(ModuleError::DuplicateDefinition(
552 decl.linkage_name(id).into_owned(),
553 ));
554 }
555
556 let align = alignment
557 .max(self.isa.function_alignment().minimum as u64)
558 .max(self.isa.symbol_alignment());
559
560 let blob = self.compiled_functions[id].insert(CompiledBlob::new(
561 &mut *self.memory,
562 bytes,
563 align,
564 relocs.to_owned(),
565 #[cfg(feature = "wasmtime-unwinder")]
566 None,
567 JITMemoryKind::Executable,
568 )?);
569 let (ptr, size) = (blob.ptr(), blob.size());
570 self.record_function_for_perf(ptr, size, &decl.linkage_name(id));
571
572 self.functions_to_finalize.push(id);
573
574 Ok(())
575 }
576
577 fn define_data(&mut self, id: DataId, data: &DataDescription) -> ModuleResult<()> {
578 let decl = self.declarations.get_data_decl(id);
579 if !decl.linkage.is_definable() {
580 return Err(ModuleError::InvalidImportDefinition(
581 decl.linkage_name(id).into_owned(),
582 ));
583 }
584
585 if !self.compiled_data_objects[id].is_none() {
586 return Err(ModuleError::DuplicateDefinition(
587 decl.linkage_name(id).into_owned(),
588 ));
589 }
590
591 assert!(!decl.tls, "JIT doesn't yet support TLS");
592
593 let &DataDescription {
594 ref init,
595 function_decls: _,
596 data_decls: _,
597 function_relocs: _,
598 data_relocs: _,
599 custom_section: _,
600 align,
601 used: _,
602 } = data;
603
604 let (align, kind) = if decl.writable {
605 (
606 align.unwrap_or(WRITABLE_DATA_ALIGNMENT),
607 JITMemoryKind::Writable,
608 )
609 } else {
610 (
611 align.unwrap_or(READONLY_DATA_ALIGNMENT),
612 JITMemoryKind::ReadOnly,
613 )
614 };
615
616 let pointer_reloc = match self.isa.triple().pointer_width().unwrap() {
617 PointerWidth::U16 => panic!(),
618 PointerWidth::U32 => Reloc::Abs4,
619 PointerWidth::U64 => Reloc::Abs8,
620 };
621 let relocs = data.all_relocs(pointer_reloc).collect::<Vec<_>>();
622
623 self.compiled_data_objects[id] = Some(match *init {
624 Init::Uninitialized => {
625 panic!("data is not initialized yet");
626 }
627 Init::Zeros { size } => CompiledBlob::new_zeroed(
628 &mut *self.memory,
629 size.max(1),
630 align,
631 relocs,
632 #[cfg(feature = "wasmtime-unwinder")]
633 None,
634 kind,
635 )?,
636 Init::Bytes { ref contents } => CompiledBlob::new(
637 &mut *self.memory,
638 if contents.is_empty() {
639 &[0]
645 } else {
646 &contents[..]
647 },
648 align,
649 relocs,
650 #[cfg(feature = "wasmtime-unwinder")]
651 None,
652 kind,
653 )?,
654 });
655
656 self.data_objects_to_finalize.push(id);
657
658 Ok(())
659 }
660}
661
662#[cfg(not(windows))]
663fn lookup_with_dlsym(name: &str) -> Option<*const u8> {
664 let c_str = CString::new(name).unwrap();
665 let c_str_ptr = c_str.as_ptr();
666 let sym = unsafe { libc::dlsym(libc::RTLD_DEFAULT, c_str_ptr) };
667 if sym.is_null() {
668 None
669 } else {
670 Some(sym as *const u8)
671 }
672}
673
674#[cfg(windows)]
675fn lookup_with_dlsym(name: &str) -> Option<*const u8> {
676 use std::os::windows::io::RawHandle;
677 use std::ptr;
678 use windows_sys::Win32::Foundation::HMODULE;
679 use windows_sys::Win32::System::LibraryLoader;
680
681 const UCRTBASE: &[u8] = b"ucrtbase.dll\0";
682
683 let c_str = CString::new(name).unwrap();
684 let c_str_ptr = c_str.as_ptr();
685
686 unsafe {
687 let handles = [
688 ptr::null_mut(),
690 LibraryLoader::GetModuleHandleA(UCRTBASE.as_ptr()) as RawHandle,
692 ];
693
694 for handle in &handles {
695 let addr = LibraryLoader::GetProcAddress(*handle as HMODULE, c_str_ptr.cast());
696 match addr {
697 None => continue,
698 Some(addr) => return Some(addr as *const u8),
699 }
700 }
701
702 None
703 }
704}
705
706fn use_bti(isa_flags: &Vec<settings::Value>) -> bool {
707 isa_flags
708 .iter()
709 .find(|&f| f.name == "use_bti")
710 .map_or(false, |f| f.as_bool().unwrap_or(false))
711}
712
713const _ASSERT_JIT_MODULE_IS_SEND: () = {
714 const fn assert_is_send<T: Send>() {}
715 assert_is_send::<JITModule>();
716};