Skip to main content

subetha_cxc/
shared_reservoir_sampler.rs

1//! `SharedReservoirSampler<T>` - cross-process uniform random
2//! sampling via Vitter's Algorithm R.
3//!
4//! Maintains `k` reservoir slots. After processing N items, each
5//! has equal probability `k/N` of being in the reservoir.
6//!
7//! # Algorithm
8//!
9//! For item n (1-indexed):
10//! - If n <= k: slot[n-1] = item.
11//! - Else: j = uniform(1..=n). If j <= k: slot[j-1] = item.
12//!
13//! # Safety
14//!
15//! - No spin loops, no CAS retry loops, no Drop guards.
16//! - Bounded capacity at create.
17//! - `total_seen.fetch_add` is monotonic (no underflow).
18//! - Concurrent races on `slot[j]` for the SAME `j` after capacity
19//!   is exceeded just keep one of the racers' values; statistically
20//!   the uniform-sampling property is preserved because both values
21//!   were equally eligible.
22//!
23//! # SeqLock per slot
24//!
25//! For T larger than 8 bytes, concurrent writes risk tearing. We use
26//! a SeqLock cell (version + payload) per slot so readers spin on
27//! odd version and re-read on version change. Same protocol as
28//! [`SharedCell`](crate::SharedCell).
29
30use 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    /// Obtain the sampler at `path`, initializing an empty one if the
123    /// path does not yet exist and attaching to it if it does.
124    /// Attaching leaves the live sample set and `total_seen` in place;
125    /// a region built with a different capacity or payload type is a
126    /// `LayoutMismatch`. The in-place [`reset`](Self::reset) restarts
127    /// a live sampler.
128    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    /// Lay out an empty sampler: config first, magic last, because
145    /// attachers spin on it. The zeroed region is already the empty
146    /// slot array and `total_seen` 0.
147    ///
148    /// # Safety
149    /// `ptr` addresses at least `reservoir_file_size(capacity)`
150    /// writable zeroed bytes.
151    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    /// Wrap an initialized region, refusing one built with a different
161    /// capacity or payload type.
162    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    /// SeqLock-write a payload into a slot.
213    fn write_slot(&self, idx: usize, value: T) {
214        let slot = self.slot(idx);
215        slot.version.fetch_add(1, Ordering::AcqRel); // odd
216        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); // even
230    }
231
232    /// SeqLock-read a payload from a slot.
233    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    /// Record a value. Returns the slot index it landed in, or
260    /// None if the value was rejected (the reservoir kept its
261    /// existing sample for this position).
262    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; // 1-indexed item number
265        let k = self.capacity as u64;
266        let r = if n <= k {
267            // Reservoir not full yet; always accept.
268            let idx = (n - 1) as usize;
269            self.write_slot(idx, value);
270            Some(idx)
271        } else {
272            // Vitter R: j = uniform(1..=n); if j <= k, accept at slot j-1.
273            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 }, // 2 = rejected
285        );
286        r
287    }
288
289    /// Snapshot the current reservoir. Returns min(total_seen, k)
290    /// slots filled so far; the rest are unused.
291    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    /// Reset to empty (total_seen = 0; slots become invalid but
300    /// not zeroed - next record overwrites them).
301    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    /// A second create attaches with the live sample set in place; the
353    /// in-place reset is what restarts sampling.
354    #[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        // After capacity, each new item has probability 5/n of acceptance.
382        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        // Expect most to be rejected since acceptance probability drops to
388        // ~5%. Allow generous bounds for randomness.
389        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        // After 3 records, snapshot has 3 items.
399        for i in 0..3u32 { r.record(i); }
400        assert_eq!(r.snapshot().len(), 3);
401        // After 100 records, snapshot has 10 (capacity).
402        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        // For capacity=1 and N=100, each item should appear in the
410        // reservoir with probability 1/100. Over 1000 trials, each
411        // item should appear ~10 times. Bin into 10 buckets and
412        // check no bucket is hugely off.
413        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        // Each bucket should be ~100. Allow [50, 200] for stochastic noise.
431        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}