use crate::rand::gen_random;
pub struct IdGenerator {}
impl IdGenerator {
pub fn generate() -> u64 {
let random_bytes: [u8; 8] = gen_random();
u64::from_le_bytes(random_bytes)
}
}
pub use cubecl_common::stream_id::StreamId;
use core::hash::{BuildHasher, Hasher};
use alloc::str::FromStr;
use data_encoding::BASE32_DNSSEC;
use serde::{Deserialize, Serialize};
type DefaultHashBuilder = core::hash::BuildHasherDefault<ahash::AHasher>;
#[derive(Debug, Hash, PartialEq, Eq, Clone, Copy, PartialOrd, Ord, Serialize, Deserialize)]
pub struct ParamId {
value: u64,
}
impl From<u64> for ParamId {
fn from(value: u64) -> Self {
Self { value }
}
}
impl Default for ParamId {
fn default() -> Self {
Self::new()
}
}
impl ParamId {
pub fn new() -> Self {
Self {
value: IdGenerator::generate(),
}
}
pub fn val(&self) -> u64 {
self.value
}
}
impl FromStr for ParamId {
type Err = &'static str;
fn from_str(encoded: &str) -> Result<Self, Self::Err> {
let u64_id: Option<u64> = match BASE32_DNSSEC.decode(encoded.as_bytes()) {
Ok(bytes) => {
let mut buffer = [0u8; 8];
buffer[..bytes.len()].copy_from_slice(&bytes);
Some(u64::from_le_bytes(buffer))
}
Err(_) => match uuid::Uuid::try_parse(encoded) {
Ok(id) => {
let mut hasher = DefaultHashBuilder::default().build_hasher();
hasher.write(id.as_bytes());
Some(hasher.finish())
}
Err(_) => None,
},
};
u64_id.map(Self::from).ok_or("Invalid id.")
}
}
impl core::fmt::Display for ParamId {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let encoded = BASE32_DNSSEC.encode(&self.value.to_le_bytes());
f.write_str(&encoded)
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::collections::BTreeSet;
use alloc::string::ToString;
#[cfg(feature = "std")]
use dashmap::DashSet; #[cfg(feature = "std")]
use std::{sync::Arc, thread};
#[test]
fn uniqueness_test() {
const IDS_CNT: usize = 10_000;
let mut set: BTreeSet<u64> = BTreeSet::new();
for _i in 0..IDS_CNT {
assert!(set.insert(IdGenerator::generate()));
}
assert_eq!(set.len(), IDS_CNT);
}
#[cfg(feature = "std")]
#[test]
fn thread_safety_test() {
const NUM_THREADS: usize = 10;
const NUM_REPEATS: usize = 1_000;
const EXPECTED_TOTAL_IDS: usize = NUM_THREADS * NUM_REPEATS;
let set: Arc<DashSet<u64>> = Arc::new(DashSet::new());
let mut handles = vec![];
for _ in 0..NUM_THREADS {
let set = set.clone();
let handle = thread::spawn(move || {
for _i in 0..NUM_REPEATS {
assert!(set.insert(IdGenerator::generate()));
}
});
handles.push(handle);
}
for handle in handles {
handle.join().unwrap();
}
assert_eq!(set.len(), EXPECTED_TOTAL_IDS);
}
#[test]
fn param_serde_try_deserialize() {
let val = ParamId::from(123456u64);
let deserialized = ParamId::from_str(&val.to_string()).unwrap();
assert_eq!(val, deserialized);
assert_eq!(ParamId::from_str("invalid_id"), Err("Invalid id."));
}
#[test]
fn param_serde_deserialize() {
let val = ParamId::from(123456u64);
let deserialized = ParamId::from_str(&val.to_string()).unwrap();
assert_eq!(val, deserialized);
}
#[test]
fn param_serde_deserialize_legacy() {
let legacy_val = [45u8; 6];
let param_id = ParamId::from_str(&BASE32_DNSSEC.encode(&legacy_val)).unwrap();
assert_eq!(param_id.val().to_le_bytes()[0..6], legacy_val);
assert_eq!(param_id.val().to_le_bytes()[6..], [0, 0]);
}
#[test]
fn param_serde_deserialize_legacy_uuid() {
let legacy_id = "30b82c23-788d-4d63-a743-ada258d5f13c";
let param_id1 = ParamId::from_str(legacy_id).unwrap();
let param_id2 = ParamId::from_str(legacy_id).unwrap();
assert_eq!(param_id1, param_id2);
}
#[test]
#[should_panic = "Invalid id."]
fn param_serde_deserialize_invalid_id() {
let invalid_uuid = "30b82c23-788d-4d63-ada258d5f13c";
let _ = ParamId::from_str(invalid_uuid).unwrap();
}
}