use std::collections::{HashMap, HashSet};
use std::sync::atomic::{AtomicBool, AtomicPtr, AtomicU32, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock, RwLock, Weak};
use cljrs_ir::IrFunction;
use cljrs_value::Value;
pub const DEFAULT_JIT_THRESHOLD: u32 = 1_000;
pub const DEFAULT_IR_THRESHOLD: u32 = 50;
static IR_THRESHOLD_OVERRIDE: AtomicU32 = AtomicU32::new(0);
pub fn set_ir_threshold(t: u32) {
IR_THRESHOLD_OVERRIDE.store(t, Ordering::Relaxed);
}
pub fn ir_threshold() -> u32 {
let v = IR_THRESHOLD_OVERRIDE.load(Ordering::Relaxed);
if v != 0 {
return v;
}
match std::env::var("CLJRS_IR_THRESHOLD")
.ok()
.and_then(|s| s.parse::<u32>().ok())
{
Some(0) => u32::MAX,
Some(t) => t,
None => DEFAULT_IR_THRESHOLD,
}
}
static JIT_THRESHOLD_OVERRIDE: AtomicU32 = AtomicU32::new(0);
pub fn set_jit_threshold(t: u32) {
JIT_THRESHOLD_OVERRIDE.store(t, Ordering::Relaxed);
}
pub fn jit_threshold() -> u32 {
let v = JIT_THRESHOLD_OVERRIDE.load(Ordering::Relaxed);
if v != 0 {
return v;
}
std::env::var("CLJRS_JIT_THRESHOLD")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.unwrap_or(DEFAULT_JIT_THRESHOLD)
}
pub struct JitEntry {
pub invocation_count: AtomicU32,
pub lower_queued: AtomicBool,
pub compile_queued: AtomicBool,
pub native_fn_ptr: AtomicPtr<()>,
pub epoch: AtomicU64,
pub arg_profile: Mutex<Vec<u8>>,
pub deopt_count: AtomicU32,
}
impl JitEntry {
fn new() -> Self {
Self {
invocation_count: AtomicU32::new(0),
lower_queued: AtomicBool::new(false),
compile_queued: AtomicBool::new(false),
native_fn_ptr: AtomicPtr::new(std::ptr::null_mut()),
epoch: AtomicU64::new(0),
arg_profile: Mutex::new(Vec::new()),
deopt_count: AtomicU32::new(0),
}
}
}
pub const PROFILE_LONG: u8 = 1;
pub const PROFILE_DOUBLE: u8 = 2;
pub const PROFILE_OTHER: u8 = 0x80;
#[inline]
fn profile_tag(v: &Value) -> u8 {
match v {
Value::Long(_) => PROFILE_LONG,
Value::Double(_) => PROFILE_DOUBLE,
_ => PROFILE_OTHER,
}
}
pub fn arg_type_profile(arity_id: u64) -> Option<Vec<u8>> {
let guard = JIT_TABLE.read().unwrap();
let entry = guard.as_ref()?.get(&arity_id)?.clone();
drop(guard);
let prof = entry.arg_profile.lock().unwrap();
if prof.is_empty() {
None
} else {
Some(prof.clone())
}
}
unsafe impl Send for JitEntry {}
unsafe impl Sync for JitEntry {}
static JIT_TABLE: RwLock<Option<HashMap<u64, Arc<JitEntry>>>> = RwLock::new(None);
fn get_or_create_entry(arity_id: u64) -> Arc<JitEntry> {
{
let guard = JIT_TABLE.read().unwrap();
if let Some(cache) = guard.as_ref()
&& let Some(e) = cache.get(&arity_id)
{
return e.clone();
}
}
let mut guard = JIT_TABLE.write().unwrap();
let cache = guard.get_or_insert_with(HashMap::new);
cache
.entry(arity_id)
.or_insert_with(|| Arc::new(JitEntry::new()))
.clone()
}
pub fn get_native_fn(arity_id: u64) -> Option<(*const (), u64)> {
let guard = JIT_TABLE.read().unwrap();
let cache = guard.as_ref()?;
let entry = cache.get(&arity_id)?;
let ptr = entry.native_fn_ptr.load(Ordering::Acquire);
if ptr.is_null() {
None
} else {
let epoch = entry.epoch.load(Ordering::Acquire);
Some((ptr as *const (), epoch))
}
}
pub fn store_native_fn(arity_id: u64, ptr: *const (), epoch: u64) {
let entry = get_or_create_entry(arity_id);
entry.epoch.store(epoch, Ordering::Release);
entry.native_fn_ptr.store(ptr as *mut (), Ordering::Release);
}
pub fn take_native_epoch(arity_id: u64) -> Option<u64> {
let mut guard = JIT_TABLE.write().unwrap();
let cache = guard.as_mut()?;
let entry = cache.remove(&arity_id)?;
let ptr = entry
.native_fn_ptr
.swap(std::ptr::null_mut(), Ordering::AcqRel);
if ptr.is_null() {
None
} else {
Some(entry.epoch.load(Ordering::Acquire))
}
}
type StaleEpochFn = fn(u64);
static STALE_EPOCH_HOOK: OnceLock<StaleEpochFn> = OnceLock::new();
pub fn set_stale_epoch_hook(f: StaleEpochFn) {
let _ = STALE_EPOCH_HOOK.set(f);
}
pub fn stale_native_code(arity_id: u64) {
let mut epochs = Vec::new();
if let Some(epoch) = take_native_epoch(arity_id) {
epochs.push(epoch);
}
epochs.extend(take_osr_epochs(arity_id));
if let Some(hook) = STALE_EPOCH_HOOK.get() {
for epoch in epochs {
hook(epoch);
}
}
}
type EnqueueFn = Box<dyn Fn(u64, Arc<IrFunction>) + Send + Sync + 'static>;
static ENQUEUE_HOOK: OnceLock<EnqueueFn> = OnceLock::new();
pub fn set_enqueue_hook(f: impl Fn(u64, Arc<IrFunction>) + Send + Sync + 'static) {
let _ = ENQUEUE_HOOK.set(Box::new(f));
}
pub fn record_call(arity_id: u64, ir_func: Arc<IrFunction>, profile_args: &[Value]) {
let entry = get_or_create_entry(arity_id);
let count = entry.invocation_count.fetch_add(1, Ordering::Relaxed) + 1;
if !entry.compile_queued.load(Ordering::Relaxed) {
let n_params = ir_func.params.len();
let mut prof = entry.arg_profile.lock().unwrap();
if prof.len() < n_params {
prof.resize(n_params, 0);
}
for (i, slot) in prof.iter_mut().enumerate().take(n_params) {
*slot |= profile_args
.get(i)
.map(profile_tag)
.unwrap_or(PROFILE_OTHER);
}
}
if count < jit_threshold() {
return;
}
if entry.compile_queued.swap(true, Ordering::AcqRel) {
return;
}
if let Some(hook) = ENQUEUE_HOOK.get() {
cljrs_logging::feat_debug!("jit", "enqueue arity_id={} (count={})", arity_id, count);
hook(arity_id, ir_func);
}
}
static BOOTSTRAP_ARITY_WATERMARK: AtomicU64 = AtomicU64::new(0);
pub fn set_bootstrap_arity_watermark(w: u64) {
BOOTSTRAP_ARITY_WATERMARK.store(w, Ordering::Relaxed);
}
pub fn is_bootstrap_arity(arity_id: u64) -> bool {
arity_id < BOOTSTRAP_ARITY_WATERMARK.load(Ordering::Relaxed)
}
pub fn record_interp_call(arity_id: u64) -> bool {
let threshold = ir_threshold();
if threshold == u32::MAX {
return false;
}
let entry = get_or_create_entry(arity_id);
let count = entry.invocation_count.fetch_add(1, Ordering::Relaxed) + 1;
count >= threshold && !entry.lower_queued.load(Ordering::Relaxed)
}
pub fn compile_queued(arity_id: u64) -> bool {
let guard = JIT_TABLE.read().unwrap();
guard
.as_ref()
.and_then(|c| c.get(&arity_id))
.is_some_and(|e| e.compile_queued.load(Ordering::Relaxed))
}
pub fn lower_queued(arity_id: u64) -> bool {
let guard = JIT_TABLE.read().unwrap();
guard
.as_ref()
.and_then(|c| c.get(&arity_id))
.is_some_and(|e| e.lower_queued.load(Ordering::Relaxed))
}
pub fn mark_lower_queued(arity_id: u64) {
let entry = get_or_create_entry(arity_id);
entry.lower_queued.store(true, Ordering::Relaxed);
}
pub fn clear_lower_queued(arity_id: u64) {
let guard = JIT_TABLE.read().unwrap();
if let Some(entry) = guard.as_ref().and_then(|c| c.get(&arity_id)) {
entry.lower_queued.store(false, Ordering::Relaxed);
}
}
pub fn on_ir_published(arity_id: u64) {
let entry = get_or_create_entry(arity_id);
entry.invocation_count.store(0, Ordering::Relaxed);
}
pub fn evict_entry_if_cold(arity_id: u64) -> bool {
let mut guard = JIT_TABLE.write().unwrap();
let Some(cache) = guard.as_mut() else {
return false;
};
let Some(entry) = cache.get(&arity_id) else {
return false;
};
if !entry.native_fn_ptr.load(Ordering::Acquire).is_null()
|| entry.compile_queued.load(Ordering::Relaxed)
{
return false;
}
cache.remove(&arity_id);
true
}
pub fn stale_osr_code(arity_id: u64) {
let epochs = take_osr_epochs(arity_id);
if let Some(hook) = STALE_EPOCH_HOOK.get() {
for epoch in epochs {
hook(epoch);
}
}
}
type PendingExceptionFn = fn() -> Option<Value>;
static PENDING_EXCEPTION_HOOK: OnceLock<PendingExceptionFn> = OnceLock::new();
pub fn set_pending_exception_hook(f: PendingExceptionFn) {
let _ = PENDING_EXCEPTION_HOOK.set(f);
}
pub fn take_pending_exception() -> Option<Value> {
PENDING_EXCEPTION_HOOK.get().and_then(|f| f())
}
type DeoptSentinelFn = fn() -> usize;
static DEOPT_SENTINEL_HOOK: OnceLock<DeoptSentinelFn> = OnceLock::new();
pub fn set_deopt_sentinel_hook(f: DeoptSentinelFn) {
let _ = DEOPT_SENTINEL_HOOK.set(f);
}
#[inline]
pub fn is_deopt_result(result: *const Value) -> bool {
DEOPT_SENTINEL_HOOK
.get()
.is_some_and(|f| f() == result as usize)
}
static SPEC_BANNED: RwLock<Option<HashSet<u64>>> = RwLock::new(None);
pub fn specialization_allowed(arity_id: u64) -> bool {
let guard = SPEC_BANNED.read().unwrap();
!guard.as_ref().is_some_and(|s| s.contains(&arity_id))
}
pub fn deopt_limit() -> u32 {
std::env::var("CLJRS_JIT_DEOPT_LIMIT")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.unwrap_or(10)
}
pub fn record_deopt(arity_id: u64) {
let entry = get_or_create_entry(arity_id);
let failures = entry.deopt_count.fetch_add(1, Ordering::Relaxed) + 1;
cljrs_logging::feat_debug!(
"jit",
"deopt arity_id={} (failure #{} of {})",
arity_id,
failures,
deopt_limit()
);
if failures < deopt_limit() {
return;
}
{
let mut guard = SPEC_BANNED.write().unwrap();
guard.get_or_insert_with(HashSet::new).insert(arity_id);
}
if let Some(epoch) = take_native_epoch(arity_id)
&& let Some(hook) = STALE_EPOCH_HOOK.get()
{
cljrs_logging::feat_debug!(
"jit",
"specialization discarded arity_id={} epoch={}",
arity_id,
epoch
);
hook(epoch);
}
}
#[derive(Clone)]
pub struct OsrSlot {
pub fn_ptr: *const (),
pub epoch: u64,
pub live_ins: Arc<[cljrs_ir::VarId]>,
}
unsafe impl Send for OsrSlot {}
unsafe impl Sync for OsrSlot {}
enum OsrState {
Queued,
Ready(OsrSlot),
Failed,
}
pub enum OsrPoll {
NotRequested,
Pending,
Ready(OsrSlot),
Failed,
}
static OSR_TABLE: RwLock<Option<HashMap<(u64, u32), OsrState>>> = RwLock::new(None);
type OsrEnqueueFn = Box<dyn Fn(u64, u32, Arc<IrFunction>) + Send + Sync + 'static>;
static OSR_ENQUEUE_HOOK: OnceLock<OsrEnqueueFn> = OnceLock::new();
pub fn set_osr_enqueue_hook(f: impl Fn(u64, u32, Arc<IrFunction>) + Send + Sync + 'static) {
let _ = OSR_ENQUEUE_HOOK.set(Box::new(f));
}
static OSR_THRESHOLD_OVERRIDE: AtomicU32 = AtomicU32::new(0);
pub fn set_osr_threshold(t: u32) {
OSR_THRESHOLD_OVERRIDE.store(t, Ordering::Relaxed);
}
pub fn osr_threshold() -> u32 {
let v = OSR_THRESHOLD_OVERRIDE.load(Ordering::Relaxed);
if v != 0 {
return v;
}
std::env::var("CLJRS_OSR_THRESHOLD")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.unwrap_or_else(jit_threshold)
}
pub fn osr_poll(arity_id: u64, header: u32) -> OsrPoll {
let guard = OSR_TABLE.read().unwrap();
match guard.as_ref().and_then(|m| m.get(&(arity_id, header))) {
None => OsrPoll::NotRequested,
Some(OsrState::Queued) => OsrPoll::Pending,
Some(OsrState::Ready(slot)) => OsrPoll::Ready(slot.clone()),
Some(OsrState::Failed) => OsrPoll::Failed,
}
}
pub fn osr_request(arity_id: u64, header: u32, ir_func: &IrFunction) {
let hook = OSR_ENQUEUE_HOOK.get();
{
let mut guard = OSR_TABLE.write().unwrap();
let map = guard.get_or_insert_with(HashMap::new);
match map.entry((arity_id, header)) {
std::collections::hash_map::Entry::Occupied(_) => return,
std::collections::hash_map::Entry::Vacant(slot) => {
slot.insert(if hook.is_some() {
OsrState::Queued
} else {
OsrState::Failed
});
}
}
}
if let Some(hook) = hook {
cljrs_logging::feat_debug!(
"jit",
"osr enqueue arity_id={} header=bb{}",
arity_id,
header
);
hook(arity_id, header, Arc::new(ir_func.clone()));
}
}
pub fn store_osr_fn(
arity_id: u64,
header: u32,
ptr: *const (),
epoch: u64,
live_ins: Vec<cljrs_ir::VarId>,
) {
let mut guard = OSR_TABLE.write().unwrap();
let map = guard.get_or_insert_with(HashMap::new);
map.insert(
(arity_id, header),
OsrState::Ready(OsrSlot {
fn_ptr: ptr,
epoch,
live_ins: live_ins.into(),
}),
);
}
pub fn mark_osr_failed(arity_id: u64, header: u32) {
let mut guard = OSR_TABLE.write().unwrap();
let map = guard.get_or_insert_with(HashMap::new);
map.insert((arity_id, header), OsrState::Failed);
}
pub fn take_osr_epochs(arity_id: u64) -> Vec<u64> {
let mut guard = OSR_TABLE.write().unwrap();
let Some(map) = guard.as_mut() else {
return Vec::new();
};
let keys: Vec<(u64, u32)> = map
.keys()
.filter(|(a, _)| *a == arity_id)
.copied()
.collect();
let mut epochs = Vec::new();
for key in keys {
if let Some(OsrState::Ready(slot)) = map.remove(&key) {
epochs.push(slot.epoch);
}
}
epochs
}
struct ThreadFrames {
stack: Mutex<Vec<u64>>,
}
static FRAME_REGISTRY: RwLock<Vec<Weak<ThreadFrames>>> = RwLock::new(Vec::new());
thread_local! {
static MY_FRAMES: Arc<ThreadFrames> = {
let frames = Arc::new(ThreadFrames { stack: Mutex::new(Vec::new()) });
let mut reg = FRAME_REGISTRY.write().unwrap();
reg.retain(|w| w.strong_count() > 0);
reg.push(Arc::downgrade(&frames));
frames
};
}
pub struct JitFrameGuard {
epoch: u64,
}
impl Drop for JitFrameGuard {
fn drop(&mut self) {
MY_FRAMES.with(|f| {
let mut stack = f.stack.lock().unwrap();
if let Some(pos) = stack.iter().rposition(|&e| e == self.epoch) {
stack.remove(pos);
}
});
}
}
pub fn push_jit_frame(epoch: u64) -> JitFrameGuard {
MY_FRAMES.with(|f| f.stack.lock().unwrap().push(epoch));
JitFrameGuard { epoch }
}
pub fn current_jit_epoch() -> Option<u64> {
MY_FRAMES.with(|f| f.stack.lock().unwrap().last().copied())
}
pub fn live_epochs() -> HashSet<u64> {
let mut live = HashSet::new();
let reg = FRAME_REGISTRY.read().unwrap();
for weak in reg.iter() {
if let Some(frames) = weak.upgrade() {
for &epoch in frames.stack.lock().unwrap().iter() {
live.insert(epoch);
}
}
}
live
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn dispatch_jit_call(fn_ptr: *const (), args: &[*const Value]) -> *const Value {
unsafe {
match args.len() {
0 => {
let f: unsafe extern "C" fn() -> *const Value = std::mem::transmute(fn_ptr);
f()
}
1 => {
let f: unsafe extern "C" fn(*const Value) -> *const Value =
std::mem::transmute(fn_ptr);
f(args[0])
}
2 => {
let f: unsafe extern "C" fn(*const Value, *const Value) -> *const Value =
std::mem::transmute(fn_ptr);
f(args[0], args[1])
}
3 => {
let f: unsafe extern "C" fn(
*const Value,
*const Value,
*const Value,
) -> *const Value = std::mem::transmute(fn_ptr);
f(args[0], args[1], args[2])
}
4 => {
let f: unsafe extern "C" fn(
*const Value,
*const Value,
*const Value,
*const Value,
) -> *const Value = std::mem::transmute(fn_ptr);
f(args[0], args[1], args[2], args[3])
}
5 => {
let f: unsafe extern "C" fn(
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
) -> *const Value = std::mem::transmute(fn_ptr);
f(args[0], args[1], args[2], args[3], args[4])
}
6 => {
let f: unsafe extern "C" fn(
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
) -> *const Value = std::mem::transmute(fn_ptr);
f(args[0], args[1], args[2], args[3], args[4], args[5])
}
7 => {
let f: unsafe extern "C" fn(
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
) -> *const Value = std::mem::transmute(fn_ptr);
f(
args[0], args[1], args[2], args[3], args[4], args[5], args[6],
)
}
8 => {
let f: unsafe extern "C" fn(
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
*const Value,
) -> *const Value = std::mem::transmute(fn_ptr);
f(
args[0], args[1], args[2], args[3], args[4], args[5], args[6], args[7],
)
}
n => panic!("JIT dispatch: unsupported arity {n} (max 8 in Phase 10.1)"),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn epoch_round_trips_through_store_get_take() {
let id = 0xF100_0001;
let ptr = 0x1234usize as *const ();
store_native_fn(id, ptr, 777);
assert_eq!(get_native_fn(id), Some((ptr, 777)));
assert_eq!(take_native_epoch(id), Some(777));
assert_eq!(get_native_fn(id), None);
assert_eq!(take_native_epoch(id), None);
}
#[test]
fn frame_guard_marks_epoch_live_then_clears() {
let e = 0xBEEF_0001;
assert!(!live_epochs().contains(&e));
let guard = push_jit_frame(e);
assert!(live_epochs().contains(&e));
drop(guard);
assert!(!live_epochs().contains(&e));
}
#[test]
fn osr_slot_round_trips_and_rebind_takes_epochs() {
use cljrs_ir::VarId;
let id = 0xF200_0001;
assert!(matches!(osr_poll(id, 1), OsrPoll::NotRequested));
let ptr = 0x5678usize as *const ();
store_osr_fn(id, 1, ptr, 901, vec![VarId(3), VarId(4)]);
match osr_poll(id, 1) {
OsrPoll::Ready(slot) => {
assert_eq!(slot.fn_ptr, ptr);
assert_eq!(slot.epoch, 901);
assert_eq!(&*slot.live_ins, &[VarId(3), VarId(4)]);
}
_ => panic!("expected Ready"),
}
mark_osr_failed(id, 7);
assert!(matches!(osr_poll(id, 7), OsrPoll::Failed));
let epochs = take_osr_epochs(id);
assert_eq!(epochs, vec![901]);
assert!(matches!(osr_poll(id, 1), OsrPoll::NotRequested));
assert!(matches!(osr_poll(id, 7), OsrPoll::NotRequested));
}
#[test]
fn osr_request_without_hook_marks_failed() {
let id = 0xF200_0002;
let ir = IrFunction::new(None, None);
osr_request(id, 2, &ir);
assert!(matches!(osr_poll(id, 2), OsrPoll::Failed));
}
#[test]
fn record_interp_call_warm_lifecycle() {
set_ir_threshold(3);
let id = 0xF300_0001;
assert!(!record_interp_call(id));
assert!(!record_interp_call(id));
assert!(record_interp_call(id));
assert!(record_interp_call(id));
mark_lower_queued(id);
assert!(lower_queued(id));
assert!(!record_interp_call(id));
clear_lower_queued(id);
assert!(record_interp_call(id));
on_ir_published(id);
assert!(!record_interp_call(id));
set_ir_threshold(u32::MAX);
for _ in 0..10 {
assert!(!record_interp_call(id));
}
set_ir_threshold(0); }
#[test]
fn evict_entry_if_cold_respects_native_and_queued() {
let id = 0xF300_0002;
mark_lower_queued(id); store_native_fn(id, 0x4242usize as *const (), 555);
assert!(!evict_entry_if_cold(id));
assert_eq!(take_native_epoch(id), Some(555));
mark_lower_queued(id);
assert!(evict_entry_if_cold(id));
assert!(!lower_queued(id));
assert!(!evict_entry_if_cold(id)); }
#[test]
fn nested_frames_pop_in_lifo_order() {
let a = 0xBEEF_1001;
let b = 0xBEEF_1002;
let ga = push_jit_frame(a);
let gb = push_jit_frame(b);
let live = live_epochs();
assert!(live.contains(&a) && live.contains(&b));
drop(gb);
assert!(live_epochs().contains(&a));
assert!(!live_epochs().contains(&b));
drop(ga);
assert!(!live_epochs().contains(&a));
}
}