use base64::Engine;
use std::collections::HashMap;
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::sync::Mutex;
pub(crate) struct RidGenerator {
db_counter: AtomicU32,
coll_counters: Mutex<HashMap<u32, u32>>,
doc_counter: AtomicU64,
}
impl RidGenerator {
pub fn new() -> Self {
Self {
db_counter: AtomicU32::new(1),
coll_counters: Mutex::new(HashMap::new()),
doc_counter: AtomicU64::new(1),
}
}
pub fn next_database_rid(&self) -> (u32, String) {
let id = self.db_counter.fetch_add(1, Ordering::SeqCst);
let bytes = id.to_be_bytes();
(id, encode_rid(&bytes))
}
pub fn next_collection_rid(&self, db_id: u32) -> (u32, String) {
let coll_id = {
let mut counters = self.coll_counters.lock().unwrap();
let entry = counters.entry(db_id).or_insert(0);
*entry += 1;
*entry
};
let coll_with_high_bit = coll_id | 0x80000000;
let mut bytes = [0u8; 8];
bytes[..4].copy_from_slice(&db_id.to_be_bytes());
bytes[4..8].copy_from_slice(&coll_with_high_bit.to_be_bytes());
(coll_id, encode_rid(&bytes))
}
pub fn next_document_rid(&self, db_id: u32, coll_id: u32) -> (u64, String) {
let doc_id = self.doc_counter.fetch_add(1, Ordering::SeqCst);
let doc_with_type = doc_id << 4; let coll_with_high_bit = coll_id | 0x80000000;
let mut bytes = [0u8; 16];
bytes[..4].copy_from_slice(&db_id.to_be_bytes());
bytes[4..8].copy_from_slice(&coll_with_high_bit.to_be_bytes());
bytes[8..16].copy_from_slice(&doc_with_type.to_be_bytes());
(doc_id, encode_rid(&bytes))
}
pub fn forget_database(&self, db_id: u32) {
self.coll_counters.lock().unwrap().remove(&db_id);
}
pub fn next_pkrange_rid(&self, db_id: u32, coll_id: u32, pkrange_id: u32) -> String {
let pkr_id = (pkrange_id as u64) << 4 | 0x05; let coll_with_high_bit = coll_id | 0x80000000;
let mut bytes = [0u8; 16];
bytes[..4].copy_from_slice(&db_id.to_be_bytes());
bytes[4..8].copy_from_slice(&coll_with_high_bit.to_be_bytes());
bytes[8..16].copy_from_slice(&pkr_id.to_be_bytes());
encode_rid(&bytes)
}
}
fn encode_rid(bytes: &[u8]) -> String {
let encoded = base64::engine::general_purpose::STANDARD.encode(bytes);
encoded.replace('/', "-")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn database_rid_4_bytes() {
let gen = RidGenerator::new();
let (id, rid) = gen.next_database_rid();
assert_eq!(id, 1);
let decoded_str = rid.replace('-', "/");
let bytes = base64::engine::general_purpose::STANDARD
.decode(&decoded_str)
.unwrap();
assert_eq!(bytes.len(), 4);
let db_id = u32::from_be_bytes(bytes[..4].try_into().unwrap());
assert_eq!(db_id, 1);
}
#[test]
fn collection_rid_encodes_parent() {
let gen = RidGenerator::new();
let (db_id, _db_rid) = gen.next_database_rid();
let (coll_id, coll_rid) = gen.next_collection_rid(db_id);
assert_eq!(coll_id, 1);
let decoded_str = coll_rid.replace('-', "/");
let bytes = base64::engine::general_purpose::STANDARD
.decode(&decoded_str)
.unwrap();
assert_eq!(bytes.len(), 8);
let parent = u32::from_be_bytes(bytes[..4].try_into().unwrap());
assert_eq!(parent, db_id);
let coll = u32::from_be_bytes(bytes[4..8].try_into().unwrap());
assert!(coll & 0x80000000 != 0);
assert_eq!(coll & 0x7FFFFFFF, coll_id);
}
#[test]
fn document_rid_encodes_hierarchy() {
let gen = RidGenerator::new();
let (db_id, _) = gen.next_database_rid();
let (coll_id, _) = gen.next_collection_rid(db_id);
let (doc_id, doc_rid) = gen.next_document_rid(db_id, coll_id);
assert_eq!(doc_id, 1);
let decoded_str = doc_rid.replace('-', "/");
let bytes = base64::engine::general_purpose::STANDARD
.decode(&decoded_str)
.unwrap();
assert_eq!(bytes.len(), 16);
let parent_db = u32::from_be_bytes(bytes[..4].try_into().unwrap());
assert_eq!(parent_db, db_id);
let doc_raw = u64::from_be_bytes(bytes[8..16].try_into().unwrap());
assert_eq!(doc_raw & 0x0F, 0x00);
}
#[test]
fn pkrange_rid_type_nibble() {
let gen = RidGenerator::new();
let (db_id, _) = gen.next_database_rid();
let (coll_id, _) = gen.next_collection_rid(db_id);
let pkr_rid = gen.next_pkrange_rid(db_id, coll_id, 0);
let decoded_str = pkr_rid.replace('-', "/");
let bytes = base64::engine::general_purpose::STANDARD
.decode(&decoded_str)
.unwrap();
assert_eq!(bytes.len(), 16);
let pkr_raw = u64::from_be_bytes(bytes[8..16].try_into().unwrap());
assert_eq!(pkr_raw & 0x0F, 0x05);
}
#[test]
fn monotonic_ids() {
let gen = RidGenerator::new();
let (id1, _) = gen.next_database_rid();
let (id2, _) = gen.next_database_rid();
assert_eq!(id1, 1);
assert_eq!(id2, 2);
}
}