use shape_vm::bytecode::FunctionHash;
use std::collections::HashMap;
use crate::optimizer::Tier2CacheKey;
#[derive(Debug, Clone)]
pub struct CacheEntry {
pub code_ptr: *const u8,
pub function_hash: FunctionHash,
pub schema_version: u32,
pub feedback_epoch: u32,
pub dependencies: Vec<FunctionHash>,
pub tier2_key: Option<Tier2CacheKey>,
}
unsafe impl Send for CacheEntry {}
unsafe impl Sync for CacheEntry {}
pub struct JitCodeCache {
entries: HashMap<FunctionHash, CacheEntry>,
dependents: HashMap<FunctionHash, Vec<FunctionHash>>,
}
unsafe impl Send for JitCodeCache {}
unsafe impl Sync for JitCodeCache {}
impl JitCodeCache {
pub fn new() -> Self {
Self {
entries: HashMap::new(),
dependents: HashMap::new(),
}
}
pub fn with_capacity(capacity: usize) -> Self {
Self {
entries: HashMap::with_capacity(capacity),
dependents: HashMap::new(),
}
}
pub fn get(&self, hash: &FunctionHash) -> Option<*const u8> {
self.entries.get(hash).map(|e| e.code_ptr)
}
pub fn insert(&mut self, hash: FunctionHash, ptr: *const u8) {
self.remove_dependency_edges(&hash);
self.entries.insert(
hash,
CacheEntry {
code_ptr: ptr,
function_hash: hash,
schema_version: 0,
feedback_epoch: 0,
dependencies: Vec::new(),
tier2_key: None,
},
);
}
pub fn insert_entry(&mut self, entry: CacheEntry) {
let hash = entry.function_hash;
self.remove_dependency_edges(&hash);
for dep in &entry.dependencies {
self.dependents.entry(*dep).or_default().push(hash);
}
self.entries.insert(hash, entry);
}
pub fn invalidate_by_dependency(&mut self, changed_hash: &FunctionHash) -> Vec<FunctionHash> {
let mut invalidated = Vec::new();
let mut worklist = vec![*changed_hash];
while let Some(current) = worklist.pop() {
if let Some(deps) = self.dependents.remove(¤t) {
for dep_hash in deps {
if self.entries.remove(&dep_hash).is_some() {
invalidated.push(dep_hash);
worklist.push(dep_hash);
}
}
}
}
for inv in &invalidated {
self.remove_dependency_edges(inv);
}
invalidated
}
pub fn get_entry(&self, hash: &FunctionHash) -> Option<&CacheEntry> {
self.entries.get(hash)
}
pub fn contains(&self, hash: &FunctionHash) -> bool {
self.entries.contains_key(hash)
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn clear(&mut self) {
self.entries.clear();
self.dependents.clear();
}
fn remove_dependency_edges(&mut self, hash: &FunctionHash) {
if let Some(entry) = self.entries.get(hash) {
let deps: Vec<FunctionHash> = entry.dependencies.clone();
for dep in &deps {
if let Some(rev) = self.dependents.get_mut(dep) {
rev.retain(|h| h != hash);
if rev.is_empty() {
self.dependents.remove(dep);
}
}
}
}
}
}
impl Default for JitCodeCache {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_cache() {
let cache = JitCodeCache::new();
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
assert!(cache.get(&FunctionHash::ZERO).is_none());
}
#[test]
fn insert_and_get() {
let mut cache = JitCodeCache::new();
let hash = FunctionHash([0xAB; 32]);
let fake_ptr = 0xDEAD_BEEF_usize as *const u8;
cache.insert(hash, fake_ptr);
assert_eq!(cache.len(), 1);
assert!(!cache.is_empty());
assert!(cache.contains(&hash));
assert_eq!(cache.get(&hash), Some(fake_ptr));
}
#[test]
fn missing_hash_returns_none() {
let mut cache = JitCodeCache::new();
let hash_a = FunctionHash([1u8; 32]);
let hash_b = FunctionHash([2u8; 32]);
cache.insert(hash_a, 0x1 as *const u8);
assert!(cache.get(&hash_b).is_none());
assert!(!cache.contains(&hash_b));
}
#[test]
fn overwrite_entry() {
let mut cache = JitCodeCache::new();
let hash = FunctionHash([0xCC; 32]);
let ptr1 = 0x1000_usize as *const u8;
let ptr2 = 0x2000_usize as *const u8;
cache.insert(hash, ptr1);
assert_eq!(cache.get(&hash), Some(ptr1));
cache.insert(hash, ptr2);
assert_eq!(cache.get(&hash), Some(ptr2));
assert_eq!(cache.len(), 1);
}
#[test]
fn clear_removes_all() {
let mut cache = JitCodeCache::new();
cache.insert(FunctionHash([1; 32]), 0x1 as *const u8);
cache.insert(FunctionHash([2; 32]), 0x2 as *const u8);
assert_eq!(cache.len(), 2);
cache.clear();
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
}
#[test]
fn with_capacity() {
let cache = JitCodeCache::with_capacity(64);
assert!(cache.is_empty());
}
fn make_entry(hash: FunctionHash, ptr: usize, deps: Vec<FunctionHash>) -> CacheEntry {
CacheEntry {
code_ptr: ptr as *const u8,
function_hash: hash,
schema_version: 1,
feedback_epoch: 1,
dependencies: deps,
tier2_key: None,
}
}
#[test]
fn test_insert_entry_with_dependencies() {
let mut cache = JitCodeCache::new();
let callee = FunctionHash([0x01; 32]);
let caller = FunctionHash([0x02; 32]);
cache.insert_entry(make_entry(callee, 0x1000, vec![]));
cache.insert_entry(make_entry(caller, 0x2000, vec![callee]));
assert_eq!(cache.len(), 2);
assert!(cache.contains(&callee));
assert!(cache.contains(&caller));
let entry = cache.get_entry(&caller).unwrap();
assert_eq!(entry.schema_version, 1);
assert_eq!(entry.feedback_epoch, 1);
assert_eq!(entry.dependencies, vec![callee]);
}
#[test]
fn test_invalidate_by_dependency() {
let mut cache = JitCodeCache::new();
let callee = FunctionHash([0x01; 32]);
let caller_a = FunctionHash([0x02; 32]);
let caller_b = FunctionHash([0x03; 32]);
let unrelated = FunctionHash([0x04; 32]);
cache.insert_entry(make_entry(callee, 0x1000, vec![]));
cache.insert_entry(make_entry(caller_a, 0x2000, vec![callee]));
cache.insert_entry(make_entry(caller_b, 0x3000, vec![callee]));
cache.insert_entry(make_entry(unrelated, 0x4000, vec![]));
assert_eq!(cache.len(), 4);
let mut invalidated = cache.invalidate_by_dependency(&callee);
invalidated.sort_by_key(|h| h.0);
assert_eq!(invalidated.len(), 2);
assert!(invalidated.contains(&caller_a));
assert!(invalidated.contains(&caller_b));
assert!(cache.contains(&callee));
assert!(cache.contains(&unrelated));
assert!(!cache.contains(&caller_a));
assert!(!cache.contains(&caller_b));
assert_eq!(cache.len(), 2);
}
#[test]
fn test_invalidate_cascading() {
let mut cache = JitCodeCache::new();
let c = FunctionHash([0x01; 32]);
let b = FunctionHash([0x02; 32]);
let a = FunctionHash([0x03; 32]);
cache.insert_entry(make_entry(c, 0x1000, vec![]));
cache.insert_entry(make_entry(b, 0x2000, vec![c]));
cache.insert_entry(make_entry(a, 0x3000, vec![b]));
assert_eq!(cache.len(), 3);
let mut invalidated = cache.invalidate_by_dependency(&c);
invalidated.sort_by_key(|h| h.0);
assert_eq!(invalidated.len(), 2);
assert!(invalidated.contains(&b));
assert!(invalidated.contains(&a));
assert!(cache.contains(&c));
assert!(!cache.contains(&b));
assert!(!cache.contains(&a));
assert_eq!(cache.len(), 1);
}
#[test]
fn test_get_entry_returns_metadata() {
let mut cache = JitCodeCache::new();
let hash = FunctionHash([0xAA; 32]);
let dep = FunctionHash([0xBB; 32]);
cache.insert_entry(CacheEntry {
code_ptr: 0x5000 as *const u8,
function_hash: hash,
schema_version: 42,
feedback_epoch: 7,
dependencies: vec![dep],
tier2_key: None,
});
let entry = cache.get_entry(&hash).unwrap();
assert_eq!(entry.code_ptr, 0x5000 as *const u8);
assert_eq!(entry.function_hash, hash);
assert_eq!(entry.schema_version, 42);
assert_eq!(entry.feedback_epoch, 7);
assert_eq!(entry.dependencies, vec![dep]);
assert_eq!(cache.get(&hash), Some(0x5000 as *const u8));
assert!(cache.get_entry(&FunctionHash([0xFF; 32])).is_none());
}
#[test]
fn test_tier2_cache_key_stored_in_entry() {
let mut cache = JitCodeCache::new();
let root = FunctionHash([0x10; 32]);
let inlined_callee = FunctionHash([0x20; 32]);
let key = Tier2CacheKey::with_versions(
root.0,
vec![inlined_callee.0],
1, 5, 3, );
cache.insert_entry(CacheEntry {
code_ptr: 0x8000 as *const u8,
function_hash: root,
schema_version: 5,
feedback_epoch: 3,
dependencies: vec![inlined_callee],
tier2_key: Some(key.clone()),
});
let entry = cache.get_entry(&root).unwrap();
let stored_key = entry.tier2_key.as_ref().unwrap();
assert_eq!(stored_key.root_hash, root.0);
assert_eq!(stored_key.inlined_hashes, vec![inlined_callee.0]);
assert_eq!(stored_key.schema_version, 5);
assert_eq!(stored_key.feedback_epoch, 3);
assert_eq!(stored_key.compiler_version, 1);
let key_no_versions = Tier2CacheKey::new(root.0, vec![inlined_callee.0], 1);
assert_ne!(stored_key.combined_hash(), key_no_versions.combined_hash());
}
#[test]
fn test_invalidate_with_tier2_entries() {
let mut cache = JitCodeCache::new();
let callee = FunctionHash([0x01; 32]);
let optimized = FunctionHash([0x02; 32]);
let key = Tier2CacheKey::with_versions(optimized.0, vec![callee.0], 1, 0, 0);
cache.insert_entry(CacheEntry {
code_ptr: 0x1000 as *const u8,
function_hash: callee,
schema_version: 0,
feedback_epoch: 0,
dependencies: vec![],
tier2_key: None,
});
cache.insert_entry(CacheEntry {
code_ptr: 0x2000 as *const u8,
function_hash: optimized,
schema_version: 0,
feedback_epoch: 0,
dependencies: vec![callee],
tier2_key: Some(key),
});
let invalidated = cache.invalidate_by_dependency(&callee);
assert_eq!(invalidated.len(), 1);
assert_eq!(invalidated[0], optimized);
assert!(!cache.contains(&optimized));
assert!(cache.contains(&callee));
}
}