use std::num::NonZeroU32;
use std::sync::atomic::{AtomicU32, Ordering};
use serde::{Deserialize, Serialize};
macro_rules! define_id_type {
($name:ident) => {
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize,
)]
#[serde(transparent)]
pub struct $name(NonZeroU32);
impl $name {
pub fn new(val: u32) -> Option<Self> {
NonZeroU32::new(val).map(Self)
}
pub fn to_raw(self) -> u32 {
self.0.get()
}
}
impl From<NonZeroU32> for $name {
fn from(val: NonZeroU32) -> Self {
Self(val)
}
}
};
}
define_id_type!(FileId);
define_id_type!(SymbolId);
define_id_type!(SnapshotId);
define_id_type!(DataNodeId);
#[derive(Debug)]
pub struct IdGenerator<T> {
counter: AtomicU32,
_marker: std::marker::PhantomData<T>,
}
impl<T> IdGenerator<T> {
pub fn new() -> Self {
Self {
counter: AtomicU32::new(1),
_marker: std::marker::PhantomData,
}
}
pub fn with_start(start: u32) -> Self {
Self {
counter: AtomicU32::new(start.max(1)),
_marker: std::marker::PhantomData,
}
}
pub fn next(&self) -> T
where
T: From<NonZeroU32>,
{
let val = self.counter.fetch_add(1, Ordering::SeqCst);
let nz = NonZeroU32::new(val)
.expect("IdGenerator exhausted u32 space (allocated > u32::MAX ids)");
T::from(nz)
}
}
impl<T> Default for IdGenerator<T> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashSet;
use std::sync::Arc;
#[test]
fn file_id_sequential() {
let idgen = IdGenerator::<FileId>::new();
assert_eq!(idgen.next(), FileId::new(1).unwrap());
assert_eq!(idgen.next(), FileId::new(2).unwrap());
assert_eq!(idgen.next(), FileId::new(3).unwrap());
}
#[test]
fn symbol_id_sequential() {
let idgen = IdGenerator::<SymbolId>::new();
assert_eq!(idgen.next(), SymbolId::new(1).unwrap());
assert_eq!(idgen.next(), SymbolId::new(2).unwrap());
assert_eq!(idgen.next(), SymbolId::new(3).unwrap());
}
#[test]
fn snapshot_id_sequential() {
let idgen = IdGenerator::<SnapshotId>::new();
assert_eq!(idgen.next(), SnapshotId::new(1).unwrap());
assert_eq!(idgen.next(), SnapshotId::new(2).unwrap());
assert_eq!(idgen.next(), SnapshotId::new(3).unwrap());
}
#[test]
fn id_generator_thread_safe() {
let idgen = Arc::new(IdGenerator::<SymbolId>::new());
let mut handles = Vec::new();
for _ in 0..4 {
let idgen = Arc::clone(&idgen);
handles.push(std::thread::spawn(move || {
let mut ids = Vec::with_capacity(250);
for _ in 0..250 {
ids.push(idgen.next());
}
ids
}));
}
let all_ids: Vec<SymbolId> = handles
.into_iter()
.flat_map(|h| h.join().unwrap_or_default())
.collect();
let unique: HashSet<SymbolId> = all_ids.iter().copied().collect();
assert_eq!(unique.len(), 1000, "expected 1000 unique IDs");
assert!(
unique
.iter()
.all(|id| id.to_raw() >= 1 && id.to_raw() <= 1000),
"all IDs must be in range 1..=1000"
);
}
#[test]
fn file_id_serde_roundtrip() {
let original = FileId::new(42).unwrap();
let json = serde_json::to_string(&original).unwrap();
let roundtrip: FileId = serde_json::from_str(&json).unwrap();
assert_eq!(original, roundtrip);
}
#[test]
fn symbol_id_serde_roundtrip() {
let original = SymbolId::new(99).unwrap();
let json = serde_json::to_string(&original).unwrap();
let roundtrip: SymbolId = serde_json::from_str(&json).unwrap();
assert_eq!(original, roundtrip);
}
#[test]
fn file_id_to_raw() {
assert_eq!(FileId::new(7).unwrap().to_raw(), 7);
}
#[test]
fn id_generator_default() {
let idgen = IdGenerator::<FileId>::default();
assert_eq!(idgen.next(), FileId::new(1).unwrap());
}
#[test]
fn id_generator_with_start_zero_safeguard() {
let idgen = IdGenerator::<SymbolId>::with_start(0);
let first = idgen.next();
assert_eq!(
first.to_raw(),
1,
"with_start(0) must sanitize to start ID 1"
);
}
#[test]
fn zero_is_invalid() {
assert!(FileId::new(0).is_none());
assert!(SymbolId::new(0).is_none());
assert!(SnapshotId::new(0).is_none());
assert!(DataNodeId::new(0).is_none());
}
#[test]
fn zero_rejected_on_deserialize() {
let res: Result<FileId, _> = serde_json::from_str("0");
assert!(res.is_err(), "deserializing 0 must fail");
}
#[test]
fn option_id_is_niche_optimized() {
assert_eq!(
std::mem::size_of::<Option<FileId>>(),
std::mem::size_of::<FileId>()
);
assert_eq!(std::mem::size_of::<Option<FileId>>(), 4);
assert_eq!(std::mem::size_of::<Option<SymbolId>>(), 4);
assert_eq!(std::mem::size_of::<Option<SnapshotId>>(), 4);
assert_eq!(std::mem::size_of::<Option<DataNodeId>>(), 4);
}
}