use std::{
collections::hash_map::DefaultHasher,
fmt,
hash::{Hash, Hasher},
num::NonZero,
sync::atomic::{AtomicU64, Ordering},
};
use num_bigint::BigInt;
use num_traits::ToPrimitive;
use strum::EnumCount;
use crate::{heap::HeapId, intern::StaticStrings};
#[derive(Clone, Copy, PartialEq, Eq)]
#[repr(transparent)]
pub(crate) struct HashValue(NonZero<u64>);
impl HashValue {
#[inline]
#[must_use]
pub const fn new(hash: u64) -> Self {
Self(match NonZero::new(!hash) {
Some(nz) => nz,
None => NonZero::<u64>::MIN,
})
}
#[inline]
#[must_use]
pub const fn raw(self) -> u64 {
!self.0.get()
}
}
impl Hash for HashValue {
#[inline]
fn hash<H: Hasher>(&self, state: &mut H) {
self.raw().hash(state);
}
}
impl fmt::Debug for HashValue {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("HashValue").field(&self.raw()).finish()
}
}
impl serde::Serialize for HashValue {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.raw().serialize(serializer)
}
}
impl<'de> serde::Deserialize<'de> for HashValue {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(Self::new(u64::deserialize(deserializer)?))
}
}
#[inline]
pub(crate) fn hash_one(value: impl Hash) -> HashValue {
let mut hasher = DefaultHasher::new();
value.hash(&mut hasher);
HashValue::new(hasher.finish())
}
#[inline]
pub(crate) fn identity_hash(id: HeapId) -> HashValue {
hash_one(id)
}
#[inline]
pub(crate) fn hash_python_str(s: &str) -> HashValue {
let mut hasher = DefaultHasher::new();
s.hash(&mut hasher);
HashValue::new(hasher.finish())
}
#[inline]
pub(crate) fn hash_python_bytes(b: &[u8]) -> HashValue {
let mut hasher = DefaultHasher::new();
b.hash(&mut hasher);
HashValue::new(hasher.finish())
}
#[inline]
pub(crate) fn hash_python_long_int(bi: &BigInt) -> HashValue {
if let Some(i) = bi.to_i64() {
HashValue::new(i.cast_unsigned())
} else {
let mut hasher = DefaultHasher::new();
let (sign, bytes) = bi.to_bytes_le();
sign.hash(&mut hasher);
bytes.hash(&mut hasher);
HashValue::new(hasher.finish())
}
}
#[derive(Debug, Clone)]
pub(crate) struct WithHash<T> {
value: T,
hash: HashValue,
}
impl<T> WithHash<T> {
#[inline]
pub fn value(&self) -> &T {
&self.value
}
#[inline]
pub fn hash(&self) -> HashValue {
self.hash
}
}
impl WithHash<String> {
#[inline]
pub fn for_str(value: String) -> Self {
let hash = hash_python_str(&value);
Self { value, hash }
}
}
impl WithHash<Vec<u8>> {
#[inline]
pub fn for_bytes(value: Vec<u8>) -> Self {
let hash = hash_python_bytes(&value);
Self { value, hash }
}
}
impl WithHash<BigInt> {
#[inline]
pub fn for_long_int(value: BigInt) -> Self {
let hash = hash_python_long_int(&value);
Self { value, hash }
}
}
impl<T: serde::Serialize> serde::Serialize for WithHash<T> {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.value.serialize(serializer)
}
}
impl<'de> serde::Deserialize<'de> for WithHash<String> {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(Self::for_str(String::deserialize(deserializer)?))
}
}
impl<'de> serde::Deserialize<'de> for WithHash<Vec<u8>> {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(Self::for_bytes(Vec::<u8>::deserialize(deserializer)?))
}
}
impl<'de> serde::Deserialize<'de> for WithHash<BigInt> {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(Self::for_long_int(BigInt::deserialize(deserializer)?))
}
}
pub(crate) struct LazyHashTable<const N: usize> {
cells: [AtomicU64; N],
}
impl<const N: usize> LazyHashTable<N> {
pub const fn new() -> Self {
Self {
cells: [const { AtomicU64::new(0) }; N],
}
}
#[inline]
pub fn get_or_compute(&self, index: usize, compute: impl FnOnce() -> HashValue) -> HashValue {
if let Some(stored) = NonZero::new(self.cells[index].load(Ordering::Relaxed)) {
HashValue(stored)
} else {
let h = compute();
self.cells[index].store(h.0.get(), Ordering::Relaxed);
h
}
}
}
pub(crate) static ASCII_HASHES: LazyHashTable<128> = LazyHashTable::new();
pub(crate) static STATIC_HASHES: LazyHashTable<{ StaticStrings::COUNT }> = LazyHashTable::new();