Skip to main content

provide_telemetry/
sampling.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 provide.io llc
2// SPDX-License-Identifier: Apache-2.0
3// SPDX-Comment: Part of provide-telemetry.
4//
5
6use std::collections::BTreeMap;
7use std::sync::{Mutex, OnceLock};
8
9use crate::errors::TelemetryError;
10use crate::health::increment_dropped;
11
12#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
13pub enum Signal {
14    Logs,
15    Traces,
16    Metrics,
17}
18
19#[derive(Clone, Debug, PartialEq)]
20pub struct SamplingPolicy {
21    pub default_rate: f64,
22    pub overrides: BTreeMap<String, f64>,
23}
24
25impl Default for SamplingPolicy {
26    fn default() -> Self {
27        Self {
28            default_rate: 1.0,
29            overrides: BTreeMap::new(),
30        }
31    }
32}
33
34static POLICIES: OnceLock<Mutex<BTreeMap<Signal, SamplingPolicy>>> = OnceLock::new();
35
36fn policies() -> &'static Mutex<BTreeMap<Signal, SamplingPolicy>> {
37    POLICIES.get_or_init(|| {
38        Mutex::new(BTreeMap::from([
39            (Signal::Logs, SamplingPolicy::default()),
40            (Signal::Traces, SamplingPolicy::default()),
41            (Signal::Metrics, SamplingPolicy::default()),
42        ]))
43    })
44}
45
46pub fn set_sampling_policy(
47    signal: Signal,
48    policy: SamplingPolicy,
49) -> Result<SamplingPolicy, TelemetryError> {
50    let normalized = SamplingPolicy {
51        default_rate: policy.default_rate.clamp(0.0, 1.0),
52        overrides: policy
53            .overrides
54            .into_iter()
55            .map(|(key, rate)| (key, rate.clamp(0.0, 1.0)))
56            .collect(),
57    };
58    policies()
59        .lock()
60        .expect("sampling policy lock poisoned")
61        .insert(signal, normalized.clone());
62    Ok(normalized)
63}
64
65pub fn get_sampling_policy(signal: Signal) -> Result<SamplingPolicy, TelemetryError> {
66    policies()
67        .lock()
68        .expect("sampling policy lock poisoned")
69        .get(&signal)
70        .cloned()
71        .ok_or_else(|| TelemetryError::new("unknown signal"))
72}
73
74pub fn should_sample(signal: Signal, key: Option<&str>) -> Result<bool, TelemetryError> {
75    let policy = get_sampling_policy(signal)?;
76    let rate = key
77        .and_then(|value| policy.overrides.get(value).copied())
78        .unwrap_or(policy.default_rate);
79
80    if rate >= 1.0 {
81        return Ok(true);
82    }
83    if rate <= 0.0 {
84        increment_dropped(signal, 1);
85        return Ok(false);
86    }
87
88    let keep = key
89        .map(|value| {
90            let total = value
91                .as_bytes()
92                .iter()
93                .fold(0u64, |acc, byte| acc + u64::from(*byte));
94            let normalized = (total % 100) as f64 / 100.0;
95            normalized < rate
96        })
97        .unwrap_or(rate >= 0.5);
98    if !keep {
99        increment_dropped(signal, 1);
100    }
101    Ok(keep)
102}
103
104pub fn _reset_sampling_for_tests() {
105    *policies().lock().expect("sampling policy lock poisoned") = BTreeMap::from([
106        (Signal::Logs, SamplingPolicy::default()),
107        (Signal::Traces, SamplingPolicy::default()),
108        (Signal::Metrics, SamplingPolicy::default()),
109    ]);
110}
111
112#[cfg(test)]
113mod tests {
114    use super::*;
115    use crate::health::{_reset_health_for_tests, get_health_snapshot};
116    use crate::testing::acquire_test_state_lock;
117
118    #[test]
119    fn sampling_test_boundary_rates_and_reset_helper() {
120        let _guard = acquire_test_state_lock();
121        _reset_sampling_for_tests();
122        set_sampling_policy(
123            Signal::Logs,
124            SamplingPolicy {
125                default_rate: 0.0,
126                overrides: BTreeMap::new(),
127            },
128        )
129        .expect("policy should set");
130        assert!(!should_sample(Signal::Logs, None).expect("sampling should work"));
131
132        set_sampling_policy(
133            Signal::Logs,
134            SamplingPolicy {
135                default_rate: 1.0,
136                overrides: BTreeMap::new(),
137            },
138        )
139        .expect("policy should set");
140        assert!(should_sample(Signal::Logs, None).expect("sampling should work"));
141
142        set_sampling_policy(
143            Signal::Logs,
144            SamplingPolicy {
145                default_rate: 0.25,
146                overrides: BTreeMap::from([("special".to_string(), 0.75)]),
147            },
148        )
149        .expect("policy should set");
150        _reset_sampling_for_tests();
151        let reset = get_sampling_policy(Signal::Logs).expect("policy should exist");
152        assert_eq!(reset.default_rate, 1.0);
153        assert!(reset.overrides.is_empty());
154    }
155
156    #[test]
157    fn sampling_test_keyed_hashing_hits_threshold_edges() {
158        let _guard = acquire_test_state_lock();
159        _reset_sampling_for_tests();
160        _reset_health_for_tests();
161        let before = get_health_snapshot().dropped_logs;
162        set_sampling_policy(
163            Signal::Logs,
164            SamplingPolicy {
165                default_rate: 0.5,
166                overrides: BTreeMap::from([("edge".to_string(), 0.5)]),
167            },
168        )
169        .expect("policy should set");
170
171        let after_first_keep = get_health_snapshot().dropped_logs;
172        assert!(should_sample(Signal::Logs, Some("1")).expect("sampling should work"));
173        assert_eq!(get_health_snapshot().dropped_logs, after_first_keep);
174
175        assert!(!should_sample(Signal::Logs, Some("11")).expect("sampling should work"));
176        let after_first_drop = get_health_snapshot().dropped_logs;
177        assert_eq!(after_first_drop - before, 1);
178
179        assert!(!should_sample(Signal::Logs, Some("2")).expect("sampling should work"));
180        let after_second_drop = get_health_snapshot().dropped_logs;
181        assert_eq!(after_second_drop - before, 2);
182
183        assert!(should_sample(Signal::Logs, None).expect("sampling should work"));
184        let after = get_health_snapshot().dropped_logs;
185        assert_eq!(after - before, 2);
186    }
187}