use std::collections::HashMap;
use std::mem;
use cranelift_codegen::ir::{self, AbiParam, InstBuilder, types};
use cranelift_codegen::settings::{self, Configurable};
use cranelift_frontend::{FunctionBuilder, FunctionBuilderContext};
use cranelift_jit::{JITBuilder, JITModule};
use cranelift_module::{Linkage, Module};
use crate::ast::PolydatNode;
use crate::ast::SlotShape;
use super::kernels::{JitCore, JitKernelPushPull, JitKernelRaw};
extern "C" fn jit_xxh3_hash(value: u64) -> u64 {
guarded(|| xxhash_rust::xxh3::xxh3_64(&value.to_le_bytes()))
}
extern "C" fn jit_interleave(a: u64, b: u64) -> u64 {
guarded(|| {
let mut result: u64 = 0;
for i in 0..32 {
result |= ((a >> i) & 1) << (2 * i);
result |= ((b >> i) & 1) << (2 * i + 1);
}
result
})
}
extern "C" fn jit_lut_sample(input_bits: u64, lut_ptr: u64, lut_len: u64) -> u64 {
guarded(|| {
let u = f64::from_bits(input_bits).clamp(0.0, 1.0);
let n = (lut_len - 1) as f64;
let pos = u * n;
let idx = (pos as usize).min(lut_len as usize - 2);
let frac = pos - idx as f64;
let result = unsafe {
let ptr = lut_ptr as *const f64;
let a = *ptr.add(idx);
let b = *ptr.add(idx + 1);
a * (1.0 - frac) + b * frac
};
result.to_bits()
})
}
extern "C" fn jit_shuffle(input: u64, feedback: u64, size: u64, min: u64) -> u64 {
guarded(|| crate::numeric::permute::shuffle_bounded(input, feedback, size, min))
}
extern "C" fn jit_sin(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).sin().to_bits())
}
extern "C" fn jit_cos(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).cos().to_bits())
}
extern "C" fn jit_tan(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).tan().to_bits())
}
extern "C" fn jit_asin(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).asin().to_bits())
}
extern "C" fn jit_acos(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).acos().to_bits())
}
extern "C" fn jit_atan(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).atan().to_bits())
}
extern "C" fn jit_sqrt(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).sqrt().to_bits())
}
extern "C" fn jit_abs_f64(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).abs().to_bits())
}
extern "C" fn jit_ln(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).ln().to_bits())
}
extern "C" fn jit_exp(bits: u64) -> u64 {
guarded(|| f64::from_bits(bits).exp().to_bits())
}
extern "C" fn jit_floor_base10(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
floor_pow10(x)
};
r.to_bits()
})
}
extern "C" fn jit_ceiling_base10(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
let lo = floor_pow10(x);
if lo == x { lo } else { lo * 10.0 }
};
r.to_bits()
})
}
extern "C" fn jit_closest_base10(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
let lo = floor_pow10(x);
let hi = if lo == x { lo } else { lo * 10.0 };
pick_closest(x, lo, hi)
};
r.to_bits()
})
}
extern "C" fn jit_floor_decade(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
let base = floor_pow10(x);
(x / base).floor() * base
};
r.to_bits()
})
}
extern "C" fn jit_ceiling_decade(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
let base = floor_pow10(x);
(x / base).ceil() * base
};
r.to_bits()
})
}
extern "C" fn jit_closest_decade(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
let base = floor_pow10(x);
(x / base).round() * base
};
r.to_bits()
})
}
extern "C" fn jit_floor_binomial(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
floor_pow2(x)
};
r.to_bits()
})
}
extern "C" fn jit_ceiling_binomial(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
let lo = floor_pow2(x);
if lo == x { lo } else { lo * 2.0 }
};
r.to_bits()
})
}
extern "C" fn jit_closest_binomial(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
let lo = floor_pow2(x);
let hi = if lo == x { lo } else { lo * 2.0 };
pick_closest(x, lo, hi)
};
r.to_bits()
})
}
extern "C" fn jit_floor_fibonacci(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
floor_fibonacci_val(x)
};
r.to_bits()
})
}
extern "C" fn jit_ceiling_fibonacci(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
ceiling_fibonacci_val(x)
};
r.to_bits()
})
}
extern "C" fn jit_closest_fibonacci(bits: u64) -> u64 {
guarded(|| {
use crate::numeric::round_numbers::*;
let x = f64::from_bits(bits);
let r = if !positive_finite(x) {
0.0
} else {
pick_closest(x, floor_fibonacci_val(x), ceiling_fibonacci_val(x))
};
r.to_bits()
})
}
extern "C" fn jit_atan2(y_bits: u64, x_bits: u64) -> u64 {
guarded(|| {
f64::from_bits(y_bits)
.atan2(f64::from_bits(x_bits))
.to_bits()
})
}
extern "C" fn jit_pow(base_bits: u64, exp_bits: u64) -> u64 {
guarded(|| {
f64::from_bits(base_bits)
.powf(f64::from_bits(exp_bits))
.to_bits()
})
}
extern "C" fn jit_round_nearest(x_bits: u64, iv_bits: u64) -> u64 {
guarded(|| {
let x = f64::from_bits(x_bits);
let interval = f64::from_bits(iv_bits);
let r = if !(interval.is_finite() && interval > 0.0) {
x
} else {
(x / interval).round() * interval
};
r.to_bits()
})
}
extern "C" fn jit_round_floor(x_bits: u64, iv_bits: u64) -> u64 {
guarded(|| {
let x = f64::from_bits(x_bits);
let interval = f64::from_bits(iv_bits);
let r = if !(interval.is_finite() && interval > 0.0) {
x
} else {
(x / interval).floor() * interval
};
r.to_bits()
})
}
extern "C" fn jit_round_ceiling(x_bits: u64, iv_bits: u64) -> u64 {
guarded(|| {
let x = f64::from_bits(x_bits);
let interval = f64::from_bits(iv_bits);
let r = if !(interval.is_finite() && interval > 0.0) {
x
} else {
(x / interval).ceil() * interval
};
r.to_bits()
})
}
extern "C" fn jit_pcg(input: u64, seed: u64, stream: u64) -> u64 {
guarded(|| {
let inc = 2u64.wrapping_mul(stream).wrapping_add(1);
crate::numeric::pcg::pcg_seek(seed, inc, input)
})
}
extern "C" fn jit_pcg_stream(input: u64, stream: u64, seed: u64) -> u64 {
guarded(|| {
let inc = 2u64.wrapping_mul(stream).wrapping_add(1);
crate::numeric::pcg::pcg_seek(seed, inc, input)
})
}
extern "C" fn jit_n_of(input: u64, n: u64, m: u64) -> u64 {
guarded(|| {
if m == 0 {
return 0;
}
crate::numeric::n_of_m::n_of_m_eval(input, n, m)
})
}
extern "C" fn jit_cycle_walk(pos: u64, range: u64, seed: u64, inc: u64) -> u64 {
guarded(|| {
let stream = inc.saturating_sub(1) / 2;
let state = crate::numeric::pcg::build_cycle_walk_state(range, seed, stream);
crate::numeric::pcg::cycle_walk_inner(
pos,
range,
state.half_bits,
state.half_mask,
&state.round_keys,
)
})
}
extern "C" fn jit_perlin_1d(input: u64, perm_ptr: u64, freq_bits: u64) -> u64 {
guarded(|| {
let perm = unsafe { &*(perm_ptr as *const crate::numeric::noise::PermTable) };
let freq = f64::from_bits(freq_bits);
let r = crate::numeric::noise::perlin_1d_algo(perm, input as f64 * freq);
r.to_bits()
})
}
extern "C" fn jit_perlin_2d(x: u64, y: u64, perm_ptr: u64, freq_bits: u64) -> u64 {
guarded(|| {
let perm = unsafe { &*(perm_ptr as *const crate::numeric::noise::PermTable) };
let freq = f64::from_bits(freq_bits);
let r = crate::numeric::noise::perlin_2d_algo(perm, x as f64 * freq, y as f64 * freq);
r.to_bits()
})
}
extern "C" fn jit_simplex_2d(x: u64, y: u64, perm_ptr: u64, freq_bits: u64) -> u64 {
guarded(|| {
let perm = unsafe { &*(perm_ptr as *const crate::numeric::noise::PermTable) };
let freq = f64::from_bits(freq_bits);
let r = crate::numeric::noise::simplex_2d_algo(perm, x as f64 * freq, y as f64 * freq);
r.to_bits()
})
}
extern "C" fn jit_fractal_noise_1d(input: u64, perm_ptr: u64, freq_bits: u64, octaves: u64) -> u64 {
guarded(|| {
let perm = unsafe { &*(perm_ptr as *const crate::numeric::noise::PermTable) };
let freq = f64::from_bits(freq_bits);
let r = crate::numeric::noise::fbm_1d(perm, input as f64, freq, octaves as u32);
r.to_bits()
})
}
extern "C" fn jit_fractal_noise_2d(
x: u64,
y: u64,
perm_ptr: u64,
freq_bits: u64,
octaves: u64,
) -> u64 {
guarded(|| {
let perm = unsafe { &*(perm_ptr as *const crate::numeric::noise::PermTable) };
let freq = f64::from_bits(freq_bits);
let r = crate::numeric::noise::fbm_2d(perm, x as f64, y as f64, freq, octaves as u32);
r.to_bits()
})
}
extern "C" fn jit_thread_id() -> u64 {
guarded(|| {
let id = std::thread::current().id();
let id_str = format!("{id:?}");
let num = id_str.trim_start_matches("ThreadId(").trim_end_matches(')');
num.parse().unwrap_or(0)
})
}
extern "C" fn jit_current_epoch_millis() -> u64 {
guarded(|| {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_millis() as u64
})
}
#[repr(C, align(16))]
struct JitJmpBuf([u8; 512]);
#[cfg(not(windows))]
unsafe extern "C" {
fn _setjmp(env: *mut JitJmpBuf) -> i32;
fn _longjmp(env: *mut JitJmpBuf, val: i32) -> !;
}
#[cfg(windows)]
unsafe extern "C" {
fn _setjmp(env: *mut JitJmpBuf, frame: *mut std::ffi::c_void) -> i32;
#[link_name = "longjmp"]
fn _longjmp(env: *mut JitJmpBuf, val: i32) -> !;
}
use std::cell::{Cell, RefCell};
thread_local! {
static JIT_JMP_BUF: Cell<Option<*mut JitJmpBuf>> = const { Cell::new(None) };
static JIT_VIOLATION_MSG: RefCell<Option<String>> = const { RefCell::new(None) };
}
fn jit_violation_longjmp(msg: String) -> ! {
JIT_VIOLATION_MSG.with(|m| *m.borrow_mut() = Some(msg.clone()));
let buf_ptr: Option<*mut JitJmpBuf> = JIT_JMP_BUF.with(|b| b.get());
match buf_ptr {
Some(ptr) => unsafe { _longjmp(ptr, 1) },
None => {
let mut err = std::io::stderr().lock();
use std::io::Write;
let _ = writeln!(err, "{msg}");
let _ = err.flush();
std::process::abort();
}
}
}
struct JmpBufGuard {
prev: Option<*mut JitJmpBuf>,
}
impl Drop for JmpBufGuard {
fn drop(&mut self) {
JIT_JMP_BUF.with(|b| b.set(self.prev));
}
}
pub(crate) fn invoke_with_catch<F: FnOnce()>(f: F) {
use std::mem::MaybeUninit;
let mut buf: MaybeUninit<JitJmpBuf> = MaybeUninit::uninit();
let buf_ptr = buf.as_mut_ptr();
let prev: Option<*mut JitJmpBuf> = JIT_JMP_BUF.with(|b| b.replace(Some(buf_ptr)));
let _guard = JmpBufGuard { prev };
#[cfg(not(windows))]
let jmpval = unsafe { _setjmp(buf_ptr) };
#[cfg(windows)]
let jmpval = unsafe { _setjmp(buf_ptr, std::ptr::null_mut()) };
if jmpval == 0 {
f();
} else {
let msg = JIT_VIOLATION_MSG
.with(|m| m.borrow_mut().take())
.unwrap_or_else(|| "JIT predicate violation (no message)".into());
std::panic::resume_unwind(Box::new(msg));
}
}
extern "C" fn jit_is_positive_fail(value: u64, name_ptr: u64, name_len: u64) -> u64 {
let name = if name_ptr != 0 {
unsafe {
std::str::from_utf8_unchecked(std::slice::from_raw_parts(
name_ptr as *const u8,
name_len as usize,
))
}
} else {
"value"
};
jit_violation_longjmp(format!(
"is_positive({name}): value must be > 0, got {value}"
));
}
extern "C" fn jit_in_range_fail(value: u64, lo: u64, hi: u64) -> u64 {
jit_violation_longjmp(format!("in_range: value {value} outside [{lo}, {hi}]"));
}
extern "C" fn jit_div_zero_fail(kind: u64) -> u64 {
jit_violation_longjmp(
if kind == 0 {
"attempt to divide by zero"
} else {
"attempt to calculate the remainder with a divisor of zero"
}
.to_string(),
);
}
extern "C" fn jit_f64_mod(a_bits: u64, b_bits: u64) -> u64 {
let (a, b) = (f64::from_bits(a_bits), f64::from_bits(b_bits));
(if b != 0.0 { a % b } else { 0.0 }).to_bits()
}
extern "C" fn jit_is_one_of_fail(value: u64, set_ptr: u64, set_len: u64) -> u64 {
let msg = if set_ptr != 0 {
let set = unsafe { std::slice::from_raw_parts(set_ptr as *const u64, set_len as usize) };
format!("is_one_of: value {value} not in allowed set {set:?}")
} else {
format!("is_one_of: value {value} not in allowed set [..]")
};
jit_violation_longjmp(msg);
}
extern "C" fn jit_weighted_pick(
input: u64,
values_ptr: u64,
biases_ptr: u64,
primaries_ptr: u64,
aliases_ptr: u64,
n: u64,
) -> u64 {
guarded(|| {
let n = n as usize;
let slot = (input as usize) % n;
let bias_test = ((input >> 32) as f64) / (u32::MAX as f64);
unsafe {
let biases = std::slice::from_raw_parts(biases_ptr as *const f64, n);
let primaries = std::slice::from_raw_parts(primaries_ptr as *const u64, n);
let aliases = std::slice::from_raw_parts(aliases_ptr as *const u64, n);
let values = std::slice::from_raw_parts(values_ptr as *const u64, n);
let index = if bias_test < biases[slot] {
primaries[slot]
} else {
aliases[slot]
};
values[index as usize]
}
})
}
fn guarded<T>(body: impl FnOnce() -> T) -> T {
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(body)) {
Ok(v) => v,
Err(payload) => {
let msg = payload
.downcast_ref::<String>()
.cloned()
.or_else(|| payload.downcast_ref::<&str>().map(|s| s.to_string()))
.unwrap_or_else(|| "panic in a compiled helper".to_string());
jit_violation_longjmp(msg)
}
}
}
#[derive(Clone)]
pub struct SlotKitRef(pub std::sync::Arc<crate::ast::CompiledSlotKit>);
impl SlotKitRef {
fn new(kit: crate::ast::CompiledSlotKit) -> Self {
SlotKitRef(std::sync::Arc::new(kit))
}
}
impl std::fmt::Debug for SlotKitRef {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"SlotKitRef({:p}, {} scratch)",
std::sync::Arc::as_ptr(&self.0),
self.0.scratch.len()
)
}
}
impl PartialEq for SlotKitRef {
fn eq(&self, other: &Self) -> bool {
std::sync::Arc::ptr_eq(&self.0, &other.0)
}
}
extern "C" fn jit_slot_call(
kit: *const crate::ast::CompiledSlotKit,
inputs: *const u64,
n_in: u64,
outputs: *mut u64,
n_out: u64,
scratch: *mut crate::ast::ScratchBuf,
base: u64,
n_scratch: u64,
) {
guarded(|| unsafe {
let kit = &*kit;
let ins = std::slice::from_raw_parts(inputs, n_in as usize);
let outs = std::slice::from_raw_parts_mut(outputs, n_out as usize);
let sc = std::slice::from_raw_parts_mut(scratch.add(base as usize), n_scratch as usize);
(kit.op)(ins, outs, sc)
})
}
impl JitOp {
pub(crate) fn slot_kit(&self) -> Option<&SlotKitRef> {
match self {
JitOp::SlotCall { kit, .. } | JitOp::Convert { kit, .. } => Some(kit),
_ => None,
}
}
pub(crate) fn scratch_elems(&self) -> &[crate::ast::ScratchElem] {
const STR_ENTRY: [crate::ast::ScratchElem; 1] = [crate::ast::ScratchElem::Str];
const F32_ENTRY: [crate::ast::ScratchElem; 1] = [crate::ast::ScratchElem::F32];
match self {
JitOp::SlotCall { kit, .. } | JitOp::Convert { kit, .. } => &kit.0.scratch,
JitOp::U64ToStr { .. }
| JitOp::I64ToStr { .. }
| JitOp::F64ToStr { .. }
| JitOp::StrConcat { .. }
| JitOp::JsonToStr { .. } => &STR_ENTRY,
JitOp::VecProduce { .. } => &F32_ENTRY,
_ => &[],
}
}
pub(crate) fn place_scratch(&mut self, base: usize) {
match self {
JitOp::SlotCall { scratch_base, .. }
| JitOp::Convert { scratch_base, .. }
| JitOp::U64ToStr { scratch_base }
| JitOp::I64ToStr { scratch_base }
| JitOp::F64ToStr { scratch_base }
| JitOp::StrConcat { scratch_base }
| JitOp::JsonToStr { scratch_base }
| JitOp::VecProduce { scratch_base, .. } => *scratch_base = base,
_ => {}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VecProducer {
Add,
Scale,
Norm,
HashVec,
XxHash3Vec,
RegToVec,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VecReducer {
Dot,
L2,
Cosine,
LidMle,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RegLaneRead {
F32,
I16,
I64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RegProducer {
WithLaneF32,
GatherF32,
VecToRegF32,
MulI8,
}
unsafe fn vec_f32_of<'a>(ptr: u64, len: u64) -> &'a [f32] {
if len == 0 {
&[]
} else {
unsafe { std::slice::from_raw_parts(ptr as usize as *const f32, len as usize) }
}
}
unsafe fn write_f32_entry(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
f: impl FnOnce(&mut Vec<f32>),
) {
unsafe {
let entry = &mut *scratch.add(base as usize);
let crate::ast::ScratchBuf::F32(v) = entry else {
panic!("a vector lowering's scratch entry is not an f32 vector");
};
f(v);
*buffer.add(out_slot as usize) = v.as_ptr() as usize as u64;
*buffer.add(out_slot as usize + 1) = v.len() as u64;
}
}
macro_rules! vec_producer {
($name:ident, |$out:ident, $w0:ident, $w1:ident, $w2:ident, $w3:ident| $body:expr) => {
extern "C" fn $name(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
$w0: u64,
$w1: u64,
$w2: u64,
$w3: u64,
) {
guarded(|| unsafe {
let _ = ($w2, $w3);
write_f32_entry(scratch, base, buffer, out_slot, |$out| $body)
})
}
};
}
vec_producer!(jit_vec_add, |out, a_ptr, a_len, b_ptr, b_len| {
let (a, b) = (vec_f32_of(a_ptr, a_len), vec_f32_of(b_ptr, b_len));
crate::numeric::vector::check_lens("vec_add", a.len(), b.len());
crate::numeric::vector::add_f32_into(a, b, out)
});
vec_producer!(jit_vec_scale, |out, a_ptr, a_len, k_bits, _z| {
let a = vec_f32_of(a_ptr, a_len);
crate::numeric::vector::scale_f32_into(a, f64::from_bits(k_bits) as f32, out)
});
vec_producer!(jit_vec_norm, |out, a_ptr, a_len, _y, _z| {
crate::numeric::vector::norm_f32_into(vec_f32_of(a_ptr, a_len), out)
});
vec_producer!(jit_hash_vec, |out, seed, dim, _y, _z| {
crate::numeric::vector::hash_vec_into(seed, dim, out)
});
vec_producer!(jit_xxhash3_vec, |out, seed, dim, _y, _z| {
crate::numeric::vector::xxhash3_vec_into(seed, dim, out)
});
vec_producer!(jit_reg_to_vec_f32, |out, lo, hi, _y, _z| {
out.clear();
out.extend_from_slice(&crate::ast::Bits128([lo, hi]).lanes_f32())
});
macro_rules! vec_reducer {
($name:ident, |$w0:ident, $w1:ident, $w2:ident, $w3:ident| $body:expr) => {
extern "C" fn $name($w0: u64, $w1: u64, $w2: u64, $w3: u64) -> u64 {
guarded(|| unsafe {
let _ = ($w2, $w3);
let r: f64 = $body;
r.to_bits()
})
}
};
}
vec_reducer!(jit_vec_dot, |a_ptr, a_len, b_ptr, b_len| {
let (a, b) = (vec_f32_of(a_ptr, a_len), vec_f32_of(b_ptr, b_len));
crate::numeric::vector::check_lens("vec_dot", a.len(), b.len());
crate::numeric::vector::dot_f32(a, b) as f64
});
vec_reducer!(jit_vec_l2, |a_ptr, a_len, b_ptr, b_len| {
let (a, b) = (vec_f32_of(a_ptr, a_len), vec_f32_of(b_ptr, b_len));
crate::numeric::vector::check_lens("vec_l2", a.len(), b.len());
(crate::numeric::vector::l2sq_f32(a, b) as f64).sqrt()
});
vec_reducer!(jit_vec_cosine, |a_ptr, a_len, b_ptr, b_len| {
let (a, b) = (vec_f32_of(a_ptr, a_len), vec_f32_of(b_ptr, b_len));
crate::numeric::vector::check_lens("vec_cosine", a.len(), b.len());
crate::numeric::vector::cosine_f32(a, b)
});
vec_reducer!(jit_lid_mle, |d_ptr, d_len, k_bits, _z| {
crate::numeric::vector::lid_mle_of(vec_f32_of(d_ptr, d_len), f64::from_bits(k_bits))
});
extern "C" fn jit_reg_lane_f32(lo: u64, hi: u64, i: u64) -> u64 {
guarded(|| crate::numeric::register::lane_f32(crate::ast::Bits128([lo, hi]), i).to_bits())
}
extern "C" fn jit_reg_lane_i16(lo: u64, hi: u64, i: u64) -> u64 {
guarded(|| crate::numeric::register::lane_i16(crate::ast::Bits128([lo, hi]), i) as i64 as u64)
}
extern "C" fn jit_reg_lane_i64(lo: u64, hi: u64, i: u64) -> u64 {
guarded(|| crate::numeric::register::lane_i64(crate::ast::Bits128([lo, hi]), i) as u64)
}
macro_rules! reg_producer {
($name:ident, |$w0:ident, $w1:ident, $w2:ident, $w3:ident| $body:expr) => {
extern "C" fn $name(
buffer: *mut u64,
out_slot: u64,
$w0: u64,
$w1: u64,
$w2: u64,
$w3: u64,
) {
guarded(|| unsafe {
let _ = ($w2, $w3);
let r: crate::ast::Bits128 = $body;
*buffer.add(out_slot as usize) = r.0[0];
*buffer.add(out_slot as usize + 1) = r.0[1];
})
}
};
}
reg_producer!(jit_reg_with_lane_f32, |lo, hi, i, v_bits| {
crate::numeric::register::with_lane_f32(
crate::ast::Bits128([lo, hi]),
i,
f64::from_bits(v_bits),
)
});
reg_producer!(jit_reg_gather_f32, |v_ptr, v_len, offset, _z| {
crate::numeric::register::gather_f32(vec_f32_of(v_ptr, v_len), offset)
});
reg_producer!(jit_vec_to_reg_f32, |v_ptr, v_len, _y, _z| {
crate::numeric::register::to_reg_f32(vec_f32_of(v_ptr, v_len))
});
reg_producer!(jit_reg_mul_i8, |a_lo, a_hi, b_lo, b_hi| {
crate::numeric::register::mul_i8(
crate::ast::Bits128([a_lo, a_hi]),
crate::ast::Bits128([b_lo, b_hi]),
)
});
unsafe fn write_str_entry(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
f: impl FnOnce(&mut Vec<u8>),
) {
unsafe {
let entry = &mut *scratch.add(base as usize);
let crate::ast::ScratchBuf::Str(v) = entry else {
panic!("a string lowering's scratch entry is not a string");
};
v.clear();
f(v);
*buffer.add(out_slot as usize) = v.as_ptr() as usize as u64;
*buffer.add(out_slot as usize + 1) = v.len() as u64;
}
}
extern "C" fn jit_u64_to_str(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
value: u64,
) {
use std::io::Write;
guarded(|| unsafe {
write_str_entry(scratch, base, buffer, out_slot, |v| {
write!(v, "{value}").expect("a vector accepts every write")
})
})
}
extern "C" fn jit_i64_to_str(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
value: u64,
) {
use std::io::Write;
guarded(|| unsafe {
write_str_entry(scratch, base, buffer, out_slot, |v| {
write!(v, "{}", value as i64).expect("a vector accepts every write")
})
})
}
extern "C" fn jit_f64_to_str(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
bits: u64,
) {
use std::io::Write;
guarded(|| unsafe {
write_str_entry(scratch, base, buffer, out_slot, |v| {
write!(v, "{}", f64::from_bits(bits)).expect("a vector accepts every write")
})
})
}
extern "C" fn jit_str_concat(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
pairs: *const u64,
n: u64,
) {
guarded(|| unsafe {
let words = std::slice::from_raw_parts(pairs, 2 * n as usize);
write_str_entry(scratch, base, buffer, out_slot, |v| {
for pair in words.as_chunks::<2>().0 {
let bytes =
std::slice::from_raw_parts(pair[0] as usize as *const u8, pair[1] as usize);
v.extend_from_slice(bytes);
}
})
})
}
extern "C" fn jit_json_to_str(
scratch: *mut crate::ast::ScratchBuf,
base: u64,
buffer: *mut u64,
out_slot: u64,
ptr: u64,
len: u64,
) {
guarded(|| unsafe {
let pair = [ptr, len];
let value = crate::derive_support::ref_value(&pair);
let json = match value {
crate::ast::Value::Json(j) => j.as_ref(),
other => panic!("expected Json wire, got {other:?}"),
};
write_str_entry(scratch, base, buffer, out_slot, |v| {
serde_json::to_writer(v, json).expect("a vector accepts every write")
})
})
}
#[derive(Debug, Clone, Copy, PartialEq)]
enum Scalar {
Unsigned(u32),
Signed(u32),
Bool,
F32,
F64,
}
impl Scalar {
fn of(t: crate::ast::PortType) -> Option<Self> {
use crate::ast::PortType as P;
Some(match t {
P::U8 => Self::Unsigned(8),
P::U16 => Self::Unsigned(16),
P::U32 => Self::Unsigned(32),
P::U64 => Self::Unsigned(64),
P::I8 => Self::Signed(8),
P::I16 => Self::Signed(16),
P::I32 => Self::Signed(32),
P::I64 => Self::Signed(64),
P::Bool => Self::Bool,
P::F32 => Self::F32,
P::F64 => Self::F64,
_ => return None,
})
}
fn int_range(self) -> Option<(i128, i128)> {
match self {
Self::Unsigned(b) => Some((0, (1i128 << b) - 1)),
Self::Signed(b) => Some((-(1i128 << (b - 1)), (1i128 << (b - 1)) - 1)),
Self::Bool => Some((0, 1)),
Self::F32 | Self::F64 => None,
}
}
}
fn conversion_op(node: &dyn PolydatNode) -> Option<JitOp> {
use crate::ast::Slot;
let meta = node.meta();
let [Slot::Wire(input)] = meta.ins.as_slice() else {
return None;
};
let [output] = meta.outs.as_slice() else {
return None;
};
let (from, to) = (input.typ, output.typ);
Scalar::of(from)?;
Scalar::of(to)?;
let canonical = crate::compile::assembly::boundary_adapter(from, to)?;
if canonical.meta().name != meta.name {
return None;
}
let op = node.compiled_u64()?;
Some(JitOp::Convert {
from,
to,
kit: SlotKitRef::new(crate::ast::CompiledSlotKit {
scratch: Vec::new(),
op: Box::new(move |inputs, outputs, _| op(inputs, outputs)),
}),
scratch_base: 0,
})
}
pub fn classify_node_typed(node: &dyn PolydatNode, wire_types: &[crate::ast::PortType]) -> JitOp {
use crate::ast::PortType as PT;
let is_ref = |t: &crate::ast::PortType| t.slot_color() == crate::ast::SlotColor::Ref2;
let vec_produce = |kind: VecProducer| JitOp::VecProduce {
kind,
scratch_base: 0,
};
let ref_copy = |ty: crate::ast::PortType| {
crate::compile::assembly::ref_copy_kit(ty)
.map(|kit| JitOp::SlotCall {
kit: SlotKitRef::new(kit),
scratch_base: 0,
})
.unwrap_or(JitOp::Fallback)
};
let meta = node.meta();
let named = match meta.name.as_str() {
"__u64_to_string" => JitOp::U64ToStr { scratch_base: 0 },
"__i64_to_string" => JitOp::I64ToStr { scratch_base: 0 },
"__f64_to_string" => JitOp::F64ToStr { scratch_base: 0 },
"json_to_str" if wire_types == [crate::ast::PortType::Json] => {
JitOp::JsonToStr { scratch_base: 0 }
}
"str_concat"
if !wire_types.is_empty()
&& wire_types.iter().all(|t| *t == crate::ast::PortType::Str) =>
{
JitOp::StrConcat { scratch_base: 0 }
}
"vec_add" if wire_types == [PT::VecF32, PT::VecF32] => vec_produce(VecProducer::Add),
"vec_scale" if wire_types == [PT::VecF32, PT::F64] => vec_produce(VecProducer::Scale),
"vec_norm" if wire_types == [PT::VecF32] => vec_produce(VecProducer::Norm),
"hash_vec" if wire_types == [PT::U64, PT::U64] => vec_produce(VecProducer::HashVec),
"xxhash3_vec" if wire_types == [PT::U64, PT::U64] => vec_produce(VecProducer::XxHash3Vec),
"reg_to_vec_f32" if wire_types == [PT::RegF32x4] => vec_produce(VecProducer::RegToVec),
"vec_dot" if wire_types == [PT::VecF32, PT::VecF32] => JitOp::VecReduce(VecReducer::Dot),
"vec_l2" if wire_types == [PT::VecF32, PT::VecF32] => JitOp::VecReduce(VecReducer::L2),
"vec_cosine" if wire_types == [PT::VecF32, PT::VecF32] => {
JitOp::VecReduce(VecReducer::Cosine)
}
"lid_mle" if wire_types == [PT::VecF32, PT::F64] => JitOp::VecReduce(VecReducer::LidMle),
"reg_lane_f32" if wire_types == [PT::RegF32x4, PT::U64] => JitOp::RegLane(RegLaneRead::F32),
"reg_lane_i16" if wire_types == [PT::RegI16x8, PT::U64] => JitOp::RegLane(RegLaneRead::I16),
"reg_lane_i64" if wire_types == [PT::RegI64x2, PT::U64] => JitOp::RegLane(RegLaneRead::I64),
"reg_with_lane_f32" if wire_types == [PT::RegF32x4, PT::U64, PT::F64] => {
JitOp::RegProduce(RegProducer::WithLaneF32)
}
"reg_gather_f32" if wire_types == [PT::VecF32, PT::U64] => {
JitOp::RegProduce(RegProducer::GatherF32)
}
"vec_to_reg_f32" if wire_types == [PT::VecF32] => {
JitOp::RegProduce(RegProducer::VecToRegF32)
}
"reg_mul_i8" if wire_types == [PT::RegI8x16, PT::RegI8x16] => {
JitOp::RegProduce(RegProducer::MulI8)
}
"reg_dot_f32" if wire_types == [PT::RegF32x4, PT::RegF32x4] => JitOp::RegDotF32,
n if n.starts_with("__port_") || n == "default_or" => match meta.outs.first() {
Some(o) if is_ref(&o.typ) => return ref_copy(o.typ),
_ => JitOp::Identity,
},
"select" | "select_u64" if wire_types.iter().skip(1).any(|t| t.slot_width() != 1) => {
JitOp::Fallback
}
_ if wire_types.iter().any(is_ref) || meta.outs.iter().any(|o| is_ref(&o.typ)) => {
JitOp::Fallback
}
_ => classify_node(node),
};
if !matches!(named, JitOp::Fallback) {
return named;
}
if let Some(kit) = node.compiled_slot(
wire_types,
crate::compile::select::Engine::Native(crate::compile::select::Provenance::Auto),
) {
return JitOp::SlotCall {
kit: SlotKitRef::new(kit),
scratch_base: 0,
};
}
if let Some(op) = node.compiled_u64() {
return JitOp::SlotCall {
kit: SlotKitRef::new(crate::ast::CompiledSlotKit {
scratch: Vec::new(),
op: Box::new(move |inputs, outputs, _| op(inputs, outputs)),
}),
scratch_base: 0,
};
}
JitOp::Fallback
}
#[derive(Debug, Clone, PartialEq)]
pub enum JitOp {
Identity,
AddConst(u64),
MulConst(u64),
DivConst(u64),
ModConst(u64),
ClampConst(u64, u64),
Interleave,
MixedRadixConst(Vec<u64>),
Hash,
SplitMix64,
ShuffleConst(u64, u64, u64),
UnitInterval,
F64ToU64,
RoundToU64,
FloorToU64,
CeilToU64,
ClampF64Const(u64, u64), LerpConst(u64, u64), ScaleRangeConst(u64, u64), QuantizeConst(u64), DiscretizeConst(u64, u64), LutSampleConst(u64, u64), WeightedPickConst(u64, u64, u64, u64, u64),
MathUnary(u8),
MathBinary(u8),
U64Add2,
U64Sub2,
U64Mul2,
U64Div2,
U64Mod2,
U64And,
U64Or,
U64Xor,
U64Shl,
U64Shr,
U64Not,
ToF64,
F64Add,
F64Sub,
F64Mul,
F64Div,
F64Mod,
U64DivWire,
U64ModWire,
SlotCall {
kit: SlotKitRef,
scratch_base: usize,
},
Convert {
from: crate::ast::PortType,
to: crate::ast::PortType,
kit: SlotKitRef,
scratch_base: usize,
},
U64ToStr {
scratch_base: usize,
},
I64ToStr {
scratch_base: usize,
},
F64ToStr {
scratch_base: usize,
},
StrConcat {
scratch_base: usize,
},
JsonToStr {
scratch_base: usize,
},
VecProduce {
kind: VecProducer,
scratch_base: usize,
},
VecReduce(VecReducer),
RegLane(RegLaneRead),
RegProduce(RegProducer),
RegDotF32,
RegShuffleConst([u8; 16]),
IsPositiveCheck {
name_ptr: u64,
name_len: u64,
},
InRangeCheck(u64, u64),
IsOneOfCheck {
allowed: Vec<u64>,
set_ptr: u64,
set_len: u64,
},
RegBinOp(u8, u8),
RegCopy,
RegSplat(u8),
U64Cmp(ir::condcodes::IntCC),
F64Cmp(ir::condcodes::FloatCC),
SelectU64,
SelectF64,
I64ToF64,
ToBool,
ConstU64(u64),
ConstF64(u64),
HashRangeConst(u64),
HashIntervalConst(u64, u64),
InvLerpConst(u64, u64),
RemapConst(u64, u64, u64, u64),
EpochOffsetConst(u64),
EpochScaleConst(u64),
ThreadId,
CurrentEpochMillis,
Perlin1dConst(u64, u64),
Perlin2dConst(u64, u64),
Simplex2dConst(u64, u64),
FractalNoise1dConst(u64, u64, u64),
FractalNoise2dConst(u64, u64, u64),
VariadicSum,
VariadicProduct,
VariadicMin,
VariadicMax,
CheckedAdd,
CheckedSub,
CheckedMul,
CeilToMultiple,
MultiplesAtLeast,
FairCoin,
BlendConst(u64),
LfsrStepConst(u64),
PcgConst(u64, u64),
PcgStreamConst(u64),
CycleWalkConst(u64, u64, u64),
UnfairCoinConst(u64),
CoinFlipConst(u64),
ChanceConst(u64),
NOfConst(u64, u64),
Fallback,
}
pub fn classify_node(node: &dyn PolydatNode) -> JitOp {
if let Some(op) = conversion_op(node) {
return op;
}
let name = node.meta().name.as_str();
let consts = node.jit_constants();
match name {
"identity" => JitOp::Identity,
"hash" | "splitmix64" | "scatter" => JitOp::SplitMix64,
"fair_coin" => JitOp::FairCoin,
"unfair_coin" => {
if let Some(&p) = consts.first() {
JitOp::UnfairCoinConst(p)
} else {
JitOp::Fallback
}
}
"chance" => {
if let Some(&p) = consts.first() {
JitOp::ChanceConst(p)
} else {
JitOp::Fallback
}
}
"xxhash3" | "xxh3" => JitOp::Hash,
"hash_range" => {
if let Some(&c) = consts.first() {
JitOp::HashRangeConst(c)
} else {
JitOp::Fallback
}
}
"hash_interval" => {
if consts.len() >= 2 {
JitOp::HashIntervalConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"add" => {
if let Some(&c) = consts.first() {
JitOp::AddConst(c)
} else {
JitOp::Fallback
}
}
"mul" => {
if let Some(&c) = consts.first() {
JitOp::MulConst(c)
} else {
JitOp::Fallback
}
}
"div" => {
if let Some(&c) = consts.first() {
JitOp::DivConst(c)
} else {
JitOp::Fallback
}
}
"mod" => {
if let Some(&c) = consts.first() {
JitOp::ModConst(c)
} else {
JitOp::Fallback
}
}
"clamp" => {
if consts.len() >= 2 && consts[0] <= consts[1] {
JitOp::ClampConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"interleave" => JitOp::Interleave,
"mixed_radix" => {
if consts.is_empty() {
JitOp::Fallback
} else {
JitOp::MixedRadixConst(consts)
}
}
"shuffle" => {
if consts.len() >= 3 {
JitOp::ShuffleConst(consts[0], consts[1], consts[2])
} else {
JitOp::Fallback
}
}
"unit_interval" => JitOp::UnitInterval,
"f64_to_u64" => JitOp::F64ToU64,
"round_to_u64" => JitOp::RoundToU64,
"floor_to_u64" => JitOp::FloorToU64,
"ceil_to_u64" => JitOp::CeilToU64,
"clamp_f64" => {
if consts.len() >= 2 {
JitOp::ClampF64Const(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"lerp" => {
if consts.len() >= 2 {
JitOp::LerpConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"scale_range" => {
if consts.len() >= 2 {
JitOp::ScaleRangeConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"quantize" => {
if let Some(&c) = consts.first() {
JitOp::QuantizeConst(c)
} else {
JitOp::Fallback
}
}
"discretize" => {
if consts.len() >= 2 {
JitOp::DiscretizeConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"lut_sample" | "dist_normal" | "icd_normal" | "dist_exponential" | "icd_exponential"
| "dist_uniform" | "dist_pareto" | "dist_zipf" | "dist_empirical" => {
if consts.len() >= 2 {
JitOp::LutSampleConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"sin" => JitOp::MathUnary(0),
"cos" => JitOp::MathUnary(1),
"tan" => JitOp::MathUnary(2),
"asin" => JitOp::MathUnary(3),
"acos" => JitOp::MathUnary(4),
"atan" => JitOp::MathUnary(5),
"sqrt" => JitOp::MathUnary(6),
"abs_f64" => JitOp::MathUnary(7),
"ln" => JitOp::MathUnary(8),
"exp" => JitOp::MathUnary(9),
"floor_base10" => JitOp::MathUnary(10),
"ceiling_base10" => JitOp::MathUnary(11),
"closest_base10" => JitOp::MathUnary(12),
"floor_decade" => JitOp::MathUnary(13),
"ceiling_decade" => JitOp::MathUnary(14),
"closest_decade" => JitOp::MathUnary(15),
"floor_binomial" => JitOp::MathUnary(16),
"ceiling_binomial" => JitOp::MathUnary(17),
"closest_binomial" => JitOp::MathUnary(18),
"floor_fibonacci" => JitOp::MathUnary(19),
"ceiling_fibonacci" => JitOp::MathUnary(20),
"closest_fibonacci" => JitOp::MathUnary(21),
"atan2" => JitOp::MathBinary(0),
"pow" => JitOp::MathBinary(1),
"round_nearest" => JitOp::MathBinary(2),
"round_floor" => JitOp::MathBinary(3),
"round_ceiling" => JitOp::MathBinary(4),
"to_f64" => JitOp::ToF64,
"u64_add" => JitOp::U64Add2,
"u64_sub" => JitOp::U64Sub2,
"u64_mul" => JitOp::U64Mul2,
"u64_div" => JitOp::U64Div2,
"u64_mod" => JitOp::U64Mod2,
"u64_and" => JitOp::U64And,
"u64_or" => JitOp::U64Or,
"u64_xor" => JitOp::U64Xor,
"u64_shl" => JitOp::U64Shl,
"u64_shr" => JitOp::U64Shr,
"u64_not" => JitOp::U64Not,
"reg_add_i8" => JitOp::RegBinOp(0, 0),
"reg_sub_i8" => JitOp::RegBinOp(0, 1),
"reg_shuffle_bytes" => {
let mut mask = [0u8; 16];
if consts.len() == 16 && consts.iter().all(|&m| m < 16) {
for (m, &c) in mask.iter_mut().zip(consts.iter()) {
*m = c as u8;
}
JitOp::RegShuffleConst(mask)
} else {
JitOp::Fallback
}
}
"reg_add_i16" => JitOp::RegBinOp(1, 0),
"reg_sub_i16" => JitOp::RegBinOp(1, 1),
"reg_mul_i16" => JitOp::RegBinOp(1, 2),
"reg_add_i32" => JitOp::RegBinOp(2, 0),
"reg_sub_i32" => JitOp::RegBinOp(2, 1),
"reg_mul_i32" => JitOp::RegBinOp(2, 2),
"reg_add_i64" => JitOp::RegBinOp(3, 0),
"reg_sub_i64" => JitOp::RegBinOp(3, 1),
"reg_mul_i64" => JitOp::RegBinOp(3, 2),
"reg_add_f32" => JitOp::RegBinOp(4, 0),
"reg_sub_f32" => JitOp::RegBinOp(4, 1),
"reg_mul_f32" => JitOp::RegBinOp(4, 2),
"reg_add_f64" => JitOp::RegBinOp(5, 0),
"reg_sub_f64" => JitOp::RegBinOp(5, 1),
"reg_mul_f64" => JitOp::RegBinOp(5, 2),
"__reg_view_raw" | "__reg_view_i8x16" | "__reg_view_i16x8" | "__reg_view_i32x4"
| "__reg_view_i64x2" | "__reg_view_f16x8" | "__reg_view_f32x4" | "__reg_view_f64x2" => {
JitOp::RegCopy
}
"reg_splat_i8" => JitOp::RegSplat(0),
"reg_splat_i16" => JitOp::RegSplat(1),
"reg_splat_i32" => JitOp::RegSplat(2),
"reg_splat_i64" => JitOp::RegSplat(3),
"reg_splat_f32" => JitOp::RegSplat(4),
"reg_splat_f64" => JitOp::RegSplat(5),
"f64_add" => JitOp::F64Add,
"f64_sub" => JitOp::F64Sub,
"f64_mul" => JitOp::F64Mul,
"f64_div" => JitOp::F64Div,
"f64_mod" => JitOp::F64Mod,
"u64_eq" => JitOp::U64Cmp(ir::condcodes::IntCC::Equal),
"u64_ne" => JitOp::U64Cmp(ir::condcodes::IntCC::NotEqual),
"u64_lt" => JitOp::U64Cmp(ir::condcodes::IntCC::UnsignedLessThan),
"u64_le" => JitOp::U64Cmp(ir::condcodes::IntCC::UnsignedLessThanOrEqual),
"u64_gt" => JitOp::U64Cmp(ir::condcodes::IntCC::UnsignedGreaterThan),
"u64_ge" => JitOp::U64Cmp(ir::condcodes::IntCC::UnsignedGreaterThanOrEqual),
"f64_eq" => JitOp::F64Cmp(ir::condcodes::FloatCC::Equal),
"f64_ne" => JitOp::F64Cmp(ir::condcodes::FloatCC::NotEqual),
"f64_lt" => JitOp::F64Cmp(ir::condcodes::FloatCC::LessThan),
"f64_le" => JitOp::F64Cmp(ir::condcodes::FloatCC::LessThanOrEqual),
"f64_gt" => JitOp::F64Cmp(ir::condcodes::FloatCC::GreaterThan),
"f64_ge" => JitOp::F64Cmp(ir::condcodes::FloatCC::GreaterThanOrEqual),
"select_u64" | "select" => JitOp::SelectU64,
"select_f64" => JitOp::SelectF64,
"div_wire" => JitOp::U64DivWire,
"mod_wire" => JitOp::U64ModWire,
"ceil_to_multiple" => JitOp::CeilToMultiple,
"multiples_at_least" => JitOp::MultiplesAtLeast,
"checked_add" => JitOp::CheckedAdd,
"checked_sub" => JitOp::CheckedSub,
"checked_mul" => JitOp::CheckedMul,
"sum" => JitOp::VariadicSum,
"product" => JitOp::VariadicProduct,
"min" => JitOp::VariadicMin,
"max" => JitOp::VariadicMax,
"blend" => {
if let Some(&c) = consts.first() {
JitOp::BlendConst(c)
} else {
JitOp::Fallback
}
}
"lfsr_step" => {
if let Some(&fb) = consts.first() {
JitOp::LfsrStepConst(fb)
} else {
JitOp::Fallback
}
}
"pcg" => {
if consts.len() >= 2 {
JitOp::PcgConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"pcg_stream" => {
if let Some(&seed) = consts.first() {
JitOp::PcgStreamConst(seed)
} else {
JitOp::Fallback
}
}
"n_of" => {
if consts.len() >= 2 {
JitOp::NOfConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"cycle_walk" => {
if consts.len() >= 3 {
JitOp::CycleWalkConst(consts[0], consts[1], consts[2])
} else {
JitOp::Fallback
}
}
"coin_flip" => {
if let Some(&threshold) = consts.first() {
JitOp::CoinFlipConst(threshold)
} else {
JitOp::Fallback
}
}
"default_or" => JitOp::Identity,
"const_u64" | "const_bool" => {
if let Some(&c) = consts.first() {
JitOp::ConstU64(c)
} else {
JitOp::Fallback
}
}
"const_f64" => {
if let Some(&c) = consts.first() {
JitOp::ConstF64(c)
} else {
JitOp::Fallback
}
}
"inv_lerp" => {
if consts.len() >= 2 {
JitOp::InvLerpConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"remap" => {
if consts.len() >= 4 {
JitOp::RemapConst(consts[0], consts[1], consts[2], consts[3])
} else {
JitOp::Fallback
}
}
"epoch_offset" => {
if let Some(&c) = consts.first() {
JitOp::EpochOffsetConst(c)
} else {
JitOp::Fallback
}
}
"epoch_scale" => {
if let Some(&c) = consts.first() {
JitOp::EpochScaleConst(c)
} else {
JitOp::Fallback
}
}
"thread_id" => JitOp::ThreadId,
"current_epoch_millis" => JitOp::CurrentEpochMillis,
"perlin_1d" => {
if consts.len() >= 2 {
JitOp::Perlin1dConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"perlin_2d" => {
if consts.len() >= 2 {
JitOp::Perlin2dConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"simplex_2d" => {
if consts.len() >= 2 {
JitOp::Simplex2dConst(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"fractal_noise_1d" => {
if consts.len() >= 3 {
JitOp::FractalNoise1dConst(consts[0], consts[1], consts[2])
} else {
JitOp::Fallback
}
}
"fractal_noise_2d" => {
if consts.len() >= 3 {
JitOp::FractalNoise2dConst(consts[0], consts[1], consts[2])
} else {
JitOp::Fallback
}
}
"trunc_u64" => JitOp::F64ToU64,
"round_u64" => JitOp::RoundToU64,
"weighted_pick" => {
if consts.len() >= 5 {
JitOp::WeightedPickConst(consts[0], consts[1], consts[2], consts[3], consts[4])
} else {
JitOp::Fallback
}
}
"is_positive" => {
let name = node.meta().ins.iter().find_map(|slot| match slot {
crate::ast::Slot::Const {
name,
value: crate::ast::ConstValue::Str(v),
} if name == "name" => Some(v),
_ => None,
});
match name {
Some(v) => JitOp::IsPositiveCheck {
name_ptr: v.as_ptr() as u64,
name_len: v.len() as u64,
},
None => JitOp::IsPositiveCheck {
name_ptr: 0,
name_len: 0,
},
}
}
"in_range" => {
if consts.len() >= 2 {
JitOp::InRangeCheck(consts[0], consts[1])
} else {
JitOp::Fallback
}
}
"is_one_of" => {
if consts.is_empty() {
JitOp::Fallback
} else {
let set = node.meta().ins.iter().find_map(|slot| match slot {
crate::ast::Slot::Const {
name,
value: crate::ast::ConstValue::VecU64(v),
} if name == "allowed" => Some(v),
_ => None,
});
let (set_ptr, set_len) = match set {
Some(v) => (v.as_ptr() as u64, v.len() as u64),
None => (0, 0),
};
JitOp::IsOneOfCheck {
allowed: consts,
set_ptr,
set_len,
}
}
}
_ => JitOp::Fallback,
}
}
#[doc(hidden)]
pub fn compile_jit_raw(
coord_count: usize,
total_slots: usize,
steps: Vec<(JitOp, Vec<usize>, Vec<usize>)>,
output_map: HashMap<String, usize>,
nodes: Vec<Box<dyn PolydatNode>>,
) -> Result<JitKernelRaw, String> {
let alone = vec![false; steps.len()];
compile_jit_raw_with(
coord_count,
total_slots,
steps,
output_map,
nodes,
crate::compile::externs::Externs::coordinates_only(coord_count),
super::kernels::ScratchPlan::default(),
Vec::new(),
alone,
)
}
fn pure_units(
steps: &[(JitOp, Vec<usize>, Vec<usize>)],
total_slots: usize,
alone: &[bool],
volatile: &[usize],
unset_read: &[usize],
) -> crate::compile::fusion_units::UnitPlan {
let mut producer = vec![usize::MAX; total_slots + 1];
for (i, (_, _, outs)) in steps.iter().enumerate() {
for &s in outs {
if s < producer.len() {
producer[s] = i;
}
}
}
let preds: Vec<Vec<usize>> = steps
.iter()
.map(|(_, ins, _)| {
let mut p: Vec<usize> = ins
.iter()
.filter_map(|&s| producer.get(s).copied().filter(|&p| p != usize::MAX))
.collect();
p.sort_unstable();
p.dedup();
p
})
.collect();
let inputs_read: Vec<Vec<usize>> = steps
.iter()
.map(|(_, ins, _)| {
ins.iter()
.copied()
.filter(|&s| producer.get(s).is_some_and(|&p| p == usize::MAX))
.collect()
})
.collect();
let fusible: Vec<bool> = (0..steps.len())
.map(|i| !alone.get(i).copied().unwrap_or(false))
.collect();
let mut class = vec![0u64; steps.len()];
for &v in volatile {
if v < class.len() {
class[v] = 1;
}
}
let reads: Vec<Vec<usize>> = inputs_read
.iter()
.map(|ins| {
ins.iter()
.copied()
.filter(|s| unset_read.binary_search(s).is_ok())
.collect()
})
.collect();
let class = crate::compile::fusion_units::refine_by_externs(&preds, &reads, &class);
let rank: Vec<usize> = (0..steps.len()).collect();
crate::compile::fusion_units::plan_units(&preds, &inputs_read, &fusible, &class, &rank, &|_| {
false
})
}
fn unit_dependents(
input_dependents: Vec<Vec<usize>>,
plan: &crate::compile::fusion_units::UnitPlan,
) -> Vec<Vec<usize>> {
input_dependents
.into_iter()
.map(|steps| {
let mut units: Vec<usize> = steps.iter().map(|&s| plan.unit_of[s]).collect();
units.sort_unstable();
units.dedup();
units
})
.collect()
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn compile_jit_raw_with(
coord_count: usize,
total_slots: usize,
steps: Vec<(JitOp, Vec<usize>, Vec<usize>)>,
output_map: HashMap<String, usize>,
nodes: Vec<Box<dyn PolydatNode>>,
externs: crate::compile::externs::Externs,
scratch: super::kernels::ScratchPlan,
volatile: Vec<usize>,
alone: Vec<bool>,
) -> Result<JitKernelRaw, String> {
let plan = pure_units(
&steps,
total_slots,
&alone,
&volatile,
&externs.unset_read_slots(),
);
let (_, entry, code) = compile_jit_impl(&steps, Some(&plan.units), Some(total_slots))?;
let cones =
super::kernels::ConePlan::new(&steps, total_slots, &plan, output_map.values().copied());
let mut core = JitCore::new(
total_slots,
coord_count,
output_map,
code,
nodes,
scratch,
volatile,
entry,
cones,
);
core.set_externs(externs);
core.engine =
crate::compile::select::Engine::PureNative(crate::compile::select::Provenance::Raw);
Ok(JitKernelRaw { core })
}
pub(crate) type JitSegmentCode = (NativeFn, super::kernels::JitCode);
pub(crate) fn compile_jit_entry(
steps: &[(JitOp, Vec<usize>, Vec<usize>)],
tracker: Option<usize>,
) -> Result<JitSegmentCode, String> {
let (raw_fn, _, code) = compile_jit_impl(steps, None, tracker)?;
Ok((raw_fn, code))
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn compile_jit_push_pull(
coord_count: usize,
total_slots: usize,
steps: Vec<(JitOp, Vec<usize>, Vec<usize>)>,
output_map: HashMap<String, usize>,
nodes: Vec<Box<dyn PolydatNode>>,
input_dependents: Vec<Vec<usize>>,
externs: crate::compile::externs::Externs,
scratch: super::kernels::ScratchPlan,
volatile: Vec<usize>,
alone: Vec<bool>,
) -> Result<JitKernelPushPull, String> {
let buffer_len = total_slots;
let plan = pure_units(
&steps,
total_slots,
&alone,
&volatile,
&externs.unset_read_slots(),
);
let (_, entry, code) = compile_jit_impl(&steps, Some(&plan.units), Some(total_slots))?;
let step_outs: Vec<&[usize]> = steps.iter().map(|(_, _, o)| o.as_slice()).collect();
let slot_provenance =
crate::compile::slot_provenance(coord_count, buffer_len, &step_outs, &input_dependents);
let cones =
super::kernels::ConePlan::new(&steps, total_slots, &plan, output_map.values().copied());
let input_dependents = unit_dependents(input_dependents, &plan);
let mut core = JitCore::new(
total_slots,
coord_count,
output_map,
code,
nodes,
scratch,
volatile,
entry,
cones,
);
core.set_externs(externs);
Ok(JitKernelPushPull {
core,
input_dependents,
slot_provenance,
changed_mask: crate::kernel::ProvMask::all_below(coord_count),
force_run: false,
})
}
pub type NativeFn = unsafe fn(*const u64, *mut u64, *mut crate::ast::ScratchBuf);
pub type NativeDispatchFn =
unsafe fn(*const u64, *mut u64, *mut crate::ast::ScratchBuf, *const u32, u64, *mut u8);
type JitCompiled = (NativeFn, NativeDispatchFn, super::kernels::JitCode);
type JitEntry = (NativeFn, NativeDispatchFn, bool);
type JitFunctionSpec<'a> = (
&'a [(JitOp, Vec<usize>, Vec<usize>)],
Option<&'a [Vec<usize>]>,
);
fn compile_jit_impl(
steps: &[(JitOp, Vec<usize>, Vec<usize>)],
dispatch: Option<&[Vec<usize>]>,
tracker: Option<usize>,
) -> Result<JitCompiled, String> {
let (entries, code) = compile_jit_module(&[(steps, dispatch)], tracker)?;
let (straight_fn, dispatch_fn, _) = entries[0];
Ok((straight_fn, dispatch_fn, code))
}
pub(crate) type JitStep = (JitOp, Vec<usize>, Vec<usize>);
pub(crate) fn compile_jit_entries(
batches: &[&[JitStep]],
tracker: Option<usize>,
) -> Result<(Vec<(NativeFn, bool)>, super::kernels::JitCode), String> {
let specs: Vec<JitFunctionSpec> = batches.iter().map(|&b| (b, None)).collect();
let (entries, code) = compile_jit_module(&specs, tracker)?;
Ok((
entries
.into_iter()
.map(|(f, _, fallible)| (f, fallible))
.collect(),
code,
))
}
fn compile_jit_module(
functions: &[JitFunctionSpec],
tracker: Option<usize>,
) -> Result<(Vec<JitEntry>, super::kernels::JitCode), String> {
let mut flag_builder = settings::builder();
flag_builder.set("opt_level", "speed").unwrap();
flag_builder.set("unwind_info", "true").unwrap();
flag_builder.set("preserve_frame_pointers", "true").unwrap();
let isa = super::host_isa::build_host_isa(flag_builder)?;
let mut jit_builder = JITBuilder::with_isa(isa, cranelift_module::default_libcall_names());
jit_builder.symbol("jit_xxh3_hash", jit_xxh3_hash as *const u8);
jit_builder.symbol("jit_interleave", jit_interleave as *const u8);
jit_builder.symbol("jit_shuffle", jit_shuffle as *const u8);
jit_builder.symbol("jit_lut_sample", jit_lut_sample as *const u8);
jit_builder.symbol("jit_weighted_pick", jit_weighted_pick as *const u8);
jit_builder.symbol("jit_pcg", jit_pcg as *const u8);
jit_builder.symbol("jit_pcg_stream", jit_pcg_stream as *const u8);
jit_builder.symbol("jit_n_of", jit_n_of as *const u8);
jit_builder.symbol("jit_cycle_walk", jit_cycle_walk as *const u8);
jit_builder.symbol("jit_perlin_1d", jit_perlin_1d as *const u8);
jit_builder.symbol("jit_perlin_2d", jit_perlin_2d as *const u8);
jit_builder.symbol("jit_simplex_2d", jit_simplex_2d as *const u8);
jit_builder.symbol("jit_fractal_noise_1d", jit_fractal_noise_1d as *const u8);
jit_builder.symbol("jit_fractal_noise_2d", jit_fractal_noise_2d as *const u8);
jit_builder.symbol("jit_thread_id", jit_thread_id as *const u8);
jit_builder.symbol(
"jit_current_epoch_millis",
jit_current_epoch_millis as *const u8,
);
jit_builder.symbol("jit_is_positive_fail", jit_is_positive_fail as *const u8);
jit_builder.symbol("jit_in_range_fail", jit_in_range_fail as *const u8);
jit_builder.symbol("jit_is_one_of_fail", jit_is_one_of_fail as *const u8);
jit_builder.symbol("jit_slot_call", jit_slot_call as *const u8);
jit_builder.symbol("jit_u64_to_str", jit_u64_to_str as *const u8);
jit_builder.symbol("jit_i64_to_str", jit_i64_to_str as *const u8);
jit_builder.symbol("jit_f64_to_str", jit_f64_to_str as *const u8);
jit_builder.symbol("jit_str_concat", jit_str_concat as *const u8);
jit_builder.symbol("jit_json_to_str", jit_json_to_str as *const u8);
jit_builder.symbol("jit_vec_add", jit_vec_add as *const u8);
jit_builder.symbol("jit_vec_scale", jit_vec_scale as *const u8);
jit_builder.symbol("jit_vec_norm", jit_vec_norm as *const u8);
jit_builder.symbol("jit_hash_vec", jit_hash_vec as *const u8);
jit_builder.symbol("jit_xxhash3_vec", jit_xxhash3_vec as *const u8);
jit_builder.symbol("jit_reg_to_vec_f32", jit_reg_to_vec_f32 as *const u8);
jit_builder.symbol("jit_vec_dot", jit_vec_dot as *const u8);
jit_builder.symbol("jit_vec_l2", jit_vec_l2 as *const u8);
jit_builder.symbol("jit_vec_cosine", jit_vec_cosine as *const u8);
jit_builder.symbol("jit_lid_mle", jit_lid_mle as *const u8);
jit_builder.symbol("jit_reg_lane_f32", jit_reg_lane_f32 as *const u8);
jit_builder.symbol("jit_reg_lane_i16", jit_reg_lane_i16 as *const u8);
jit_builder.symbol("jit_reg_lane_i64", jit_reg_lane_i64 as *const u8);
jit_builder.symbol("jit_reg_with_lane_f32", jit_reg_with_lane_f32 as *const u8);
jit_builder.symbol("jit_reg_gather_f32", jit_reg_gather_f32 as *const u8);
jit_builder.symbol("jit_vec_to_reg_f32", jit_vec_to_reg_f32 as *const u8);
jit_builder.symbol("jit_reg_mul_i8", jit_reg_mul_i8 as *const u8);
jit_builder.symbol("jit_sin", jit_sin as *const u8);
jit_builder.symbol("jit_cos", jit_cos as *const u8);
jit_builder.symbol("jit_tan", jit_tan as *const u8);
jit_builder.symbol("jit_asin", jit_asin as *const u8);
jit_builder.symbol("jit_acos", jit_acos as *const u8);
jit_builder.symbol("jit_atan", jit_atan as *const u8);
jit_builder.symbol("jit_sqrt", jit_sqrt as *const u8);
jit_builder.symbol("jit_abs_f64", jit_abs_f64 as *const u8);
jit_builder.symbol("jit_ln", jit_ln as *const u8);
jit_builder.symbol("jit_exp", jit_exp as *const u8);
jit_builder.symbol("jit_floor_base10", jit_floor_base10 as *const u8);
jit_builder.symbol("jit_ceiling_base10", jit_ceiling_base10 as *const u8);
jit_builder.symbol("jit_closest_base10", jit_closest_base10 as *const u8);
jit_builder.symbol("jit_floor_decade", jit_floor_decade as *const u8);
jit_builder.symbol("jit_ceiling_decade", jit_ceiling_decade as *const u8);
jit_builder.symbol("jit_closest_decade", jit_closest_decade as *const u8);
jit_builder.symbol("jit_floor_binomial", jit_floor_binomial as *const u8);
jit_builder.symbol("jit_ceiling_binomial", jit_ceiling_binomial as *const u8);
jit_builder.symbol("jit_closest_binomial", jit_closest_binomial as *const u8);
jit_builder.symbol("jit_floor_fibonacci", jit_floor_fibonacci as *const u8);
jit_builder.symbol("jit_ceiling_fibonacci", jit_ceiling_fibonacci as *const u8);
jit_builder.symbol("jit_closest_fibonacci", jit_closest_fibonacci as *const u8);
jit_builder.symbol("jit_atan2", jit_atan2 as *const u8);
jit_builder.symbol("jit_pow", jit_pow as *const u8);
jit_builder.symbol("jit_round_nearest", jit_round_nearest as *const u8);
jit_builder.symbol("jit_round_floor", jit_round_floor as *const u8);
jit_builder.symbol("jit_round_ceiling", jit_round_ceiling as *const u8);
jit_builder.symbol("jit_f64_mod", jit_f64_mod as *const u8);
jit_builder.symbol("jit_div_zero_fail", jit_div_zero_fail as *const u8);
let mut module = JITModule::new(jit_builder);
let hash_func_id = {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64));
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_xxh3_hash", Linkage::Import, &sig)
.map_err(|e| format!("declare hash: {e}"))?
};
let interleave_func_id = {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64));
sig.params.push(AbiParam::new(types::I64));
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_interleave", Linkage::Import, &sig)
.map_err(|e| format!("declare interleave: {e}"))?
};
let shuffle_func_id = {
let mut sig = module.make_signature();
for _ in 0..4 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_shuffle", Linkage::Import, &sig)
.map_err(|e| format!("declare shuffle: {e}"))?
};
let lut_sample_func_id = {
let mut sig = module.make_signature();
for _ in 0..3 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_lut_sample", Linkage::Import, &sig)
.map_err(|e| format!("declare lut_sample: {e}"))?
};
let weighted_pick_func_id = {
let mut sig = module.make_signature();
for _ in 0..6 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_weighted_pick", Linkage::Import, &sig)
.map_err(|e| format!("declare weighted_pick: {e}"))?
};
let pcg_func_id = {
let mut sig = module.make_signature();
for _ in 0..3 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_pcg", Linkage::Import, &sig)
.map_err(|e| format!("declare pcg: {e}"))?
};
let pcg_stream_func_id = {
let mut sig = module.make_signature();
for _ in 0..3 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_pcg_stream", Linkage::Import, &sig)
.map_err(|e| format!("declare pcg_stream: {e}"))?
};
let n_of_func_id = {
let mut sig = module.make_signature();
for _ in 0..3 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_n_of", Linkage::Import, &sig)
.map_err(|e| format!("declare n_of: {e}"))?
};
let cycle_walk_func_id = {
let mut sig = module.make_signature();
for _ in 0..4 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_cycle_walk", Linkage::Import, &sig)
.map_err(|e| format!("declare cycle_walk: {e}"))?
};
let perlin_1d_func_id = {
let mut sig = module.make_signature();
for _ in 0..3 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_perlin_1d", Linkage::Import, &sig)
.map_err(|e| format!("declare perlin_1d: {e}"))?
};
let perlin_2d_func_id = {
let mut sig = module.make_signature();
for _ in 0..4 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_perlin_2d", Linkage::Import, &sig)
.map_err(|e| format!("declare perlin_2d: {e}"))?
};
let simplex_2d_func_id = {
let mut sig = module.make_signature();
for _ in 0..4 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_simplex_2d", Linkage::Import, &sig)
.map_err(|e| format!("declare simplex_2d: {e}"))?
};
let fractal_noise_1d_func_id = {
let mut sig = module.make_signature();
for _ in 0..4 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_fractal_noise_1d", Linkage::Import, &sig)
.map_err(|e| format!("declare fractal_noise_1d: {e}"))?
};
let fractal_noise_2d_func_id = {
let mut sig = module.make_signature();
for _ in 0..5 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_fractal_noise_2d", Linkage::Import, &sig)
.map_err(|e| format!("declare fractal_noise_2d: {e}"))?
};
let thread_id_func_id = {
let mut sig = module.make_signature();
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_thread_id", Linkage::Import, &sig)
.map_err(|e| format!("declare thread_id: {e}"))?
};
let current_epoch_millis_func_id = {
let mut sig = module.make_signature();
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_current_epoch_millis", Linkage::Import, &sig)
.map_err(|e| format!("declare current_epoch_millis: {e}"))?
};
let math_unary_names = [
"jit_sin",
"jit_cos",
"jit_tan",
"jit_asin",
"jit_acos",
"jit_atan",
"jit_sqrt",
"jit_abs_f64",
"jit_ln",
"jit_exp",
"jit_floor_base10",
"jit_ceiling_base10",
"jit_closest_base10",
"jit_floor_decade",
"jit_ceiling_decade",
"jit_closest_decade",
"jit_floor_binomial",
"jit_ceiling_binomial",
"jit_closest_binomial",
"jit_floor_fibonacci",
"jit_ceiling_fibonacci",
"jit_closest_fibonacci",
];
let mut math_unary_ids = Vec::new();
for name in &math_unary_names {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64));
sig.returns.push(AbiParam::new(types::I64));
math_unary_ids.push(
module
.declare_function(name, Linkage::Import, &sig)
.map_err(|e| format!("declare {name}: {e}"))?,
);
}
let is_positive_fail_id = {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64));
sig.params.push(AbiParam::new(types::I64));
sig.params.push(AbiParam::new(types::I64));
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_is_positive_fail", Linkage::Import, &sig)
.map_err(|e| format!("declare is_positive_fail: {e}"))?
};
let in_range_fail_id = {
let mut sig = module.make_signature();
for _ in 0..3 {
sig.params.push(AbiParam::new(types::I64));
}
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_in_range_fail", Linkage::Import, &sig)
.map_err(|e| format!("declare in_range_fail: {e}"))?
};
let is_one_of_fail_id = {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64));
sig.params.push(AbiParam::new(types::I64));
sig.params.push(AbiParam::new(types::I64));
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_is_one_of_fail", Linkage::Import, &sig)
.map_err(|e| format!("declare is_one_of_fail: {e}"))?
};
let math_binary_names = [
"jit_atan2",
"jit_pow",
"jit_round_nearest",
"jit_round_floor",
"jit_round_ceiling",
"jit_f64_mod",
];
const F64_MOD_HELPER: usize = 5;
let div_zero_fail_id = {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64));
sig.returns.push(AbiParam::new(types::I64));
module
.declare_function("jit_div_zero_fail", Linkage::Import, &sig)
.map_err(|e| format!("declare div_zero_fail: {e}"))?
};
let mut math_binary_ids = Vec::new();
for name in &math_binary_names {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64));
sig.params.push(AbiParam::new(types::I64));
sig.returns.push(AbiParam::new(types::I64));
math_binary_ids.push(
module
.declare_function(name, Linkage::Import, &sig)
.map_err(|e| format!("declare {name}: {e}"))?,
);
}
let slot_call_id = {
let mut sig = module.make_signature();
for _ in 0..8 {
sig.params.push(AbiParam::new(types::I64));
}
module
.declare_function("jit_slot_call", Linkage::Import, &sig)
.map_err(|e| format!("declare jit_slot_call: {e}"))?
};
let mut declare_str = |name: &str, args: usize| -> Result<cranelift_module::FuncId, String> {
let mut sig = module.make_signature();
for _ in 0..args {
sig.params.push(AbiParam::new(types::I64));
}
module
.declare_function(name, Linkage::Import, &sig)
.map_err(|e| format!("declare {name}: {e}"))
};
let u64_to_str_id = declare_str("jit_u64_to_str", 5)?;
let i64_to_str_id = declare_str("jit_i64_to_str", 5)?;
let f64_to_str_id = declare_str("jit_f64_to_str", 5)?;
let str_concat_id = declare_str("jit_str_concat", 6)?;
let json_to_str_id = declare_str("jit_json_to_str", 6)?;
let mut declare_words =
|name: &str, args: usize, returns: bool| -> Result<cranelift_module::FuncId, String> {
let mut sig = module.make_signature();
for _ in 0..args {
sig.params.push(AbiParam::new(types::I64));
}
if returns {
sig.returns.push(AbiParam::new(types::I64));
}
module
.declare_function(name, Linkage::Import, &sig)
.map_err(|e| format!("declare {name}: {e}"))
};
let vec_producer_ids = [
(VecProducer::Add, declare_words("jit_vec_add", 8, false)?),
(
VecProducer::Scale,
declare_words("jit_vec_scale", 8, false)?,
),
(VecProducer::Norm, declare_words("jit_vec_norm", 8, false)?),
(
VecProducer::HashVec,
declare_words("jit_hash_vec", 8, false)?,
),
(
VecProducer::XxHash3Vec,
declare_words("jit_xxhash3_vec", 8, false)?,
),
(
VecProducer::RegToVec,
declare_words("jit_reg_to_vec_f32", 8, false)?,
),
];
let vec_reducer_ids = [
(VecReducer::Dot, declare_words("jit_vec_dot", 4, true)?),
(VecReducer::L2, declare_words("jit_vec_l2", 4, true)?),
(
VecReducer::Cosine,
declare_words("jit_vec_cosine", 4, true)?,
),
(VecReducer::LidMle, declare_words("jit_lid_mle", 4, true)?),
];
let reg_lane_ids = [
(
RegLaneRead::F32,
declare_words("jit_reg_lane_f32", 3, true)?,
),
(
RegLaneRead::I16,
declare_words("jit_reg_lane_i16", 3, true)?,
),
(
RegLaneRead::I64,
declare_words("jit_reg_lane_i64", 3, true)?,
),
];
let reg_producer_ids = [
(
RegProducer::WithLaneF32,
declare_words("jit_reg_with_lane_f32", 6, false)?,
),
(
RegProducer::GatherF32,
declare_words("jit_reg_gather_f32", 6, false)?,
),
(
RegProducer::VecToRegF32,
declare_words("jit_vec_to_reg_f32", 6, false)?,
),
(
RegProducer::MulI8,
declare_words("jit_reg_mul_i8", 6, false)?,
),
];
let mut defined: Vec<(cranelift_module::FuncId, bool)> = Vec::with_capacity(functions.len());
for (function_idx, &(steps, dispatch)) in functions.iter().enumerate() {
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); if dispatch.is_some() {
sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); }
let func_id = module
.declare_function(
&format!("polydat_kernel_{function_idx}"),
Linkage::Local,
&sig,
)
.map_err(|e| format!("declare kernel: {e}"))?;
let mut ctx = module.make_context();
ctx.func.signature = sig;
let mut fb_ctx = FunctionBuilderContext::new();
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let block = builder.create_block();
builder.append_block_params_for_function_params(block);
builder.switch_to_block(block);
builder.seal_block(block);
let _coords_ptr = builder.block_params(block)[0];
let buffer_ptr = builder.block_params(block)[1];
let scratch_ptr = builder.block_params(block)[2];
let hash_func_ref = module.declare_func_in_func(hash_func_id, builder.func);
let interleave_func_ref = module.declare_func_in_func(interleave_func_id, builder.func);
let shuffle_func_ref = module.declare_func_in_func(shuffle_func_id, builder.func);
let lut_sample_func_ref = module.declare_func_in_func(lut_sample_func_id, builder.func);
let weighted_pick_func_ref =
module.declare_func_in_func(weighted_pick_func_id, builder.func);
let is_positive_fail_ref =
module.declare_func_in_func(is_positive_fail_id, builder.func);
let in_range_fail_ref = module.declare_func_in_func(in_range_fail_id, builder.func);
let div_zero_fail_ref = module.declare_func_in_func(div_zero_fail_id, builder.func);
let is_one_of_fail_ref = module.declare_func_in_func(is_one_of_fail_id, builder.func);
let slot_call_ref = module.declare_func_in_func(slot_call_id, builder.func);
let u64_to_str_ref = module.declare_func_in_func(u64_to_str_id, builder.func);
let i64_to_str_ref = module.declare_func_in_func(i64_to_str_id, builder.func);
let f64_to_str_ref = module.declare_func_in_func(f64_to_str_id, builder.func);
let str_concat_ref = module.declare_func_in_func(str_concat_id, builder.func);
let json_to_str_ref = module.declare_func_in_func(json_to_str_id, builder.func);
let vec_producer_refs: Vec<(VecProducer, ir::FuncRef)> = vec_producer_ids
.iter()
.map(|(k, id)| (*k, module.declare_func_in_func(*id, builder.func)))
.collect();
let vec_reducer_refs: Vec<(VecReducer, ir::FuncRef)> = vec_reducer_ids
.iter()
.map(|(k, id)| (*k, module.declare_func_in_func(*id, builder.func)))
.collect();
let reg_lane_refs: Vec<(RegLaneRead, ir::FuncRef)> = reg_lane_ids
.iter()
.map(|(k, id)| (*k, module.declare_func_in_func(*id, builder.func)))
.collect();
let reg_producer_refs: Vec<(RegProducer, ir::FuncRef)> = reg_producer_ids
.iter()
.map(|(k, id)| (*k, module.declare_func_in_func(*id, builder.func)))
.collect();
let pcg_func_ref = module.declare_func_in_func(pcg_func_id, builder.func);
let pcg_stream_func_ref = module.declare_func_in_func(pcg_stream_func_id, builder.func);
let n_of_func_ref = module.declare_func_in_func(n_of_func_id, builder.func);
let cycle_walk_func_ref = module.declare_func_in_func(cycle_walk_func_id, builder.func);
let perlin_1d_func_ref = module.declare_func_in_func(perlin_1d_func_id, builder.func);
let perlin_2d_func_ref = module.declare_func_in_func(perlin_2d_func_id, builder.func);
let simplex_2d_func_ref = module.declare_func_in_func(simplex_2d_func_id, builder.func);
let fractal_noise_1d_func_ref =
module.declare_func_in_func(fractal_noise_1d_func_id, builder.func);
let fractal_noise_2d_func_ref =
module.declare_func_in_func(fractal_noise_2d_func_id, builder.func);
let thread_id_func_ref = module.declare_func_in_func(thread_id_func_id, builder.func);
let current_epoch_millis_func_ref =
module.declare_func_in_func(current_epoch_millis_func_id, builder.func);
let math_unary_refs: Vec<_> = math_unary_ids
.iter()
.map(|id| module.declare_func_in_func(*id, builder.func))
.collect();
let math_binary_refs: Vec<_> = math_binary_ids
.iter()
.map(|id| module.declare_func_in_func(*id, builder.func))
.collect();
let everything: [Vec<usize>; 1] = [(0..steps.len()).collect()];
let schedule: &[Vec<usize>] = dispatch.unwrap_or(&everything);
let dispatcher = dispatch.map(|units| {
let list_ptr = builder.block_params(block)[3];
let list_len = builder.block_params(block)[4];
let clean_ptr = builder.block_params(block)[5];
let at = builder.create_sized_stack_slot(ir::StackSlotData::new(
ir::StackSlotKind::ExplicitSlot,
8,
3,
));
let zero = builder.ins().iconst(types::I64, 0);
builder.ins().stack_store(zero, at, 0);
let head = builder.create_block();
let fetch = builder.create_block();
let dispatch_unit = builder.create_block();
let skip = builder.create_block();
let exit = builder.create_block();
let unit_blocks: Vec<ir::Block> =
units.iter().map(|_| builder.create_block()).collect();
builder.ins().jump(head, &[]);
builder.switch_to_block(head);
let i = builder.ins().stack_load(types::I64, at, 0);
let done = builder.ins().icmp(
ir::condcodes::IntCC::UnsignedGreaterThanOrEqual,
i,
list_len,
);
builder.ins().brif(done, exit, &[], fetch, &[]);
builder.switch_to_block(fetch);
builder.seal_block(fetch);
let offset = builder.ins().ishl_imm(i, 2);
let addr = builder.ins().iadd(list_ptr, offset);
let unit = builder
.ins()
.load(types::I32, ir::MemFlags::trusted(), addr, 0);
let unit_wide = builder.ins().uextend(types::I64, unit);
let flag_addr = builder.ins().iadd(clean_ptr, unit_wide);
let flag = builder
.ins()
.load(types::I8, ir::MemFlags::trusted(), flag_addr, 0);
builder.ins().brif(flag, skip, &[], dispatch_unit, &[]);
builder.switch_to_block(skip);
builder.seal_block(skip);
let next = builder.ins().iadd_imm(i, 1);
builder.ins().stack_store(next, at, 0);
builder.ins().jump(head, &[]);
builder.switch_to_block(dispatch_unit);
builder.seal_block(dispatch_unit);
let default = builder.func.dfg.block_call(exit, &[]);
let targets: Vec<ir::BlockCall> = unit_blocks
.iter()
.map(|&b| builder.func.dfg.block_call(b, &[]))
.collect();
let table = builder.create_jump_table(ir::JumpTableData::new(default, &targets));
builder.ins().br_table(unit, table);
(at, head, exit, unit_blocks, clean_ptr)
});
for (unit_idx, members) in schedule.iter().enumerate() {
if let Some((_, _, _, unit_blocks, _)) = &dispatcher {
builder.switch_to_block(unit_blocks[unit_idx]);
builder.seal_block(unit_blocks[unit_idx]);
}
for &step_idx in members {
let (jit_op, input_slots, output_slots) = &steps[step_idx];
let tracker_store = tracker.map(|t| {
let idx = builder.ins().iconst(types::I64, step_idx as i64);
let inst = store_slot(&mut builder, buffer_ptr, t, idx);
(inst, builder.func.dfg.num_insts())
});
match jit_op {
JitOp::Identity => {
for (&i, &o) in input_slots.iter().zip(output_slots.iter()) {
let val = load_slot(&mut builder, buffer_ptr, i);
store_slot(&mut builder, buffer_ptr, o, val);
}
}
JitOp::AddConst(c) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_val = builder.ins().iconst(types::I64, *c as i64);
let result = builder.ins().iadd(val, c_val);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::MulConst(c) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_val = builder.ins().iconst(types::I64, *c as i64);
let result = builder.ins().imul(val, c_val);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::DivConst(c) | JitOp::ModConst(c) => {
let is_div = matches!(jit_op, JitOp::DivConst(_));
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
if *c == 0 {
let kind =
builder.ins().iconst(types::I64, if is_div { 0 } else { 1 });
let _ = builder.ins().call(div_zero_fail_ref, &[kind]);
store_slot(&mut builder, buffer_ptr, output_slots[0], val);
} else {
let c_val = builder.ins().iconst(types::I64, *c as i64);
let result = if is_div {
builder.ins().udiv(val, c_val)
} else {
builder.ins().urem(val, c_val)
};
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
}
JitOp::U64DivWire | JitOp::U64ModWire => {
let is_div = matches!(jit_op, JitOp::U64DivWire);
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let zero = builder.ins().iconst(types::I64, 0);
let is_zero = builder.ins().icmp(ir::condcodes::IntCC::Equal, b, zero);
let fail_block = builder.create_block();
let ok_block = builder.create_block();
builder.ins().brif(is_zero, fail_block, &[], ok_block, &[]);
builder.switch_to_block(fail_block);
builder.seal_block(fail_block);
let kind = builder.ins().iconst(types::I64, if is_div { 0 } else { 1 });
let _ = builder.ins().call(div_zero_fail_ref, &[kind]);
builder.ins().jump(ok_block, &[]);
builder.switch_to_block(ok_block);
builder.seal_block(ok_block);
let result = if is_div {
builder.ins().udiv(a, b)
} else {
builder.ins().urem(a, b)
};
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::ClampConst(min, max) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let min_val = builder.ins().iconst(types::I64, *min as i64);
let max_val = builder.ins().iconst(types::I64, *max as i64);
let clamped_lo = builder.ins().umax(val, min_val);
let clamped = builder.ins().umin(clamped_lo, max_val);
store_slot(&mut builder, buffer_ptr, output_slots[0], clamped);
}
JitOp::Interleave => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let call = builder.ins().call(interleave_func_ref, &[a, b]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::MixedRadixConst(radixes) => {
let mut remainder = load_slot(&mut builder, buffer_ptr, input_slots[0]);
for (i, &radix) in radixes.iter().enumerate() {
if radix == 0 {
store_slot(
&mut builder,
buffer_ptr,
output_slots[i],
remainder,
);
} else {
let r = builder.ins().iconst(types::I64, radix as i64);
let digit = builder.ins().urem(remainder, r);
store_slot(&mut builder, buffer_ptr, output_slots[i], digit);
remainder = builder.ins().udiv(remainder, r);
}
}
}
JitOp::Hash => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let call = builder.ins().call(hash_func_ref, &[val]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::SplitMix64 => {
let x0 = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_gamma = builder
.ins()
.iconst(types::I64, 0x9e3779b97f4a7c15u64 as i64);
let x1 = builder.ins().iadd(x0, c_gamma);
let s30 = builder.ins().ushr_imm(x1, 30);
let x2 = builder.ins().bxor(x1, s30);
let c_m1 = builder
.ins()
.iconst(types::I64, 0xbf58476d1ce4e5b9u64 as i64);
let x3 = builder.ins().imul(x2, c_m1);
let s27 = builder.ins().ushr_imm(x3, 27);
let x4 = builder.ins().bxor(x3, s27);
let c_m2 = builder
.ins()
.iconst(types::I64, 0x94d049bb133111ebu64 as i64);
let x5 = builder.ins().imul(x4, c_m2);
let s31 = builder.ins().ushr_imm(x5, 31);
let result = builder.ins().bxor(x5, s31);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::FairCoin => {
let x0 = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_gamma = builder
.ins()
.iconst(types::I64, 0x9e3779b97f4a7c15u64 as i64);
let x1 = builder.ins().iadd(x0, c_gamma);
let s30 = builder.ins().ushr_imm(x1, 30);
let x2 = builder.ins().bxor(x1, s30);
let c_m1 = builder
.ins()
.iconst(types::I64, 0xbf58476d1ce4e5b9u64 as i64);
let x3 = builder.ins().imul(x2, c_m1);
let s27 = builder.ins().ushr_imm(x3, 27);
let x4 = builder.ins().bxor(x3, s27);
let c_m2 = builder
.ins()
.iconst(types::I64, 0x94d049bb133111ebu64 as i64);
let x5 = builder.ins().imul(x4, c_m2);
let s31 = builder.ins().ushr_imm(x5, 31);
let h = builder.ins().bxor(x5, s31);
let one = builder.ins().iconst(types::I64, 1);
let result = builder.ins().band(h, one);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::CoinFlipConst(threshold) => {
let x = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let thr = builder.ins().iconst(types::I64, *threshold as i64);
let cmp =
builder
.ins()
.icmp(ir::condcodes::IntCC::UnsignedLessThan, x, thr);
let zero = builder.ins().iconst(types::I64, 0);
let one = builder.ins().iconst(types::I64, 1);
let result = builder.ins().select(cmp, one, zero);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::UnfairCoinConst(p_bits) => {
let x0 = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_gamma = builder
.ins()
.iconst(types::I64, 0x9e3779b97f4a7c15u64 as i64);
let x1 = builder.ins().iadd(x0, c_gamma);
let s30 = builder.ins().ushr_imm(x1, 30);
let x2 = builder.ins().bxor(x1, s30);
let c_m1 = builder
.ins()
.iconst(types::I64, 0xbf58476d1ce4e5b9u64 as i64);
let x3 = builder.ins().imul(x2, c_m1);
let s27 = builder.ins().ushr_imm(x3, 27);
let x4 = builder.ins().bxor(x3, s27);
let c_m2 = builder
.ins()
.iconst(types::I64, 0x94d049bb133111ebu64 as i64);
let x5 = builder.ins().imul(x4, c_m2);
let s31 = builder.ins().ushr_imm(x5, 31);
let h = builder.ins().bxor(x5, s31);
let fval = builder.ins().fcvt_from_uint(types::F64, h);
let max_f = builder.ins().f64const(u64::MAX as f64);
let unit = builder.ins().fdiv(fval, max_f);
let p_f = builder.ins().f64const(f64::from_bits(*p_bits));
let cmp =
builder
.ins()
.fcmp(ir::condcodes::FloatCC::LessThan, unit, p_f);
let zero = builder.ins().iconst(types::I64, 0);
let one = builder.ins().iconst(types::I64, 1);
let result = builder.ins().select(cmp, one, zero);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::ChanceConst(p_bits) => {
let x0 = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_gamma = builder
.ins()
.iconst(types::I64, 0x9e3779b97f4a7c15u64 as i64);
let x1 = builder.ins().iadd(x0, c_gamma);
let s30 = builder.ins().ushr_imm(x1, 30);
let x2 = builder.ins().bxor(x1, s30);
let c_m1 = builder
.ins()
.iconst(types::I64, 0xbf58476d1ce4e5b9u64 as i64);
let x3 = builder.ins().imul(x2, c_m1);
let s27 = builder.ins().ushr_imm(x3, 27);
let x4 = builder.ins().bxor(x3, s27);
let c_m2 = builder
.ins()
.iconst(types::I64, 0x94d049bb133111ebu64 as i64);
let x5 = builder.ins().imul(x4, c_m2);
let s31 = builder.ins().ushr_imm(x5, 31);
let h = builder.ins().bxor(x5, s31);
let fval = builder.ins().fcvt_from_uint(types::F64, h);
let max_f = builder.ins().f64const(u64::MAX as f64);
let unit = builder.ins().fdiv(fval, max_f);
let p_f = builder.ins().f64const(f64::from_bits(*p_bits));
let cmp =
builder
.ins()
.fcmp(ir::condcodes::FloatCC::LessThan, unit, p_f);
let zero_bits =
builder.ins().iconst(types::I64, 0.0_f64.to_bits() as i64);
let one_bits =
builder.ins().iconst(types::I64, 1.0_f64.to_bits() as i64);
let result = builder.ins().select(cmp, one_bits, zero_bits);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::ShuffleConst(feedback, size, min) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let fb = builder.ins().iconst(types::I64, *feedback as i64);
let sz = builder.ins().iconst(types::I64, *size as i64);
let mn = builder.ins().iconst(types::I64, *min as i64);
let call = builder.ins().call(shuffle_func_ref, &[val, fb, sz, mn]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::UnitInterval => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let fval = builder.ins().fcvt_from_uint(types::F64, val);
let max_f = builder.ins().f64const(u64::MAX as f64);
let result = builder.ins().fdiv(fval, max_f);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::F64ToU64 => {
let fval = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let result = builder.ins().fcvt_to_uint_sat(types::I64, fval);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::RoundToU64 => {
let fval = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let rounded = round_half_away(&mut builder, fval);
let result = builder.ins().fcvt_to_uint_sat(types::I64, rounded);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::FloorToU64 => {
let fval = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let floored = builder.ins().floor(fval);
let result = builder.ins().fcvt_to_uint_sat(types::I64, floored);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::CeilToU64 => {
let fval = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let ceiled = builder.ins().ceil(fval);
let result = builder.ins().fcvt_to_uint_sat(types::I64, ceiled);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::ClampF64Const(min_bits, max_bits) => {
let fval = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let fmin = builder.ins().f64const(f64::from_bits(*min_bits));
let fmax = builder.ins().f64const(f64::from_bits(*max_bits));
let clamped = clamp_ir(&mut builder, fval, fmin, fmax);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], clamped);
}
JitOp::LerpConst(a_bits, b_bits) => {
let t = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let a = builder.ins().f64const(f64::from_bits(*a_bits));
let b = builder.ins().f64const(f64::from_bits(*b_bits));
let diff = builder.ins().fsub(b, a);
let scaled = builder.ins().fmul(t, diff);
let result = builder.ins().fadd(a, scaled);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::ScaleRangeConst(min_bits, range_bits) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let fval = builder.ins().fcvt_from_uint(types::F64, val);
let max_f = builder.ins().f64const(u64::MAX as f64);
let t = builder.ins().fdiv(fval, max_f);
let fmin = builder.ins().f64const(f64::from_bits(*min_bits));
let frange = builder.ins().f64const(f64::from_bits(*range_bits));
let scaled = builder.ins().fmul(t, frange);
let result = builder.ins().fadd(fmin, scaled);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::QuantizeConst(step_bits) => {
let fval = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let step = builder.ins().f64const(f64::from_bits(*step_bits));
let divided = builder.ins().fdiv(fval, step);
let rounded = round_half_away(&mut builder, divided);
let result = builder.ins().fmul(rounded, step);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::LutSampleConst(lut_ptr, lut_len) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let ptr_val = builder.ins().iconst(types::I64, *lut_ptr as i64);
let len_val = builder.ins().iconst(types::I64, *lut_len as i64);
let call = builder
.ins()
.call(lut_sample_func_ref, &[input, ptr_val, len_val]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::DiscretizeConst(range_bits, buckets) => {
let fval = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let range = f64::from_bits(*range_bits);
let fzero = builder.ins().f64const(0.0);
let frange = builder.ins().f64const(range);
let fbuckets = builder.ins().f64const(*buckets as f64);
let clamped = clamp_ir(&mut builder, fval, fzero, frange);
let divided = builder.ins().fdiv(clamped, frange);
let scaled = builder.ins().fmul(divided, fbuckets);
let as_u64 = builder.ins().fcvt_to_uint_sat(types::I64, scaled);
let max_bucket =
builder.ins().iconst(types::I64, (*buckets - 1) as i64);
let result = builder.ins().umin(as_u64, max_bucket);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::WeightedPickConst(
values_ptr,
biases_ptr,
primaries_ptr,
aliases_ptr,
n,
) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let v_ptr = builder.ins().iconst(types::I64, *values_ptr as i64);
let b_ptr = builder.ins().iconst(types::I64, *biases_ptr as i64);
let p_ptr = builder.ins().iconst(types::I64, *primaries_ptr as i64);
let a_ptr = builder.ins().iconst(types::I64, *aliases_ptr as i64);
let n_val = builder.ins().iconst(types::I64, *n as i64);
let call = builder.ins().call(
weighted_pick_func_ref,
&[input, v_ptr, b_ptr, p_ptr, a_ptr, n_val],
);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::MathUnary(idx) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let func_ref = math_unary_refs[*idx as usize];
let call = builder.ins().call(func_ref, &[input]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::MathBinary(idx) => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let func_ref = math_binary_refs[*idx as usize];
let call = builder.ins().call(func_ref, &[a, b]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::ToF64 => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let fval = builder.ins().fcvt_from_uint(types::F64, val);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], fval);
}
JitOp::RegBinOp(lane, arith) => {
let vt = reg_lane_type(*lane);
let a = load_reg128(&mut builder, buffer_ptr, input_slots[0], vt);
let b = load_reg128(&mut builder, buffer_ptr, input_slots[2], vt);
let is_float = matches!(*lane, 4 | 5);
let r = match (arith, is_float) {
(0, false) => builder.ins().iadd(a, b),
(1, false) => builder.ins().isub(a, b),
(2, false) => builder.ins().imul(a, b),
(0, true) => builder.ins().fadd(a, b),
(1, true) => builder.ins().fsub(a, b),
(2, true) => builder.ins().fmul(a, b),
_ => unreachable!("RegBinOp arith index out of range"),
};
store_reg128(&mut builder, buffer_ptr, output_slots[0], r);
}
JitOp::RegCopy => {
let v =
load_reg128(&mut builder, buffer_ptr, input_slots[0], types::I64X2);
store_reg128(&mut builder, buffer_ptr, output_slots[0], v);
}
JitOp::RegSplat(lane) => {
let vt = reg_lane_type(*lane);
let scalar = match *lane {
0 => {
let v = load_slot(&mut builder, buffer_ptr, input_slots[0]);
builder.ins().ireduce(types::I8, v)
}
1 => {
let v = load_slot(&mut builder, buffer_ptr, input_slots[0]);
builder.ins().ireduce(types::I16, v)
}
2 => {
let v = load_slot(&mut builder, buffer_ptr, input_slots[0]);
builder.ins().ireduce(types::I32, v)
}
3 => load_slot(&mut builder, buffer_ptr, input_slots[0]),
4 => {
let f = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
builder.ins().fdemote(types::F32, f)
}
5 => load_slot_f64(&mut builder, buffer_ptr, input_slots[0]),
_ => unreachable!("RegSplat lane index out of range"),
};
let v = builder.ins().splat(vt, scalar);
store_reg128(&mut builder, buffer_ptr, output_slots[0], v);
}
JitOp::U64Add2 => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().iadd(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Sub2 => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().isub(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Mul2 => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().imul(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Div2 => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let zero = builder.ins().iconst(types::I64, 0);
let is_zero = builder.ins().icmp(ir::condcodes::IntCC::Equal, b, zero);
let div_block = builder.create_block();
let merge_block = builder.create_block();
builder.append_block_param(merge_block, types::I64);
builder
.ins()
.brif(is_zero, merge_block, &[zero], div_block, &[]);
builder.switch_to_block(div_block);
builder.seal_block(div_block);
let div_result = builder.ins().udiv(a, b);
builder.ins().jump(merge_block, &[div_result]);
builder.switch_to_block(merge_block);
builder.seal_block(merge_block);
let result = builder.block_params(merge_block)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Mod2 => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let zero = builder.ins().iconst(types::I64, 0);
let is_zero = builder.ins().icmp(ir::condcodes::IntCC::Equal, b, zero);
let rem_block = builder.create_block();
let merge_block = builder.create_block();
builder.append_block_param(merge_block, types::I64);
builder
.ins()
.brif(is_zero, merge_block, &[zero], rem_block, &[]);
builder.switch_to_block(rem_block);
builder.seal_block(rem_block);
let rem_result = builder.ins().urem(a, b);
builder.ins().jump(merge_block, &[rem_result]);
builder.switch_to_block(merge_block);
builder.seal_block(merge_block);
let result = builder.block_params(merge_block)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64And => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().band(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Or => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().bor(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Xor => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().bxor(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Shl => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().ishl(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Shr => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().ushr(a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::U64Not => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let result = builder.ins().bnot(a);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::F64Add => {
let a = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot_f64(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().fadd(a, b);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::F64Sub => {
let a = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot_f64(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().fsub(a, b);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::F64Mul => {
let a = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot_f64(&mut builder, buffer_ptr, input_slots[1]);
let result = builder.ins().fmul(a, b);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::F64Div => {
let a = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot_f64(&mut builder, buffer_ptr, input_slots[1]);
let zero = builder.ins().f64const(0.0);
let is_zero =
builder.ins().fcmp(ir::condcodes::FloatCC::Equal, b, zero);
let div_result = builder.ins().fdiv(a, b);
let result = builder.ins().select(is_zero, zero, div_result);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::F64Mod => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let call = builder
.ins()
.call(math_binary_refs[F64_MOD_HELPER], &[a, b]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::IsPositiveCheck { name_ptr, name_len } => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let zero = builder.ins().iconst(types::I64, 0);
let is_zero =
builder.ins().icmp(ir::condcodes::IntCC::Equal, val, zero);
let fail_block = builder.create_block();
let ok_block = builder.create_block();
builder.ins().brif(is_zero, fail_block, &[], ok_block, &[]);
builder.switch_to_block(fail_block);
builder.seal_block(fail_block);
let np = builder.ins().iconst(types::I64, *name_ptr as i64);
let nl = builder.ins().iconst(types::I64, *name_len as i64);
let _ = builder.ins().call(is_positive_fail_ref, &[val, np, nl]);
builder.ins().jump(ok_block, &[]);
builder.switch_to_block(ok_block);
builder.seal_block(ok_block);
store_slot(&mut builder, buffer_ptr, output_slots[0], val);
}
JitOp::InRangeCheck(lo, hi) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let lo_v = builder.ins().iconst(types::I64, *lo as i64);
let hi_v = builder.ins().iconst(types::I64, *hi as i64);
let below = builder.ins().icmp(
ir::condcodes::IntCC::UnsignedLessThan,
val,
lo_v,
);
let above = builder.ins().icmp(
ir::condcodes::IntCC::UnsignedGreaterThan,
val,
hi_v,
);
let out_of_range = builder.ins().bor(below, above);
let fail_block = builder.create_block();
let ok_block = builder.create_block();
builder
.ins()
.brif(out_of_range, fail_block, &[], ok_block, &[]);
builder.switch_to_block(fail_block);
builder.seal_block(fail_block);
let _ = builder.ins().call(in_range_fail_ref, &[val, lo_v, hi_v]);
builder.ins().jump(ok_block, &[]);
builder.switch_to_block(ok_block);
builder.seal_block(ok_block);
store_slot(&mut builder, buffer_ptr, output_slots[0], val);
}
JitOp::IsOneOfCheck {
allowed,
set_ptr,
set_len,
} => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let mut any_match = builder.ins().iconst(types::I8, 0);
for allow in allowed.iter() {
let c = builder.ins().iconst(types::I64, *allow as i64);
let eq = builder.ins().icmp(ir::condcodes::IntCC::Equal, val, c);
any_match = builder.ins().bor(any_match, eq);
}
let fail_block = builder.create_block();
let ok_block = builder.create_block();
builder
.ins()
.brif(any_match, ok_block, &[], fail_block, &[]);
builder.switch_to_block(fail_block);
builder.seal_block(fail_block);
let sp = builder.ins().iconst(types::I64, *set_ptr as i64);
let sl = builder.ins().iconst(types::I64, *set_len as i64);
let _ = builder.ins().call(is_one_of_fail_ref, &[val, sp, sl]);
builder.ins().jump(ok_block, &[]);
builder.switch_to_block(ok_block);
builder.seal_block(ok_block);
store_slot(&mut builder, buffer_ptr, output_slots[0], val);
}
JitOp::U64Cmp(cc) => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let cmp = builder.ins().icmp(*cc, a, b);
let zero = builder.ins().iconst(types::I64, 0);
let one = builder.ins().iconst(types::I64, 1);
let result = builder.ins().select(cmp, one, zero);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::F64Cmp(cc) => {
let a = load_slot_f64(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot_f64(
&mut builder,
buffer_ptr,
if input_slots.len() > 1 {
input_slots[1]
} else {
input_slots[0]
},
);
let cmp = builder.ins().fcmp(*cc, a, b);
let zero = builder.ins().iconst(types::I64, 0);
let one = builder.ins().iconst(types::I64, 1);
let result = builder.ins().select(cmp, one, zero);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::SelectU64 => {
let cond = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let a = load_slot(
&mut builder,
buffer_ptr,
if input_slots.len() > 1 {
input_slots[1]
} else {
input_slots[0]
},
);
let b = load_slot(
&mut builder,
buffer_ptr,
if input_slots.len() > 2 {
input_slots[2]
} else {
input_slots[0]
},
);
let zero = builder.ins().iconst(types::I64, 0);
let is_nonzero =
builder
.ins()
.icmp(ir::condcodes::IntCC::NotEqual, cond, zero);
let result = builder.ins().select(is_nonzero, a, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::SelectF64 => {
let cond = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let a = load_slot_f64(
&mut builder,
buffer_ptr,
if input_slots.len() > 1 {
input_slots[1]
} else {
input_slots[0]
},
);
let b = load_slot_f64(
&mut builder,
buffer_ptr,
if input_slots.len() > 2 {
input_slots[2]
} else {
input_slots[0]
},
);
let zero = builder.ins().iconst(types::I64, 0);
let is_nonzero =
builder
.ins()
.icmp(ir::condcodes::IntCC::NotEqual, cond, zero);
let result = builder.ins().select(is_nonzero, a, b);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::I64ToF64 => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let fval = builder.ins().fcvt_from_sint(types::F64, val);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], fval);
}
JitOp::ToBool => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let zero = builder.ins().iconst(types::I64, 0);
let one = builder.ins().iconst(types::I64, 1);
let cmp = builder
.ins()
.icmp(ir::condcodes::IntCC::NotEqual, val, zero);
let result = builder.ins().select(cmp, one, zero);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::ConstU64(v) | JitOp::ConstF64(v) => {
let result = builder.ins().iconst(types::I64, *v as i64);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::HashRangeConst(max) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_gamma = builder
.ins()
.iconst(types::I64, 0x9e3779b97f4a7c15u64 as i64);
let x1 = builder.ins().iadd(input, c_gamma);
let s30 = builder.ins().ushr_imm(x1, 30);
let x2 = builder.ins().bxor(x1, s30);
let c_m1 = builder
.ins()
.iconst(types::I64, 0xbf58476d1ce4e5b9u64 as i64);
let x3 = builder.ins().imul(x2, c_m1);
let s27 = builder.ins().ushr_imm(x3, 27);
let x4 = builder.ins().bxor(x3, s27);
let c_m2 = builder
.ins()
.iconst(types::I64, 0x94d049bb133111ebu64 as i64);
let x5 = builder.ins().imul(x4, c_m2);
let s31 = builder.ins().ushr_imm(x5, 31);
let h = builder.ins().bxor(x5, s31);
if *max == 0 {
let zero = builder.ins().iconst(types::I64, 0);
store_slot(&mut builder, buffer_ptr, output_slots[0], zero);
} else {
let m = builder.ins().iconst(types::I64, *max as i64);
let rem = builder.ins().urem(h, m);
store_slot(&mut builder, buffer_ptr, output_slots[0], rem);
}
}
JitOp::HashIntervalConst(min_bits, max_bits) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let c_gamma = builder
.ins()
.iconst(types::I64, 0x9e3779b97f4a7c15u64 as i64);
let x1 = builder.ins().iadd(input, c_gamma);
let s30 = builder.ins().ushr_imm(x1, 30);
let x2 = builder.ins().bxor(x1, s30);
let c_m1 = builder
.ins()
.iconst(types::I64, 0xbf58476d1ce4e5b9u64 as i64);
let x3 = builder.ins().imul(x2, c_m1);
let s27 = builder.ins().ushr_imm(x3, 27);
let x4 = builder.ins().bxor(x3, s27);
let c_m2 = builder
.ins()
.iconst(types::I64, 0x94d049bb133111ebu64 as i64);
let x5 = builder.ins().imul(x4, c_m2);
let s31 = builder.ins().ushr_imm(x5, 31);
let h = builder.ins().bxor(x5, s31);
let h_f = builder.ins().fcvt_from_uint(types::F64, h);
let denom = builder.ins().f64const(u64::MAX as f64);
let unit = builder.ins().fdiv(h_f, denom);
let min_f = f64::from_bits(*min_bits);
let max_f = f64::from_bits(*max_bits);
let span = builder.ins().f64const(max_f - min_f);
let min_val = builder.ins().f64const(min_f);
let scaled = builder.ins().fmul(unit, span);
let res_f = builder.ins().fadd(min_val, scaled);
let res = builder
.ins()
.bitcast(types::I64, ir::MemFlags::new(), res_f);
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::InvLerpConst(a_bits, b_bits) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let in_f =
builder
.ins()
.bitcast(types::F64, ir::MemFlags::new(), input);
let a_f = f64::from_bits(*a_bits);
let b_f = f64::from_bits(*b_bits);
let a_val = builder.ins().f64const(a_f);
let inv_span = builder.ins().f64const(1.0 / (b_f - a_f));
let diff = builder.ins().fsub(in_f, a_val);
let t = builder.ins().fmul(diff, inv_span);
let zero = builder.ins().f64const(0.0);
let one = builder.ins().f64const(1.0);
let res_f = clamp_ir(&mut builder, t, zero, one);
let res = builder
.ins()
.bitcast(types::I64, ir::MemFlags::new(), res_f);
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::RemapConst(in_min_bits, in_max_bits, out_min_bits, out_max_bits) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let in_f =
builder
.ins()
.bitcast(types::F64, ir::MemFlags::new(), input);
let in_min = f64::from_bits(*in_min_bits);
let in_max = f64::from_bits(*in_max_bits);
let out_min = f64::from_bits(*out_min_bits);
let out_max = f64::from_bits(*out_max_bits);
let in_span_val = builder.ins().f64const(in_max - in_min);
let in_min_val = builder.ins().f64const(in_min);
let out_min_val = builder.ins().f64const(out_min);
let out_span_val = builder.ins().f64const(out_max - out_min);
let diff = builder.ins().fsub(in_f, in_min_val);
let t = builder.ins().fdiv(diff, in_span_val);
let scaled = builder.ins().fmul(t, out_span_val);
let res_f = builder.ins().fadd(out_min_val, scaled);
let res = builder
.ins()
.bitcast(types::I64, ir::MemFlags::new(), res_f);
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::EpochOffsetConst(base) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = builder.ins().iconst(types::I64, *base as i64);
let res = builder.ins().iadd(val, b);
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::EpochScaleConst(factor) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let f = builder.ins().iconst(types::I64, *factor as i64);
let res = builder.ins().imul(val, f);
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::ThreadId => {
let call = builder.ins().call(thread_id_func_ref, &[]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::CurrentEpochMillis => {
let call = builder.ins().call(current_epoch_millis_func_ref, &[]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::Perlin1dConst(perm_ptr, freq_bits) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let p = builder.ins().iconst(types::I64, *perm_ptr as i64);
let fb = builder.ins().iconst(types::I64, *freq_bits as i64);
let call = builder.ins().call(perlin_1d_func_ref, &[input, p, fb]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::Perlin2dConst(perm_ptr, freq_bits) => {
let x = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let y = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let p = builder.ins().iconst(types::I64, *perm_ptr as i64);
let fb = builder.ins().iconst(types::I64, *freq_bits as i64);
let call = builder.ins().call(perlin_2d_func_ref, &[x, y, p, fb]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::Simplex2dConst(perm_ptr, freq_bits) => {
let x = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let y = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let p = builder.ins().iconst(types::I64, *perm_ptr as i64);
let fb = builder.ins().iconst(types::I64, *freq_bits as i64);
let call = builder.ins().call(simplex_2d_func_ref, &[x, y, p, fb]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::FractalNoise1dConst(perm_ptr, freq_bits, octaves) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let p = builder.ins().iconst(types::I64, *perm_ptr as i64);
let fb = builder.ins().iconst(types::I64, *freq_bits as i64);
let oct = builder.ins().iconst(types::I64, *octaves as i64);
let call = builder
.ins()
.call(fractal_noise_1d_func_ref, &[input, p, fb, oct]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::FractalNoise2dConst(perm_ptr, freq_bits, octaves) => {
let x = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let y = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let p = builder.ins().iconst(types::I64, *perm_ptr as i64);
let fb = builder.ins().iconst(types::I64, *freq_bits as i64);
let oct = builder.ins().iconst(types::I64, *octaves as i64);
let call = builder
.ins()
.call(fractal_noise_2d_func_ref, &[x, y, p, fb, oct]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::CycleWalkConst(range, seed, inc) => {
let pos = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let r = builder.ins().iconst(types::I64, *range as i64);
let s = builder.ins().iconst(types::I64, *seed as i64);
let i = builder.ins().iconst(types::I64, *inc as i64);
let call = builder.ins().call(cycle_walk_func_ref, &[pos, r, s, i]);
let res = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], res);
}
JitOp::VariadicSum => {
if input_slots.is_empty() {
let zero = builder.ins().iconst(types::I64, 0);
store_slot(&mut builder, buffer_ptr, output_slots[0], zero);
} else {
let mut acc = load_slot(&mut builder, buffer_ptr, input_slots[0]);
for &slot in &input_slots[1..] {
let v = load_slot(&mut builder, buffer_ptr, slot);
acc = builder.ins().iadd(acc, v);
}
store_slot(&mut builder, buffer_ptr, output_slots[0], acc);
}
}
JitOp::VariadicProduct => {
if input_slots.is_empty() {
let one = builder.ins().iconst(types::I64, 1);
store_slot(&mut builder, buffer_ptr, output_slots[0], one);
} else {
let mut acc = load_slot(&mut builder, buffer_ptr, input_slots[0]);
for &slot in &input_slots[1..] {
let v = load_slot(&mut builder, buffer_ptr, slot);
acc = builder.ins().imul(acc, v);
}
store_slot(&mut builder, buffer_ptr, output_slots[0], acc);
}
}
JitOp::VariadicMin => {
if input_slots.is_empty() {
let ident = builder.ins().iconst(types::I64, u64::MAX as i64);
store_slot(&mut builder, buffer_ptr, output_slots[0], ident);
} else {
let mut acc = load_slot(&mut builder, buffer_ptr, input_slots[0]);
for &slot in &input_slots[1..] {
let v = load_slot(&mut builder, buffer_ptr, slot);
let cmp = builder.ins().icmp(
ir::condcodes::IntCC::UnsignedLessThan,
v,
acc,
);
acc = builder.ins().select(cmp, v, acc);
}
store_slot(&mut builder, buffer_ptr, output_slots[0], acc);
}
}
JitOp::VariadicMax => {
if input_slots.is_empty() {
let zero = builder.ins().iconst(types::I64, 0);
store_slot(&mut builder, buffer_ptr, output_slots[0], zero);
} else {
let mut acc = load_slot(&mut builder, buffer_ptr, input_slots[0]);
for &slot in &input_slots[1..] {
let v = load_slot(&mut builder, buffer_ptr, slot);
let cmp = builder.ins().icmp(
ir::condcodes::IntCC::UnsignedGreaterThan,
v,
acc,
);
acc = builder.ins().select(cmp, v, acc);
}
store_slot(&mut builder, buffer_ptr, output_slots[0], acc);
}
}
JitOp::CeilToMultiple => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let m = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let zero = builder.ins().iconst(types::I64, 0);
let one = builder.ins().iconst(types::I64, 1);
let is_zero = builder.ins().icmp(ir::condcodes::IntCC::Equal, m, zero);
let calc_block = builder.create_block();
let merge_block = builder.create_block();
builder.append_block_param(merge_block, types::I64);
builder
.ins()
.brif(is_zero, merge_block, &[val], calc_block, &[]);
builder.switch_to_block(calc_block);
builder.seal_block(calc_block);
let div = div_ceil(&mut builder, val, m, one);
let high = builder.ins().umulhi(div, m);
let low = builder.ins().imul(div, m);
let zero_hi = builder.ins().iconst(types::I64, 0);
let overflows =
builder
.ins()
.icmp(ir::condcodes::IntCC::NotEqual, high, zero_hi);
let max = builder.ins().iconst(types::I64, -1);
let mul = builder.ins().select(overflows, max, low);
builder.ins().jump(merge_block, &[mul]);
builder.switch_to_block(merge_block);
builder.seal_block(merge_block);
let result = builder.block_params(merge_block)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::CheckedAdd => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let sum = builder.ins().iadd(a, b);
let is_overflow =
builder
.ins()
.icmp(ir::condcodes::IntCC::UnsignedLessThan, sum, a);
let zero = builder.ins().iconst(types::I64, 0);
let result = builder.ins().select(is_overflow, zero, sum);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::CheckedSub => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let is_lt =
builder
.ins()
.icmp(ir::condcodes::IntCC::UnsignedLessThan, a, b);
let diff = builder.ins().isub(a, b);
let zero = builder.ins().iconst(types::I64, 0);
let result = builder.ins().select(is_lt, zero, diff);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::CheckedMul => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let prod = builder.ins().imul(a, b);
let zero = builder.ins().iconst(types::I64, 0);
let a_is_zero =
builder.ins().icmp(ir::condcodes::IntCC::Equal, a, zero);
let div_block = builder.create_block();
let merge_block = builder.create_block();
builder.append_block_param(merge_block, types::I64);
builder
.ins()
.brif(a_is_zero, merge_block, &[zero], div_block, &[]);
builder.switch_to_block(div_block);
builder.seal_block(div_block);
let div = builder.ins().udiv(prod, a);
let ok = builder.ins().icmp(ir::condcodes::IntCC::Equal, div, b);
let mul_res = builder.ins().select(ok, prod, zero);
builder.ins().jump(merge_block, &[mul_res]);
builder.switch_to_block(merge_block);
builder.seal_block(merge_block);
let result = builder.block_params(merge_block)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::MultiplesAtLeast => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let m = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let zero = builder.ins().iconst(types::I64, 0);
let one = builder.ins().iconst(types::I64, 1);
let is_zero = builder.ins().icmp(ir::condcodes::IntCC::Equal, m, zero);
let calc_block = builder.create_block();
let merge_block = builder.create_block();
builder.append_block_param(merge_block, types::I64);
builder
.ins()
.brif(is_zero, merge_block, &[zero], calc_block, &[]);
builder.switch_to_block(calc_block);
builder.seal_block(calc_block);
let div = div_ceil(&mut builder, val, m, one);
builder.ins().jump(merge_block, &[div]);
builder.switch_to_block(merge_block);
builder.seal_block(merge_block);
let result = builder.block_params(merge_block)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::BlendConst(mix_bits) => {
let a = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let b = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let fa = builder.ins().bitcast(types::F64, ir::MemFlags::new(), a);
let fb = builder.ins().bitcast(types::F64, ir::MemFlags::new(), b);
let mix_f64 = f64::from_bits(*mix_bits);
let mix_val = builder.ins().f64const(mix_f64);
let one = builder.ins().f64const(1.0);
let one_minus_mix = builder.ins().fsub(one, mix_val);
let a_part = builder.ins().fmul(fa, one_minus_mix);
let b_part = builder.ins().fmul(fb, mix_val);
let sum = builder.ins().fadd(a_part, b_part);
let result =
builder.ins().bitcast(types::I64, ir::MemFlags::new(), sum);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::LfsrStepConst(feedback) => {
let val = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let feedback = builder.ins().iconst(types::I64, *feedback as i64);
let one = builder.ins().iconst(types::I64, 1);
let zero = builder.ins().iconst(types::I64, 0);
let shifted = builder.ins().ushr(val, one);
let lsb = builder.ins().band(val, one);
let is_odd =
builder
.ins()
.icmp(ir::condcodes::IntCC::NotEqual, lsb, zero);
let fb_mask = builder.ins().select(is_odd, feedback, zero);
let result = builder.ins().bxor(shifted, fb_mask);
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::PcgConst(seed, stream) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let s = builder.ins().iconst(types::I64, *seed as i64);
let st = builder.ins().iconst(types::I64, *stream as i64);
let call = builder.ins().call(pcg_func_ref, &[input, s, st]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::PcgStreamConst(seed) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let st = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let s = builder.ins().iconst(types::I64, *seed as i64);
let call = builder.ins().call(pcg_stream_func_ref, &[input, st, s]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::NOfConst(n, m) => {
let input = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let n_val = builder.ins().iconst(types::I64, *n as i64);
let m_val = builder.ins().iconst(types::I64, *m as i64);
let call = builder.ins().call(n_of_func_ref, &[input, n_val, m_val]);
let result = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], result);
}
JitOp::SlotCall { kit, scratch_base } => {
emit_slot_call(
&mut builder,
buffer_ptr,
scratch_ptr,
slot_call_ref,
kit,
*scratch_base,
input_slots,
output_slots,
);
}
JitOp::Convert {
from,
to,
kit,
scratch_base,
} => {
emit_conversion(
&mut builder,
buffer_ptr,
input_slots[0],
output_slots[0],
*from,
*to,
|builder| {
emit_slot_call(
builder,
buffer_ptr,
scratch_ptr,
slot_call_ref,
kit,
*scratch_base,
input_slots,
output_slots,
)
},
);
}
JitOp::U64ToStr { scratch_base }
| JitOp::I64ToStr { scratch_base }
| JitOp::F64ToStr { scratch_base } => {
let func = match jit_op {
JitOp::U64ToStr { .. } => u64_to_str_ref,
JitOp::I64ToStr { .. } => i64_to_str_ref,
_ => f64_to_str_ref,
};
let value = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let base_v = builder.ins().iconst(types::I64, *scratch_base as i64);
let out_v = builder.ins().iconst(types::I64, output_slots[0] as i64);
builder
.ins()
.call(func, &[scratch_ptr, base_v, buffer_ptr, out_v, value]);
}
JitOp::JsonToStr { scratch_base } => {
let ptr = load_slot(&mut builder, buffer_ptr, input_slots[0]);
let len = load_slot(&mut builder, buffer_ptr, input_slots[1]);
let base_v = builder.ins().iconst(types::I64, *scratch_base as i64);
let out_v = builder.ins().iconst(types::I64, output_slots[0] as i64);
builder.ins().call(
json_to_str_ref,
&[scratch_ptr, base_v, buffer_ptr, out_v, ptr, len],
);
}
JitOp::StrConcat { scratch_base } => {
let n_words = input_slots.len();
let frame = builder.create_sized_stack_slot(ir::StackSlotData::new(
ir::StackSlotKind::ExplicitSlot,
(n_words.max(1) * 8) as u32,
3,
));
for (k, &s) in input_slots.iter().enumerate() {
let v = load_slot(&mut builder, buffer_ptr, s);
builder.ins().stack_store(v, frame, (k * 8) as i32);
}
let pairs_ptr = builder.ins().stack_addr(types::I64, frame, 0);
let n_v = builder.ins().iconst(types::I64, (n_words / 2) as i64);
let base_v = builder.ins().iconst(types::I64, *scratch_base as i64);
let out_v = builder.ins().iconst(types::I64, output_slots[0] as i64);
builder.ins().call(
str_concat_ref,
&[scratch_ptr, base_v, buffer_ptr, out_v, pairs_ptr, n_v],
);
}
JitOp::VecProduce { kind, scratch_base } => {
let func = func_of(&vec_producer_refs, *kind);
let base_v = builder.ins().iconst(types::I64, *scratch_base as i64);
let out_v = builder.ins().iconst(types::I64, output_slots[0] as i64);
let words = load_words(&mut builder, buffer_ptr, input_slots, 4);
let mut args = vec![scratch_ptr, base_v, buffer_ptr, out_v];
args.extend(words);
builder.ins().call(func, &args);
}
JitOp::VecReduce(kind) => {
let func = func_of(&vec_reducer_refs, *kind);
let words = load_words(&mut builder, buffer_ptr, input_slots, 4);
let call = builder.ins().call(func, &words);
let bits = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], bits);
}
JitOp::RegLane(kind) => {
let func = func_of(®_lane_refs, *kind);
let words = load_words(&mut builder, buffer_ptr, input_slots, 3);
let call = builder.ins().call(func, &words);
let word = builder.inst_results(call)[0];
store_slot(&mut builder, buffer_ptr, output_slots[0], word);
}
JitOp::RegProduce(kind) => {
let func = func_of(®_producer_refs, *kind);
let out_v = builder.ins().iconst(types::I64, output_slots[0] as i64);
let words = load_words(&mut builder, buffer_ptr, input_slots, 4);
let mut args = vec![buffer_ptr, out_v];
args.extend(words);
builder.ins().call(func, &args);
}
JitOp::RegDotF32 => {
let a =
load_reg128(&mut builder, buffer_ptr, input_slots[0], types::F32X4);
let b =
load_reg128(&mut builder, buffer_ptr, input_slots[2], types::F32X4);
let p = builder.ins().fmul(a, b);
let p0 = builder.ins().extractlane(p, 0);
let p1 = builder.ins().extractlane(p, 1);
let p2 = builder.ins().extractlane(p, 2);
let p3 = builder.ins().extractlane(p, 3);
let s01 = builder.ins().fadd(p0, p1);
let s23 = builder.ins().fadd(p2, p3);
let s = builder.ins().fadd(s01, s23);
let wide = builder.ins().fpromote(types::F64, s);
store_slot_f64(&mut builder, buffer_ptr, output_slots[0], wide);
}
JitOp::RegShuffleConst(mask) => {
let x =
load_reg128(&mut builder, buffer_ptr, input_slots[0], types::I8X16);
let imm = builder
.func
.dfg
.immediates
.push(ir::ConstantData::from(&mask[..]));
let r = builder.ins().shuffle(x, x, imm);
store_reg128(&mut builder, buffer_ptr, output_slots[0], r);
}
JitOp::Fallback => {
}
}
if let Some((inst, mark)) = tracker_store {
let calls = (mark..builder.func.dfg.num_insts()).any(|i| {
builder.func.dfg.insts[ir::Inst::from_u32(i as u32)]
.opcode()
.is_call()
});
if !calls {
builder.func.layout.remove_inst(inst);
}
}
}
if let Some((at, head, _, _, clean_ptr)) = &dispatcher {
let one = builder.ins().iconst(types::I8, 1);
builder
.ins()
.store(ir::MemFlags::trusted(), one, *clean_ptr, unit_idx as i32);
let i = builder.ins().stack_load(types::I64, *at, 0);
let next = builder.ins().iadd_imm(i, 1);
builder.ins().stack_store(next, *at, 0);
builder.ins().jump(*head, &[]);
}
}
if let Some((_, head, exit, _, _)) = dispatcher {
builder.seal_block(head);
builder.switch_to_block(exit);
builder.seal_block(exit);
}
builder.ins().return_(&[]);
builder.finalize();
}
let fallible = ctx.func.layout.blocks().any(|block| {
ctx.func
.layout
.block_insts(block)
.any(|inst| ctx.func.dfg.insts[inst].opcode().is_call())
});
module
.define_function(func_id, &mut ctx)
.map_err(|e| format!("define function: {e}"))?;
module.clear_context(&mut ctx);
defined.push((func_id, fallible));
}
module
.finalize_definitions()
.map_err(|e| format!("finalize: {e}"))?;
let entries: Vec<JitEntry> = defined
.iter()
.map(|&(func_id, fallible)| {
let code_ptr = module.get_finalized_function(func_id);
let straight_fn: NativeFn = unsafe { mem::transmute(code_ptr) };
let dispatch_fn: NativeDispatchFn = unsafe { mem::transmute(code_ptr) };
(straight_fn, dispatch_fn, fallible)
})
.collect();
let kits: Vec<SlotKitRef> = functions
.iter()
.flat_map(|(steps, _)| steps.iter())
.filter_map(|(op, _, _)| op.slot_kit().cloned())
.collect();
let any_fallible = defined.iter().any(|&(_, f)| f);
let code = super::kernels::JitCode::new(module, kits, any_fallible);
Ok((entries, code))
}
#[allow(clippy::too_many_arguments)]
fn emit_slot_call(
builder: &mut FunctionBuilder,
buffer_ptr: ir::Value,
scratch_ptr: ir::Value,
slot_call_ref: ir::FuncRef,
kit: &SlotKitRef,
scratch_base: usize,
input_slots: &[usize],
output_slots: &[usize],
) {
let n_in = input_slots.len();
let n_out = output_slots.len();
let frame = |builder: &mut FunctionBuilder, n: usize| {
builder.create_sized_stack_slot(ir::StackSlotData::new(
ir::StackSlotKind::ExplicitSlot,
(n.max(1) * 8) as u32,
3,
))
};
let in_frame = frame(builder, n_in);
let out_frame = frame(builder, n_out);
for (k, &s) in input_slots.iter().enumerate() {
let v = load_slot(builder, buffer_ptr, s);
builder.ins().stack_store(v, in_frame, (k * 8) as i32);
}
let kit_ptr = builder
.ins()
.iconst(types::I64, std::sync::Arc::as_ptr(&kit.0) as usize as i64);
let in_ptr = builder.ins().stack_addr(types::I64, in_frame, 0);
let n_in_v = builder.ins().iconst(types::I64, n_in as i64);
let out_ptr = builder.ins().stack_addr(types::I64, out_frame, 0);
let n_out_v = builder.ins().iconst(types::I64, n_out as i64);
let base_v = builder.ins().iconst(types::I64, scratch_base as i64);
let n_sc_v = builder.ins().iconst(types::I64, kit.0.scratch.len() as i64);
builder.ins().call(
slot_call_ref,
&[
kit_ptr,
in_ptr,
n_in_v,
out_ptr,
n_out_v,
scratch_ptr,
base_v,
n_sc_v,
],
);
for (k, &s) in output_slots.iter().enumerate() {
let v = builder
.ins()
.stack_load(types::I64, out_frame, (k * 8) as i32);
store_slot(builder, buffer_ptr, s, v);
}
}
fn emit_conversion(
builder: &mut FunctionBuilder,
buffer_ptr: ir::Value,
in_slot: usize,
out_slot: usize,
from: crate::ast::PortType,
to: crate::ast::PortType,
slow: impl FnOnce(&mut FunctionBuilder),
) {
use Scalar::{Bool, F32, F64, Signed, Unsigned};
use ir::condcodes::{FloatCC, IntCC};
let (Some(src), Some(dst)) = (Scalar::of(from), Scalar::of(to)) else {
slow(builder);
return;
};
let raw = load_slot(builder, buffer_ptr, in_slot);
let float_of = |builder: &mut FunctionBuilder| -> ir::Value {
match src {
F32 => {
let bits = builder.ins().ireduce(types::I32, raw);
let x = builder.ins().bitcast(types::F32, ir::MemFlags::new(), bits);
builder.ins().fpromote(types::F64, x)
}
_ => builder.ins().bitcast(types::F64, ir::MemFlags::new(), raw),
}
};
let store_float = |builder: &mut FunctionBuilder, x: ir::Value| {
let word = if dst == F32 {
let bits = builder.ins().bitcast(types::I32, ir::MemFlags::new(), x);
builder.ins().uextend(types::I64, bits)
} else {
builder.ins().bitcast(types::I64, ir::MemFlags::new(), x)
};
store_slot(builder, buffer_ptr, out_slot, word);
};
match (src, dst) {
(_, Bool) => {
let truth = match src {
F32 | F64 => {
let x = float_of(builder);
let zero = builder.ins().f64const(0.0);
let nonzero = builder.ins().fcmp(FloatCC::NotEqual, x, zero);
let ordered = builder.ins().fcmp(FloatCC::Ordered, x, x);
builder.ins().band(nonzero, ordered)
}
_ => builder.ins().icmp_imm(IntCC::NotEqual, raw, 0),
};
let word = builder.ins().uextend(types::I64, truth);
store_slot(builder, buffer_ptr, out_slot, word);
}
(s, d) if s.int_range().is_some() && d.int_range().is_some() => {
let (smin, smax) = s.int_range().expect("an integer");
let (dmin, dmax) = d.int_range().expect("an integer");
let mut fits = Vec::new();
if dmin > smin {
fits.push(builder.ins().icmp_imm(
IntCC::SignedGreaterThanOrEqual,
raw,
dmin as i64,
));
}
if dmax < smax {
let cc = if matches!(s, Signed(_)) {
IntCC::SignedLessThanOrEqual
} else {
IntCC::UnsignedLessThanOrEqual
};
fits.push(builder.ins().icmp_imm(cc, raw, dmax as u64 as i64));
}
branch_on(
builder,
fits,
|b| {
store_slot(b, buffer_ptr, out_slot, raw);
},
slow,
);
}
(s, F32 | F64) if s.int_range().is_some() => {
let ty = if dst == F32 { types::F32 } else { types::F64 };
let x = if matches!(s, Signed(_)) {
builder.ins().fcvt_from_sint(ty, raw)
} else {
builder.ins().fcvt_from_uint(ty, raw)
};
store_float(builder, x);
}
(F32, F64) => {
let x = float_of(builder);
store_float(builder, x);
}
(F64, F32) => {
let x = float_of(builder);
let narrow = builder.ins().fdemote(types::F32, x);
store_float(builder, narrow);
}
(F32 | F64, d) => {
let (lo, hi) = match d {
Unsigned(b) => (0.0, 2f64.powi(b as i32)),
Signed(b) => (-(2f64.powi(b as i32 - 1)), 2f64.powi(b as i32 - 1)),
_ => {
slow(builder);
return;
}
};
let x = float_of(builder);
let lo_v = builder.ins().f64const(lo);
let hi_v = builder.ins().f64const(hi);
let above = builder.ins().fcmp(FloatCC::GreaterThanOrEqual, x, lo_v);
let below = builder.ins().fcmp(FloatCC::LessThan, x, hi_v);
branch_on(
builder,
vec![above, below],
|b| {
let word = if matches!(d, Signed(_)) {
b.ins().fcvt_to_sint_sat(types::I64, x)
} else {
b.ins().fcvt_to_uint_sat(types::I64, x)
};
store_slot(b, buffer_ptr, out_slot, word);
},
slow,
);
}
_ => slow(builder),
}
}
fn branch_on(
builder: &mut FunctionBuilder,
conds: Vec<ir::Value>,
fast: impl FnOnce(&mut FunctionBuilder),
slow: impl FnOnce(&mut FunctionBuilder),
) {
let mut conds = conds.into_iter();
let Some(first) = conds.next() else {
fast(builder);
return;
};
let mut ok = first;
for c in conds {
ok = builder.ins().band(ok, c);
}
let fast_block = builder.create_block();
let slow_block = builder.create_block();
let done = builder.create_block();
builder.ins().brif(ok, fast_block, &[], slow_block, &[]);
builder.switch_to_block(fast_block);
builder.seal_block(fast_block);
fast(builder);
builder.ins().jump(done, &[]);
builder.switch_to_block(slow_block);
builder.seal_block(slow_block);
slow(builder);
builder.ins().jump(done, &[]);
builder.switch_to_block(done);
builder.seal_block(done);
}
fn load_slot(builder: &mut FunctionBuilder, buffer_ptr: ir::Value, slot: usize) -> ir::Value {
let offset = (slot * 8) as i32;
builder
.ins()
.load(types::I64, ir::MemFlags::trusted(), buffer_ptr, offset)
}
fn store_slot(
builder: &mut FunctionBuilder,
buffer_ptr: ir::Value,
slot: usize,
value: ir::Value,
) -> ir::Inst {
let offset = (slot * 8) as i32;
builder
.ins()
.store(ir::MemFlags::trusted(), value, buffer_ptr, offset)
}
fn reg_lane_type(lane: u8) -> ir::Type {
match lane {
0 => types::I8X16,
1 => types::I16X8,
2 => types::I32X4,
3 => types::I64X2,
4 => types::F32X4,
5 => types::F64X2,
_ => unreachable!("register lane index out of range"),
}
}
fn load_reg128(
builder: &mut FunctionBuilder,
buffer_ptr: ir::Value,
first_slot: usize,
vt: ir::Type,
) -> ir::Value {
let offset = (first_slot * 8) as i32;
builder
.ins()
.load(vt, ir::MemFlags::new(), buffer_ptr, offset)
}
fn store_reg128(
builder: &mut FunctionBuilder,
buffer_ptr: ir::Value,
first_slot: usize,
value: ir::Value,
) {
let offset = (first_slot * 8) as i32;
builder
.ins()
.store(ir::MemFlags::new(), value, buffer_ptr, offset);
}
fn round_half_away(builder: &mut FunctionBuilder, x: ir::Value) -> ir::Value {
let t = builder.ins().trunc(x);
let frac = builder.ins().fsub(x, t);
let mag = builder.ins().fabs(frac);
let half = builder.ins().f64const(0.5);
let reaches = builder
.ins()
.fcmp(ir::condcodes::FloatCC::GreaterThanOrEqual, mag, half);
let one = builder.ins().f64const(1.0);
let step = builder.ins().fcopysign(one, x);
let up = builder.ins().fadd(t, step);
builder.ins().select(reaches, up, t)
}
fn clamp_ir(
builder: &mut FunctionBuilder,
x: ir::Value,
lo: ir::Value,
hi: ir::Value,
) -> ir::Value {
let below = builder.ins().fcmp(ir::condcodes::FloatCC::LessThan, x, lo);
let above = builder
.ins()
.fcmp(ir::condcodes::FloatCC::GreaterThan, x, hi);
let capped = builder.ins().select(above, hi, x);
builder.ins().select(below, lo, capped)
}
fn div_ceil(
builder: &mut FunctionBuilder,
val: ir::Value,
m: ir::Value,
one: ir::Value,
) -> ir::Value {
let q = builder.ins().udiv(val, m);
let r = builder.ins().urem(val, m);
let zero = builder.ins().iconst(types::I64, 0);
let inexact = builder.ins().icmp(ir::condcodes::IntCC::NotEqual, r, zero);
let q1 = builder.ins().iadd(q, one);
builder.ins().select(inexact, q1, q)
}
fn func_of<K: PartialEq + Copy>(refs: &[(K, ir::FuncRef)], key: K) -> ir::FuncRef {
refs.iter()
.find(|(k, _)| *k == key)
.map(|(_, r)| *r)
.expect("every helper of the group is declared")
}
fn load_words(
builder: &mut FunctionBuilder,
buffer_ptr: ir::Value,
input_slots: &[usize],
n: usize,
) -> Vec<ir::Value> {
(0..n)
.map(|k| match input_slots.get(k) {
Some(&s) => load_slot(builder, buffer_ptr, s),
None => builder.ins().iconst(types::I64, 0),
})
.collect()
}
fn load_slot_f64(builder: &mut FunctionBuilder, buffer_ptr: ir::Value, slot: usize) -> ir::Value {
let i64_val = load_slot(builder, buffer_ptr, slot);
builder
.ins()
.bitcast(types::F64, ir::MemFlags::new(), i64_val)
}
fn store_slot_f64(
builder: &mut FunctionBuilder,
buffer_ptr: ir::Value,
slot: usize,
value: ir::Value,
) {
let i64_val = builder
.ins()
.bitcast(types::I64, ir::MemFlags::new(), value);
store_slot(builder, buffer_ptr, slot, i64_val);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn every_scalar_conversion_lowers_by_its_types() {
let mut missing = Vec::new();
for &from in crate::ast::PortType::ALL {
for &to in crate::ast::PortType::ALL {
if from == to || Scalar::of(from).is_none() || Scalar::of(to).is_none() {
continue;
}
let Some(node) = crate::compile::assembly::boundary_adapter(from, to) else {
continue;
};
match classify_node(node.as_ref()) {
JitOp::Convert { from: f, to: t, .. } if f == from && t == to => {}
other => missing.push(format!(
"{} ({from:?} -> {to:?}) classified as {other:?}",
node.meta().name
)),
}
}
}
assert!(missing.is_empty(), "{}", missing.join("\n"));
}
#[test]
fn jit_identity() {
let steps = vec![(JitOp::Identity, vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[42]);
assert_eq!(kernel.get("out"), 42);
}
#[test]
fn jit_add_const() {
let steps = vec![(JitOp::AddConst(100), vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[5]);
assert_eq!(kernel.get("out"), 105);
}
#[test]
fn jit_mul_const() {
let steps = vec![(JitOp::MulConst(7), vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[6]);
assert_eq!(kernel.get("out"), 42);
}
#[test]
fn jit_mod_const() {
let steps = vec![(JitOp::ModConst(100), vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[542]);
assert_eq!(kernel.get("out"), 42);
}
#[test]
fn jit_hash() {
let steps = vec![(JitOp::Hash, vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[42]);
let v1 = kernel.get("out");
let expected = xxhash_rust::xxh3::xxh3_64(&42u64.to_le_bytes());
assert_eq!(v1, expected);
}
#[test]
fn jit_hash_deterministic() {
let steps = vec![(JitOp::Hash, vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[42]);
let v1 = kernel.get("out");
kernel.eval(&[42]);
let v2 = kernel.get("out");
assert_eq!(v1, v2);
}
#[test]
fn jit_chain_hash_mod() {
let steps = vec![
(JitOp::Hash, vec![0], vec![1]), (JitOp::ModConst(1_000_000), vec![1], vec![2]), ];
let mut output_map = HashMap::new();
output_map.insert("user_id".into(), 2);
let mut kernel = compile_jit_raw(1, 3, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[42]);
let uid = kernel.get("user_id");
assert!(uid < 1_000_000, "got {uid}");
}
#[test]
fn jit_clamp_const() {
let steps = vec![(JitOp::ClampConst(10, 50), vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[5]);
assert_eq!(kernel.get("out"), 10);
kernel.eval(&[30]);
assert_eq!(kernel.get("out"), 30);
kernel.eval(&[100]);
assert_eq!(kernel.get("out"), 50); }
#[test]
fn jit_interleave() {
let steps = vec![(JitOp::Interleave, vec![0, 1], vec![2])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 2);
let mut kernel = compile_jit_raw(2, 3, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[0b101, 0b010]);
assert_eq!(kernel.get("out"), 0b01_10_01);
}
#[test]
fn jit_mixed_radix() {
let steps = vec![(
JitOp::MixedRadixConst(vec![100, 1000, 0]),
vec![0],
vec![1, 2, 3],
)];
let mut output_map = HashMap::new();
output_map.insert("d0".into(), 1);
output_map.insert("d1".into(), 2);
output_map.insert("d2".into(), 3);
let mut kernel = compile_jit_raw(1, 4, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[4_201_337]);
assert_eq!(kernel.get("d0"), 37);
assert_eq!(kernel.get("d1"), 13);
assert_eq!(kernel.get("d2"), 42);
}
#[test]
fn jit_unit_interval() {
let steps = vec![(JitOp::UnitInterval, vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[0]);
let v = f64::from_bits(kernel.get("out"));
assert!((v - 0.0).abs() < 1e-10);
kernel.eval(&[u64::MAX]);
let v = f64::from_bits(kernel.get("out"));
assert!((v - 1.0).abs() < 1e-10);
}
#[test]
fn jit_f64_to_u64() {
let steps = vec![(JitOp::F64ToU64, vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[3.7f64.to_bits()]);
assert_eq!(kernel.get("out"), 3); }
#[test]
fn jit_round_to_u64() {
let steps = vec![(JitOp::RoundToU64, vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[3.7f64.to_bits()]);
assert_eq!(kernel.get("out"), 4);
kernel.eval(&[3.2f64.to_bits()]);
assert_eq!(kernel.get("out"), 3);
}
#[test]
fn jit_clamp_f64() {
let steps = vec![(
JitOp::ClampF64Const(0.0f64.to_bits(), 1.0f64.to_bits()),
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[(-0.5f64).to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 0.0);
kernel.eval(&[0.5f64.to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 0.5);
kernel.eval(&[1.5f64.to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 1.0);
}
#[test]
fn jit_lerp() {
let steps = vec![(
JitOp::LerpConst(10.0f64.to_bits(), 20.0f64.to_bits()),
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[0.0f64.to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 10.0);
kernel.eval(&[1.0f64.to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 20.0);
kernel.eval(&[0.5f64.to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 15.0);
}
#[test]
fn jit_scale_range() {
let steps = vec![(
JitOp::ScaleRangeConst(10.0f64.to_bits(), 10.0f64.to_bits()),
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[0]);
let v = f64::from_bits(kernel.get("out"));
assert!((v - 10.0).abs() < 0.001);
kernel.eval(&[u64::MAX]);
let v = f64::from_bits(kernel.get("out"));
assert!((v - 20.0).abs() < 0.001);
}
#[test]
fn jit_quantize() {
let steps = vec![(JitOp::QuantizeConst(10.0f64.to_bits()), vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[13.0f64.to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 10.0);
kernel.eval(&[17.0f64.to_bits()]);
assert_eq!(f64::from_bits(kernel.get("out")), 20.0);
}
#[test]
fn jit_discretize() {
let steps = vec![(
JitOp::DiscretizeConst(100.0f64.to_bits(), 10),
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[0.0f64.to_bits()]);
assert_eq!(kernel.get("out"), 0);
kernel.eval(&[55.0f64.to_bits()]);
assert_eq!(kernel.get("out"), 5);
kernel.eval(&[99.0f64.to_bits()]);
assert_eq!(kernel.get("out"), 9);
kernel.eval(&[200.0f64.to_bits()]);
assert_eq!(kernel.get("out"), 9);
}
#[test]
fn jit_chain_unit_interval_lerp() {
let steps = vec![
(JitOp::UnitInterval, vec![0], vec![1]),
(
JitOp::LerpConst(100.0f64.to_bits(), 200.0f64.to_bits()),
vec![1],
vec![2],
),
];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 2);
let mut kernel = compile_jit_raw(1, 3, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[0]);
let v = f64::from_bits(kernel.get("out"));
assert!((v - 100.0).abs() < 0.001);
kernel.eval(&[u64::MAX]);
let v = f64::from_bits(kernel.get("out"));
assert!((v - 200.0).abs() < 0.001);
}
#[test]
fn jit_multi_step_chain() {
let steps = vec![
(JitOp::AddConst(10), vec![0], vec![1]),
(JitOp::MulConst(3), vec![1], vec![2]),
(JitOp::ModConst(100), vec![2], vec![3]),
];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 3);
let mut kernel = compile_jit_raw(1, 4, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[5]);
assert_eq!(kernel.get("out"), 45);
}
#[test]
fn jit_is_positive_check_passes_positive() {
let steps = vec![(
JitOp::IsPositiveCheck {
name_ptr: 0,
name_len: 0,
},
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[42]);
assert_eq!(kernel.get("out"), 42);
kernel.eval(&[u64::MAX]);
assert_eq!(kernel.get("out"), u64::MAX);
}
#[test]
fn jit_in_range_check_passes_interior() {
let steps = vec![(JitOp::InRangeCheck(10, 100), vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[50]);
assert_eq!(kernel.get("out"), 50);
kernel.eval(&[10]);
assert_eq!(kernel.get("out"), 10);
kernel.eval(&[100]);
assert_eq!(kernel.get("out"), 100);
}
#[test]
fn jit_is_one_of_check_passes_allowed_values() {
let steps = vec![(
JitOp::IsOneOfCheck {
allowed: vec![1, 2, 3, 5, 8],
set_ptr: 0,
set_len: 0,
},
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
for v in [1u64, 2, 3, 5, 8] {
kernel.eval(&[v]);
assert_eq!(kernel.get("out"), v);
}
}
#[test]
fn jit_is_one_of_check_accepts_single_element_allow_list() {
let steps = vec![(
JitOp::IsOneOfCheck {
allowed: vec![42],
set_ptr: 0,
set_len: 0,
},
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
kernel.eval(&[42]);
assert_eq!(kernel.get("out"), 42);
}
fn extract_panic_msg(payload: Box<dyn std::any::Any + Send + 'static>) -> String {
payload
.downcast_ref::<String>()
.cloned()
.or_else(|| payload.downcast_ref::<&str>().map(|s| s.to_string()))
.unwrap_or_else(|| "(non-string panic)".into())
}
#[test]
fn jit_is_positive_violation_is_catchable() {
let steps = vec![(
JitOp::IsPositiveCheck {
name_ptr: 0,
name_len: 0,
},
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
let err = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| kernel.eval(&[0])))
.expect_err("JIT violation should panic");
assert!(extract_panic_msg(err).contains("must be > 0"));
}
#[test]
fn jit_in_range_violation_is_catchable() {
let steps = vec![(JitOp::InRangeCheck(10, 100), vec![0], vec![1])];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
let err = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| kernel.eval(&[5])))
.expect_err("below-range should panic");
assert!(extract_panic_msg(err).contains("outside [10, 100]"));
let err = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| kernel.eval(&[500])))
.expect_err("above-range should panic");
assert!(extract_panic_msg(err).contains("outside [10, 100]"));
}
#[test]
fn jit_is_one_of_violation_is_catchable() {
let steps = vec![(
JitOp::IsOneOfCheck {
allowed: vec![1, 3, 5],
set_ptr: 0,
set_len: 0,
},
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
let err = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| kernel.eval(&[2])))
.expect_err("disallowed value should panic");
assert!(extract_panic_msg(err).contains("not in allowed set"));
}
#[test]
fn invoke_with_catch_restores_slot_after_foreign_panic() {
let caught = std::panic::catch_unwind(|| {
invoke_with_catch(|| panic!("foreign panic"));
});
assert!(caught.is_err(), "foreign panic should propagate out");
let steps = vec![(
JitOp::IsPositiveCheck {
name_ptr: 0,
name_len: 0,
},
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
let err = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| kernel.eval(&[0])))
.expect_err("JIT violation should panic cleanly after foreign panic");
assert!(extract_panic_msg(err).contains("must be > 0"));
kernel.eval(&[42]);
assert_eq!(kernel.get("out"), 42);
}
#[test]
fn jit_kernel_survives_multiple_violations() {
let steps = vec![(
JitOp::IsPositiveCheck {
name_ptr: 0,
name_len: 0,
},
vec![0],
vec![1],
)];
let mut output_map = HashMap::new();
output_map.insert("out".into(), 1);
let mut kernel = compile_jit_raw(1, 2, steps, output_map, Vec::new()).unwrap();
for _ in 0..3 {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| kernel.eval(&[0])))
.expect_err("violation should still panic");
}
kernel.eval(&[42]);
assert_eq!(kernel.get("out"), 42);
}
const LADDER: &str = "input cycle: u64\ninput tenant_seed: u64\ninput operation_seed: u64\n\
seeded_cycle := u64_add(cycle, tenant_seed)\nidentity_entropy := hash(seeded_cycle)\n\
account_id := mul(identity_entropy, 10000000)\n\
route_seed := u64_add(account_id, operation_seed)\nroute_entropy := hash(route_seed)\n\
shard := mul(route_entropy, 64)\n\
payload_seed := u64_add(identity_entropy, operation_seed)\n\
payload_entropy := hash(payload_seed)\npayload_class := mul(payload_entropy, 8)\n\
token_seed := u64_add(route_entropy, payload_entropy)\nevent_token := hash(token_seed)\n";
fn unit_counts(src: &str) -> (usize, usize) {
let asm = || crate::dsl::compile::compile_polydat_to_assembler(src).unwrap();
let pure = asm()
.try_compile_pure_jit_raw()
.unwrap()
.core
.cones
.unit_count();
let native = crate::dsl::compile::compile_polydat_with(
src,
crate::compile::select::Engine::Native(crate::compile::select::Provenance::Raw),
)
.unwrap()
.plan()
.native_segments;
(pure, native)
}
#[test]
fn only_an_extern_splits_a_unit() {
assert_eq!(unit_counts(LADDER), (1, 1));
let extern_seed = LADDER.replace(
"input operation_seed: u64",
"extern operation_seed: u64 = 7",
);
assert_eq!(unit_counts(&extern_seed), (2, 2));
let shape = "input x: u64\nMODE\nt := hash(x)\na := hash(t)\n\
b := u64_add(t, mode)\nc := add(b, 1)\n";
assert_eq!(
unit_counts(&shape.replace("MODE", "input mode: u64")),
(1, 1)
);
assert_eq!(
unit_counts(&shape.replace("MODE", "extern mode: u64 = 3")),
(2, 2)
);
}
}