use std::ops::{Index, IndexMut};
use std::sync::RwLock;
use jstd::registry::{Identified, Identifier, Registry};
use rustc_hash::FxHashMap as HashMap;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use crate::{
types::TypeId,
value::literal::{Literal, LiteralId},
};
fn mask_to_size(value: u64, size: usize) -> u64 {
if size >= 8 {
value
} else {
value & ((1u64 << (size * 8)) - 1)
}
}
pub struct Interner<Id: Identifier, T> {
inner: RwLock<Registry<Id, T>>,
}
impl<Id: Identifier, T> Default for Interner<Id, T> {
fn default() -> Self {
Self {
inner: RwLock::new(Registry::default()),
}
}
}
impl<Id: Identifier, T: Clone> Clone for Interner<Id, T> {
fn clone(&self) -> Self {
Self {
inner: RwLock::new(self.read().clone()),
}
}
}
impl<Id: Identifier, T> Interner<Id, T> {
fn read(&self) -> std::sync::RwLockReadGuard<'_, Registry<Id, T>> {
self.inner.read().expect("interner RwLock poisoned")
}
pub fn push(&self, value: T) -> Id {
self.inner
.write()
.expect("interner RwLock poisoned")
.push(value)
}
pub fn len(&self) -> usize {
self.read().len()
}
pub fn is_empty(&self) -> bool {
self.read().is_empty()
}
}
impl<Id: Identifier, T: Clone> Interner<Id, T> {
pub fn iter(&self) -> impl Iterator<Item = Identified<Id, T>> {
self.read()
.iter()
.map(|item| Identified::new(item.id, item.inner.clone()))
.collect::<Vec<_>>()
.into_iter()
}
}
impl<Id: Identifier, T> Index<Id> for Interner<Id, T> {
type Output = T;
fn index(&self, id: Id) -> &T {
let ptr: *const T = {
let reg = self.read();
®[id] as *const T
};
unsafe { &*ptr }
}
}
impl<Id: Identifier, T> IndexMut<Id> for Interner<Id, T> {
fn index_mut(&mut self, id: Id) -> &mut T {
&mut self.inner.get_mut().expect("interner RwLock poisoned")[id]
}
}
impl<Id: Identifier, T: Serialize> Serialize for Interner<Id, T> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.read().serialize(serializer)
}
}
impl<'de, Id: Identifier, T: Deserialize<'de>> Deserialize<'de> for Interner<Id, T> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(Self {
inner: RwLock::new(Registry::deserialize(deserializer)?),
})
}
}
#[derive(Default)]
pub struct LiteralInterner {
inner: RwLock<LiteralPool>,
}
#[derive(Default, Clone)]
struct LiteralPool {
literals: Registry<LiteralId, Literal>,
cache: HashMap<(u64, TypeId), LiteralId>,
}
impl Clone for LiteralInterner {
fn clone(&self) -> Self {
Self {
inner: RwLock::new(self.read().clone()),
}
}
}
impl LiteralInterner {
fn read(&self) -> std::sync::RwLockReadGuard<'_, LiteralPool> {
self.inner.read().expect("literal interner RwLock poisoned")
}
pub fn get_or_make_typed_literal(&self, value: u64, type_id: TypeId, size: usize) -> LiteralId {
let value = mask_to_size(value, size);
if let Some(&id) = self.read().cache.get(&(value, type_id)) {
return id;
}
let mut pool = self
.inner
.write()
.expect("literal interner RwLock poisoned");
if let Some(&id) = pool.cache.get(&(value, type_id)) {
return id;
}
let id = pool.literals.push(Literal {
value,
type_id,
symbolic: None,
});
pool.cache.insert((value, type_id), id);
id
}
pub fn push_literal(&self, literal: Literal) -> LiteralId {
self.inner
.write()
.expect("literal interner RwLock poisoned")
.literals
.push(literal)
}
pub fn len(&self) -> usize {
self.read().literals.len()
}
pub fn is_empty(&self) -> bool {
self.read().literals.is_empty()
}
pub fn iter(&self) -> impl Iterator<Item = Identified<LiteralId, Literal>> {
self.read()
.literals
.iter()
.map(|item| Identified::new(item.id, item.inner.clone()))
.collect::<Vec<_>>()
.into_iter()
}
}
impl Index<LiteralId> for LiteralInterner {
type Output = Literal;
fn index(&self, id: LiteralId) -> &Literal {
let ptr: *const Literal = {
let pool = self.read();
&pool.literals[id] as *const Literal
};
unsafe { &*ptr }
}
}
impl IndexMut<LiteralId> for LiteralInterner {
fn index_mut(&mut self, id: LiteralId) -> &mut Literal {
&mut self
.inner
.get_mut()
.expect("literal interner RwLock poisoned")
.literals[id]
}
}
impl Serialize for LiteralInterner {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.read().literals.serialize(serializer)
}
}
impl<'de> Deserialize<'de> for LiteralInterner {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let literals = Registry::<LiteralId, Literal>::deserialize(deserializer)?;
let mut cache: HashMap<(u64, TypeId), LiteralId> = HashMap::default();
for item in literals.iter() {
if item.inner.symbolic.is_none() {
cache
.entry((item.inner.value, item.inner.type_id))
.or_insert(item.id);
}
}
Ok(Self {
inner: RwLock::new(LiteralPool { literals, cache }),
})
}
}