pliron_llvm/llvm_sys/
lljit.rs1use bitflags::bitflags;
48use core::{marker::PhantomData, mem, ops::Deref, ptr};
49
50use llvm_sys::{
51 core::LLVMGetModuleContext,
52 orc2::{
53 LLVMJITEvaluatedSymbol, LLVMJITSymbolFlags, LLVMJITSymbolGenericFlags,
54 LLVMOrcAbsoluteSymbols, LLVMOrcCSymbolMapPair,
55 LLVMOrcCreateNewThreadSafeContextFromLLVMContext, LLVMOrcCreateNewThreadSafeModule,
56 LLVMOrcDisposeThreadSafeContext, lljit,
57 },
58};
59
60use crate::llvm_sys::{
61 core::{LLVMContext, LLVMModule, handle_err},
62 cstr_to_string,
63 target::initialize_native,
64 to_c_str,
65};
66
67bitflags! {
68 #[derive(PartialEq, Eq, Clone, Debug, Hash, Copy)]
69 pub struct JITSymbolGenericFlags: u8 {
70 const JITSymbolGenericFlagsNone = 0;
71 const JITSymbolGenericFlagsExported = 1;
72 const JITSymbolGenericFlagsWeak = 2;
73 const JITSymbolGenericFlagsCallable = 4;
74 const JITSymbolGenericFlagsMaterializationSideEffectsOnly = 8;
75 }
76}
77
78impl From<LLVMJITSymbolGenericFlags> for JITSymbolGenericFlags {
79 fn from(value: LLVMJITSymbolGenericFlags) -> Self {
80 let mut flags = JITSymbolGenericFlags::empty();
81 if (value as u8) & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsExported as u8) != 0
82 {
83 flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsExported;
84 }
85 if (value as u8) & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsWeak as u8) != 0 {
86 flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsWeak;
87 }
88 if (value as u8) & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsCallable as u8) != 0
89 {
90 flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsCallable;
91 }
92 if (value as u8)
93 & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsMaterializationSideEffectsOnly
94 as u8)
95 != 0
96 {
97 flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsMaterializationSideEffectsOnly;
98 }
99 flags
100 }
101}
102
103impl From<JITSymbolGenericFlags> for u8 {
104 fn from(value: JITSymbolGenericFlags) -> Self {
105 let mut flags = LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsNone as u8;
106 if value.contains(JITSymbolGenericFlags::JITSymbolGenericFlagsExported) {
107 flags |= LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsExported as u8;
108 }
109 if value.contains(JITSymbolGenericFlags::JITSymbolGenericFlagsWeak) {
110 flags |= LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsWeak as u8;
111 }
112 if value.contains(JITSymbolGenericFlags::JITSymbolGenericFlagsCallable) {
113 flags |= LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsCallable as u8;
114 }
115 if value
116 .contains(JITSymbolGenericFlags::JITSymbolGenericFlagsMaterializationSideEffectsOnly)
117 {
118 flags |=
119 LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsMaterializationSideEffectsOnly
120 as u8;
121 }
122 flags
123 }
124}
125
126pub struct LLVMLLJIT(lljit::LLVMOrcLLJITRef);
130
131impl LLVMLLJIT {
132 pub fn new_with_default_builder() -> Result<Self, String> {
134 unsafe {
135 let mut jit = mem::MaybeUninit::uninit();
136 let err = lljit::LLVMOrcCreateLLJIT(jit.as_mut_ptr(), ptr::null_mut());
137 handle_err(err)?;
138 Ok(LLVMLLJIT(jit.assume_init()))
139 }
140 }
141
142 pub fn add_module(&self, context: LLVMContext, module: LLVMModule) -> Result<(), String> {
144 unsafe {
145 assert_eq!(
146 context.inner_ref(),
147 LLVMGetModuleContext(module.inner_ref())
148 );
149 let tsctx = LLVMOrcCreateNewThreadSafeContextFromLLVMContext(context.inner_ref());
150 mem::forget(context);
152 let tsm = LLVMOrcCreateNewThreadSafeModule(module.inner_ref(), tsctx);
153 let main_jd = lljit::LLVMOrcLLJITGetMainJITDylib(self.0);
154 let err = lljit::LLVMOrcLLJITAddLLVMIRModule(self.0, main_jd, tsm);
155 LLVMOrcDisposeThreadSafeContext(tsctx);
156 mem::forget(module);
158 handle_err(err)
159 }
160 }
161
162 pub fn lookup_symbol(&self, name: &str) -> Result<u64, String> {
164 unsafe {
165 let mut addr = mem::MaybeUninit::uninit();
166 let err = lljit::LLVMOrcLLJITLookup(self.0, addr.as_mut_ptr(), to_c_str(name).as_ptr());
167 handle_err(err)?;
168 Ok(addr.assume_init())
169 }
170 }
171
172 pub fn get_triple_string(&self) -> String {
174 unsafe {
175 let triple_ptr = lljit::LLVMOrcLLJITGetTripleString(self.0);
176 cstr_to_string(triple_ptr).unwrap()
177 }
178 }
179
180 pub fn add_symbol_mapping(
182 &self,
183 name: &str,
184 addr: u64,
185 flags: JITSymbolGenericFlags,
186 ) -> Result<(), String> {
187 let symbol_pool_ref =
188 unsafe { lljit::LLVMOrcLLJITMangleAndIntern(self.0, to_c_str(name).as_ptr()) };
189
190 let jit_evaluated_symbol = LLVMJITEvaluatedSymbol {
191 Address: addr,
192 Flags: LLVMJITSymbolFlags {
193 GenericFlags: flags.into(),
194 TargetFlags: 0,
195 },
196 };
197
198 let mut symbol_pair = LLVMOrcCSymbolMapPair {
199 Name: symbol_pool_ref,
200 Sym: jit_evaluated_symbol,
201 };
202
203 let materialization_unit = unsafe { LLVMOrcAbsoluteSymbols(&mut symbol_pair as *mut _, 1) };
204 let main_dylib = unsafe { lljit::LLVMOrcLLJITGetMainJITDylib(self.0) };
205
206 let res =
207 unsafe { llvm_sys::orc2::LLVMOrcJITDylibDefine(main_dylib, materialization_unit) };
208 handle_err(res)
209 }
210}
211
212impl Drop for LLVMLLJIT {
213 fn drop(&mut self) {
214 unsafe {
215 let err = lljit::LLVMOrcDisposeLLJIT(self.0);
216 if let Err(err) = handle_err(err) {
217 panic!("Error disposing LLJIT: {}", err);
218 }
219 }
220 }
221}
222
223pub struct SimpleJIT {
255 jit: LLVMLLJIT,
256}
257
258pub struct JitSymbol<'jit, F: Copy> {
264 f: F,
265 _jit: PhantomData<&'jit SimpleJIT>,
266}
267
268impl<'jit, F: Copy> Deref for JitSymbol<'jit, F> {
269 type Target = F;
270
271 fn deref(&self) -> &F {
272 &self.f
273 }
274}
275
276impl SimpleJIT {
277 pub fn new(context: LLVMContext, module: LLVMModule) -> Result<Self, String> {
281 initialize_native()?;
282 let jit = LLVMLLJIT::new_with_default_builder()?;
283 jit.add_module(context, module)?;
284 Ok(SimpleJIT { jit })
285 }
286
287 pub unsafe fn lookup_symbol<'jit, T: Copy>(
294 &'jit self,
295 symbol_name: &str,
296 ) -> Result<JitSymbol<'jit, T>, String> {
297 const { assert!(mem::size_of::<T>() == mem::size_of::<u64>()) };
298 let addr = self.jit.lookup_symbol(symbol_name)?;
299 assert!(addr != 0);
300 let f = unsafe { mem::transmute_copy::<u64, T>(&addr) };
301 Ok(JitSymbol {
302 f,
303 _jit: PhantomData,
304 })
305 }
306}