1use std::cell::Cell;
31use std::fs::{File, OpenOptions};
32use std::marker::PhantomData;
33use std::mem::size_of;
34use std::path::Path;
35use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
36
37use memmap2::{MmapMut, MmapOptions};
38
39pub const RESERVOIR_MAGIC: u64 = 0x4150_5253_4D50_4C31;
40pub const RESERVOIR_SLOT_PAYLOAD: usize = 56;
41
42#[repr(C, align(64))]
43pub struct ReservoirHeader {
44 pub magic: u64,
45 pub capacity: u32,
46 pub slot_size: u32,
47 pub total_seen: AtomicU64,
48 _pad: [u8; 40],
49}
50
51#[repr(C, align(64))]
52pub struct ReservoirSlot {
53 pub version: AtomicU32,
54 _pad: [u8; 4],
55 pub payload: [u8; RESERVOIR_SLOT_PAYLOAD],
56}
57
58const _: () = {
59 assert!(size_of::<ReservoirHeader>() == 64);
60 assert!(size_of::<ReservoirSlot>() == 64);
61};
62
63#[derive(Debug, Clone, Copy, PartialEq, Eq)]
64pub enum ReservoirError {
65 PayloadTooLarge,
66 LayoutMismatch,
67 IoError(std::io::ErrorKind),
68}
69
70impl From<std::io::Error> for ReservoirError {
71 fn from(e: std::io::Error) -> Self { Self::IoError(e.kind()) }
72}
73
74pub fn reservoir_file_size(capacity: usize) -> usize {
75 size_of::<ReservoirHeader>() + capacity * size_of::<ReservoirSlot>()
76}
77
78thread_local! {
79 static RNG_STATE: Cell<u64> = Cell::new({
80 let t = std::time::SystemTime::now()
81 .duration_since(std::time::UNIX_EPOCH)
82 .map(|d| d.as_nanos() as u64)
83 .unwrap_or(1);
84 let mix = t.wrapping_mul(0x9E37_79B9_7F4A_7C15);
85 if mix == 0 { 1 } else { mix }
86 });
87}
88
89#[inline]
90fn next_random_u64() -> u64 {
91 RNG_STATE.with(|s| {
92 let mut x = s.get();
93 x ^= x << 13;
94 x ^= x >> 7;
95 x ^= x << 17;
96 s.set(x);
97 x
98 })
99}
100
101pub struct SharedReservoirSampler<T: Copy + 'static> {
102 _file: File,
103 mmap: MmapMut,
104 capacity: usize,
105 _phantom: PhantomData<T>,
106 header_sidecar: subetha_core::HandshakeHeader,
107 ring_sidecar: Box<subetha_core::ObservationRing>,
108}
109
110unsafe impl<T: Copy + Send + 'static> Send for SharedReservoirSampler<T> {}
111unsafe impl<T: Copy + Sync + 'static> Sync for SharedReservoirSampler<T> {}
112
113impl<T: Copy + Send + Sync + 'static> subetha_sidecar::AdaptiveInstance for SharedReservoirSampler<T> {
114 fn header(&self) -> &subetha_core::HandshakeHeader { &self.header_sidecar }
115 fn ring(&self) -> &subetha_core::ObservationRing { &self.ring_sidecar }
116 fn make_policy(&self) -> Box<dyn subetha_sidecar::Policy> {
117 Box::new(subetha_sidecar::NoMigrationPolicy)
118 }
119}
120
121impl<T: Copy + 'static> SharedReservoirSampler<T> {
122 pub fn create(
129 path: impl AsRef<Path>, capacity: usize,
130 ) -> Result<Self, ReservoirError> {
131 if size_of::<T>() > RESERVOIR_SLOT_PAYLOAD {
132 return Err(ReservoirError::PayloadTooLarge);
133 }
134 assert!(capacity >= 1);
135 let (file, mmap) = crate::mmf_attach::create_or_attach(
136 path.as_ref(),
137 reservoir_file_size(capacity),
138 |ptr| unsafe { Self::init_region(ptr, capacity) },
139 |ptr| unsafe { (*(ptr as *const ReservoirHeader)).magic == RESERVOIR_MAGIC },
140 )?;
141 Self::from_region(file, mmap, capacity)
142 }
143
144 unsafe fn init_region(ptr: *mut u8, capacity: usize) {
152 let hdr = ptr as *mut ReservoirHeader;
153 unsafe {
154 (*hdr).capacity = capacity as u32;
155 (*hdr).slot_size = size_of::<T>() as u32;
156 std::ptr::write_volatile(&raw mut (*hdr).magic, RESERVOIR_MAGIC);
157 }
158 }
159
160 fn from_region(
163 file: File,
164 mmap: MmapMut,
165 capacity: usize,
166 ) -> Result<Self, ReservoirError> {
167 let hdr = unsafe { &*(mmap.as_ptr() as *const ReservoirHeader) };
168 if hdr.magic != RESERVOIR_MAGIC
169 || hdr.capacity != capacity as u32
170 || hdr.slot_size != size_of::<T>() as u32
171 {
172 return Err(ReservoirError::LayoutMismatch);
173 }
174 Ok(Self {
175 _file: file, mmap, capacity, _phantom: PhantomData,
176 header_sidecar: subetha_core::HandshakeHeader::new(),
177 ring_sidecar: Box::new(subetha_core::ObservationRing::new()),
178 })
179 }
180
181 pub fn open(
182 path: impl AsRef<Path>, expected_capacity: usize,
183 ) -> Result<Self, ReservoirError> {
184 if size_of::<T>() > RESERVOIR_SLOT_PAYLOAD {
185 return Err(ReservoirError::PayloadTooLarge);
186 }
187 let total = reservoir_file_size(expected_capacity);
188 let file = OpenOptions::new().read(true).write(true).open(path.as_ref())?;
189 if file.metadata()?.len() < total as u64 {
190 return Err(ReservoirError::LayoutMismatch);
191 }
192 let mmap = unsafe { MmapOptions::new().len(total).map_mut(&file)? };
193 Self::from_region(file, mmap, expected_capacity)
194 }
195
196 #[inline]
197 pub fn capacity(&self) -> usize { self.capacity }
198
199 pub fn total_seen(&self) -> u64 {
200 self.header().total_seen.load(Ordering::Acquire)
201 }
202
203 fn header(&self) -> &ReservoirHeader {
204 unsafe { &*(self.mmap.as_ptr() as *const ReservoirHeader) }
205 }
206
207 fn slot(&self, i: usize) -> &ReservoirSlot {
208 let base = unsafe { self.mmap.as_ptr().add(size_of::<ReservoirHeader>()) };
209 unsafe { &*(base.add(i * size_of::<ReservoirSlot>()) as *const ReservoirSlot) }
210 }
211
212 fn write_slot(&self, idx: usize, value: T) {
214 let slot = self.slot(idx);
215 slot.version.fetch_add(1, Ordering::AcqRel); let dst = unsafe {
217 let base = self.mmap.as_ptr().add(size_of::<ReservoirHeader>())
218 .add(idx * size_of::<ReservoirSlot>())
219 .add(std::mem::offset_of!(ReservoirSlot, payload));
220 base as *mut u8
221 };
222 unsafe {
223 std::ptr::copy_nonoverlapping(
224 &value as *const T as *const u8,
225 dst,
226 size_of::<T>(),
227 );
228 }
229 slot.version.fetch_add(1, Ordering::AcqRel); }
231
232 fn read_slot(&self, idx: usize) -> T {
234 let slot = self.slot(idx);
235 loop {
236 let v1 = slot.version.load(Ordering::Acquire);
237 if v1 & 1 != 0 {
238 std::hint::spin_loop();
239 continue;
240 }
241 let mut out = std::mem::MaybeUninit::<T>::uninit();
242 let src = unsafe {
243 self.mmap.as_ptr().add(size_of::<ReservoirHeader>())
244 .add(idx * size_of::<ReservoirSlot>())
245 .add(std::mem::offset_of!(ReservoirSlot, payload))
246 };
247 unsafe {
248 std::ptr::copy_nonoverlapping(
249 src, out.as_mut_ptr() as *mut u8, size_of::<T>(),
250 );
251 }
252 let v2 = slot.version.load(Ordering::Acquire);
253 if v1 == v2 {
254 return unsafe { out.assume_init() };
255 }
256 }
257 }
258
259 pub fn record(&self, value: T) -> Option<usize> {
263 let prev = self.header().total_seen.fetch_add(1, Ordering::AcqRel);
264 let n = prev + 1; let k = self.capacity as u64;
266 let r = if n <= k {
267 let idx = (n - 1) as usize;
269 self.write_slot(idx, value);
270 Some(idx)
271 } else {
272 let j = (next_random_u64() % n) + 1;
274 if j <= k {
275 let idx = (j - 1) as usize;
276 self.write_slot(idx, value);
277 Some(idx)
278 } else {
279 None
280 }
281 };
282 self.ring_sidecar.push_op(
283 crate::sidecar_ops::reservoir::OP_RECORD,
284 if r.is_none() { 2 } else { 0 }, );
286 r
287 }
288
289 pub fn snapshot(&self) -> Vec<T> {
292 let filled = (self.total_seen() as usize).min(self.capacity);
293 let v: Vec<T> = (0..filled).map(|i| self.read_slot(i)).collect();
294 self.ring_sidecar
295 .push_op(crate::sidecar_ops::reservoir::OP_SNAPSHOT, 0);
296 v
297 }
298
299 pub fn reset(&self) {
302 self.header().total_seen.store(0, Ordering::Release);
303 }
304
305 pub fn flush(&self) -> Result<(), ReservoirError> {
306 self.mmap.flush()?;
307 Ok(())
308 }
309 pub fn flush_async(&self) -> Result<(), ReservoirError> {
310 self.mmap.flush_async()?;
311 Ok(())
312 }
313}
314
315#[cfg(test)]
316mod tests {
317 use super::*;
318 use std::sync::Arc;
319 use std::thread;
320
321 fn tmp(name: &str) -> std::path::PathBuf {
322 let mut p = std::env::temp_dir();
323 let pid = std::process::id();
324 p.push(format!("subetha-reservoir-{name}-{pid}.bin"));
325 p
326 }
327
328 #[test]
329 fn create_initial_state_is_empty() {
330 let p = tmp("init");
331 let r: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 10).unwrap();
332 assert_eq!(r.capacity(), 10);
333 assert_eq!(r.total_seen(), 0);
334 assert_eq!(r.snapshot(), Vec::<u32>::new());
335 std::fs::remove_file(&p).ok();
336 }
337
338 #[test]
339 fn first_k_items_always_accepted() {
340 let p = tmp("first-k");
341 let r: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 5).unwrap();
342 for i in 0..5u32 {
343 let idx = r.record(i);
344 assert_eq!(idx, Some(i as usize));
345 }
346 let snap = r.snapshot();
347 assert_eq!(snap, vec![0, 1, 2, 3, 4]);
348 assert_eq!(r.total_seen(), 5);
349 std::fs::remove_file(&p).ok();
350 }
351
352 #[test]
355 fn second_create_attaches_and_keeps_samples() {
356 let p = tmp("attach");
357 std::fs::remove_file(&p).ok();
358 let r: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 5).unwrap();
359 for i in 0..3u32 { r.record(i); }
360
361 let r2: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 5).unwrap();
362 assert_eq!(r2.total_seen(), 3, "attach restarted a live sampler");
363 assert_eq!(r2.snapshot().len(), 3);
364 assert!(matches!(
365 SharedReservoirSampler::<u32>::create(&p, 4),
366 Err(ReservoirError::LayoutMismatch),
367 ));
368
369 r2.reset();
370 assert_eq!(r.total_seen(), 0, "reset did not restart for every handle");
371 drop(r);
372 drop(r2);
373 std::fs::remove_file(&p).ok();
374 }
375
376 #[test]
377 fn after_capacity_some_items_rejected() {
378 let p = tmp("after-cap");
379 let r: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 5).unwrap();
380 for i in 0..5u32 { r.record(i); }
381 let mut accepted = 0;
383 let mut rejected = 0;
384 for i in 5..100u32 {
385 if r.record(i).is_some() { accepted += 1; } else { rejected += 1; }
386 }
387 assert!(rejected > 50, "expected mostly rejections; got {accepted} accepted, {rejected} rejected");
390 assert_eq!(r.total_seen(), 100);
391 std::fs::remove_file(&p).ok();
392 }
393
394 #[test]
395 fn snapshot_always_has_correct_length() {
396 let p = tmp("snap-len");
397 let r: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 10).unwrap();
398 for i in 0..3u32 { r.record(i); }
400 assert_eq!(r.snapshot().len(), 3);
401 for i in 3..100u32 { r.record(i); }
403 assert_eq!(r.snapshot().len(), 10);
404 std::fs::remove_file(&p).ok();
405 }
406
407 #[test]
408 fn uniform_distribution_over_many_trials() {
409 let p = tmp("uniform");
414 let n_trials = 1000;
415 let n_items = 100u32;
416 let mut counts = [0u32; 10];
417 for trial in 0..n_trials {
418 let path = std::env::temp_dir().join(
419 format!("subetha-reservoir-uniform-{trial}-{}.bin", std::process::id()),
420 );
421 let r: SharedReservoirSampler<u32>
422 = SharedReservoirSampler::create(&path, 1).unwrap();
423 for i in 0..n_items { r.record(i); }
424 let snap = r.snapshot();
425 let kept = snap[0];
426 let bucket = (kept * 10 / n_items) as usize;
427 counts[bucket.min(9)] += 1;
428 std::fs::remove_file(&path).ok();
429 }
430 for (i, &c) in counts.iter().enumerate() {
432 assert!((30..=200).contains(&c),
433 "bucket {i} count {c} is way out of expected ~100");
434 }
435 let _p = p;
436 }
437
438 #[test]
439 fn reset_clears_count() {
440 let p = tmp("reset");
441 let r: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 5).unwrap();
442 for i in 0..10u32 { r.record(i); }
443 assert_eq!(r.total_seen(), 10);
444 r.reset();
445 assert_eq!(r.total_seen(), 0);
446 assert_eq!(r.snapshot(), Vec::<u32>::new());
447 std::fs::remove_file(&p).ok();
448 }
449
450 #[test]
451 fn cross_handle_visibility() {
452 let p = tmp("cross-handle");
453 let w: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 5).unwrap();
454 let rdr: SharedReservoirSampler<u32> = SharedReservoirSampler::open(&p, 5).unwrap();
455 for i in 0..5u32 { w.record(i); }
456 let snap = rdr.snapshot();
457 assert_eq!(snap, vec![0, 1, 2, 3, 4]);
458 assert_eq!(rdr.total_seen(), 5);
459 std::fs::remove_file(&p).ok();
460 }
461
462 #[test]
463 fn payload_too_large_rejected() {
464 #[allow(dead_code)]
465 struct Big([u8; RESERVOIR_SLOT_PAYLOAD + 1]);
466 impl Copy for Big {}
467 impl Clone for Big { fn clone(&self) -> Self { *self } }
468 let p = tmp("too-large");
469 assert_eq!(
470 SharedReservoirSampler::<Big>::create(&p, 4).err(),
471 Some(ReservoirError::PayloadTooLarge)
472 );
473 std::fs::remove_file(&p).ok();
474 }
475
476 #[test]
477 fn concurrent_recorders_count_correctly() {
478 let p = tmp("concurrent");
479 let r: Arc<SharedReservoirSampler<u32>>
480 = Arc::new(SharedReservoirSampler::create(&p, 10).unwrap());
481 let n_threads = 4;
482 let per_thread = 100;
483 let mut handles = vec![];
484 for t in 0..n_threads as u32 {
485 let r = r.clone();
486 handles.push(thread::spawn(move || {
487 for i in 0..per_thread as u32 {
488 r.record(t * 1000 + i);
489 }
490 }));
491 }
492 for h in handles { h.join().unwrap(); }
493 assert_eq!(r.total_seen() as usize, n_threads * per_thread);
494 let snap = r.snapshot();
495 assert_eq!(snap.len(), 10);
496 std::fs::remove_file(&p).ok();
497 }
498
499 #[test]
500 fn struct_payload_round_trip() {
501 #[derive(Clone, Copy, Debug, PartialEq)]
502 #[repr(C)]
503 struct LogEntry { ts: u64, code: u32, severity: u32 }
504 let p = tmp("struct");
505 let r: SharedReservoirSampler<LogEntry>
506 = SharedReservoirSampler::create(&p, 3).unwrap();
507 let e1 = LogEntry { ts: 100, code: 1, severity: 1 };
508 let e2 = LogEntry { ts: 200, code: 2, severity: 2 };
509 let e3 = LogEntry { ts: 300, code: 3, severity: 3 };
510 r.record(e1);
511 r.record(e2);
512 r.record(e3);
513 let snap = r.snapshot();
514 assert_eq!(snap, vec![e1, e2, e3]);
515 std::fs::remove_file(&p).ok();
516 }
517
518 #[test]
519 fn disk_persistence_survives_reopen() {
520 let p = tmp("disk");
521 {
522 let r: SharedReservoirSampler<u32> = SharedReservoirSampler::create(&p, 5).unwrap();
523 for i in 0..5u32 { r.record(i); }
524 r.flush().unwrap();
525 }
526 let r2: SharedReservoirSampler<u32> = SharedReservoirSampler::open(&p, 5).unwrap();
527 assert_eq!(r2.total_seen(), 5);
528 assert_eq!(r2.snapshot(), vec![0, 1, 2, 3, 4]);
529 std::fs::remove_file(&p).ok();
530 }
531}