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(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 std::{mem::MaybeUninit, ptr};
49
50use llvm_sys::orc2::{
51    LLVMJITEvaluatedSymbol, LLVMJITSymbolFlags, LLVMJITSymbolGenericFlags, LLVMOrcAbsoluteSymbols,
52    LLVMOrcCSymbolMapPair, LLVMOrcCreateNewThreadSafeContext, LLVMOrcCreateNewThreadSafeModule,
53    LLVMOrcDisposeThreadSafeContext, LLVMOrcDisposeThreadSafeModule, lljit,
54};
55
56use crate::llvm_sys::{
57    core::{LLVMModule, handle_err},
58    cstr_to_string, to_c_str,
59};
60
61bitflags! {
62     #[derive(PartialEq, Eq, Clone, Debug, Hash, Copy)]
63    pub struct JITSymbolGenericFlags: u8 {
64        const JITSymbolGenericFlagsNone = 0;
65        const JITSymbolGenericFlagsExported = 1;
66        const JITSymbolGenericFlagsWeak = 2;
67        const JITSymbolGenericFlagsCallable = 4;
68        const JITSymbolGenericFlagsMaterializationSideEffectsOnly = 8;
69    }
70}
71
72impl From<LLVMJITSymbolGenericFlags> for JITSymbolGenericFlags {
73    fn from(value: LLVMJITSymbolGenericFlags) -> Self {
74        let mut flags = JITSymbolGenericFlags::empty();
75        if (value as u8) & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsExported as u8) != 0
76        {
77            flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsExported;
78        }
79        if (value as u8) & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsWeak as u8) != 0 {
80            flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsWeak;
81        }
82        if (value as u8) & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsCallable as u8) != 0
83        {
84            flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsCallable;
85        }
86        if (value as u8)
87            & (LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsMaterializationSideEffectsOnly
88                as u8)
89            != 0
90        {
91            flags |= JITSymbolGenericFlags::JITSymbolGenericFlagsMaterializationSideEffectsOnly;
92        }
93        flags
94    }
95}
96
97impl From<JITSymbolGenericFlags> for u8 {
98    fn from(value: JITSymbolGenericFlags) -> Self {
99        let mut flags = LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsNone as u8;
100        if value.contains(JITSymbolGenericFlags::JITSymbolGenericFlagsExported) {
101            flags |= LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsExported as u8;
102        }
103        if value.contains(JITSymbolGenericFlags::JITSymbolGenericFlagsWeak) {
104            flags |= LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsWeak as u8;
105        }
106        if value.contains(JITSymbolGenericFlags::JITSymbolGenericFlagsCallable) {
107            flags |= LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsCallable as u8;
108        }
109        if value
110            .contains(JITSymbolGenericFlags::JITSymbolGenericFlagsMaterializationSideEffectsOnly)
111        {
112            flags |=
113                LLVMJITSymbolGenericFlags::LLVMJITSymbolGenericFlagsMaterializationSideEffectsOnly
114                    as u8;
115        }
116        flags
117    }
118}
119
120pub struct LLVMLLJIT(lljit::LLVMOrcLLJITRef);
121
122impl LLVMLLJIT {
123    /// Create a new LLJIT instance with default settings.
124    pub fn new_with_default_builder() -> Result<Self, String> {
125        unsafe {
126            let mut jit = MaybeUninit::uninit();
127            let err = lljit::LLVMOrcCreateLLJIT(jit.as_mut_ptr(), ptr::null_mut());
128            handle_err(err)?;
129            Ok(LLVMLLJIT(jit.assume_init()))
130        }
131    }
132
133    /// Add an [LLVMModule] to the JIT's main JITDylib, in its own thread-safe context.
134    pub fn add_module(&self, module: LLVMModule) -> Result<(), String> {
135        unsafe {
136            let tsctx = LLVMOrcCreateNewThreadSafeContext();
137            let tsm = LLVMOrcCreateNewThreadSafeModule(module.inner_ref(), tsctx);
138            let main_jd = lljit::LLVMOrcLLJITGetMainJITDylib(self.0);
139            let err = lljit::LLVMOrcLLJITAddLLVMIRModule(self.0, main_jd, tsm);
140            // The underlying LLVMContext will be kept alive by our ThreadSafeModule
141            // (See OrcV2CBindingsBasicUsage.c)
142            LLVMOrcDisposeThreadSafeContext(tsctx);
143            // Ownership of the module has been transferred to the JIT
144            std::mem::forget(module);
145            handle_err(err).inspect_err(|_| {
146                // Dispose of the ThreadSafeModule on error
147                LLVMOrcDisposeThreadSafeModule(tsm);
148            })
149        }
150    }
151
152    /// Lookup a symbol in the JIT.
153    pub fn lookup_symbol(&self, name: &str) -> Result<u64, String> {
154        unsafe {
155            let mut addr = MaybeUninit::uninit();
156            let err = lljit::LLVMOrcLLJITLookup(self.0, addr.as_mut_ptr(), to_c_str(name).as_ptr());
157            handle_err(err)?;
158            Ok(addr.assume_init())
159        }
160    }
161
162    /// Get the target triple string for this JIT instance.
163    pub fn get_triple_string(&self) -> String {
164        unsafe {
165            let triple_ptr = lljit::LLVMOrcLLJITGetTripleString(self.0);
166            cstr_to_string(triple_ptr).unwrap()
167        }
168    }
169
170    /// Add a symbol mapping to the JIT's main DyLib
171    pub fn add_symbol_mapping(
172        &self,
173        name: &str,
174        addr: u64,
175        flags: JITSymbolGenericFlags,
176    ) -> Result<(), String> {
177        let symbol_pool_ref =
178            unsafe { lljit::LLVMOrcLLJITMangleAndIntern(self.0, to_c_str(name).as_ptr()) };
179
180        let jit_evaluated_symbol = LLVMJITEvaluatedSymbol {
181            Address: addr,
182            Flags: LLVMJITSymbolFlags {
183                GenericFlags: flags.into(),
184                TargetFlags: 0,
185            },
186        };
187
188        let mut symbol_pair = LLVMOrcCSymbolMapPair {
189            Name: symbol_pool_ref,
190            Sym: jit_evaluated_symbol,
191        };
192
193        let materialization_unit = unsafe { LLVMOrcAbsoluteSymbols(&mut symbol_pair as *mut _, 1) };
194        let main_dylib = unsafe { lljit::LLVMOrcLLJITGetMainJITDylib(self.0) };
195
196        let res =
197            unsafe { llvm_sys::orc2::LLVMOrcJITDylibDefine(main_dylib, materialization_unit) };
198        handle_err(res)
199    }
200}
201
202impl Drop for LLVMLLJIT {
203    fn drop(&mut self) {
204        unsafe {
205            let err = lljit::LLVMOrcDisposeLLJIT(self.0);
206            if let Err(err) = handle_err(err) {
207                panic!("Error disposing LLJIT: {}", err);
208            }
209        }
210    }
211}