pliron_llvm/llvm_sys/
lljit.rs1use 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 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 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 LLVMOrcDisposeThreadSafeContext(tsctx);
143 std::mem::forget(module);
145 handle_err(err).inspect_err(|_| {
146 LLVMOrcDisposeThreadSafeModule(tsm);
148 })
149 }
150 }
151
152 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 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 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}