Skip to main content

rill_ml/models/
naive_bayes.rs

1//! Online Naive Bayes classifiers.
2//!
3//! Three variants are provided:
4//! - [`GaussianNaiveBayes`]: for continuous features (assumes Gaussian distribution)
5//! - [`BernoulliNaiveBayes`]: for binary features (0/1)
6//! - [`MultinomialNaiveBayes`]: for count features (non-negative integers)
7//!
8//! All three implement [`OnlineBinaryClassifier`] for binary classification.
9//! Multi-class support may be added in a future version.
10
11use crate::error::{
12    RillError, checked_finite_add, checked_increment, ensure_finite, validate_features,
13};
14use crate::loss::log_loss::sigmoid;
15#[cfg(feature = "serde")]
16use crate::persistence::ValidateState;
17use crate::traits::OnlineBinaryClassifier;
18
19/// Configuration for Naive Bayes classifiers.
20#[derive(Debug, Clone)]
21#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
22#[non_exhaustive]
23pub struct NaiveBayesConfig {
24    /// Laplace smoothing parameter (alpha). Must be `> 0`.
25    ///
26    /// Default: `1.0` (standard Laplace smoothing).
27    pub alpha: f64,
28}
29
30impl Default for NaiveBayesConfig {
31    fn default() -> Self {
32        Self { alpha: 1.0 }
33    }
34}
35
36fn validate_config(config: &NaiveBayesConfig) -> Result<(), RillError> {
37    ensure_finite("alpha", config.alpha)?;
38    if config.alpha <= 0.0 {
39        return Err(RillError::InvalidParameter {
40            name: "alpha",
41            value: config.alpha,
42        });
43    }
44    Ok(())
45}
46
47/// Compare two non-negative finite floats for approximate equality using a
48/// combination of absolute and relative tolerance, returning `true` when
49/// they are "close enough" to be plausibly the same quantity.
50///
51/// This is used by [`MultinomialNaiveBayes`] to validate that
52/// `sum(feature_sums)` matches the cached `total` despite floating-point
53/// accumulation order. Strict equality would reject legitimate round-off;
54/// an unbounded tolerance would accept clearly inconsistent state.
55///
56/// - `atol` is the absolute floor: differences below it are always accepted.
57/// - `rtol` scales with the magnitude of the larger value.
58///
59/// Either `a` or `b` being non-finite makes the comparison return `false`;
60/// callers should still run [`ensure_finite`] separately to produce a
61/// precise diagnostic.
62fn approx_equal_non_negative(a: f64, b: f64, atol: f64, rtol: f64) -> bool {
63    if !a.is_finite() || !b.is_finite() {
64        return false;
65    }
66    let abs_diff = (a - b).abs();
67    if abs_diff <= atol {
68        return true;
69    }
70    let larger = a.abs().max(b.abs());
71    abs_diff <= rtol * larger
72}
73
74/// Tolerances used to verify Multinomial NB cache consistency.
75///
76/// `total_false` / `total_true` cache the sum of the corresponding
77/// `feature_sums_*` vector. Floating-point accumulation order can introduce
78/// small discrepancies, so a strict equality check would reject legitimate
79/// state. These tolerances are tight enough to reject grossly inconsistent
80/// (e.g. off-by-orders-of-magnitude) malicious state while tolerating
81/// normal round-off.
82const MULTINOMIAL_TOTAL_ATOL: f64 = 1e-9;
83const MULTINOMIAL_TOTAL_RTOL: f64 = 1e-6;
84
85/// Validate that a log-domain value is admissible.
86///
87/// `-Infinity` is allowed (it represents probability 0 for an impossible
88/// event, e.g. a class with zero training samples). `NaN` and
89/// `+Infinity` are rejected because they indicate an arithmetic breakdown
90/// (such as `-Infinity - (-Infinity)`) that would propagate into the final
91/// probability as `NaN`.
92fn ensure_log_domain(field: &'static str, value: f64) -> Result<(), RillError> {
93    if value.is_nan() || (value.is_infinite() && value > 0.0) {
94        return Err(RillError::NonFiniteValue { field, value });
95    }
96    Ok(())
97}
98
99/// Add two log-domain values, rejecting `NaN` and `+Infinity` results.
100///
101/// `-Infinity` is preserved (probability 0 + anything = probability 0),
102/// matching [`ensure_log_domain`].
103fn checked_log_add(current: f64, delta: f64, field: &'static str) -> Result<f64, RillError> {
104    let value = current + delta;
105    ensure_log_domain(field, value)?;
106    Ok(value)
107}
108
109/// Validate that features are finite and non-negative (for Bernoulli/Multinomial).
110fn validate_non_negative(feature_count: usize, features: &[f64]) -> Result<(), RillError> {
111    validate_features(feature_count, features)?;
112    for &x in features {
113        if x < 0.0 {
114            return Err(RillError::InvalidParameter {
115                name: "feature",
116                value: x,
117            });
118        }
119    }
120    Ok(())
121}
122
123// ============================================================================
124// Gaussian Naive Bayes
125// ============================================================================
126
127/// Per-class Gaussian statistics (Welford algorithm per feature).
128#[derive(Debug, Clone)]
129#[cfg_attr(feature = "serde", derive(serde::Serialize))]
130struct GaussianClassStats {
131    counts: Vec<u64>,
132    means: Vec<f64>,
133    m2s: Vec<f64>,
134    class_count: u64,
135}
136
137impl GaussianClassStats {
138    fn new(feature_count: usize) -> Self {
139        Self {
140            counts: vec![0; feature_count],
141            means: vec![0.0; feature_count],
142            m2s: vec![0.0; feature_count],
143            class_count: 0,
144        }
145    }
146
147    fn variance(&self, idx: usize) -> f64 {
148        if self.counts[idx] < 2 {
149            0.0
150        } else {
151            self.m2s[idx] / self.counts[idx] as f64
152        }
153    }
154
155    fn reset(&mut self) {
156        self.counts.fill(0);
157        self.means.fill(0.0);
158        self.m2s.fill(0.0);
159        self.class_count = 0;
160    }
161
162    /// Validate internal consistency of this class's statistics.
163    ///
164    /// Does not check that the per-feature vector lengths match the parent
165    /// model's `feature_count`; that is the parent's responsibility.
166    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
167    fn validate_invariants(&self) -> Result<(), RillError> {
168        let n = self.counts.len();
169        if self.means.len() != n || self.m2s.len() != n {
170            return Err(RillError::InvalidState(
171                "gaussian class stats: counts/means/m2s length mismatch".to_owned(),
172            ));
173        }
174        // Gaussian NB updates every feature on each learn, so every
175        // per-feature count must equal the class_count.
176        for &c in &self.counts {
177            if c != self.class_count {
178                return Err(RillError::InvalidState(format!(
179                    "gaussian class stats: feature count {c} != class_count {}",
180                    self.class_count
181                )));
182            }
183        }
184        for &m in &self.means {
185            ensure_finite("gaussian mean", m)?;
186        }
187        for &m2 in &self.m2s {
188            ensure_finite("gaussian m2", m2)?;
189            if m2 < 0.0 {
190                return Err(RillError::InvalidState(format!(
191                    "gaussian m2 must be non-negative, got {m2}"
192                )));
193            }
194        }
195        // count == 0: no samples seen → mean and m2 must be exactly 0.
196        if self.class_count == 0 {
197            for &m in &self.means {
198                if m != 0.0 {
199                    return Err(RillError::InvalidState(format!(
200                        "gaussian class_count=0 but mean={m} (must be 0)"
201                    )));
202                }
203            }
204            for &m2 in &self.m2s {
205                if m2 != 0.0 {
206                    return Err(RillError::InvalidState(format!(
207                        "gaussian class_count=0 but m2={m2} (must be 0)"
208                    )));
209                }
210            }
211        }
212        // count == 1: Welford M2 is exactly 0 after a single sample.
213        if self.class_count == 1 {
214            for &m2 in &self.m2s {
215                if m2 != 0.0 {
216                    return Err(RillError::InvalidState(format!(
217                        "gaussian class_count=1 but m2={m2} (must be 0)"
218                    )));
219                }
220            }
221        }
222        Ok(())
223    }
224}
225
226#[cfg(feature = "serde")]
227impl<'de> serde::Deserialize<'de> for GaussianClassStats {
228    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
229    where
230        D: serde::Deserializer<'de>,
231    {
232        #[derive(serde::Deserialize)]
233        struct State {
234            counts: Vec<u64>,
235            means: Vec<f64>,
236            m2s: Vec<f64>,
237            class_count: u64,
238        }
239        let s = State::deserialize(deserializer)?;
240        let stats = GaussianClassStats {
241            counts: s.counts,
242            means: s.means,
243            m2s: s.m2s,
244            class_count: s.class_count,
245        };
246        stats
247            .validate_invariants()
248            .map_err(serde::de::Error::custom)?;
249        Ok(stats)
250    }
251}
252
253/// Online Gaussian Naive Bayes classifier.
254///
255/// Assumes features are conditionally independent given the class,
256/// and each feature follows a Gaussian distribution per class.
257/// Uses Welford's algorithm for numerically stable variance updates.
258///
259/// # Examples
260///
261/// ```
262/// use rill_ml::models::GaussianNaiveBayes;
263/// use rill_ml::OnlineBinaryClassifier;
264///
265/// let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
266/// model.learn(&[1.0, 2.0], true).unwrap();
267/// model.learn(&[-1.0, -2.0], false).unwrap();
268/// let proba = model.predict_proba(&[0.5, 1.0]).unwrap();
269/// assert!(proba > 0.0 && proba < 1.0);
270/// ```
271#[derive(Debug, Clone)]
272#[cfg_attr(feature = "serde", derive(serde::Serialize))]
273pub struct GaussianNaiveBayes {
274    feature_count: usize,
275    config: NaiveBayesConfig,
276    class_false: GaussianClassStats,
277    class_true: GaussianClassStats,
278    samples_seen: u64,
279}
280
281impl GaussianNaiveBayes {
282    /// Create a new Gaussian Naive Bayes classifier.
283    ///
284    /// `feature_count` must be greater than zero.
285    pub fn new(feature_count: usize, config: NaiveBayesConfig) -> Result<Self, RillError> {
286        validate_config(&config)?;
287        if feature_count == 0 {
288            return Err(RillError::EmptyFeatures);
289        }
290        Ok(Self {
291            feature_count,
292            config,
293            class_false: GaussianClassStats::new(feature_count),
294            class_true: GaussianClassStats::new(feature_count),
295            samples_seen: 0,
296        })
297    }
298
299    /// The Laplace smoothing parameter.
300    pub const fn alpha(&self) -> f64 {
301        self.config.alpha
302    }
303
304    /// Gaussian log probability density function.
305    ///
306    /// Returns `0.0` when `variance <= 0.0` (constant or unseen feature),
307    /// otherwise the standard Gaussian log-density. The result is always
308    /// in `(-∞, 0]` for admissible inputs; callers still validate it via
309    /// [`ensure_log_domain`] to catch `NaN` from malicious state.
310    fn gaussian_log_pdf(x: f64, mean: f64, variance: f64) -> f64 {
311        if variance <= 0.0 {
312            return 0.0;
313        }
314        let sigma = variance.sqrt();
315        -0.5 * ((x - mean) / sigma).powi(2) - sigma.ln() - 0.5 * (2.0 * std::f64::consts::PI).ln()
316    }
317
318    /// Validate all invariants required for a deserialized model to be safe.
319    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
320    pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
321        if self.feature_count == 0 {
322            return Err(RillError::EmptyFeatures);
323        }
324        validate_config(&self.config)?;
325        // Per-feature vector lengths must match feature_count.
326        if self.class_false.counts.len() != self.feature_count
327            || self.class_false.means.len() != self.feature_count
328            || self.class_false.m2s.len() != self.feature_count
329        {
330            return Err(RillError::InvalidState(
331                "gaussian class_false vector length != feature_count".to_owned(),
332            ));
333        }
334        if self.class_true.counts.len() != self.feature_count
335            || self.class_true.means.len() != self.feature_count
336            || self.class_true.m2s.len() != self.feature_count
337        {
338            return Err(RillError::InvalidState(
339                "gaussian class_true vector length != feature_count".to_owned(),
340            ));
341        }
342        self.class_false.validate_invariants()?;
343        self.class_true.validate_invariants()?;
344        let total_class = self
345            .class_false
346            .class_count
347            .checked_add(self.class_true.class_count)
348            .ok_or_else(|| {
349                RillError::InvalidState(format!(
350                    "gaussian class_false({}) + class_true({}) overflow",
351                    self.class_false.class_count, self.class_true.class_count
352                ))
353            })?;
354        if total_class != self.samples_seen {
355            return Err(RillError::InvalidState(format!(
356                "gaussian samples_seen={} != class_false({}) + class_true({})",
357                self.samples_seen, self.class_false.class_count, self.class_true.class_count
358            )));
359        }
360        Ok(())
361    }
362}
363
364#[cfg(feature = "serde")]
365impl<'de> serde::Deserialize<'de> for GaussianNaiveBayes {
366    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
367    where
368        D: serde::Deserializer<'de>,
369    {
370        #[derive(serde::Deserialize)]
371        struct State {
372            feature_count: usize,
373            config: NaiveBayesConfig,
374            class_false: GaussianClassStats,
375            class_true: GaussianClassStats,
376            samples_seen: u64,
377        }
378        let s = State::deserialize(deserializer)?;
379        let model = GaussianNaiveBayes {
380            feature_count: s.feature_count,
381            config: s.config,
382            class_false: s.class_false,
383            class_true: s.class_true,
384            samples_seen: s.samples_seen,
385        };
386        model
387            .validate_invariants()
388            .map_err(serde::de::Error::custom)?;
389        Ok(model)
390    }
391}
392
393impl OnlineBinaryClassifier for GaussianNaiveBayes {
394    fn feature_count(&self) -> usize {
395        self.feature_count
396    }
397
398    fn samples_seen(&self) -> u64 {
399        self.samples_seen
400    }
401
402    fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
403        validate_features(self.feature_count, features)?;
404
405        if self.samples_seen == 0 {
406            return Ok(0.5);
407        }
408
409        let count_true = self.class_true.class_count as f64;
410        let count_false = self.class_false.class_count as f64;
411        let total = count_true + count_false;
412
413        let log_prior_true = (count_true / total).ln();
414        ensure_log_domain("nb_gaussian_log_prior_true", log_prior_true)?;
415        let log_prior_false = (count_false / total).ln();
416        ensure_log_domain("nb_gaussian_log_prior_false", log_prior_false)?;
417
418        let mut log_likelihood_true = 0.0;
419        let mut log_likelihood_false = 0.0;
420
421        for (i, &x) in features.iter().enumerate() {
422            let ll_true =
423                Self::gaussian_log_pdf(x, self.class_true.means[i], self.class_true.variance(i));
424            ensure_log_domain("nb_gaussian_log_likelihood_true", ll_true)?;
425            log_likelihood_true = checked_log_add(
426                log_likelihood_true,
427                ll_true,
428                "nb_gaussian_log_likelihood_true",
429            )?;
430            let ll_false =
431                Self::gaussian_log_pdf(x, self.class_false.means[i], self.class_false.variance(i));
432            ensure_log_domain("nb_gaussian_log_likelihood_false", ll_false)?;
433            log_likelihood_false = checked_log_add(
434                log_likelihood_false,
435                ll_false,
436                "nb_gaussian_log_likelihood_false",
437            )?;
438        }
439
440        let log_p_true = checked_log_add(
441            log_prior_true,
442            log_likelihood_true,
443            "nb_gaussian_log_p_true",
444        )?;
445        let log_p_false = checked_log_add(
446            log_prior_false,
447            log_likelihood_false,
448            "nb_gaussian_log_p_false",
449        )?;
450
451        let log_odds = log_p_true - log_p_false;
452        // `log_odds` may be `±Infinity` (one class dominates) — sigmoid maps
453        // those to 0.0 or 1.0, which is valid per the closed `[0, 1]` trait
454        // contract. Only `NaN` (both sides `-Infinity`) is unrecoverable.
455        if log_odds.is_nan() {
456            return Err(RillError::NonFiniteValue {
457                field: "nb_gaussian_log_odds",
458                value: log_odds,
459            });
460        }
461
462        let probability = sigmoid(log_odds);
463        ensure_finite("nb_gaussian_probability", probability)?;
464        if !(0.0..=1.0).contains(&probability) {
465            return Err(RillError::InvalidProbability(probability));
466        }
467        Ok(probability)
468    }
469
470    fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
471        validate_features(self.feature_count, features)?;
472
473        // Phase 1: compute the next per-feature state without mutating self.
474        // If any feature overflows or produces a non-finite value, the
475        // caller's `Err` leaves the model untouched.
476        let stats = if target {
477            &self.class_true
478        } else {
479            &self.class_false
480        };
481        let mut next_states: Vec<(usize, u64, f64, f64)> = Vec::with_capacity(features.len());
482        for (i, &x) in features.iter().enumerate() {
483            let n = checked_increment(stats.counts[i], "feature count")?;
484            let delta = x - stats.means[i];
485            ensure_finite("mean delta", delta)?;
486            let new_mean = checked_finite_add(stats.means[i], delta / n as f64, "mean")?;
487            let delta2 = x - new_mean;
488            ensure_finite("mean delta2", delta2)?;
489            let new_m2 = checked_finite_add(stats.m2s[i], delta * delta2, "m2")?;
490            next_states.push((i, n, new_mean, new_m2));
491        }
492        let new_class_count = checked_increment(stats.class_count, "class_count")?;
493        let new_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
494
495        // Phase 2: commit atomically.
496        let stats = if target {
497            &mut self.class_true
498        } else {
499            &mut self.class_false
500        };
501        for (i, n, new_mean, new_m2) in next_states {
502            stats.counts[i] = n;
503            stats.means[i] = new_mean;
504            stats.m2s[i] = new_m2;
505        }
506        stats.class_count = new_class_count;
507        self.samples_seen = new_samples_seen;
508        Ok(())
509    }
510
511    fn reset(&mut self) {
512        self.class_false.reset();
513        self.class_true.reset();
514        self.samples_seen = 0;
515    }
516}
517
518// ============================================================================
519// Bernoulli Naive Bayes
520// ============================================================================
521
522/// Online Bernoulli Naive Bayes classifier.
523///
524/// Designed for binary features (0 or 1). Uses Laplace smoothing
525/// for probability estimation.
526///
527/// # Examples
528///
529/// ```
530/// use rill_ml::models::BernoulliNaiveBayes;
531/// use rill_ml::OnlineBinaryClassifier;
532///
533/// let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
534/// model.learn(&[1.0, 0.0, 1.0], true).unwrap();
535/// model.learn(&[0.0, 1.0, 0.0], false).unwrap();
536/// let proba = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
537/// assert!(proba > 0.5);
538/// ```
539#[derive(Debug, Clone)]
540#[cfg_attr(feature = "serde", derive(serde::Serialize))]
541pub struct BernoulliNaiveBayes {
542    feature_count: usize,
543    config: NaiveBayesConfig,
544    feature_true_counts_false: Vec<u64>,
545    feature_true_counts_true: Vec<u64>,
546    class_false_count: u64,
547    class_true_count: u64,
548    samples_seen: u64,
549}
550
551impl BernoulliNaiveBayes {
552    /// Create a new Bernoulli Naive Bayes classifier.
553    ///
554    /// `feature_count` must be greater than zero.
555    pub fn new(feature_count: usize, config: NaiveBayesConfig) -> Result<Self, RillError> {
556        validate_config(&config)?;
557        if feature_count == 0 {
558            return Err(RillError::EmptyFeatures);
559        }
560        Ok(Self {
561            feature_count,
562            config,
563            feature_true_counts_false: vec![0; feature_count],
564            feature_true_counts_true: vec![0; feature_count],
565            class_false_count: 0,
566            class_true_count: 0,
567            samples_seen: 0,
568        })
569    }
570
571    /// Compute log P(x_i | class) for a single Bernoulli feature.
572    fn log_bernoulli(x: f64, p: f64) -> f64 {
573        x * p.ln() + (1.0 - x) * (1.0 - p).ln()
574    }
575
576    /// Validate all invariants required for a deserialized model to be safe.
577    ///
578    /// Bernoulli NB tracks, per class, how many training samples had each
579    /// feature "present" (`x > 0.5`). Each per-feature count must be at most
580    /// the corresponding class count, and `samples_seen` must equal the sum
581    /// of the two class counts. Lengths must match `feature_count`.
582    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
583    pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
584        if self.feature_count == 0 {
585            return Err(RillError::EmptyFeatures);
586        }
587        validate_config(&self.config)?;
588        if self.feature_true_counts_false.len() != self.feature_count
589            || self.feature_true_counts_true.len() != self.feature_count
590        {
591            return Err(RillError::InvalidState(
592                "bernoulli feature count vector length != feature_count".to_owned(),
593            ));
594        }
595        if let Some(total) = self.class_false_count.checked_add(self.class_true_count) {
596            if total != self.samples_seen {
597                return Err(RillError::InvalidState(format!(
598                    "bernoulli samples_seen={} != class_false({}) + class_true({})",
599                    self.samples_seen, self.class_false_count, self.class_true_count
600                )));
601            }
602        } else {
603            return Err(RillError::InvalidState(format!(
604                "bernoulli class_false({}) + class_true({}) overflow",
605                self.class_false_count, self.class_true_count
606            )));
607        }
608        for &c in &self.feature_true_counts_false {
609            if c > self.class_false_count {
610                return Err(RillError::InvalidState(format!(
611                    "bernoulli feature_true_counts_false entry {c} > class_false_count {}",
612                    self.class_false_count
613                )));
614            }
615        }
616        for &c in &self.feature_true_counts_true {
617            if c > self.class_true_count {
618                return Err(RillError::InvalidState(format!(
619                    "bernoulli feature_true_counts_true entry {c} > class_true_count {}",
620                    self.class_true_count
621                )));
622            }
623        }
624        Ok(())
625    }
626}
627
628#[cfg(feature = "serde")]
629impl<'de> serde::Deserialize<'de> for BernoulliNaiveBayes {
630    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
631    where
632        D: serde::Deserializer<'de>,
633    {
634        #[derive(serde::Deserialize)]
635        struct State {
636            feature_count: usize,
637            config: NaiveBayesConfig,
638            feature_true_counts_false: Vec<u64>,
639            feature_true_counts_true: Vec<u64>,
640            class_false_count: u64,
641            class_true_count: u64,
642            samples_seen: u64,
643        }
644        let s = State::deserialize(deserializer)?;
645        let model = BernoulliNaiveBayes {
646            feature_count: s.feature_count,
647            config: s.config,
648            feature_true_counts_false: s.feature_true_counts_false,
649            feature_true_counts_true: s.feature_true_counts_true,
650            class_false_count: s.class_false_count,
651            class_true_count: s.class_true_count,
652            samples_seen: s.samples_seen,
653        };
654        model
655            .validate_invariants()
656            .map_err(serde::de::Error::custom)?;
657        Ok(model)
658    }
659}
660
661impl OnlineBinaryClassifier for BernoulliNaiveBayes {
662    fn feature_count(&self) -> usize {
663        self.feature_count
664    }
665
666    fn samples_seen(&self) -> u64 {
667        self.samples_seen
668    }
669
670    fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
671        validate_non_negative(self.feature_count, features)?;
672
673        if self.samples_seen == 0 {
674            return Ok(0.5);
675        }
676
677        let count_true = self.class_true_count as f64;
678        let count_false = self.class_false_count as f64;
679        let total = count_true + count_false;
680
681        let log_prior_true = (count_true / total).ln();
682        ensure_log_domain("nb_bernoulli_log_prior_true", log_prior_true)?;
683        let log_prior_false = (count_false / total).ln();
684        ensure_log_domain("nb_bernoulli_log_prior_false", log_prior_false)?;
685
686        let mut log_likelihood_true = 0.0;
687        let mut log_likelihood_false = 0.0;
688
689        for (i, &x) in features.iter().enumerate() {
690            let p_true = (self.feature_true_counts_true[i] as f64 + self.config.alpha)
691                / (count_true + 2.0 * self.config.alpha);
692            let p_false = (self.feature_true_counts_false[i] as f64 + self.config.alpha)
693                / (count_false + 2.0 * self.config.alpha);
694            let ll_true = Self::log_bernoulli(x, p_true);
695            ensure_log_domain("nb_bernoulli_log_likelihood_true", ll_true)?;
696            log_likelihood_true = checked_log_add(
697                log_likelihood_true,
698                ll_true,
699                "nb_bernoulli_log_likelihood_true",
700            )?;
701            let ll_false = Self::log_bernoulli(x, p_false);
702            ensure_log_domain("nb_bernoulli_log_likelihood_false", ll_false)?;
703            log_likelihood_false = checked_log_add(
704                log_likelihood_false,
705                ll_false,
706                "nb_bernoulli_log_likelihood_false",
707            )?;
708        }
709
710        let log_p_true = checked_log_add(
711            log_prior_true,
712            log_likelihood_true,
713            "nb_bernoulli_log_p_true",
714        )?;
715        let log_p_false = checked_log_add(
716            log_prior_false,
717            log_likelihood_false,
718            "nb_bernoulli_log_p_false",
719        )?;
720
721        let log_odds = log_p_true - log_p_false;
722        // See `GaussianNaiveBayes::predict_proba`: `±Infinity` is valid
723        // (sigmoid maps to 0/1), only `NaN` is unrecoverable.
724        if log_odds.is_nan() {
725            return Err(RillError::NonFiniteValue {
726                field: "nb_bernoulli_log_odds",
727                value: log_odds,
728            });
729        }
730
731        let probability = sigmoid(log_odds);
732        ensure_finite("nb_bernoulli_probability", probability)?;
733        if !(0.0..=1.0).contains(&probability) {
734            return Err(RillError::InvalidProbability(probability));
735        }
736        Ok(probability)
737    }
738
739    fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
740        validate_non_negative(self.feature_count, features)?;
741
742        // Phase 1: compute the next per-feature counts without mutating self.
743        // If any counter overflows, the caller's `Err` leaves the model
744        // untouched.
745        let mut next_feature_counts = if target {
746            self.feature_true_counts_true.clone()
747        } else {
748            self.feature_true_counts_false.clone()
749        };
750        for (i, &x) in features.iter().enumerate() {
751            if x > 0.5 {
752                next_feature_counts[i] =
753                    checked_increment(next_feature_counts[i], "feature_true_count")?;
754            }
755        }
756        let next_class_count = if target {
757            checked_increment(self.class_true_count, "class_true_count")?
758        } else {
759            checked_increment(self.class_false_count, "class_false_count")?
760        };
761        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
762
763        // Phase 2: commit atomically.
764        if target {
765            self.feature_true_counts_true = next_feature_counts;
766            self.class_true_count = next_class_count;
767        } else {
768            self.feature_true_counts_false = next_feature_counts;
769            self.class_false_count = next_class_count;
770        }
771        self.samples_seen = next_samples_seen;
772        Ok(())
773    }
774
775    fn reset(&mut self) {
776        self.feature_true_counts_false.fill(0);
777        self.feature_true_counts_true.fill(0);
778        self.class_false_count = 0;
779        self.class_true_count = 0;
780        self.samples_seen = 0;
781    }
782}
783
784// ============================================================================
785// Multinomial Naive Bayes
786// ============================================================================
787
788/// Online Multinomial Naive Bayes classifier.
789///
790/// Designed for count features (non-negative values). Uses Laplace
791/// smoothing. Commonly used for text classification with word counts.
792///
793/// # Examples
794///
795/// ```
796/// use rill_ml::models::MultinomialNaiveBayes;
797/// use rill_ml::OnlineBinaryClassifier;
798///
799/// let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
800/// model.learn(&[2.0, 1.0, 0.0], true).unwrap();
801/// model.learn(&[0.0, 1.0, 3.0], false).unwrap();
802/// let proba = model.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
803/// assert!(proba > 0.0 && proba < 1.0);
804/// ```
805#[derive(Debug, Clone)]
806#[cfg_attr(feature = "serde", derive(serde::Serialize))]
807pub struct MultinomialNaiveBayes {
808    feature_count: usize,
809    config: NaiveBayesConfig,
810    feature_sums_false: Vec<f64>,
811    feature_sums_true: Vec<f64>,
812    total_false: f64,
813    total_true: f64,
814    class_false_count: u64,
815    class_true_count: u64,
816    samples_seen: u64,
817}
818
819impl MultinomialNaiveBayes {
820    /// Create a new Multinomial Naive Bayes classifier.
821    ///
822    /// `feature_count` must be greater than zero.
823    pub fn new(feature_count: usize, config: NaiveBayesConfig) -> Result<Self, RillError> {
824        validate_config(&config)?;
825        if feature_count == 0 {
826            return Err(RillError::EmptyFeatures);
827        }
828        Ok(Self {
829            feature_count,
830            config,
831            feature_sums_false: vec![0.0; feature_count],
832            feature_sums_true: vec![0.0; feature_count],
833            total_false: 0.0,
834            total_true: 0.0,
835            class_false_count: 0,
836            class_true_count: 0,
837            samples_seen: 0,
838        })
839    }
840
841    /// Validate all invariants required for a deserialized model to be safe.
842    ///
843    /// Multinomial NB caches `total_false` / `total_true` as the sum of the
844    /// corresponding `feature_sums_*` vector. Floating-point accumulation
845    /// order can introduce small discrepancies, so the cache is checked
846    /// against the vector sum with a tight combined absolute / relative
847    /// tolerance (see [`MULTINOMIAL_TOTAL_ATOL`] / [`MULTINOMIAL_TOTAL_RTOL`]).
848    /// Grossly inconsistent (e.g. off-by-orders-of-magnitude) state is still
849    /// rejected.
850    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
851    pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
852        if self.feature_count == 0 {
853            return Err(RillError::EmptyFeatures);
854        }
855        validate_config(&self.config)?;
856        if self.feature_sums_false.len() != self.feature_count
857            || self.feature_sums_true.len() != self.feature_count
858        {
859            return Err(RillError::InvalidState(
860                "multinomial feature sum vector length != feature_count".to_owned(),
861            ));
862        }
863        for &s in &self.feature_sums_false {
864            ensure_finite("multinomial feature_sums_false", s)?;
865            if s < 0.0 {
866                return Err(RillError::InvalidState(format!(
867                    "multinomial feature_sums_false entry {s} is negative"
868                )));
869            }
870        }
871        for &s in &self.feature_sums_true {
872            ensure_finite("multinomial feature_sums_true", s)?;
873            if s < 0.0 {
874                return Err(RillError::InvalidState(format!(
875                    "multinomial feature_sums_true entry {s} is negative"
876                )));
877            }
878        }
879        ensure_finite("multinomial total_false", self.total_false)?;
880        if self.total_false < 0.0 {
881            return Err(RillError::InvalidState(format!(
882                "multinomial total_false is negative: {}",
883                self.total_false
884            )));
885        }
886        ensure_finite("multinomial total_true", self.total_true)?;
887        if self.total_true < 0.0 {
888            return Err(RillError::InvalidState(format!(
889                "multinomial total_true is negative: {}",
890                self.total_true
891            )));
892        }
893        if let Some(total) = self.class_false_count.checked_add(self.class_true_count) {
894            if total != self.samples_seen {
895                return Err(RillError::InvalidState(format!(
896                    "multinomial samples_seen={} != class_false({}) + class_true({})",
897                    self.samples_seen, self.class_false_count, self.class_true_count
898                )));
899            }
900        } else {
901            return Err(RillError::InvalidState(format!(
902                "multinomial class_false({}) + class_true({}) overflow",
903                self.class_false_count, self.class_true_count
904            )));
905        }
906        // Cache consistency: sum(feature_sums_*) must match total_* within
907        // tolerance. Use Vec::iter().sum::<f64>() as an independent
908        // accumulation order from the learn() hot path.
909        let summed_false: f64 = self.feature_sums_false.iter().sum();
910        let summed_true: f64 = self.feature_sums_true.iter().sum();
911        if !approx_equal_non_negative(
912            summed_false,
913            self.total_false,
914            MULTINOMIAL_TOTAL_ATOL,
915            MULTINOMIAL_TOTAL_RTOL,
916        ) {
917            return Err(RillError::InvalidState(format!(
918                "multinomial sum(feature_sums_false)={summed_false} mismatches total_false={} beyond tolerance",
919                self.total_false
920            )));
921        }
922        if !approx_equal_non_negative(
923            summed_true,
924            self.total_true,
925            MULTINOMIAL_TOTAL_ATOL,
926            MULTINOMIAL_TOTAL_RTOL,
927        ) {
928            return Err(RillError::InvalidState(format!(
929                "multinomial sum(feature_sums_true)={summed_true} mismatches total_true={} beyond tolerance",
930                self.total_true
931            )));
932        }
933        Ok(())
934    }
935}
936
937#[cfg(feature = "serde")]
938impl<'de> serde::Deserialize<'de> for MultinomialNaiveBayes {
939    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
940    where
941        D: serde::Deserializer<'de>,
942    {
943        #[derive(serde::Deserialize)]
944        struct State {
945            feature_count: usize,
946            config: NaiveBayesConfig,
947            feature_sums_false: Vec<f64>,
948            feature_sums_true: Vec<f64>,
949            total_false: f64,
950            total_true: f64,
951            class_false_count: u64,
952            class_true_count: u64,
953            samples_seen: u64,
954        }
955        let s = State::deserialize(deserializer)?;
956        let model = MultinomialNaiveBayes {
957            feature_count: s.feature_count,
958            config: s.config,
959            feature_sums_false: s.feature_sums_false,
960            feature_sums_true: s.feature_sums_true,
961            total_false: s.total_false,
962            total_true: s.total_true,
963            class_false_count: s.class_false_count,
964            class_true_count: s.class_true_count,
965            samples_seen: s.samples_seen,
966        };
967        model
968            .validate_invariants()
969            .map_err(serde::de::Error::custom)?;
970        Ok(model)
971    }
972}
973
974#[cfg(feature = "serde")]
975impl ValidateState for GaussianNaiveBayes {
976    fn validate_state(&self) -> Result<(), RillError> {
977        GaussianNaiveBayes::validate_invariants(self)
978    }
979}
980
981#[cfg(feature = "serde")]
982impl ValidateState for BernoulliNaiveBayes {
983    fn validate_state(&self) -> Result<(), RillError> {
984        BernoulliNaiveBayes::validate_invariants(self)
985    }
986}
987
988#[cfg(feature = "serde")]
989impl ValidateState for MultinomialNaiveBayes {
990    fn validate_state(&self) -> Result<(), RillError> {
991        MultinomialNaiveBayes::validate_invariants(self)
992    }
993}
994
995impl OnlineBinaryClassifier for MultinomialNaiveBayes {
996    fn feature_count(&self) -> usize {
997        self.feature_count
998    }
999
1000    fn samples_seen(&self) -> u64 {
1001        self.samples_seen
1002    }
1003
1004    fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
1005        validate_non_negative(self.feature_count, features)?;
1006
1007        if self.samples_seen == 0 {
1008            return Ok(0.5);
1009        }
1010
1011        let count_true = self.class_true_count as f64;
1012        let count_false = self.class_false_count as f64;
1013        let total = count_true + count_false;
1014
1015        let log_prior_true = (count_true / total).ln();
1016        ensure_log_domain("nb_multinomial_log_prior_true", log_prior_true)?;
1017        let log_prior_false = (count_false / total).ln();
1018        ensure_log_domain("nb_multinomial_log_prior_false", log_prior_false)?;
1019
1020        let denom_true = self.total_true + self.config.alpha * self.feature_count as f64;
1021        let denom_false = self.total_false + self.config.alpha * self.feature_count as f64;
1022
1023        let mut log_likelihood_true = 0.0;
1024        let mut log_likelihood_false = 0.0;
1025
1026        for (i, &x) in features.iter().enumerate() {
1027            let p_true = (self.feature_sums_true[i] + self.config.alpha) / denom_true;
1028            let p_false = (self.feature_sums_false[i] + self.config.alpha) / denom_false;
1029            let ll_true = x * p_true.ln();
1030            ensure_log_domain("nb_multinomial_log_likelihood_true", ll_true)?;
1031            log_likelihood_true = checked_log_add(
1032                log_likelihood_true,
1033                ll_true,
1034                "nb_multinomial_log_likelihood_true",
1035            )?;
1036            let ll_false = x * p_false.ln();
1037            ensure_log_domain("nb_multinomial_log_likelihood_false", ll_false)?;
1038            log_likelihood_false = checked_log_add(
1039                log_likelihood_false,
1040                ll_false,
1041                "nb_multinomial_log_likelihood_false",
1042            )?;
1043        }
1044
1045        let log_p_true = checked_log_add(
1046            log_prior_true,
1047            log_likelihood_true,
1048            "nb_multinomial_log_p_true",
1049        )?;
1050        let log_p_false = checked_log_add(
1051            log_prior_false,
1052            log_likelihood_false,
1053            "nb_multinomial_log_p_false",
1054        )?;
1055
1056        let log_odds = log_p_true - log_p_false;
1057        // See `GaussianNaiveBayes::predict_proba`: `±Infinity` is valid
1058        // (sigmoid maps to 0/1), only `NaN` is unrecoverable.
1059        if log_odds.is_nan() {
1060            return Err(RillError::NonFiniteValue {
1061                field: "nb_multinomial_log_odds",
1062                value: log_odds,
1063            });
1064        }
1065
1066        let probability = sigmoid(log_odds);
1067        ensure_finite("nb_multinomial_probability", probability)?;
1068        if !(0.0..=1.0).contains(&probability) {
1069            return Err(RillError::InvalidProbability(probability));
1070        }
1071        Ok(probability)
1072    }
1073
1074    fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
1075        validate_non_negative(self.feature_count, features)?;
1076
1077        // Phase 1: compute the next per-feature sums and total without
1078        // mutating self. If any sum overflows, the caller's `Err` leaves the
1079        // model untouched.
1080        let mut next_feature_sums = if target {
1081            self.feature_sums_true.clone()
1082        } else {
1083            self.feature_sums_false.clone()
1084        };
1085        let mut next_total = if target {
1086            self.total_true
1087        } else {
1088            self.total_false
1089        };
1090        for (i, &x) in features.iter().enumerate() {
1091            next_feature_sums[i] = checked_finite_add(next_feature_sums[i], x, "feature_sum")?;
1092            next_total = checked_finite_add(next_total, x, "total")?;
1093        }
1094        let next_class_count = if target {
1095            checked_increment(self.class_true_count, "class_true_count")?
1096        } else {
1097            checked_increment(self.class_false_count, "class_false_count")?
1098        };
1099        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
1100
1101        // Phase 2: commit atomically.
1102        if target {
1103            self.feature_sums_true = next_feature_sums;
1104            self.total_true = next_total;
1105            self.class_true_count = next_class_count;
1106        } else {
1107            self.feature_sums_false = next_feature_sums;
1108            self.total_false = next_total;
1109            self.class_false_count = next_class_count;
1110        }
1111        self.samples_seen = next_samples_seen;
1112        Ok(())
1113    }
1114
1115    fn reset(&mut self) {
1116        self.feature_sums_false.fill(0.0);
1117        self.feature_sums_true.fill(0.0);
1118        self.total_false = 0.0;
1119        self.total_true = 0.0;
1120        self.class_false_count = 0;
1121        self.class_true_count = 0;
1122        self.samples_seen = 0;
1123    }
1124}
1125
1126#[cfg(test)]
1127mod tests {
1128    use super::*;
1129    use rand::SeedableRng;
1130
1131    // ====================
1132    // GaussianNaiveBayes
1133    // ====================
1134
1135    #[test]
1136    fn gaussian_cold_start_returns_0_5() {
1137        let model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1138        let p = model.predict_proba(&[1.0, 2.0]).unwrap();
1139        assert!((p - 0.5).abs() < 1e-12);
1140    }
1141
1142    #[test]
1143    fn gaussian_learn_separable_data() {
1144        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1145        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1146        for _ in 0..200 {
1147            let x1 = 2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1148            let x2 = 2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1149            model.learn(&[x1, x2], true).unwrap();
1150            let x1 = -2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1151            let x2 = -2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1152            model.learn(&[x1, x2], false).unwrap();
1153        }
1154        let p_pos = model.predict_proba(&[2.0, 2.0]).unwrap();
1155        let p_neg = model.predict_proba(&[-2.0, -2.0]).unwrap();
1156        assert!(p_pos > 0.7, "p_pos = {p_pos}");
1157        assert!(p_neg < 0.3, "p_neg = {p_neg}");
1158    }
1159
1160    #[test]
1161    fn gaussian_dimension_mismatch_rejected() {
1162        let mut model = GaussianNaiveBayes::new(3, Default::default()).unwrap();
1163        assert!(model.predict_proba(&[1.0, 2.0]).is_err());
1164        assert!(model.learn(&[1.0, 2.0], true).is_err());
1165    }
1166
1167    #[test]
1168    fn gaussian_non_finite_rejected() {
1169        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1170        assert!(model.learn(&[f64::NAN, 1.0], true).is_err());
1171        assert!(model.learn(&[1.0, f64::INFINITY], true).is_err());
1172    }
1173
1174    #[test]
1175    fn gaussian_reset_clears_state() {
1176        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1177        model.learn(&[1.0, 2.0], true).unwrap();
1178        model.learn(&[-1.0, -2.0], false).unwrap();
1179        model.reset();
1180        assert_eq!(model.samples_seen(), 0);
1181        assert!((model.predict_proba(&[1.0, 2.0]).unwrap() - 0.5).abs() < 1e-12);
1182    }
1183
1184    #[test]
1185    fn gaussian_invalid_alpha_rejected() {
1186        assert!(GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: 0.0 }).is_err());
1187        assert!(GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: -1.0 }).is_err());
1188        assert!(GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: f64::NAN }).is_err());
1189    }
1190
1191    #[test]
1192    fn gaussian_predict_does_not_update_state() {
1193        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1194        model.learn(&[1.0, 2.0], true).unwrap();
1195        let before = model.samples_seen();
1196        let _ = model.predict_proba(&[0.5, 0.5]).unwrap();
1197        assert_eq!(model.samples_seen(), before);
1198    }
1199
1200    #[cfg(feature = "serde")]
1201    #[test]
1202    fn gaussian_serde_roundtrip() {
1203        let mut model = GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: 0.5 }).unwrap();
1204        model.learn(&[1.0, 2.0], true).unwrap();
1205        model.learn(&[1.5, 2.5], true).unwrap();
1206        model.learn(&[-1.0, -2.0], false).unwrap();
1207        model.learn(&[-1.5, -2.5], false).unwrap();
1208        let json = serde_json::to_string(&model).unwrap();
1209        let restored: GaussianNaiveBayes = serde_json::from_str(&json).unwrap();
1210        assert_eq!(restored.samples_seen(), model.samples_seen());
1211        assert_eq!(restored.feature_count(), model.feature_count());
1212        let p1 = model.predict_proba(&[0.5, 0.5]).unwrap();
1213        let p2 = restored.predict_proba(&[0.5, 0.5]).unwrap();
1214        assert!((p1 - p2).abs() < 1e-12);
1215    }
1216
1217    #[test]
1218    fn gaussian_predict_proba_in_range() {
1219        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1220        model.learn(&[1.0, 2.0], true).unwrap();
1221        model.learn(&[3.0, 4.0], true).unwrap();
1222        model.learn(&[-1.0, -2.0], false).unwrap();
1223        model.learn(&[-3.0, -4.0], false).unwrap();
1224        let p = model.predict_proba(&[0.5, 1.0]).unwrap();
1225        assert!(p > 0.0 && p < 1.0, "p = {p}");
1226    }
1227
1228    #[test]
1229    fn gaussian_zero_features_rejected() {
1230        assert!(GaussianNaiveBayes::new(0, Default::default()).is_err());
1231    }
1232
1233    #[test]
1234    fn gaussian_learns_gaussian_distribution() {
1235        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1236        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1237        for _ in 0..500 {
1238            let x1 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1239            let x2 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1240            model.learn(&[x1, x2], true).unwrap();
1241            let x1 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1242            let x2 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1243            model.learn(&[x1, x2], false).unwrap();
1244        }
1245        let mut correct = 0;
1246        let total = 100;
1247        for _ in 0..total {
1248            let x1 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1249            let x2 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1250            if model.predict(&[x1, x2]).unwrap() {
1251                correct += 1;
1252            }
1253            let x1 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1254            let x2 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1255            if !model.predict(&[x1, x2]).unwrap() {
1256                correct += 1;
1257            }
1258        }
1259        let accuracy = correct as f64 / (total * 2) as f64;
1260        assert!(accuracy > 0.95, "accuracy = {accuracy}");
1261    }
1262
1263    #[test]
1264    fn gaussian_single_class_predicts_that_class() {
1265        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1266        model.learn(&[1.0, 2.0], true).unwrap();
1267        model.learn(&[1.5, 2.5], true).unwrap();
1268        let p = model.predict_proba(&[1.0, 2.0]).unwrap();
1269        assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1270    }
1271
1272    // ====================
1273    // BernoulliNaiveBayes
1274    // ====================
1275
1276    #[test]
1277    fn bernoulli_cold_start_returns_0_5() {
1278        let model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1279        let p = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1280        assert!((p - 0.5).abs() < 1e-12);
1281    }
1282
1283    #[test]
1284    fn bernoulli_learn_separable_data() {
1285        let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1286        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1287        for _ in 0..200 {
1288            let f0 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.9 {
1289                1.0
1290            } else {
1291                0.0
1292            };
1293            let f1 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.1 {
1294                1.0
1295            } else {
1296                0.0
1297            };
1298            let f2 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.5 {
1299                1.0
1300            } else {
1301                0.0
1302            };
1303            model.learn(&[f0, f1, f2], true).unwrap();
1304            let f0 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.1 {
1305                1.0
1306            } else {
1307                0.0
1308            };
1309            let f1 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.9 {
1310                1.0
1311            } else {
1312                0.0
1313            };
1314            let f2 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.5 {
1315                1.0
1316            } else {
1317                0.0
1318            };
1319            model.learn(&[f0, f1, f2], false).unwrap();
1320        }
1321        let p_pos = model.predict_proba(&[1.0, 0.0, 0.0]).unwrap();
1322        let p_neg = model.predict_proba(&[0.0, 1.0, 0.0]).unwrap();
1323        assert!(p_pos > 0.7, "p_pos = {p_pos}");
1324        assert!(p_neg < 0.3, "p_neg = {p_neg}");
1325    }
1326
1327    #[test]
1328    fn bernoulli_dimension_mismatch_rejected() {
1329        let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1330        assert!(model.predict_proba(&[1.0, 0.0]).is_err());
1331        assert!(model.learn(&[1.0, 0.0], true).is_err());
1332    }
1333
1334    #[test]
1335    fn bernoulli_non_finite_rejected() {
1336        let mut model = BernoulliNaiveBayes::new(2, Default::default()).unwrap();
1337        assert!(model.learn(&[f64::NAN, 1.0], true).is_err());
1338        assert!(model.learn(&[1.0, f64::INFINITY], true).is_err());
1339    }
1340
1341    #[test]
1342    fn bernoulli_reset_clears_state() {
1343        let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1344        model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1345        model.learn(&[0.0, 1.0, 0.0], false).unwrap();
1346        model.reset();
1347        assert_eq!(model.samples_seen(), 0);
1348        assert!((model.predict_proba(&[1.0, 0.0, 1.0]).unwrap() - 0.5).abs() < 1e-12);
1349    }
1350
1351    #[test]
1352    fn bernoulli_invalid_alpha_rejected() {
1353        assert!(BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.0 }).is_err());
1354        assert!(BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: -1.0 }).is_err());
1355        assert!(BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: f64::NAN }).is_err());
1356    }
1357
1358    #[test]
1359    fn bernoulli_predict_does_not_update_state() {
1360        let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1361        model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1362        let before = model.samples_seen();
1363        let _ = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1364        assert_eq!(model.samples_seen(), before);
1365    }
1366
1367    #[cfg(feature = "serde")]
1368    #[test]
1369    fn bernoulli_serde_roundtrip() {
1370        let mut model = BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
1371        model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1372        model.learn(&[0.0, 1.0, 0.0], false).unwrap();
1373        model.learn(&[1.0, 1.0, 0.0], true).unwrap();
1374        let json = serde_json::to_string(&model).unwrap();
1375        let restored: BernoulliNaiveBayes = serde_json::from_str(&json).unwrap();
1376        assert_eq!(restored.samples_seen(), model.samples_seen());
1377        assert_eq!(restored.feature_count(), model.feature_count());
1378        let p1 = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1379        let p2 = restored.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1380        assert!((p1 - p2).abs() < 1e-12);
1381    }
1382
1383    #[test]
1384    fn bernoulli_predict_proba_in_range() {
1385        let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1386        model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1387        model.learn(&[0.0, 1.0, 0.0], false).unwrap();
1388        let p = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1389        assert!(p > 0.0 && p < 1.0, "p = {p}");
1390    }
1391
1392    #[test]
1393    fn bernoulli_zero_features_rejected() {
1394        assert!(BernoulliNaiveBayes::new(0, Default::default()).is_err());
1395    }
1396
1397    #[test]
1398    fn bernoulli_rejects_negative_values() {
1399        let mut model = BernoulliNaiveBayes::new(2, Default::default()).unwrap();
1400        assert!(model.learn(&[-1.0, 0.0], true).is_err());
1401        assert!(model.predict_proba(&[-0.5, 0.0]).is_err());
1402    }
1403
1404    // ====================
1405    // MultinomialNaiveBayes
1406    // ====================
1407
1408    #[test]
1409    fn multinomial_cold_start_returns_0_5() {
1410        let model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1411        let p = model.predict_proba(&[1.0, 2.0, 3.0]).unwrap();
1412        assert!((p - 0.5).abs() < 1e-12);
1413    }
1414
1415    #[test]
1416    fn multinomial_learn_separable_data() {
1417        let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1418        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1419        for _ in 0..200 {
1420            let f0 = rand::Rng::gen_range(&mut rng, 3.0..6.0);
1421            let f1 = rand::Rng::gen_range(&mut rng, 2.0..5.0);
1422            let f2 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
1423            model.learn(&[f0, f1, f2], true).unwrap();
1424            let f0 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
1425            let f1 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
1426            let f2 = rand::Rng::gen_range(&mut rng, 3.0..6.0);
1427            model.learn(&[f0, f1, f2], false).unwrap();
1428        }
1429        let p_pos = model.predict_proba(&[4.0, 3.0, 0.0]).unwrap();
1430        let p_neg = model.predict_proba(&[0.0, 0.0, 4.0]).unwrap();
1431        assert!(p_pos > 0.7, "p_pos = {p_pos}");
1432        assert!(p_neg < 0.3, "p_neg = {p_neg}");
1433    }
1434
1435    #[test]
1436    fn multinomial_dimension_mismatch_rejected() {
1437        let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1438        assert!(model.predict_proba(&[1.0, 2.0]).is_err());
1439        assert!(model.learn(&[1.0, 2.0], true).is_err());
1440    }
1441
1442    #[test]
1443    fn multinomial_non_finite_rejected() {
1444        let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1445        assert!(model.learn(&[f64::NAN, 1.0], true).is_err());
1446        assert!(model.learn(&[1.0, f64::INFINITY], true).is_err());
1447    }
1448
1449    #[test]
1450    fn multinomial_reset_clears_state() {
1451        let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1452        model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1453        model.learn(&[0.0, 1.0, 3.0], false).unwrap();
1454        model.reset();
1455        assert_eq!(model.samples_seen(), 0);
1456        assert!((model.predict_proba(&[1.0, 1.0, 1.0]).unwrap() - 0.5).abs() < 1e-12);
1457    }
1458
1459    #[test]
1460    fn multinomial_invalid_alpha_rejected() {
1461        assert!(MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.0 }).is_err());
1462        assert!(MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: -1.0 }).is_err());
1463        assert!(MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: f64::NAN }).is_err());
1464    }
1465
1466    #[test]
1467    fn multinomial_predict_does_not_update_state() {
1468        let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1469        model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1470        let before = model.samples_seen();
1471        let _ = model.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1472        assert_eq!(model.samples_seen(), before);
1473    }
1474
1475    #[cfg(feature = "serde")]
1476    #[test]
1477    fn multinomial_serde_roundtrip() {
1478        let mut model = MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
1479        model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1480        model.learn(&[0.0, 1.0, 3.0], false).unwrap();
1481        model.learn(&[1.0, 2.0, 1.0], true).unwrap();
1482        let json = serde_json::to_string(&model).unwrap();
1483        let restored: MultinomialNaiveBayes = serde_json::from_str(&json).unwrap();
1484        assert_eq!(restored.samples_seen(), model.samples_seen());
1485        assert_eq!(restored.feature_count(), model.feature_count());
1486        let p1 = model.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1487        let p2 = restored.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1488        assert!((p1 - p2).abs() < 1e-12);
1489    }
1490
1491    #[test]
1492    fn multinomial_predict_proba_in_range() {
1493        let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1494        model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1495        model.learn(&[0.0, 1.0, 3.0], false).unwrap();
1496        let p = model.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1497        assert!(p > 0.0 && p < 1.0, "p = {p}");
1498    }
1499
1500    #[test]
1501    fn multinomial_zero_features_rejected() {
1502        assert!(MultinomialNaiveBayes::new(0, Default::default()).is_err());
1503    }
1504
1505    #[test]
1506    fn multinomial_rejects_negative_values() {
1507        let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1508        assert!(model.learn(&[-1.0, 0.0], true).is_err());
1509        assert!(model.predict_proba(&[-0.5, 0.0]).is_err());
1510    }
1511
1512    #[test]
1513    fn multinomial_handles_all_zero_features() {
1514        let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1515        model.learn(&[0.0, 0.0, 0.0], true).unwrap();
1516        model.learn(&[0.0, 0.0, 0.0], false).unwrap();
1517        let p = model.predict_proba(&[0.0, 0.0, 0.0]).unwrap();
1518        assert!((p - 0.5).abs() < 1e-12, "p = {p}");
1519    }
1520
1521    // ====================
1522    // 4.4: probability finiteness
1523    // ====================
1524
1525    #[test]
1526    fn gaussian_predict_proba_rejects_extreme_features_causing_nan_log_odds() {
1527        // Train a balanced two-class model where both classes have small but
1528        // non-zero variance. Predicting with an extreme finite feature makes
1529        // both log-likelihoods underflow to -Infinity, which would yield
1530        // `log_odds = -Inf - (-Inf) = NaN` without the finiteness guard.
1531        let mut model = GaussianNaiveBayes::new(1, Default::default()).unwrap();
1532        model.learn(&[1.0], true).unwrap();
1533        model.learn(&[2.0], true).unwrap();
1534        model.learn(&[1.0], false).unwrap();
1535        model.learn(&[2.0], false).unwrap();
1536        let result = model.predict_proba(&[1e200]);
1537        assert!(result.is_err(), "expected Err for NaN log_odds, got Ok");
1538    }
1539
1540    #[test]
1541    fn bernoulli_predict_proba_rejects_extreme_features_causing_nan_log_odds() {
1542        // Train both classes so the feature is "present" most of the time,
1543        // giving p > 0.5 for both. `log_bernoulli(x, p)` simplifies to
1544        // `x * ln(p/(1-p)) + ln(1-p)`; with p > 0.5 the slope is positive,
1545        // so a huge `x` drives both log-likelihoods to +Infinity.
1546        // `log_odds = +Inf - (+Inf) = NaN` must be rejected.
1547        let mut model = BernoulliNaiveBayes::new(1, Default::default()).unwrap();
1548        for _ in 0..10 {
1549            model.learn(&[1.0], true).unwrap();
1550            model.learn(&[1.0], false).unwrap();
1551        }
1552        let result = model.predict_proba(&[1e308]);
1553        assert!(result.is_err(), "expected Err for NaN log_odds, got Ok");
1554    }
1555
1556    #[test]
1557    fn multinomial_predict_proba_rejects_extreme_features_causing_nan_log_odds() {
1558        // Train with large totals so the Laplace-smoothed p for the
1559        // zero-sum feature is ≈ 1/total (very small). With x = 1e308,
1560        // `x * ln(p)` overflows to -Infinity for both classes, making
1561        // log_odds = -Inf - (-Inf) = NaN.
1562        let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1563        model.learn(&[1e10, 0.0], true).unwrap();
1564        model.learn(&[0.0, 1e10], false).unwrap();
1565        let result = model.predict_proba(&[1e308, 1e308]);
1566        assert!(result.is_err(), "expected Err for NaN log_odds, got Ok");
1567    }
1568
1569    #[test]
1570    fn gaussian_single_class_predict_proba_returns_0_or_1() {
1571        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1572        model.learn(&[1.0, 2.0], true).unwrap();
1573        model.learn(&[1.5, 2.5], true).unwrap();
1574        let p = model.predict_proba(&[1.0, 2.0]).unwrap();
1575        assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1576    }
1577
1578    #[test]
1579    fn bernoulli_single_class_predict_proba_returns_0_or_1() {
1580        let mut model = BernoulliNaiveBayes::new(2, Default::default()).unwrap();
1581        model.learn(&[1.0, 0.0], true).unwrap();
1582        model.learn(&[0.0, 1.0], true).unwrap();
1583        let p = model.predict_proba(&[1.0, 0.0]).unwrap();
1584        assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1585    }
1586
1587    #[test]
1588    fn multinomial_single_class_predict_proba_returns_0_or_1() {
1589        let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1590        model.learn(&[1.0, 0.0], true).unwrap();
1591        model.learn(&[0.0, 1.0], true).unwrap();
1592        let p = model.predict_proba(&[1.0, 0.0]).unwrap();
1593        assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1594    }
1595
1596    #[test]
1597    fn gaussian_predict_proba_does_not_modify_state_on_error() {
1598        let mut model = GaussianNaiveBayes::new(1, Default::default()).unwrap();
1599        model.learn(&[1.0], true).unwrap();
1600        model.learn(&[2.0], true).unwrap();
1601        model.learn(&[1.0], false).unwrap();
1602        model.learn(&[2.0], false).unwrap();
1603        let before = model.samples_seen();
1604        let _ = model.predict_proba(&[1e200]);
1605        assert_eq!(model.samples_seen(), before);
1606    }
1607
1608    // ====================
1609    // 4.5: failure atomicity
1610    // ====================
1611
1612    #[test]
1613    fn gaussian_learn_failure_leaves_state_unchanged() {
1614        // Train one sample so feature 1 has non-zero mean; the second
1615        // learn call computes a huge m2 for feature 1 (overflow to Inf),
1616        // but feature 0's next state is valid. The old code would commit
1617        // feature 0 before feature 1 fails; the new code must reject the
1618        // entire call.
1619        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1620        model.learn(&[1.0, 1.0], true).unwrap();
1621        let before_samples = model.samples_seen();
1622        let before_class_count = model.class_true.class_count;
1623
1624        let result = model.learn(&[2.0, 1e200], true);
1625        assert!(result.is_err(), "expected overflow error");
1626
1627        // State must be unchanged.
1628        assert_eq!(model.samples_seen(), before_samples);
1629        assert_eq!(model.class_true.class_count, before_class_count);
1630        assert_eq!(model.class_true.counts, vec![1, 1]);
1631        assert_eq!(model.class_true.means, vec![1.0, 1.0]);
1632        assert_eq!(model.class_true.m2s, vec![0.0, 0.0]);
1633    }
1634
1635    #[test]
1636    fn gaussian_learn_succeeds_after_failed_attempt() {
1637        let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1638        model.learn(&[1.0, 1.0], true).unwrap();
1639        // Failed call (overflow) must not corrupt state.
1640        let _ = model.learn(&[2.0, 1e200], true);
1641        // Subsequent valid call must work.
1642        model.learn(&[2.0, 3.0], true).unwrap();
1643        assert_eq!(model.samples_seen(), 2);
1644    }
1645
1646    #[test]
1647    fn multinomial_learn_failure_leaves_state_unchanged() {
1648        // With a fresh model, learn([1e308, 1e308]) causes total to overflow
1649        // to Inf on the second feature. The old code commits feature 0 and
1650        // the first feature's contribution to total before failing; the new
1651        // code must reject the entire call.
1652        let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1653        let before_samples = model.samples_seen();
1654
1655        let result = model.learn(&[1e308, 1e308], true);
1656        assert!(result.is_err(), "expected overflow error");
1657
1658        // State must be unchanged.
1659        assert_eq!(model.samples_seen(), before_samples);
1660        assert_eq!(model.class_true_count, 0);
1661        assert_eq!(model.feature_sums_true, vec![0.0, 0.0]);
1662        assert_eq!(model.total_true, 0.0);
1663    }
1664
1665    #[test]
1666    fn multinomial_learn_succeeds_after_failed_attempt() {
1667        let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1668        let _ = model.learn(&[1e308, 1e308], true);
1669        // Subsequent valid call must work.
1670        model.learn(&[1.0, 2.0], true).unwrap();
1671        assert_eq!(model.samples_seen(), 1);
1672        assert_eq!(model.feature_sums_true, vec![1.0, 2.0]);
1673    }
1674
1675    #[cfg(feature = "serde")]
1676    #[test]
1677    fn bernoulli_learn_failure_leaves_state_unchanged() {
1678        // Construct a model where feature 1's count is at u64::MAX via serde.
1679        // A learn call that tries to increment both features must fail on
1680        // feature 1 without modifying feature 0. The class_count must be
1681        // at least as large as the per-feature count to satisfy the
1682        // deserialization invariants (`feature_true_counts_true[i] <=
1683        // class_true_count`), so `class_true_count` and `samples_seen`
1684        // sit at u64::MAX as well.
1685        let json = r#"{
1686            "feature_count": 2,
1687            "config": {"alpha": 1.0},
1688            "feature_true_counts_false": [0, 0],
1689            "feature_true_counts_true": [0, 18446744073709551615],
1690            "class_false_count": 0,
1691            "class_true_count": 18446744073709551615,
1692            "samples_seen": 18446744073709551615
1693        }"#;
1694        let mut model: BernoulliNaiveBayes = serde_json::from_str(json).unwrap();
1695        let before_samples = model.samples_seen();
1696
1697        let result = model.learn(&[1.0, 1.0], true);
1698        assert!(result.is_err(), "expected counter overflow");
1699
1700        // State must be unchanged.
1701        assert_eq!(model.samples_seen(), before_samples);
1702        assert_eq!(model.feature_true_counts_true, vec![0u64, u64::MAX]);
1703        assert_eq!(model.class_true_count, u64::MAX);
1704    }
1705
1706    #[cfg(feature = "serde")]
1707    #[test]
1708    fn gaussian_learn_failure_on_samples_seen_overflow_leaves_state_unchanged() {
1709        // Construct a valid model where the per-class counts are one below
1710        // the overflow boundary and `samples_seen` already sits at u64::MAX.
1711        // A successful feature update and class_count increment must not be
1712        // committed when `samples_seen` overflows.
1713        //
1714        // The Gaussian invariant requires every per-feature count to equal
1715        // `class_count`, so the original "class_count = u64::MAX, counts = 1"
1716        // fixture is no longer admissible. This repurposed fixture instead
1717        // pins `class_true.counts[0] == class_true.class_count == u64::MAX-1`
1718        // and `samples_seen = u64::MAX` so that the feature and class_count
1719        // increments both succeed but the samples_seen increment overflows.
1720        let max_minus_one = u64::MAX - 1;
1721        let json = format!(
1722            r#"{{
1723                "feature_count": 1,
1724                "config": {{"alpha": 1.0}},
1725                "class_false": {{
1726                    "counts": [1],
1727                    "means": [3.0],
1728                    "m2s": [0.0],
1729                    "class_count": 1
1730                }},
1731                "class_true": {{
1732                    "counts": [{max_minus_one}],
1733                    "means": [5.0],
1734                    "m2s": [0.0],
1735                    "class_count": {max_minus_one}
1736                }},
1737                "samples_seen": 18446744073709551615
1738            }}"#
1739        );
1740        let mut model: GaussianNaiveBayes = serde_json::from_str(&json).unwrap();
1741        let before_samples = model.samples_seen();
1742
1743        let result = model.learn(&[6.0], true);
1744        assert!(result.is_err(), "expected samples_seen overflow");
1745
1746        // State must be unchanged — feature 0's next state and the new
1747        // class_count were computed but not committed because samples_seen
1748        // overflowed.
1749        assert_eq!(model.samples_seen(), before_samples);
1750        assert_eq!(model.class_true.counts, vec![max_minus_one]);
1751        assert_eq!(model.class_true.means, vec![5.0]);
1752        assert_eq!(model.class_true.m2s, vec![0.0]);
1753    }
1754
1755    #[cfg(feature = "serde")]
1756    #[test]
1757    fn multinomial_learn_failure_on_class_count_overflow_leaves_state_unchanged() {
1758        // `class_true_count` sits at u64::MAX so the class_count increment
1759        // overflows. The Multinomial invariants do not couple feature sums
1760        // to `class_true_count`, so `class_true_count = u64::MAX` is
1761        // admissible as long as `samples_seen` matches (here both equal
1762        // u64::MAX, since `class_false_count = 0`).
1763        let json = r#"{
1764            "feature_count": 1,
1765            "config": {"alpha": 1.0},
1766            "feature_sums_false": [0.0],
1767            "feature_sums_true": [5.0],
1768            "total_false": 0.0,
1769            "total_true": 5.0,
1770            "class_false_count": 0,
1771            "class_true_count": 18446744073709551615,
1772            "samples_seen": 18446744073709551615
1773        }"#;
1774        let mut model: MultinomialNaiveBayes = serde_json::from_str(json).unwrap();
1775        let before_samples = model.samples_seen();
1776
1777        let result = model.learn(&[3.0], true);
1778        assert!(result.is_err(), "expected class_count overflow");
1779
1780        assert_eq!(model.samples_seen(), before_samples);
1781        assert_eq!(model.feature_sums_true, vec![5.0]);
1782        assert_eq!(model.total_true, 5.0);
1783        assert_eq!(model.class_true_count, u64::MAX);
1784    }
1785
1786    // ====================================================================
1787    // §6.1: Naive Bayes serde trust-boundary validation
1788    // ====================================================================
1789
1790    #[test]
1791    fn approx_equal_non_negative_tolerances() {
1792        // Exact equality.
1793        assert!(approx_equal_non_negative(0.0, 0.0, 1e-9, 1e-6));
1794        assert!(approx_equal_non_negative(5.0, 5.0, 1e-9, 1e-6));
1795        // Absolute tolerance floor.
1796        assert!(approx_equal_non_negative(0.0, 1e-10, 1e-9, 1e-6));
1797        assert!(!approx_equal_non_negative(0.0, 1e-8, 1e-9, 1e-6));
1798        // Relative tolerance scales with magnitude.
1799        assert!(approx_equal_non_negative(
1800            1_000_000.0,
1801            1_000_000.5,
1802            1e-9,
1803            1e-6
1804        ));
1805        assert!(!approx_equal_non_negative(
1806            1_000_000.0,
1807            1_000_002.0,
1808            1e-9,
1809            1e-6
1810        ));
1811        // Non-finite inputs always compare false.
1812        assert!(!approx_equal_non_negative(f64::NAN, 0.0, 1e-9, 1e-6));
1813        assert!(!approx_equal_non_negative(
1814            f64::INFINITY,
1815            f64::INFINITY,
1816            1e-9,
1817            1e-6
1818        ));
1819        // Order-of-magnitude mismatch must be rejected.
1820        assert!(!approx_equal_non_negative(1.0, 1e6, 1e-9, 1e-6));
1821    }
1822
1823    #[cfg(feature = "serde")]
1824    #[test]
1825    fn gaussian_serde_rejects_feature_vector_length_mismatch() {
1826        // class_true vectors have length 1 while feature_count = 2.
1827        let json = r#"{
1828            "feature_count": 2,
1829            "config": {"alpha": 1.0},
1830            "class_false": {
1831                "counts": [0, 0],
1832                "means": [0.0, 0.0],
1833                "m2s": [0.0, 0.0],
1834                "class_count": 0
1835            },
1836            "class_true": {
1837                "counts": [0],
1838                "means": [0.0],
1839                "m2s": [0.0],
1840                "class_count": 0
1841            },
1842            "samples_seen": 0
1843        }"#;
1844        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1845        assert!(
1846            result.is_err(),
1847            "feature vector length mismatch must be rejected"
1848        );
1849    }
1850
1851    #[cfg(feature = "serde")]
1852    #[test]
1853    fn gaussian_serde_rejects_negative_m2() {
1854        let json = r#"{
1855            "feature_count": 1,
1856            "config": {"alpha": 1.0},
1857            "class_false": {
1858                "counts": [0],
1859                "means": [0.0],
1860                "m2s": [0.0],
1861                "class_count": 0
1862            },
1863            "class_true": {
1864                "counts": [2],
1865                "means": [5.0],
1866                "m2s": [-1.0],
1867                "class_count": 2
1868            },
1869            "samples_seen": 2
1870        }"#;
1871        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1872        assert!(result.is_err(), "negative m2 must be rejected");
1873    }
1874
1875    #[cfg(feature = "serde")]
1876    #[test]
1877    fn gaussian_serde_rejects_count_zero_with_nonzero_state() {
1878        // class_count == 0 but means/m2s are non-zero: this state can never
1879        // arise from a legitimate learn() path.
1880        let json = r#"{
1881            "feature_count": 1,
1882            "config": {"alpha": 1.0},
1883            "class_false": {
1884                "counts": [0],
1885                "means": [0.0],
1886                "m2s": [0.0],
1887                "class_count": 0
1888            },
1889            "class_true": {
1890                "counts": [0],
1891                "means": [7.0],
1892                "m2s": [0.0],
1893                "class_count": 0
1894            },
1895            "samples_seen": 0
1896        }"#;
1897        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1898        assert!(
1899            result.is_err(),
1900            "class_count=0 with non-zero mean must be rejected"
1901        );
1902    }
1903
1904    #[cfg(feature = "serde")]
1905    #[test]
1906    fn gaussian_serde_rejects_count_one_with_nonzero_m2() {
1907        // After a single Welford update, M2 is exactly 0; a non-zero M2 at
1908        // count == 1 indicates a corrupted or malicious payload.
1909        let json = r#"{
1910            "feature_count": 1,
1911            "config": {"alpha": 1.0},
1912            "class_false": {
1913                "counts": [0],
1914                "means": [0.0],
1915                "m2s": [0.0],
1916                "class_count": 0
1917            },
1918            "class_true": {
1919                "counts": [1],
1920                "means": [5.0],
1921                "m2s": [0.25],
1922                "class_count": 1
1923            },
1924            "samples_seen": 1
1925        }"#;
1926        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1927        assert!(
1928            result.is_err(),
1929            "class_count=1 with non-zero m2 must be rejected"
1930        );
1931    }
1932
1933    #[cfg(feature = "serde")]
1934    #[test]
1935    fn gaussian_serde_rejects_samples_seen_mismatch() {
1936        // class_false.class_count + class_true.class_count = 1 + 1 = 2,
1937        // but samples_seen = 5.
1938        let json = r#"{
1939            "feature_count": 1,
1940            "config": {"alpha": 1.0},
1941            "class_false": {
1942                "counts": [1],
1943                "means": [3.0],
1944                "m2s": [0.0],
1945                "class_count": 1
1946            },
1947            "class_true": {
1948                "counts": [1],
1949                "means": [5.0],
1950                "m2s": [0.0],
1951                "class_count": 1
1952            },
1953            "samples_seen": 5
1954        }"#;
1955        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1956        assert!(result.is_err(), "samples_seen mismatch must be rejected");
1957    }
1958
1959    #[cfg(feature = "serde")]
1960    #[test]
1961    fn gaussian_serde_rejects_feature_count_not_equal_class_count() {
1962        // Gaussian NB updates every feature on each learn, so per-feature
1963        // counts must equal class_count. Here counts[0] = 3 but class_count
1964        // = 1.
1965        let json = r#"{
1966            "feature_count": 1,
1967            "config": {"alpha": 1.0},
1968            "class_false": {
1969                "counts": [0],
1970                "means": [0.0],
1971                "m2s": [0.0],
1972                "class_count": 0
1973            },
1974            "class_true": {
1975                "counts": [3],
1976                "means": [5.0],
1977                "m2s": [0.5],
1978                "class_count": 1
1979            },
1980            "samples_seen": 1
1981        }"#;
1982        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1983        assert!(
1984            result.is_err(),
1985            "per-feature count != class_count must be rejected"
1986        );
1987    }
1988
1989    #[cfg(feature = "serde")]
1990    #[test]
1991    fn gaussian_serde_rejects_invalid_alpha() {
1992        let json = r#"{
1993            "feature_count": 1,
1994            "config": {"alpha": 0.0},
1995            "class_false": {
1996                "counts": [0],
1997                "means": [0.0],
1998                "m2s": [0.0],
1999                "class_count": 0
2000            },
2001            "class_true": {
2002                "counts": [0],
2003                "means": [0.0],
2004                "m2s": [0.0],
2005                "class_count": 0
2006            },
2007            "samples_seen": 0
2008        }"#;
2009        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2010        assert!(result.is_err(), "invalid alpha must be rejected");
2011    }
2012
2013    #[cfg(feature = "serde")]
2014    #[test]
2015    fn gaussian_serde_rejects_non_finite_state() {
2016        let json = r#"{
2017            "feature_count": 1,
2018            "config": {"alpha": 1.0},
2019            "class_false": {
2020                "counts": [0],
2021                "means": [0.0],
2022                "m2s": [0.0],
2023                "class_count": 0
2024            },
2025            "class_true": {
2026                "counts": [1],
2027                "means": [NaN],
2028                "m2s": [0.0],
2029                "class_count": 1
2030            },
2031            "samples_seen": 1
2032        }"#;
2033        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2034        assert!(result.is_err(), "non-finite mean must be rejected");
2035    }
2036
2037    #[cfg(feature = "serde")]
2038    #[test]
2039    fn gaussian_serde_rejects_malicious_state_without_panic() {
2040        // A payload that, if accepted, would index out of bounds during
2041        // predict(). Validation must reject it cleanly instead of panicking.
2042        let json = r#"{
2043            "feature_count": 3,
2044            "config": {"alpha": 1.0},
2045            "class_false": {
2046                "counts": [0],
2047                "means": [0.0],
2048                "m2s": [0.0],
2049                "class_count": 0
2050            },
2051            "class_true": {
2052                "counts": [0],
2053                "means": [0.0],
2054                "m2s": [0.0],
2055                "class_count": 0
2056            },
2057            "samples_seen": 0
2058        }"#;
2059        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2060        assert!(
2061            result.is_err(),
2062            "malicious length-mismatched state must be rejected, not panicked on"
2063        );
2064    }
2065
2066    #[cfg(feature = "serde")]
2067    #[test]
2068    fn bernoulli_serde_rejects_feature_count_above_class_count() {
2069        // feature_true_counts_true[0] = 5 but class_true_count = 1: the
2070        // per-feature count cannot exceed the number of training samples
2071        // for that class.
2072        let json = r#"{
2073            "feature_count": 1,
2074            "config": {"alpha": 1.0},
2075            "feature_true_counts_false": [0],
2076            "feature_true_counts_true": [5],
2077            "class_false_count": 0,
2078            "class_true_count": 1,
2079            "samples_seen": 1
2080        }"#;
2081        let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2082        assert!(
2083            result.is_err(),
2084            "feature count > class_count must be rejected"
2085        );
2086    }
2087
2088    #[cfg(feature = "serde")]
2089    #[test]
2090    fn bernoulli_serde_rejects_samples_seen_mismatch() {
2091        let json = r#"{
2092            "feature_count": 1,
2093            "config": {"alpha": 1.0},
2094            "feature_true_counts_false": [0],
2095            "feature_true_counts_true": [0],
2096            "class_false_count": 1,
2097            "class_true_count": 1,
2098            "samples_seen": 5
2099        }"#;
2100        let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2101        assert!(result.is_err(), "samples_seen mismatch must be rejected");
2102    }
2103
2104    #[cfg(feature = "serde")]
2105    #[test]
2106    fn bernoulli_serde_rejects_feature_vector_length_mismatch() {
2107        let json = r#"{
2108            "feature_count": 2,
2109            "config": {"alpha": 1.0},
2110            "feature_true_counts_false": [0],
2111            "feature_true_counts_true": [0, 0],
2112            "class_false_count": 0,
2113            "class_true_count": 0,
2114            "samples_seen": 0
2115        }"#;
2116        let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2117        assert!(
2118            result.is_err(),
2119            "feature vector length mismatch must be rejected"
2120        );
2121    }
2122
2123    #[cfg(feature = "serde")]
2124    #[test]
2125    fn bernoulli_serde_rejects_invalid_alpha() {
2126        let json = r#"{
2127            "feature_count": 1,
2128            "config": {"alpha": -1.0},
2129            "feature_true_counts_false": [0],
2130            "feature_true_counts_true": [0],
2131            "class_false_count": 0,
2132            "class_true_count": 0,
2133            "samples_seen": 0
2134        }"#;
2135        let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2136        assert!(result.is_err(), "invalid alpha must be rejected");
2137    }
2138
2139    #[cfg(feature = "serde")]
2140    #[test]
2141    fn multinomial_serde_rejects_negative_feature_sum() {
2142        let json = r#"{
2143            "feature_count": 1,
2144            "config": {"alpha": 1.0},
2145            "feature_sums_false": [0.0],
2146            "feature_sums_true": [-3.0],
2147            "total_false": 0.0,
2148            "total_true": -3.0,
2149            "class_false_count": 0,
2150            "class_true_count": 1,
2151            "samples_seen": 1
2152        }"#;
2153        let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2154        assert!(result.is_err(), "negative feature sum must be rejected");
2155    }
2156
2157    #[cfg(feature = "serde")]
2158    #[test]
2159    fn multinomial_serde_rejects_total_mismatch() {
2160        // sum(feature_sums_true) = 1.0 + 2.0 = 3.0 but total_true = 100.0.
2161        // The discrepancy is orders of magnitude beyond tolerance.
2162        let json = r#"{
2163            "feature_count": 2,
2164            "config": {"alpha": 1.0},
2165            "feature_sums_false": [0.0, 0.0],
2166            "feature_sums_true": [1.0, 2.0],
2167            "total_false": 0.0,
2168            "total_true": 100.0,
2169            "class_false_count": 0,
2170            "class_true_count": 1,
2171            "samples_seen": 1
2172        }"#;
2173        let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2174        assert!(
2175            result.is_err(),
2176            "total mismatch beyond tolerance must be rejected"
2177        );
2178    }
2179
2180    #[cfg(feature = "serde")]
2181    #[test]
2182    fn multinomial_serde_rejects_samples_seen_mismatch() {
2183        let json = r#"{
2184            "feature_count": 1,
2185            "config": {"alpha": 1.0},
2186            "feature_sums_false": [0.0],
2187            "feature_sums_true": [0.0],
2188            "total_false": 0.0,
2189            "total_true": 0.0,
2190            "class_false_count": 1,
2191            "class_true_count": 1,
2192            "samples_seen": 5
2193        }"#;
2194        let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2195        assert!(result.is_err(), "samples_seen mismatch must be rejected");
2196    }
2197
2198    #[cfg(feature = "serde")]
2199    #[test]
2200    fn multinomial_serde_accepts_tiny_total_roundoff() {
2201        // Floating-point accumulation can introduce tiny round-off between
2202        // the cached total and the recomputed sum. The tolerance must
2203        // accept legitimate round-off (here ~1e-12 on a total of ~3.0).
2204        let json = r#"{
2205            "feature_count": 2,
2206            "config": {"alpha": 1.0},
2207            "feature_sums_false": [0.0, 0.0],
2208            "feature_sums_true": [1.0, 2.0],
2209            "total_false": 0.0,
2210            "total_true": 3.000000000001,
2211            "class_false_count": 0,
2212            "class_true_count": 1,
2213            "samples_seen": 1
2214        }"#;
2215        let model: MultinomialNaiveBayes = serde_json::from_str(json).unwrap();
2216        assert_eq!(model.samples_seen(), 1);
2217    }
2218
2219    #[cfg(feature = "serde")]
2220    #[test]
2221    fn multinomial_serde_rejects_feature_vector_length_mismatch() {
2222        let json = r#"{
2223            "feature_count": 2,
2224            "config": {"alpha": 1.0},
2225            "feature_sums_false": [0.0, 0.0],
2226            "feature_sums_true": [0.0],
2227            "total_false": 0.0,
2228            "total_true": 0.0,
2229            "class_false_count": 0,
2230            "class_true_count": 0,
2231            "samples_seen": 0
2232        }"#;
2233        let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2234        assert!(
2235            result.is_err(),
2236            "feature vector length mismatch must be rejected"
2237        );
2238    }
2239
2240    #[cfg(feature = "serde")]
2241    #[test]
2242    fn gaussian_serde_rejects_class_count_overflow() {
2243        // Two class counts near u64::MAX must not panic on addition;
2244        // the validator must reject via checked_add overflow path.
2245        let json = r#"{
2246            "feature_count": 1,
2247            "config": {"alpha": 1.0},
2248            "class_false": {
2249                "counts": [1],
2250                "means": [0.0],
2251                "m2s": [0.0],
2252                "class_count": 18446744073709551615
2253            },
2254            "class_true": {
2255                "counts": [1],
2256                "means": [0.0],
2257                "m2s": [0.0],
2258                "class_count": 18446744073709551615
2259            },
2260            "samples_seen": 18446744073709551614
2261        }"#;
2262        let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2263        assert!(
2264            result.is_err(),
2265            "u64 overflow on class_count sum must be rejected, not panic"
2266        );
2267    }
2268
2269    #[cfg(feature = "serde")]
2270    #[test]
2271    fn bernoulli_serde_rejects_class_count_overflow() {
2272        let json = r#"{
2273            "feature_count": 1,
2274            "config": {"alpha": 1.0},
2275            "feature_true_counts_false": [1],
2276            "feature_true_counts_true": [1],
2277            "class_false_count": 18446744073709551615,
2278            "class_true_count": 18446744073709551615,
2279            "samples_seen": 18446744073709551614
2280        }"#;
2281        let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2282        assert!(
2283            result.is_err(),
2284            "u64 overflow on class_count sum must be rejected, not panic"
2285        );
2286    }
2287
2288    #[cfg(feature = "serde")]
2289    #[test]
2290    fn multinomial_serde_rejects_class_count_overflow() {
2291        let json = r#"{
2292            "feature_count": 1,
2293            "config": {"alpha": 1.0},
2294            "feature_sums_false": [1.0],
2295            "feature_sums_true": [1.0],
2296            "total_false": 1.0,
2297            "total_true": 1.0,
2298            "class_false_count": 18446744073709551615,
2299            "class_true_count": 18446744073709551615,
2300            "samples_seen": 18446744073709551614
2301        }"#;
2302        let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2303        assert!(
2304            result.is_err(),
2305            "u64 overflow on class_count sum must be rejected, not panic"
2306        );
2307    }
2308
2309    #[cfg(feature = "serde")]
2310    #[test]
2311    fn naive_bayes_valid_roundtrip_preserves_prediction() {
2312        // Round-trip each variant through serde and confirm that predictions
2313        // on a held-out feature vector are bit-for-bit identical. This guards
2314        // against validation logic that is so strict it rejects legitimate
2315        // trained state, or so loose it alters the prediction path.
2316        let probe = [0.5, 1.0, 0.0];
2317
2318        let mut g = GaussianNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
2319        g.learn(&[1.0, 2.0, 0.0], true).unwrap();
2320        g.learn(&[1.5, 2.5, 1.0], true).unwrap();
2321        g.learn(&[-1.0, -2.0, 0.0], false).unwrap();
2322        g.learn(&[-1.5, -2.5, 1.0], false).unwrap();
2323        let g_json = serde_json::to_string(&g).unwrap();
2324        let g_restored: GaussianNaiveBayes = serde_json::from_str(&g_json).unwrap();
2325        assert_eq!(g_restored.samples_seen(), g.samples_seen());
2326        assert_eq!(g_restored.feature_count(), g.feature_count());
2327        let g_p1 = g.predict_proba(&probe).unwrap();
2328        let g_p2 = g_restored.predict_proba(&probe).unwrap();
2329        assert!((g_p1 - g_p2).abs() < 1e-12);
2330
2331        let mut b = BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
2332        b.learn(&[1.0, 0.0, 1.0], true).unwrap();
2333        b.learn(&[0.0, 1.0, 0.0], false).unwrap();
2334        b.learn(&[1.0, 1.0, 0.0], true).unwrap();
2335        let b_json = serde_json::to_string(&b).unwrap();
2336        let b_restored: BernoulliNaiveBayes = serde_json::from_str(&b_json).unwrap();
2337        assert_eq!(b_restored.samples_seen(), b.samples_seen());
2338        assert_eq!(b_restored.feature_count(), b.feature_count());
2339        let b_p1 = b.predict_proba(&probe).unwrap();
2340        let b_p2 = b_restored.predict_proba(&probe).unwrap();
2341        assert!((b_p1 - b_p2).abs() < 1e-12);
2342
2343        let mut m = MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
2344        m.learn(&[2.0, 1.0, 0.0], true).unwrap();
2345        m.learn(&[0.0, 1.0, 3.0], false).unwrap();
2346        m.learn(&[1.0, 2.0, 1.0], true).unwrap();
2347        let m_json = serde_json::to_string(&m).unwrap();
2348        let m_restored: MultinomialNaiveBayes = serde_json::from_str(&m_json).unwrap();
2349        assert_eq!(m_restored.samples_seen(), m.samples_seen());
2350        assert_eq!(m_restored.feature_count(), m.feature_count());
2351        let m_p1 = m.predict_proba(&probe).unwrap();
2352        let m_p2 = m_restored.predict_proba(&probe).unwrap();
2353        assert!((m_p1 - m_p2).abs() < 1e-12);
2354    }
2355}