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;
use crate::tiered::backend::JitBackend;
use crate::tiered::tiers::Tiers;
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)
}
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 deopt_limit() -> u32 {
std::env::var("CLJRS_JIT_DEOPT_LIMIT")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.unwrap_or(10)
}
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),
}
}
}
unsafe impl Send for JitEntry {}
unsafe impl Sync for JitEntry {}
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,
}
}
#[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,
}
pub struct JitState {
entries: RwLock<HashMap<u64, Arc<JitEntry>>>,
osr: RwLock<HashMap<(u64, u32), OsrState>>,
spec_banned: RwLock<HashSet<u64>>,
bootstrap_watermark: AtomicU64,
backend: OnceLock<Arc<dyn JitBackend>>,
tiers: Weak<Tiers>,
}
impl JitState {
pub(crate) fn new(tiers: Weak<Tiers>) -> Self {
Self {
entries: RwLock::new(HashMap::new()),
osr: RwLock::new(HashMap::new()),
spec_banned: RwLock::new(HashSet::new()),
bootstrap_watermark: AtomicU64::new(0),
backend: OnceLock::new(),
tiers,
}
}
pub fn install_backend(&self, backend: Arc<dyn JitBackend>) {
let _ = self.backend.set(backend);
}
pub fn backend(&self) -> Option<&Arc<dyn JitBackend>> {
self.backend.get()
}
fn entry(&self, arity_id: u64) -> Arc<JitEntry> {
{
let guard = self.entries.read().unwrap();
if let Some(e) = guard.get(&arity_id) {
return e.clone();
}
}
self.entries
.write()
.unwrap()
.entry(arity_id)
.or_insert_with(|| Arc::new(JitEntry::new()))
.clone()
}
pub fn get_native_fn(&self, arity_id: u64) -> Option<(*const (), u64)> {
let guard = self.entries.read().unwrap();
let entry = guard.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(&self, arity_id: u64, ptr: *const (), epoch: u64) {
let entry = self.entry(arity_id);
entry.epoch.store(epoch, Ordering::Release);
entry.native_fn_ptr.store(ptr as *mut (), Ordering::Release);
}
pub fn take_native_epoch(&self, arity_id: u64) -> Option<u64> {
let entry = self.entries.write().unwrap().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))
}
}
pub fn stale_native_code(&self, arity_id: u64) {
let mut epochs = Vec::new();
if let Some(epoch) = self.take_native_epoch(arity_id) {
epochs.push(epoch);
}
epochs.extend(self.take_osr_epochs(arity_id));
if let Some(backend) = self.backend() {
for epoch in epochs {
backend.mark_stale(epoch);
}
}
}
pub fn arg_type_profile(&self, arity_id: u64) -> Option<Vec<u8>> {
let entry = self.entries.read().unwrap().get(&arity_id)?.clone();
let prof = entry.arg_profile.lock().unwrap();
if prof.is_empty() {
None
} else {
Some(prof.clone())
}
}
pub fn record_call(&self, arity_id: u64, ir_func: Arc<IrFunction>, profile_args: &[Value]) {
let entry = self.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(backend) = self.backend() {
cljrs_logging::feat_debug!("jit", "enqueue arity_id={} (count={})", arity_id, count);
backend.enqueue_function(self.tiers.clone(), arity_id, ir_func);
}
}
pub fn record_interp_call(&self, arity_id: u64) -> bool {
let threshold = ir_threshold();
if threshold == u32::MAX {
return false;
}
let entry = self.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(&self, arity_id: u64) -> bool {
self.entries
.read()
.unwrap()
.get(&arity_id)
.is_some_and(|e| e.compile_queued.load(Ordering::Relaxed))
}
pub fn pins_ir(&self, arity_id: u64) -> bool {
self.get_native_fn(arity_id).is_some() || self.compile_queued(arity_id)
}
pub fn lower_queued(&self, arity_id: u64) -> bool {
self.entries
.read()
.unwrap()
.get(&arity_id)
.is_some_and(|e| e.lower_queued.load(Ordering::Relaxed))
}
pub fn mark_lower_queued(&self, arity_id: u64) {
self.entry(arity_id)
.lower_queued
.store(true, Ordering::Relaxed);
}
pub fn clear_lower_queued(&self, arity_id: u64) {
if let Some(entry) = self.entries.read().unwrap().get(&arity_id) {
entry.lower_queued.store(false, Ordering::Relaxed);
}
}
pub fn on_ir_published(&self, arity_id: u64) {
self.entry(arity_id)
.invocation_count
.store(0, Ordering::Relaxed);
}
pub fn evict_entry_if_cold(&self, arity_id: u64) -> bool {
let mut guard = self.entries.write().unwrap();
let Some(entry) = guard.get(&arity_id) else {
return false;
};
if !entry.native_fn_ptr.load(Ordering::Acquire).is_null()
|| entry.compile_queued.load(Ordering::Relaxed)
{
return false;
}
guard.remove(&arity_id);
true
}
pub fn set_bootstrap_watermark(&self, w: u64) {
self.bootstrap_watermark.store(w, Ordering::Relaxed);
}
pub fn is_bootstrap_arity(&self, arity_id: u64) -> bool {
arity_id < self.bootstrap_watermark.load(Ordering::Relaxed)
}
#[inline]
pub fn is_deopt_result(&self, result: *const Value) -> bool {
self.backend()
.is_some_and(|b| b.deopt_sentinel() == result as usize)
}
pub fn take_pending_exception(&self) -> Option<Value> {
self.backend().and_then(|b| b.take_pending_exception())
}
pub fn specialization_allowed(&self, arity_id: u64) -> bool {
!self.spec_banned.read().unwrap().contains(&arity_id)
}
pub fn record_deopt(&self, arity_id: u64) {
let entry = self.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;
}
self.spec_banned.write().unwrap().insert(arity_id);
if let Some(epoch) = self.take_native_epoch(arity_id)
&& let Some(backend) = self.backend()
{
cljrs_logging::feat_debug!(
"jit",
"specialization discarded arity_id={} epoch={}",
arity_id,
epoch
);
backend.mark_stale(epoch);
}
}
pub fn osr_poll(&self, arity_id: u64, header: u32) -> OsrPoll {
match self.osr.read().unwrap().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(&self, arity_id: u64, header: u32, ir_func: &IrFunction) {
let has_backend = self.backend().is_some();
{
let mut guard = self.osr.write().unwrap();
match guard.entry((arity_id, header)) {
std::collections::hash_map::Entry::Occupied(_) => return,
std::collections::hash_map::Entry::Vacant(slot) => {
slot.insert(if has_backend {
OsrState::Queued
} else {
OsrState::Failed
});
}
}
}
if let Some(backend) = self.backend() {
cljrs_logging::feat_debug!(
"jit",
"osr enqueue arity_id={} header=bb{}",
arity_id,
header
);
backend.enqueue_osr(
self.tiers.clone(),
arity_id,
header,
Arc::new(ir_func.clone()),
);
}
}
pub fn store_osr_fn(
&self,
arity_id: u64,
header: u32,
ptr: *const (),
epoch: u64,
live_ins: Vec<cljrs_ir::VarId>,
) {
self.osr.write().unwrap().insert(
(arity_id, header),
OsrState::Ready(OsrSlot {
fn_ptr: ptr,
epoch,
live_ins: live_ins.into(),
}),
);
}
pub fn mark_osr_failed(&self, arity_id: u64, header: u32) {
self.osr
.write()
.unwrap()
.insert((arity_id, header), OsrState::Failed);
}
pub fn take_osr_epochs(&self, arity_id: u64) -> Vec<u64> {
let mut guard = self.osr.write().unwrap();
let keys: Vec<(u64, u32)> = guard
.keys()
.filter(|(a, _)| *a == arity_id)
.copied()
.collect();
let mut epochs = Vec::new();
for key in keys {
if let Some(OsrState::Ready(slot)) = guard.remove(&key) {
epochs.push(slot.epoch);
}
}
epochs
}
pub fn stale_osr_code(&self, arity_id: u64) {
let epochs = self.take_osr_epochs(arity_id);
if let Some(backend) = self.backend() {
for epoch in epochs {
backend.mark_stale(epoch);
}
}
}
fn published_epochs(&self) -> Vec<u64> {
let mut epochs = Vec::new();
for entry in self.entries.read().unwrap().values() {
if !entry.native_fn_ptr.load(Ordering::Acquire).is_null() {
epochs.push(entry.epoch.load(Ordering::Acquire));
}
}
for state in self.osr.read().unwrap().values() {
if let OsrState::Ready(slot) = state {
epochs.push(slot.epoch);
}
}
epochs
}
}
impl Drop for JitState {
fn drop(&mut self) {
let Some(backend) = self.backend.get() else {
return;
};
for epoch in self.published_epochs() {
backend.mark_stale(epoch);
}
}
}
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::*;
use crate::tiered::tiers::Tiers;
#[derive(Default)]
struct RecordingBackend {
staled: Mutex<Vec<u64>>,
}
impl RecordingBackend {
fn staled(&self) -> Vec<u64> {
let mut v = self.staled.lock().unwrap().clone();
v.sort_unstable();
v
}
}
impl JitBackend for RecordingBackend {
fn enqueue_function(&self, _: Weak<Tiers>, _: u64, _: Arc<cljrs_ir::IrFunction>) {}
fn enqueue_osr(&self, _: Weak<Tiers>, _: u64, _: u32, _: Arc<cljrs_ir::IrFunction>) {}
fn mark_stale(&self, epoch: u64) {
self.staled.lock().unwrap().push(epoch);
}
fn take_pending_exception(&self) -> Option<Value> {
None
}
fn deopt_sentinel(&self) -> usize {
0
}
fn compile_async_arity(&self, _: &Value, _: usize, _: &mut crate::env::env::Env) {}
}
fn jit() -> Arc<Tiers> {
Tiers::new(0xF000_0000)
}
#[test]
fn epoch_round_trips_through_store_get_take() {
let t = jit();
let id = 0xF100_0001;
let ptr = 0x1234usize as *const ();
t.jit().store_native_fn(id, ptr, 777);
assert_eq!(t.jit().get_native_fn(id), Some((ptr, 777)));
assert_eq!(t.jit().take_native_epoch(id), Some(777));
assert_eq!(t.jit().get_native_fn(id), None);
assert_eq!(t.jit().take_native_epoch(id), None);
}
#[test]
fn native_code_is_per_runtime() {
let a = jit();
let b = jit();
let id = 0xF100_0002;
a.jit().store_native_fn(id, 0x1234usize as *const (), 778);
assert!(a.jit().get_native_fn(id).is_some());
assert!(b.jit().get_native_fn(id).is_none());
assert!(!b.jit().compile_queued(id));
}
#[test]
fn dropping_a_runtime_stales_its_published_code() {
use cljrs_ir::VarId;
let backend = Arc::new(RecordingBackend::default());
let t = jit();
t.jit().install_backend(backend.clone());
t.jit()
.store_native_fn(0xF300_0001, 0x1000usize as *const (), 1001);
t.jit()
.store_native_fn(0xF300_0002, 0x2000usize as *const (), 1002);
t.jit().store_osr_fn(
0xF300_0001,
4,
0x3000usize as *const (),
1003,
vec![VarId(1)],
);
t.jit().mark_osr_failed(0xF300_0002, 9);
t.jit()
.store_native_fn(0xF300_0003, 0x4000usize as *const (), 1004);
t.jit().stale_native_code(0xF300_0003);
assert_eq!(backend.staled(), vec![1004], "redefinition released 1004");
drop(t);
assert_eq!(
backend.staled(),
vec![1001, 1002, 1003, 1004],
"the drop must release every epoch still published"
);
}
#[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 t = jit();
let id = 0xF200_0001;
assert!(matches!(t.jit().osr_poll(id, 1), OsrPoll::NotRequested));
let ptr = 0x5678usize as *const ();
t.jit()
.store_osr_fn(id, 1, ptr, 901, vec![VarId(3), VarId(4)]);
match t.jit().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"),
}
t.jit().mark_osr_failed(id, 7);
assert!(matches!(t.jit().osr_poll(id, 7), OsrPoll::Failed));
let epochs = t.jit().take_osr_epochs(id);
assert_eq!(epochs, vec![901]);
assert!(matches!(t.jit().osr_poll(id, 1), OsrPoll::NotRequested));
assert!(matches!(t.jit().osr_poll(id, 7), OsrPoll::NotRequested));
}
#[test]
fn osr_request_without_backend_marks_failed() {
let t = jit();
let id = 0xF200_0002;
let ir = IrFunction::new(None, None);
t.jit().osr_request(id, 2, &ir);
assert!(matches!(t.jit().osr_poll(id, 2), OsrPoll::Failed));
}
#[test]
fn record_interp_call_warm_lifecycle() {
set_ir_threshold(3);
let t = jit();
let id = 0xF300_0001;
assert!(!t.jit().record_interp_call(id));
assert!(!t.jit().record_interp_call(id));
assert!(t.jit().record_interp_call(id));
assert!(t.jit().record_interp_call(id));
t.jit().mark_lower_queued(id);
assert!(t.jit().lower_queued(id));
assert!(!t.jit().record_interp_call(id));
t.jit().clear_lower_queued(id);
assert!(t.jit().record_interp_call(id));
t.jit().on_ir_published(id);
assert!(!t.jit().record_interp_call(id));
set_ir_threshold(u32::MAX);
for _ in 0..10 {
assert!(!t.jit().record_interp_call(id));
}
set_ir_threshold(0); }
#[test]
fn evict_entry_if_cold_respects_native_and_queued() {
let t = jit();
let id = 0xF300_0002;
t.jit().mark_lower_queued(id); t.jit().store_native_fn(id, 0x4242usize as *const (), 555);
assert!(!t.jit().evict_entry_if_cold(id));
assert_eq!(t.jit().take_native_epoch(id), Some(555));
t.jit().mark_lower_queued(id);
assert!(t.jit().evict_entry_if_cold(id));
assert!(!t.jit().lower_queued(id));
assert!(!t.jit().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));
}
}