Skip to main content

pliron_llvm/llvm_sys/
lljit.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) The pliron contributors
3
4//! Safe(r) wrappers around llvm_sys::lljit
5//!
6//! ### Example
7//!```
8//! use pliron_llvm::llvm_sys::target::initialize_native;
9//! use pliron_llvm::llvm_sys::core::{LLVMContext, LLVMModule, LLVMMemoryBuffer};
10//! use pliron_llvm::llvm_sys::lljit::{LLVMLLJIT, JITSymbolGenericFlags};
11//! fn main() -> Result<(), String> {
12//!    initialize_native()?;
13//!    let context = LLVMContext::default();
14//!
15//!    fn my_rust_adder(a: i32, b: i32) -> i32 {
16//!        a + b
17//!    }
18//!
19//!    let ir = r#"
20//!      declare i32 @my_rust_adder(i32, i32)
21//!      define i32 @add(i32 %a, i32 %b) {
22//!          %sum = call i32 @my_rust_adder(i32 %a, i32 %b)
23//!          ret i32 %sum
24//!      }"#;
25//!    let ir_mb = LLVMMemoryBuffer::from_str(ir, "test_buffer");
26//!    let module = LLVMModule::from_ir_in_memory_buffer(&context, ir_mb)?;
27//!
28//!    let jit = LLVMLLJIT::new_with_default_builder()?;
29//!    jit.add_module(context, module)?;
30//!    // Add the Rust function as a symbol mapping
31//!    let rust_adder_addr = my_rust_adder as *const () as u64;
32//!    jit.add_symbol_mapping
33//!     ("my_rust_adder", rust_adder_addr,
34//!       JITSymbolGenericFlags::JITSymbolGenericFlagsCallable
35//!         | JITSymbolGenericFlags::JITSymbolGenericFlagsExported)?;
36//!
37//!    // Get symbol address for 'add' in the LLVM module
38//!    let symbol_addr = jit.lookup_symbol("add")?;
39//!    assert!(symbol_addr != 0);
40//!
41//!    let adder = unsafe { std::mem::transmute::<u64, fn(i32, i32) -> i32>(symbol_addr) };
42//!    assert_eq!(adder(2, 3), 5);
43//!    Ok(())
44//! }
45//! ```
46
47use 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
126/// Wrapper around LLVM's LLJIT.
127///
128/// For simple uses, considering using [SimpleJIT].
129pub struct LLVMLLJIT(lljit::LLVMOrcLLJITRef);
130
131impl LLVMLLJIT {
132    /// Create a new LLJIT instance with default settings.
133    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    /// Add an [LLVMModule] (contained in [LLVMContext]) to the JIT's main JITDylib
143    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            // Ownership of the context has been transferred to the thread-safe context
151            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            // Ownership of the module has been transferred to the JIT
157            mem::forget(module);
158            handle_err(err)
159        }
160    }
161
162    /// Lookup a symbol in the JIT.
163    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    /// Get the target triple string for this JIT instance.
173    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    /// Add a symbol mapping to the JIT's main DyLib
181    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
223/// A convenience wrapper around [LLVMLLJIT].
224///
225/// It takes ownership of the [LLVMModule] to be compiled and its owning [LLVMContext].
226///
227/// ### Example
228/// ```
229/// use pliron_llvm::llvm_sys::core::{LLVMContext, LLVMMemoryBuffer, LLVMModule};
230/// use pliron_llvm::llvm_sys::lljit::SimpleJIT;
231///
232/// let context = LLVMContext::default();
233///
234/// let ir = r#"
235///   @sum_global = global i32 0
236///
237///   define i32 @add(i32 %a, i32 %b) {
238///       %sum = add i32 %a, %b
239///       store i32 %sum, ptr @sum_global
240///       ret i32 %sum
241///   }"#;
242///
243/// let module = LLVMModule::from_ir_in_str(&context, ir, None).unwrap();
244///
245/// let jit = SimpleJIT::new(context, module).unwrap();
246/// let adder = unsafe { jit.lookup_symbol::<fn(i32, i32) -> i32>("add").unwrap() };
247/// assert_eq!(adder(2, 3), 5);
248///
249/// // The global is updated as a side effect of calling `add`, and can be
250/// // looked up (as a pointer to its storage).
251/// let sum_global = unsafe { jit.lookup_symbol::<*const i32>("sum_global").unwrap() };
252/// assert_eq!(unsafe { **sum_global }, 5);
253/// ```
254pub struct SimpleJIT {
255    jit: LLVMLLJIT,
256}
257
258/// A symbol looked up in [SimpleJIT], typed as the function pointer `F`.
259///
260/// Borrows the [SimpleJIT] it was looked up from, so it can't outlive the JIT
261/// (and therefore can't outlive the compiled code it points into). Dereferences
262/// to `F`, so it can be called directly.
263pub 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    /// Create a JIT from `module` and the `context` that contains it.
278    ///
279    /// See type documentation for examples.
280    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    /// Lookup a symbol in the JIT and interpret it to be of type `F`.
288    ///
289    /// See type documentation for examples.
290    ///
291    /// # Safety
292    /// The type of the symbol must be correct.
293    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}