use std::cell::RefCell;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::OnceLock;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::level_filters::LevelFilter;
use tracing::Metadata;
thread_local! {
static LCG_STATE: RefCell<u64> = RefCell::new(seed_for_new_thread());
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub struct SampleConfig {
pub rate: f64,
pub min_level: LevelFilter,
}
impl Default for SampleConfig {
fn default() -> Self {
Self {
rate: 1.0,
min_level: LevelFilter::TRACE,
}
}
}
impl SampleConfig {
pub fn new(rate: f64) -> Self {
Self::default().with_rate(rate)
}
pub fn with_rate(mut self, rate: f64) -> Self {
self.rate = rate.clamp(0.0, 1.0);
self
}
pub fn with_min_level(mut self, min_level: LevelFilter) -> Self {
self.min_level = min_level;
self
}
}
#[derive(Debug, Clone)]
pub struct Sampler {
config: SampleConfig,
}
impl Sampler {
pub fn new(config: SampleConfig) -> Self {
Self { config }
}
pub fn should_sample(&self, meta: &Metadata) -> bool {
let level = meta.level();
if *level > self.config.min_level {
return false;
}
if self.config.rate >= 1.0 {
return true;
}
if self.config.rate <= 0.0 {
return false;
}
thread_random() < self.config.rate
}
}
impl super::gate::EventGate for Sampler {
fn allows(&self, meta: &Metadata<'_>) -> bool {
self.should_sample(meta)
}
}
fn thread_random() -> f64 {
const MULTIPLIER: u64 = 6_364_136_223_846_793_005;
const INCREMENT: u64 = 1_442_695_040_888_963_407;
const SCALE: f64 = 4_294_967_296.0;
LCG_STATE.with(|cell| {
let mut state = cell.borrow_mut();
*state = state.wrapping_mul(MULTIPLIER).wrapping_add(INCREMENT);
f64::from((*state >> 32) as u32) / SCALE
})
}
static THREAD_ORDINAL: AtomicU64 = AtomicU64::new(0);
fn seed_for_new_thread() -> u64 {
let ordinal = THREAD_ORDINAL.fetch_add(1, Ordering::Relaxed);
splitmix64(process_nonce() ^ ordinal.wrapping_mul(GOLDEN_GAMMA))
}
fn process_nonce() -> u64 {
static NONCE: OnceLock<u64> = OnceLock::new();
*NONCE.get_or_init(|| {
let wall_clock = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(GOLDEN_GAMMA);
let stack_marker = &wall_clock as *const u64 as u64;
splitmix64(wall_clock ^ stack_marker)
})
}
const GOLDEN_GAMMA: u64 = 0x9E37_79B9_7F4A_7C15;
fn splitmix64(seed: u64) -> u64 {
let mut z = seed;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_config_logs_everything() {
let cfg = SampleConfig::default();
assert_eq!(cfg.rate, 1.0);
assert_eq!(cfg.min_level, LevelFilter::TRACE);
}
#[test]
fn thread_random_in_range() {
for _ in 0..1000 {
let r = thread_random();
assert!((0.0..1.0).contains(&r), "got {r}");
}
}
#[test]
fn thread_random_approximates_uniform() {
let mut buckets = [0u32; 10];
for _ in 0..10_000 {
let r = thread_random();
let bucket = (r * 10.0) as usize;
buckets[bucket.min(9)] += 1;
}
for &count in &buckets {
assert!(count > 700, "bucket underflow: {count}");
assert!(count < 1300, "bucket overflow: {count}");
}
}
#[test]
fn sampler_clone_is_independent() {
let s1 = Sampler::new(SampleConfig::default());
let _s2 = s1.clone();
assert_eq!(s1.config.rate, 1.0);
}
fn draw_sequence() -> Vec<u64> {
(0..32).map(|_| thread_random().to_bits()).collect()
}
#[test]
fn separate_threads_draw_different_sequences() {
let first = std::thread::spawn(draw_sequence);
let second = std::thread::spawn(draw_sequence);
let first = first.join().expect("first sampling thread panicked");
let second = second.join().expect("second sampling thread panicked");
assert_ne!(
first, second,
"threads sharing an LCG seed correlate their sampling decisions"
);
}
#[test]
fn every_thread_seed_is_distinct() {
let handles: Vec<_> = (0..8)
.map(|_| std::thread::spawn(seed_for_new_thread))
.collect();
let mut seeds: Vec<u64> = handles
.into_iter()
.map(|h| h.join().expect("seeding thread panicked"))
.collect();
let drawn = seeds.len();
seeds.sort_unstable();
seeds.dedup();
assert_eq!(seeds.len(), drawn, "thread seeds collided");
}
#[test]
fn splitmix64_avalanches_adjacent_inputs() {
let (low, high) = (splitmix64(0), splitmix64(1));
let differing_bits = (low ^ high).count_ones();
assert!(
(16..=48).contains(&differing_bits),
"adjacent seeds differ in only {differing_bits} bits"
);
}
}