Skip to main content

rill_ml/
error.rs

1//! Error types for RillML.
2//!
3//! All public APIs that can receive invalid input return [`Result<_, RillError>`]
4//! instead of panicking. Only truly unrecoverable internal invariant violations
5//! use assertions.
6
7use thiserror::Error;
8
9/// The unified error type returned by all RillML public APIs.
10#[derive(Debug, Error)]
11#[non_exhaustive]
12pub enum RillError {
13    /// A feature slice did not match the expected dimension.
14    #[error("dimension mismatch: expected {expected}, got {actual}")]
15    DimensionMismatch {
16        /// The expected number of features.
17        expected: usize,
18        /// The actual number of features provided.
19        actual: usize,
20    },
21
22    /// An empty feature slice was provided where at least one feature is required.
23    #[error("empty features are not allowed")]
24    EmptyFeatures,
25
26    /// A rolling window size of zero was provided.
27    #[error("invalid window size: must be greater than zero")]
28    InvalidWindowSize,
29
30    /// A learning rate that is not strictly positive was provided.
31    #[error("invalid learning rate: {0} (must be finite and > 0)")]
32    InvalidLearningRate(f64),
33
34    /// A generic numeric parameter was invalid.
35    #[error("invalid parameter `{name}`: {value}")]
36    InvalidParameter {
37        /// The name of the parameter.
38        name: &'static str,
39        /// The invalid value.
40        value: f64,
41    },
42
43    /// A NaN or Infinity value was encountered.
44    #[error("non-finite value for `{field}`: {value}")]
45    NonFiniteValue {
46        /// Which field or quantity the bad value belongs to.
47        field: &'static str,
48        /// The offending value.
49        value: f64,
50    },
51
52    /// A probability outside `[0, 1]` was provided.
53    #[error("invalid probability: {0} (must be in [0, 1])")]
54    InvalidProbability(f64),
55
56    /// Not enough data has been observed to compute the requested quantity.
57    #[error("insufficient data to compute the requested quantity")]
58    InsufficientData,
59
60    /// A serialized snapshot used an incompatible format version.
61    #[error("incompatible state version: expected {expected}, got {actual}")]
62    IncompatibleStateVersion {
63        /// The format version the loader expects.
64        expected: u32,
65        /// The format version found in the snapshot.
66        actual: u32,
67    },
68
69    /// Sparse features were not sorted by FeatureId.
70    #[error("sparse features must be sorted by FeatureId")]
71    UnsortedFeatureIds,
72
73    /// Duplicate FeatureId encountered in sparse features.
74    #[error("duplicate feature id: {0}")]
75    DuplicateFeatureId(u64),
76
77    /// FeatureHasher dimension is invalid (must be > 0).
78    #[error("invalid hash dimension: {0} (must be > 0)")]
79    InvalidHashDimension(usize),
80
81    /// An unknown category was encountered by a categorical encoder.
82    #[error("unknown category: {0}")]
83    UnknownCategory(String),
84
85    /// A missing value (NaN) was encountered where it is not allowed.
86    #[error("missing value (NaN) at index {index}")]
87    MissingValue {
88        /// The index of the missing value.
89        index: usize,
90    },
91
92    /// A window size or buffer capacity was invalid (must be > 0).
93    #[error("invalid capacity: {0} (must be greater than zero)")]
94    InvalidCapacity(usize),
95
96    /// A significance level or probability threshold was outside `(0, 1)`.
97    #[error("invalid significance level: {0} (must be in (0, 1))")]
98    InvalidSignificanceLevel(f64),
99
100    /// A bandit arm count was invalid (must be greater than zero).
101    #[error("invalid arm count: {0} (must be greater than zero)")]
102    InvalidArmCount(usize),
103
104    /// An epsilon value for epsilon-greedy was outside `[0, 1]`.
105    #[error("invalid epsilon: {0} (must be in [0, 1])")]
106    InvalidEpsilon(f64),
107
108    /// A bandit arm index was out of range.
109    #[error("invalid arm index: {actual} (must be < {expected})")]
110    InvalidArm {
111        /// The number of arms (upper bound).
112        expected: usize,
113        /// The offending arm index.
114        actual: usize,
115    },
116
117    /// A reward value was outside the valid range for the given bandit type.
118    #[error("invalid reward: {0} (must be finite and in the valid range)")]
119    InvalidReward(f64),
120
121    /// A feature count was invalid (must be greater than zero).
122    #[error("invalid feature count: {0} (must be greater than zero)")]
123    InvalidFeatureCount(usize),
124
125    /// A restored model violated one of its internal invariants.
126    #[error("invalid model state: {0}")]
127    InvalidState(String),
128}
129
130/// Helper to validate that a value is finite.
131pub(crate) fn ensure_finite(field: &'static str, value: f64) -> Result<(), RillError> {
132    if value.is_finite() {
133        Ok(())
134    } else {
135        Err(RillError::NonFiniteValue { field, value })
136    }
137}
138
139/// Add two values without allowing a finite-input overflow to poison state.
140pub(crate) fn checked_finite_add(
141    current: f64,
142    delta: f64,
143    field: &'static str,
144) -> Result<f64, RillError> {
145    let value = current + delta;
146    ensure_finite(field, value)?;
147    Ok(value)
148}
149
150/// Increment a long-running counter without allowing wraparound.
151pub(crate) fn checked_increment(value: u64, field: &'static str) -> Result<u64, RillError> {
152    value
153        .checked_add(1)
154        .ok_or_else(|| RillError::InvalidState(format!("{field} counter overflow")))
155}
156
157/// Helper to validate a feature slice's length and finiteness.
158pub(crate) fn validate_features(expected: usize, features: &[f64]) -> Result<(), RillError> {
159    if features.is_empty() {
160        return Err(RillError::EmptyFeatures);
161    }
162    if features.len() != expected {
163        return Err(RillError::DimensionMismatch {
164            expected,
165            actual: features.len(),
166        });
167    }
168    for (i, &v) in features.iter().enumerate() {
169        if !v.is_finite() {
170            return Err(RillError::NonFiniteValue {
171                field: "feature",
172                value: features[i],
173            });
174        }
175    }
176    Ok(())
177}
178
179/// Helper to validate a single finite scalar target/label.
180pub(crate) fn ensure_finite_target(value: f64) -> Result<(), RillError> {
181    ensure_finite("target", value)
182}
183
184#[cfg(test)]
185mod tests {
186    use super::*;
187
188    #[test]
189    fn dimension_mismatch_display() {
190        let e = RillError::DimensionMismatch {
191            expected: 3,
192            actual: 2,
193        };
194        assert!(format!("{e}").contains("expected 3"));
195        assert!(format!("{e}").contains("got 2"));
196    }
197
198    #[test]
199    fn ensure_finite_passes_for_normal_values() {
200        assert!(ensure_finite("x", 1.5).is_ok());
201        assert!(ensure_finite("x", -1e10).is_ok());
202        assert!(ensure_finite("x", 0.0).is_ok());
203    }
204
205    #[test]
206    fn ensure_finite_rejects_nan_and_infinity() {
207        assert!(ensure_finite("x", f64::NAN).is_err());
208        assert!(ensure_finite("x", f64::INFINITY).is_err());
209        assert!(ensure_finite("x", f64::NEG_INFINITY).is_err());
210    }
211
212    #[test]
213    fn validate_features_checks_dimension() {
214        assert!(validate_features(3, &[1.0, 2.0, 3.0]).is_ok());
215        assert!(validate_features(3, &[1.0, 2.0]).is_err());
216    }
217
218    #[test]
219    fn validate_features_rejects_empty() {
220        assert!(matches!(
221            validate_features(0, &[]),
222            Err(RillError::EmptyFeatures)
223        ));
224    }
225
226    #[test]
227    fn validate_features_rejects_non_finite() {
228        assert!(validate_features(2, &[1.0, f64::NAN]).is_err());
229        assert!(validate_features(2, &[f64::INFINITY, 2.0]).is_err());
230    }
231
232    #[test]
233    fn invalid_arm_count_display() {
234        let e = RillError::InvalidArmCount(0);
235        assert!(format!("{e}").contains("invalid arm count"));
236        assert!(format!("{e}").contains("0"));
237    }
238
239    #[test]
240    fn invalid_epsilon_display() {
241        let e = RillError::InvalidEpsilon(-0.5);
242        assert!(format!("{e}").contains("invalid epsilon"));
243        assert!(format!("{e}").contains("-0.5"));
244    }
245
246    #[test]
247    fn invalid_arm_display() {
248        let e = RillError::InvalidArm {
249            expected: 3,
250            actual: 5,
251        };
252        assert!(format!("{e}").contains("5"));
253        assert!(format!("{e}").contains("3"));
254    }
255
256    #[test]
257    fn invalid_reward_display() {
258        let e = RillError::InvalidReward(f64::NAN);
259        assert!(format!("{e}").contains("invalid reward"));
260    }
261
262    #[test]
263    fn invalid_feature_count_display() {
264        let e = RillError::InvalidFeatureCount(0);
265        assert!(format!("{e}").contains("invalid feature count"));
266        assert!(format!("{e}").contains("0"));
267    }
268}