1use crate::bandit::stats::ArmStats;
24use crate::bandit::{
25 Bandit, checked_finite_add, checked_increment, validate_arm, validate_reward_01,
26 validate_sample_count,
27};
28use crate::error::RillError;
29#[cfg(feature = "serde")]
30use crate::persistence::ValidateState;
31use rand::Rng;
32
33#[derive(Debug, Clone, PartialEq)]
46#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
47#[non_exhaustive]
48pub struct ThompsonConfig {
49 pub alpha_prior: f64,
54
55 pub beta_prior: f64,
60}
61
62impl Default for ThompsonConfig {
63 fn default() -> Self {
64 Self {
65 alpha_prior: 1.0,
66 beta_prior: 1.0,
67 }
68 }
69}
70
71impl ThompsonConfig {
72 pub fn validate(&self) -> Result<(), RillError> {
74 if !self.alpha_prior.is_finite() || self.alpha_prior <= 0.0 {
75 return Err(RillError::InvalidParameter {
76 name: "alpha_prior",
77 value: self.alpha_prior,
78 });
79 }
80 if !self.beta_prior.is_finite() || self.beta_prior <= 0.0 {
81 return Err(RillError::InvalidParameter {
82 name: "beta_prior",
83 value: self.beta_prior,
84 });
85 }
86 Ok(())
87 }
88}
89
90#[derive(Debug, Clone)]
116#[cfg_attr(feature = "serde", derive(serde::Serialize))]
117pub struct ThompsonSampling {
118 arm_count: usize,
119 config: ThompsonConfig,
120 alphas: Vec<f64>,
122 betas: Vec<f64>,
124 pulls: Vec<u64>,
126 total_rewards: Vec<f64>,
128 samples_seen: u64,
130}
131
132impl ThompsonSampling {
133 pub fn new(arm_count: usize, config: ThompsonConfig) -> Result<Self, RillError> {
140 if arm_count == 0 {
141 return Err(RillError::InvalidArmCount(arm_count));
142 }
143 config.validate()?;
144
145 let alpha_prior = config.alpha_prior;
146 let beta_prior = config.beta_prior;
147 Ok(Self {
148 arm_count,
149 config,
150 alphas: vec![alpha_prior; arm_count],
151 betas: vec![beta_prior; arm_count],
152 pulls: vec![0; arm_count],
153 total_rewards: vec![0.0; arm_count],
154 samples_seen: 0,
155 })
156 }
157
158 pub fn alphas(&self) -> &[f64] {
160 &self.alphas
161 }
162
163 pub fn betas(&self) -> &[f64] {
165 &self.betas
166 }
167
168 pub fn pulls(&self) -> &[u64] {
170 &self.pulls
171 }
172
173 pub fn validate(&self) -> Result<(), RillError> {
177 if self.arm_count == 0 {
178 return Err(RillError::InvalidArmCount(self.arm_count));
179 }
180 self.config.validate()?;
181 if self.alphas.len() != self.arm_count
182 || self.betas.len() != self.arm_count
183 || self.pulls.len() != self.arm_count
184 || self.total_rewards.len() != self.arm_count
185 {
186 return Err(RillError::InvalidState(
187 "arm_count does not match per-arm state lengths".to_owned(),
188 ));
189 }
190 validate_sample_count(&self.pulls, self.samples_seen)?;
191
192 for arm in 0..self.arm_count {
193 let pulls = self.pulls[arm] as f64;
194 let total = self.total_rewards[arm];
195 let alpha = self.alphas[arm];
196 let beta = self.betas[arm];
197 if !total.is_finite() || total < 0.0 || total > pulls {
198 return Err(RillError::InvalidState(format!(
199 "total reward for arm {arm} is inconsistent with [0, 1] rewards"
200 )));
201 }
202 let expected_alpha = self.config.alpha_prior + total;
203 let expected_beta = self.config.beta_prior + pulls - total;
204 let alpha_tolerance = 1e-9 * expected_alpha.abs().max(1.0);
205 let beta_tolerance = 1e-9 * expected_beta.abs().max(1.0);
206 if !alpha.is_finite()
207 || !beta.is_finite()
208 || (alpha - expected_alpha).abs() > alpha_tolerance
209 || (beta - expected_beta).abs() > beta_tolerance
210 {
211 return Err(RillError::InvalidState(format!(
212 "posterior parameters for arm {arm} are inconsistent with observations"
213 )));
214 }
215 }
216 Ok(())
217 }
218
219 fn sample_beta(rng: &mut impl Rng, alpha: f64, beta: f64) -> f64 {
224 let x = Self::sample_gamma(rng, alpha);
225 let y = Self::sample_gamma(rng, beta);
226 let denom = x + y;
228 if denom <= 0.0 {
229 0.5
231 } else {
232 x / denom
233 }
234 }
235
236 fn sample_gamma(rng: &mut impl Rng, shape: f64) -> f64 {
243 if shape < 1.0 {
244 let u: f64 = rng.gen_range(1e-10..1.0);
246 let g = Self::sample_gamma(rng, shape + 1.0);
247 return g * u.powf(1.0 / shape);
248 }
249
250 let d = shape - 1.0 / 3.0;
252 let c = 1.0 / (9.0 * d).sqrt();
253
254 loop {
255 let (x, _unused) = Self::box_muller(rng);
257 let v = (1.0 + c * x).powi(3);
258 if v <= 0.0 {
259 continue;
260 }
261 let u: f64 = rng.gen_range(0.0..1.0);
262 if u < 1.0 - 0.0331 * x.powi(4) {
263 return d * v;
264 }
265 if u.ln() < 0.5 * x * x + d * (1.0 - v + v.ln()) {
266 return d * v;
267 }
268 }
269 }
270
271 fn box_muller(rng: &mut impl Rng) -> (f64, f64) {
274 let u1: f64 = rng.gen_range(1e-10..1.0);
275 let u2: f64 = rng.gen_range(0.0..1.0);
276 let mag = (-2.0 * u1.ln()).sqrt();
277 let z0 = mag * (2.0 * std::f64::consts::PI * u2).cos();
278 let z1 = mag * (2.0 * std::f64::consts::PI * u2).sin();
279 (z0, z1)
280 }
281}
282
283impl Bandit for ThompsonSampling {
284 fn arm_count(&self) -> usize {
285 self.arm_count
286 }
287
288 fn samples_seen(&self) -> u64 {
289 self.samples_seen
290 }
291
292 fn select(&self, rng: &mut impl Rng) -> Result<usize, RillError> {
293 let mut best_arm = 0usize;
294 let mut best_sample = f64::NEG_INFINITY;
295
296 for arm in 0..self.arm_count {
297 let sample = Self::sample_beta(rng, self.alphas[arm], self.betas[arm]);
298 if sample > best_sample {
299 best_sample = sample;
300 best_arm = arm;
301 }
302 }
303
304 Ok(best_arm)
305 }
306
307 fn update(&mut self, arm: usize, reward: f64) -> Result<(), RillError> {
308 validate_arm(self.arm_count, arm)?;
309 validate_reward_01(reward)?;
310
311 let next_alpha = checked_finite_add(self.alphas[arm], reward, "alpha")?;
316 let next_beta = checked_finite_add(self.betas[arm], 1.0 - reward, "beta")?;
317 let next_pulls = checked_increment(self.pulls[arm], "pulls")?;
318 let next_total = checked_finite_add(self.total_rewards[arm], reward, "total_rewards")?;
319 let next_samples = checked_increment(self.samples_seen, "samples_seen")?;
320
321 self.alphas[arm] = next_alpha;
322 self.betas[arm] = next_beta;
323 self.pulls[arm] = next_pulls;
324 self.total_rewards[arm] = next_total;
325 self.samples_seen = next_samples;
326 Ok(())
327 }
328
329 fn reset(&mut self) {
330 self.alphas.fill(self.config.alpha_prior);
331 self.betas.fill(self.config.beta_prior);
332 self.pulls.fill(0);
333 self.total_rewards.fill(0.0);
334 self.samples_seen = 0;
335 }
336
337 fn arm_stats(&self, arm: usize) -> Result<ArmStats, RillError> {
338 validate_arm(self.arm_count, arm)?;
339 ArmStats::new(self.pulls[arm], self.total_rewards[arm])
340 }
341}
342
343#[cfg(feature = "serde")]
344#[derive(serde::Deserialize)]
345struct ThompsonSamplingState {
346 arm_count: usize,
347 config: ThompsonConfig,
348 alphas: Vec<f64>,
349 betas: Vec<f64>,
350 pulls: Vec<u64>,
351 total_rewards: Vec<f64>,
352 samples_seen: u64,
353}
354
355#[cfg(feature = "serde")]
356impl<'de> serde::Deserialize<'de> for ThompsonSampling {
357 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
358 where
359 D: serde::Deserializer<'de>,
360 {
361 let state = ThompsonSamplingState::deserialize(deserializer)?;
362 let bandit = Self {
363 arm_count: state.arm_count,
364 config: state.config,
365 alphas: state.alphas,
366 betas: state.betas,
367 pulls: state.pulls,
368 total_rewards: state.total_rewards,
369 samples_seen: state.samples_seen,
370 };
371 bandit.validate().map_err(serde::de::Error::custom)?;
372 Ok(bandit)
373 }
374}
375
376#[cfg(feature = "serde")]
377impl ValidateState for ThompsonSampling {
378 fn validate_state(&self) -> Result<(), RillError> {
379 ThompsonSampling::validate(self)
380 }
381}
382
383#[cfg(test)]
384mod tests {
385 use super::*;
386 use rand::SeedableRng;
387 use rand_chacha::ChaCha8Rng;
388
389 fn make_bandit() -> ThompsonSampling {
390 ThompsonSampling::new(3, ThompsonConfig::default()).unwrap()
391 }
392
393 #[test]
394 fn rejects_zero_arm_count() {
395 let result = ThompsonSampling::new(0, ThompsonConfig::default());
396 assert!(matches!(result, Err(RillError::InvalidArmCount(0))));
397 }
398
399 #[test]
400 fn rejects_invalid_priors() {
401 for &bad in &[0.0, -1.0, f64::NAN, f64::INFINITY] {
402 let result = ThompsonSampling::new(
403 3,
404 ThompsonConfig {
405 alpha_prior: bad,
406 beta_prior: 1.0,
407 },
408 );
409 assert!(matches!(result, Err(RillError::InvalidParameter { .. })));
410
411 let result = ThompsonSampling::new(
412 3,
413 ThompsonConfig {
414 alpha_prior: 1.0,
415 beta_prior: bad,
416 },
417 );
418 assert!(matches!(result, Err(RillError::InvalidParameter { .. })));
419 }
420 }
421
422 #[test]
423 fn initial_state() {
424 let b = make_bandit();
425 assert_eq!(b.arm_count(), 3);
426 assert_eq!(b.samples_seen(), 0);
427 for &a in b.alphas() {
429 assert!((a - 1.0).abs() < 1e-12);
430 }
431 for &be in b.betas() {
432 assert!((be - 1.0).abs() < 1e-12);
433 }
434 }
435
436 #[test]
437 fn select_returns_valid_arm() {
438 let b = make_bandit();
439 let mut rng = ChaCha8Rng::seed_from_u64(42);
440 let arm = b.select(&mut rng).unwrap();
441 assert!(arm < 3);
442 }
443
444 #[test]
445 fn update_with_success_increases_alpha() {
446 let mut b = make_bandit();
447 b.update(0, 1.0).unwrap();
448 assert!((b.alphas()[0] - 2.0).abs() < 1e-12);
450 assert!((b.betas()[0] - 1.0).abs() < 1e-12);
451 }
452
453 #[test]
454 fn update_with_failure_increases_beta() {
455 let mut b = make_bandit();
456 b.update(0, 0.0).unwrap();
457 assert!((b.alphas()[0] - 1.0).abs() < 1e-12);
459 assert!((b.betas()[0] - 2.0).abs() < 1e-12);
460 }
461
462 #[test]
463 fn update_with_continuous_reward() {
464 let mut b = make_bandit();
465 b.update(0, 0.7).unwrap();
466 assert!((b.alphas()[0] - 1.7).abs() < 1e-12);
468 assert!((b.betas()[0] - 1.3).abs() < 1e-12);
469 }
470
471 #[test]
472 fn update_rejects_invalid_arm() {
473 let mut b = make_bandit();
474 assert!(b.update(3, 1.0).is_err());
475 }
476
477 #[test]
478 fn update_rejects_reward_out_of_range() {
479 let mut b = make_bandit();
480 assert!(b.update(0, 1.5).is_err());
481 assert!(b.update(0, -0.1).is_err());
482 assert!(b.update(0, f64::NAN).is_err());
483 }
484
485 #[test]
486 fn reset_clears_state() {
487 let mut b = make_bandit();
488 b.update(0, 1.0).unwrap();
489 b.update(1, 0.0).unwrap();
490 assert_eq!(b.samples_seen(), 2);
491
492 b.reset();
493 assert_eq!(b.samples_seen(), 0);
494 for &a in b.alphas() {
495 assert!((a - 1.0).abs() < 1e-12);
496 }
497 for &be in b.betas() {
498 assert!((be - 1.0).abs() < 1e-12);
499 }
500 for &p in b.pulls() {
501 assert_eq!(p, 0);
502 }
503 }
504
505 #[test]
506 fn arm_stats_after_updates() {
507 let mut b = make_bandit();
508 b.update(0, 1.0).unwrap();
509 b.update(0, 0.0).unwrap();
510 b.update(0, 1.0).unwrap();
511 let stats = b.arm_stats(0).unwrap();
512 assert_eq!(stats.pulls, 3);
513 assert!((stats.total_reward - 2.0).abs() < 1e-12);
514 }
515
516 #[test]
517 fn arm_stats_rejects_invalid_arm() {
518 let b = make_bandit();
519 assert!(b.arm_stats(5).is_err());
520 }
521
522 #[test]
523 fn finds_best_arm_in_simulation() {
524 let mut b = make_bandit();
525 let mut rng = ChaCha8Rng::seed_from_u64(42);
526
527 for _ in 0..1000 {
530 let arm = b.select(&mut rng).unwrap();
531 let p = match arm {
532 0 => 0.8,
533 1 => 0.2,
534 _ => 0.5,
535 };
536 let reward = if rng.gen_range(0.0..1.0) < p {
537 1.0
538 } else {
539 0.0
540 };
541 b.update(arm, reward).unwrap();
542 }
543
544 let stats0 = b.arm_stats(0).unwrap();
546 let stats1 = b.arm_stats(1).unwrap();
547 let stats2 = b.arm_stats(2).unwrap();
548 assert!(stats0.pulls > stats1.pulls);
549 assert!(stats0.pulls > stats2.pulls);
550 assert!(stats0.mean_reward > 0.6);
552 }
553
554 #[test]
555 fn sample_beta_returns_value_in_unit_interval() {
556 let mut rng = ChaCha8Rng::seed_from_u64(99);
557 for _ in 0..1000 {
558 let v = ThompsonSampling::sample_beta(&mut rng, 2.0, 5.0);
559 assert!((0.0..=1.0).contains(&v), "Beta sample {v} out of [0, 1]");
560 }
561 }
562
563 #[test]
564 fn sample_gamma_returns_positive_value() {
565 let mut rng = ChaCha8Rng::seed_from_u64(7);
566 for shape in &[0.5, 1.0, 2.0, 5.0, 10.0] {
567 for _ in 0..100 {
568 let v = ThompsonSampling::sample_gamma(&mut rng, *shape);
569 assert!(v > 0.0, "Gamma sample {v} not positive for shape {shape}");
570 }
571 }
572 }
573
574 #[test]
575 fn sample_gamma_mean_converges() {
576 let mut rng = ChaCha8Rng::seed_from_u64(42);
578 let shape = 5.0;
579 let n = 10000;
580 let mut sum = 0.0;
581 for _ in 0..n {
582 sum += ThompsonSampling::sample_gamma(&mut rng, shape);
583 }
584 let mean = sum / n as f64;
585 assert!(
587 (mean - shape).abs() / shape < 0.1,
588 "Gamma mean {mean} too far from {shape}"
589 );
590 }
591
592 #[test]
593 fn sample_beta_mean_converges() {
594 let mut rng = ChaCha8Rng::seed_from_u64(42);
596 let alpha = 2.0;
597 let beta = 5.0;
598 let n = 10000;
599 let mut sum = 0.0;
600 for _ in 0..n {
601 sum += ThompsonSampling::sample_beta(&mut rng, alpha, beta);
602 }
603 let mean = sum / n as f64;
604 let expected = alpha / (alpha + beta);
605 assert!(
607 (mean - expected).abs() / expected < 0.1,
608 "Beta mean {mean} too far from {expected}"
609 );
610 }
611
612 #[cfg(feature = "serde")]
613 #[test]
614 fn serde_roundtrip() {
615 let mut b = ThompsonSampling::new(
616 3,
617 ThompsonConfig {
618 alpha_prior: 2.0,
619 beta_prior: 3.0,
620 },
621 )
622 .unwrap();
623 b.update(0, 1.0).unwrap();
624 b.update(1, 0.0).unwrap();
625
626 let json = serde_json::to_string(&b).unwrap();
627 let restored: ThompsonSampling = serde_json::from_str(&json).unwrap();
628 assert_eq!(restored.arm_count(), b.arm_count());
629 assert_eq!(restored.samples_seen(), b.samples_seen());
630 assert_eq!(restored.alphas(), b.alphas());
631 assert_eq!(restored.betas(), b.betas());
632 }
633
634 #[cfg(feature = "serde")]
635 #[test]
636 fn serde_rejects_malformed_state() {
637 let json = r#"{
638 "arm_count": 2,
639 "config": {"alpha_prior": 1.0, "beta_prior": 1.0},
640 "alphas": [2.0],
641 "betas": [1.0],
642 "pulls": [1],
643 "total_rewards": [1.0],
644 "samples_seen": 1
645 }"#;
646 assert!(serde_json::from_str::<ThompsonSampling>(json).is_err());
647 }
648}