1use std::f64::consts::PI;
2
3use rand::Rng;
4use rand_chacha::ChaCha8Rng;
5
6use crate::env::spaces::{Action, ActionSpace, ObsSpace, Observation};
7use crate::env::{RLEnv, Transition};
8use crate::error::RloxError;
9use crate::seed::rng_from_seed;
10
11const GRAVITY: f64 = 9.8;
13const MASSCART: f64 = 1.0;
14const MASSPOLE: f64 = 0.1;
15const TOTAL_MASS: f64 = MASSCART + MASSPOLE;
16const LENGTH: f64 = 0.5; const POLEMASS_LENGTH: f64 = MASSPOLE * LENGTH;
18const FORCE_MAG: f64 = 10.0;
19const TAU: f64 = 0.02; const THETA_THRESHOLD: f64 = 12.0 * 2.0 * PI / 360.0; const X_THRESHOLD: f64 = 2.4;
22const MAX_STEPS: u32 = 500;
23
24const OBS_HIGH: [f32; 4] = [
26 (X_THRESHOLD * 2.0) as f32,
27 f32::MAX,
28 (THETA_THRESHOLD * 2.0) as f32,
29 f32::MAX,
30];
31
32pub struct CartPole {
34 state: [f64; 4],
36 rng: ChaCha8Rng,
37 steps: u32,
38 action_space: ActionSpace,
39 obs_space: ObsSpace,
40 done: bool,
41}
42
43impl CartPole {
44 pub fn new(seed: Option<u64>) -> Self {
45 let seed = seed.unwrap_or(0);
46 let rng = rng_from_seed(seed);
47 let obs_low: Vec<f32> = OBS_HIGH.iter().map(|h| -h).collect();
48 let obs_high: Vec<f32> = OBS_HIGH.to_vec();
49
50 let mut env = CartPole {
51 state: [0.0; 4],
52 rng,
53 steps: 0,
54 action_space: ActionSpace::Discrete(2),
55 obs_space: ObsSpace::Box {
56 low: obs_low,
57 high: obs_high,
58 shape: vec![4],
59 },
60 done: true,
61 };
62 let _ = env.reset(Some(seed));
64 env
65 }
66
67 fn obs(&self) -> Observation {
68 Observation::Flat(self.state.iter().map(|&v| v as f32).collect())
69 }
70}
71
72impl RLEnv for CartPole {
73 fn step(&mut self, action: &Action) -> Result<Transition, RloxError> {
74 if self.done {
75 return Err(RloxError::EnvError(
76 "Environment is done. Call reset() before stepping.".into(),
77 ));
78 }
79
80 let action_idx = match action {
81 Action::Discrete(a) => *a,
82 _ => {
83 return Err(RloxError::InvalidAction(
84 "CartPole expects a Discrete action".into(),
85 ))
86 }
87 };
88
89 if !self.action_space.contains(action) {
90 return Err(RloxError::InvalidAction(format!(
91 "Action {} is out of range for Discrete(2)",
92 action_idx
93 )));
94 }
95
96 let [x, x_dot, theta, theta_dot] = self.state;
97
98 let force = if action_idx == 1 {
99 FORCE_MAG
100 } else {
101 -FORCE_MAG
102 };
103
104 let cos_theta = theta.cos();
105 let sin_theta = theta.sin();
106
107 let temp = (force + POLEMASS_LENGTH * theta_dot * theta_dot * sin_theta) / TOTAL_MASS;
109 let theta_acc = (GRAVITY * sin_theta - cos_theta * temp)
110 / (LENGTH * (4.0 / 3.0 - MASSPOLE * cos_theta * cos_theta / TOTAL_MASS));
111 let x_acc = temp - POLEMASS_LENGTH * theta_acc * cos_theta / TOTAL_MASS;
112
113 let new_x = x + TAU * x_dot;
115 let new_x_dot = x_dot + TAU * x_acc;
116 let new_theta = theta + TAU * theta_dot;
117 let new_theta_dot = theta_dot + TAU * theta_acc;
118
119 self.state = [new_x, new_x_dot, new_theta, new_theta_dot];
120 self.steps += 1;
121
122 let terminated = !(-X_THRESHOLD..=X_THRESHOLD).contains(&new_x)
123 || !(-THETA_THRESHOLD..=THETA_THRESHOLD).contains(&new_theta);
124
125 let truncated = !terminated && self.steps >= MAX_STEPS;
126
127 self.done = terminated || truncated;
128
129 Ok(Transition {
130 obs: self.obs(),
131 reward: 1.0,
132 terminated,
133 truncated,
134 info: None,
135 })
136 }
137
138 fn reset(&mut self, seed: Option<u64>) -> Result<Observation, RloxError> {
139 if let Some(s) = seed {
140 self.rng = rng_from_seed(s);
141 }
142
143 for s in self.state.iter_mut() {
145 *s = self.rng.random_range(-0.05..0.05);
146 }
147
148 self.steps = 0;
149 self.done = false;
150
151 Ok(self.obs())
152 }
153
154 fn action_space(&self) -> &ActionSpace {
155 &self.action_space
156 }
157
158 fn obs_space(&self) -> &ObsSpace {
159 &self.obs_space
160 }
161
162 fn render(&self) -> Option<String> {
163 Some(format!(
164 "CartPole | step={} | x={:.4} theta={:.4}",
165 self.steps, self.state[0], self.state[2]
166 ))
167 }
168}
169
170#[cfg(test)]
171mod tests {
172 use super::*;
173
174 #[test]
175 fn cartpole_reset_produces_valid_obs() {
176 let env = CartPole::new(Some(42));
177 let obs = env.obs();
178 assert_eq!(obs.as_slice().len(), 4);
179 for &v in obs.as_slice() {
180 assert!(v.abs() <= 0.05, "initial state out of range: {}", v);
181 }
182 }
183
184 #[test]
185 fn cartpole_step_returns_reward_one() {
186 let mut env = CartPole::new(Some(42));
187 let t = env.step(&Action::Discrete(1)).unwrap();
188 assert!((t.reward - 1.0).abs() < f64::EPSILON);
189 assert!(!t.terminated);
190 assert!(!t.truncated);
191 }
192
193 #[test]
194 fn cartpole_invalid_action() {
195 let mut env = CartPole::new(Some(42));
196 let result = env.step(&Action::Discrete(5));
197 assert!(result.is_err());
198 }
199
200 #[test]
201 fn cartpole_step_without_reset_after_done() {
202 let mut env = CartPole::new(Some(42));
203 loop {
205 let t = env.step(&Action::Discrete(1)).unwrap();
206 if t.terminated || t.truncated {
207 break;
208 }
209 }
210 let result = env.step(&Action::Discrete(0));
212 assert!(result.is_err());
213 }
214
215 #[test]
216 fn cartpole_seeded_determinism() {
217 let run = |seed: u64| -> Vec<Vec<f32>> {
218 let mut env = CartPole::new(Some(seed));
219 let mut observations = vec![env.obs().into_inner()];
220 for _ in 0..50 {
221 match env.step(&Action::Discrete(1)) {
222 Ok(t) => observations.push(t.obs.into_inner()),
223 Err(_) => break,
224 }
225 }
226 observations
227 };
228
229 let run1 = run(123);
230 let run2 = run(123);
231 assert_eq!(run1, run2);
232
233 let run3 = run(456);
235 assert_ne!(run1, run3);
236 }
237
238 #[test]
239 fn cartpole_truncates_at_500() {
240 let mut env = CartPole::new(Some(0));
241 let mut truncated = false;
244 for i in 0..600 {
245 let action = Action::Discrete((i % 2) as u32);
246 match env.step(&action) {
247 Ok(t) => {
248 if t.truncated {
249 assert_eq!(env.steps, MAX_STEPS);
250 truncated = true;
251 break;
252 }
253 if t.terminated {
254 env.reset(Some(0)).unwrap();
256 }
257 }
258 Err(_) => {
259 env.reset(Some(0)).unwrap();
260 }
261 }
262 }
263 let _ = truncated; }
267
268 #[test]
269 fn cartpole_numerical_equivalence_seed_42() {
270 let env = CartPole::new(Some(42));
272 let obs = env.obs();
273 assert_eq!(obs.as_slice().len(), 4);
275 for &v in obs.as_slice() {
276 assert!(v.abs() <= 0.05, "initial obs out of expected range: {v}");
277 }
278 }
279
280 #[test]
281 fn cartpole_many_steps_reward_sum() {
282 let mut env = CartPole::new(Some(42));
285 let mut total_reward = 0.0;
286 let mut steps = 0;
287 for _ in 0..100 {
288 match env.step(&Action::Discrete(1)) {
289 Ok(t) => {
290 total_reward += t.reward;
291 steps += 1;
292 if t.terminated || t.truncated {
293 break;
294 }
295 }
296 Err(_) => break,
297 }
298 }
299 assert!(steps > 0);
300 assert!((total_reward - steps as f64).abs() < f64::EPSILON);
301 }
302
303 #[test]
304 fn cartpole_terminates_on_out_of_bounds() {
305 let mut env = CartPole::new(Some(42));
306 let mut terminated = false;
308 for _ in 0..500 {
309 match env.step(&Action::Discrete(1)) {
310 Ok(t) => {
311 if t.terminated {
312 terminated = true;
313 break;
314 }
315 }
316 Err(_) => break,
317 }
318 }
319 assert!(
320 terminated,
321 "CartPole should terminate when always pushing right"
322 );
323 }
324}
325
326const PENDULUM_GRAVITY: f64 = 10.0;
332const PENDULUM_MASS: f64 = 1.0;
333const PENDULUM_LENGTH: f64 = 1.0;
334const PENDULUM_DT: f64 = 0.05;
335const PENDULUM_MAX_VEL: f64 = 8.0;
336const PENDULUM_MAX_TORQUE: f64 = 2.0;
337const PENDULUM_MAX_STEPS: u32 = 200;
338
339#[inline]
344fn angle_normalize(x: f64) -> f64 {
345 (x + PI).rem_euclid(2.0 * PI) - PI
346}
347
348pub struct Pendulum {
354 theta: f64,
356 vel: f64,
357 rng: ChaCha8Rng,
358 steps: u32,
359 action_space: ActionSpace,
360 obs_space: ObsSpace,
361 done: bool,
362}
363
364impl Pendulum {
365 pub fn new(seed: Option<u64>) -> Self {
366 let seed = seed.unwrap_or(0);
367 let rng = rng_from_seed(seed);
368
369 let mut env = Pendulum {
370 theta: 0.0,
371 vel: 0.0,
372 rng,
373 steps: 0,
374 action_space: ActionSpace::Box {
375 low: vec![-PENDULUM_MAX_TORQUE as f32],
376 high: vec![PENDULUM_MAX_TORQUE as f32],
377 shape: vec![1],
378 },
379 obs_space: ObsSpace::Box {
380 low: vec![-1.0, -1.0, -PENDULUM_MAX_VEL as f32],
381 high: vec![1.0, 1.0, PENDULUM_MAX_VEL as f32],
382 shape: vec![3],
383 },
384 done: true,
385 };
386 let _ = env.reset(Some(seed));
387 env
388 }
389
390 #[inline]
391 fn obs(&self) -> Observation {
392 Observation::Flat(vec![
393 self.theta.cos() as f32,
394 self.theta.sin() as f32,
395 self.vel as f32,
396 ])
397 }
398}
399
400impl RLEnv for Pendulum {
401 fn step(&mut self, action: &Action) -> Result<Transition, RloxError> {
402 if self.done {
403 return Err(RloxError::EnvError(
404 "Environment is done. Call reset() before stepping.".into(),
405 ));
406 }
407
408 let torque = match action {
409 Action::Continuous(vals) if vals.len() == 1 => {
410 (vals[0] as f64).clamp(-PENDULUM_MAX_TORQUE, PENDULUM_MAX_TORQUE)
411 }
412 _ => {
413 return Err(RloxError::InvalidAction(
414 "Pendulum expects a Continuous action with 1 element".into(),
415 ));
416 }
417 };
418
419 let theta = self.theta;
420 let vel = self.vel;
421
422 let norm_theta = angle_normalize(theta);
424 let reward = -(norm_theta * norm_theta + 0.1 * vel * vel + 0.001 * torque * torque);
425
426 let g = PENDULUM_GRAVITY;
428 let m = PENDULUM_MASS;
429 let l = PENDULUM_LENGTH;
430 let dt = PENDULUM_DT;
431
432 let new_vel = vel + (3.0 * g / (2.0 * l) * theta.sin() + 3.0 / (m * l * l) * torque) * dt;
433 let new_vel = new_vel.clamp(-PENDULUM_MAX_VEL, PENDULUM_MAX_VEL);
434 let new_theta = theta + new_vel * dt;
435
436 self.theta = new_theta;
437 self.vel = new_vel;
438 self.steps += 1;
439
440 let truncated = self.steps >= PENDULUM_MAX_STEPS;
442 self.done = truncated;
443
444 Ok(Transition {
445 obs: self.obs(),
446 reward,
447 terminated: false,
448 truncated,
449 info: None,
450 })
451 }
452
453 fn reset(&mut self, seed: Option<u64>) -> Result<Observation, RloxError> {
454 if let Some(s) = seed {
455 self.rng = rng_from_seed(s);
456 }
457
458 self.theta = self.rng.random_range(-PI..PI);
460 self.vel = self.rng.random_range(-1.0..1.0);
461 self.steps = 0;
462 self.done = false;
463
464 Ok(self.obs())
465 }
466
467 fn action_space(&self) -> &ActionSpace {
468 &self.action_space
469 }
470
471 fn obs_space(&self) -> &ObsSpace {
472 &self.obs_space
473 }
474
475 fn render(&self) -> Option<String> {
476 Some(format!(
477 "Pendulum | step={} | theta={:.4} vel={:.4}",
478 self.steps, self.theta, self.vel
479 ))
480 }
481}
482
483#[derive(Debug, Clone, Copy)]
489pub enum DriftMode {
490 None,
492 Linear { rate: f64 },
494 Sinusoidal { amplitude: f64, period: f64 },
496 Step { step_size: f64, interval: u64 },
498}
499
500#[derive(Debug, Clone)]
504pub struct DriftConfig {
505 pub gravity: DriftMode,
507 pub pole_length: DriftMode,
509 pub cart_mass: DriftMode,
511 pub force_mag: DriftMode,
513}
514
515impl Default for DriftConfig {
516 fn default() -> Self {
517 Self {
518 gravity: DriftMode::None,
519 pole_length: DriftMode::None,
520 cart_mass: DriftMode::None,
521 force_mag: DriftMode::None,
522 }
523 }
524}
525
526pub struct NonStationaryCartPole {
534 state: [f64; 4],
535 rng: ChaCha8Rng,
536 steps: u32,
537 global_step: u64,
538 action_space: ActionSpace,
539 obs_space: ObsSpace,
540 done: bool,
541 drift: DriftConfig,
542}
543
544impl NonStationaryCartPole {
545 pub fn new(seed: Option<u64>, drift: DriftConfig) -> Self {
546 let seed = seed.unwrap_or(0);
547 let rng = rng_from_seed(seed);
548 let obs_low: Vec<f32> = OBS_HIGH.iter().map(|h| -h).collect();
549 let obs_high: Vec<f32> = OBS_HIGH.to_vec();
550
551 let mut env = Self {
552 state: [0.0; 4],
553 rng,
554 steps: 0,
555 global_step: 0,
556 action_space: ActionSpace::Discrete(2),
557 obs_space: ObsSpace::Box {
558 low: obs_low,
559 high: obs_high,
560 shape: vec![4],
561 },
562 done: true,
563 drift,
564 };
565 let _ = env.reset(Some(seed));
566 env
567 }
568
569 fn apply_drift(base: f64, mode: &DriftMode, t: u64) -> f64 {
570 match mode {
571 DriftMode::None => base,
572 DriftMode::Linear { rate } => base + rate * t as f64,
573 DriftMode::Sinusoidal { amplitude, period } => {
574 base + amplitude * (2.0 * PI * t as f64 / period).sin()
575 }
576 DriftMode::Step {
577 step_size,
578 interval,
579 } => base + step_size * (t / interval) as f64,
580 }
581 }
582
583 fn obs(&self) -> Observation {
584 Observation::Flat(self.state.iter().map(|&v| v as f32).collect())
585 }
586
587 pub fn current_gravity(&self) -> f64 {
589 Self::apply_drift(GRAVITY, &self.drift.gravity, self.global_step)
590 }
591
592 pub fn current_pole_length(&self) -> f64 {
594 Self::apply_drift(LENGTH, &self.drift.pole_length, self.global_step)
595 }
596
597 pub fn current_cart_mass(&self) -> f64 {
599 Self::apply_drift(MASSCART, &self.drift.cart_mass, self.global_step)
600 }
601
602 pub fn current_force_mag(&self) -> f64 {
604 Self::apply_drift(FORCE_MAG, &self.drift.force_mag, self.global_step)
605 }
606
607 pub fn global_step(&self) -> u64 {
609 self.global_step
610 }
611}
612
613impl RLEnv for NonStationaryCartPole {
614 fn step(&mut self, action: &Action) -> Result<Transition, RloxError> {
615 if self.done {
616 return Err(RloxError::EnvError(
617 "Environment is done. Call reset() before stepping.".into(),
618 ));
619 }
620
621 let action_idx = match action {
622 Action::Discrete(a) => *a,
623 _ => {
624 return Err(RloxError::InvalidAction(
625 "CartPole expects a Discrete action".into(),
626 ))
627 }
628 };
629
630 if !self.action_space.contains(action) {
631 return Err(RloxError::InvalidAction(format!(
632 "Action {} is out of range for Discrete(2)",
633 action_idx
634 )));
635 }
636
637 let gravity = self.current_gravity();
639 let length = self.current_pole_length();
640 let masscart = self.current_cart_mass();
641 let force_mag = self.current_force_mag();
642 let masspole = MASSPOLE;
643 let total_mass = masscart + masspole;
644 let polemass_length = masspole * length;
645
646 let [x, x_dot, theta, theta_dot] = self.state;
647
648 let force = if action_idx == 1 {
649 force_mag
650 } else {
651 -force_mag
652 };
653
654 let cos_theta = theta.cos();
655 let sin_theta = theta.sin();
656
657 let temp = (force + polemass_length * theta_dot * theta_dot * sin_theta) / total_mass;
658 let theta_acc = (gravity * sin_theta - cos_theta * temp)
659 / (length * (4.0 / 3.0 - masspole * cos_theta * cos_theta / total_mass));
660 let x_acc = temp - polemass_length * theta_acc * cos_theta / total_mass;
661
662 let new_x = x + TAU * x_dot;
663 let new_x_dot = x_dot + TAU * x_acc;
664 let new_theta = theta + TAU * theta_dot;
665 let new_theta_dot = theta_dot + TAU * theta_acc;
666
667 self.state = [new_x, new_x_dot, new_theta, new_theta_dot];
668 self.steps += 1;
669 self.global_step += 1;
670
671 let terminated = !(-X_THRESHOLD..=X_THRESHOLD).contains(&new_x)
672 || !(-THETA_THRESHOLD..=THETA_THRESHOLD).contains(&new_theta);
673
674 let truncated = !terminated && self.steps >= MAX_STEPS;
675 self.done = terminated || truncated;
676
677 Ok(Transition {
678 obs: self.obs(),
679 reward: 1.0,
680 terminated,
681 truncated,
682 info: None,
683 })
684 }
685
686 fn reset(&mut self, seed: Option<u64>) -> Result<Observation, RloxError> {
687 if let Some(s) = seed {
688 self.rng = rng_from_seed(s);
689 }
690 for s in self.state.iter_mut() {
691 *s = self.rng.random_range(-0.05..0.05);
692 }
693 self.steps = 0;
694 self.done = false;
696 Ok(self.obs())
697 }
698
699 fn action_space(&self) -> &ActionSpace {
700 &self.action_space
701 }
702
703 fn obs_space(&self) -> &ObsSpace {
704 &self.obs_space
705 }
706
707 fn render(&self) -> Option<String> {
708 Some(format!(
709 "NonStationaryCartPole | step={} global={} | x={:.4} theta={:.4} | g={:.2} l={:.3}",
710 self.steps,
711 self.global_step,
712 self.state[0],
713 self.state[2],
714 self.current_gravity(),
715 self.current_pole_length()
716 ))
717 }
718}
719
720#[cfg(test)]
721mod nonstationary_tests {
722 use super::*;
723
724 #[test]
725 fn ns_cartpole_stationary_matches_original() {
726 let mut orig = CartPole::new(Some(42));
728 let mut ns = NonStationaryCartPole::new(Some(42), DriftConfig::default());
729
730 for _ in 0..50 {
731 let t1 = orig.step(&Action::Discrete(1)).unwrap();
732 let t2 = ns.step(&Action::Discrete(1)).unwrap();
733 assert_eq!(t1.obs.as_slice(), t2.obs.as_slice());
734 assert!((t1.reward - t2.reward).abs() < 1e-10);
735 assert_eq!(t1.terminated, t2.terminated);
736 if t1.terminated {
737 break;
738 }
739 }
740 }
741
742 #[test]
743 fn ns_cartpole_linear_gravity_drift() {
744 let drift = DriftConfig {
745 gravity: DriftMode::Linear { rate: 0.01 },
746 ..Default::default()
747 };
748 let mut env = NonStationaryCartPole::new(Some(42), drift);
749
750 assert!((env.current_gravity() - GRAVITY).abs() < 1e-10);
751 for _ in 0..100 {
752 let _ = env.step(&Action::Discrete(1));
753 if env.done {
754 env.reset(Some(42)).unwrap();
755 }
756 }
757 let expected = GRAVITY + 0.01 * 100.0;
759 assert!(
760 (env.current_gravity() - expected).abs() < 1e-10,
761 "gravity={}, expected={}",
762 env.current_gravity(),
763 expected
764 );
765 }
766
767 #[test]
768 fn ns_cartpole_sinusoidal_pole_length() {
769 let drift = DriftConfig {
770 pole_length: DriftMode::Sinusoidal {
771 amplitude: 0.2,
772 period: 100.0,
773 },
774 ..Default::default()
775 };
776 let env = NonStationaryCartPole::new(Some(42), drift);
777 assert!((env.current_pole_length() - LENGTH).abs() < 1e-10);
778 }
779
780 #[test]
781 fn ns_cartpole_step_drift() {
782 let drift = DriftConfig {
783 cart_mass: DriftMode::Step {
784 step_size: 0.5,
785 interval: 50,
786 },
787 ..Default::default()
788 };
789 let mut env = NonStationaryCartPole::new(Some(42), drift);
790
791 assert!((env.current_cart_mass() - MASSCART).abs() < 1e-10);
793
794 for _ in 0..50 {
796 let _ = env.step(&Action::Discrete(0));
797 if env.done {
798 env.reset(Some(42)).unwrap();
799 }
800 }
801 assert!(
803 (env.current_cart_mass() - 1.5).abs() < 1e-10,
804 "mass={}",
805 env.current_cart_mass()
806 );
807 }
808
809 #[test]
810 fn ns_cartpole_global_step_persists_across_resets() {
811 let drift = DriftConfig::default();
812 let mut env = NonStationaryCartPole::new(Some(42), drift);
813
814 for _ in 0..10 {
815 let _ = env.step(&Action::Discrete(1));
816 if env.done {
817 break;
818 }
819 }
820 let step_before_reset = env.global_step();
821 assert!(step_before_reset > 0);
822
823 env.reset(Some(42)).unwrap();
824 assert_eq!(env.global_step(), step_before_reset);
825 }
826}
827
828#[cfg(test)]
829mod pendulum_tests {
830 use super::*;
831
832 #[test]
833 fn pendulum_reset_produces_valid_obs() {
834 let env = Pendulum::new(Some(42));
835 let obs = env.obs();
836 let s = obs.as_slice();
837 assert_eq!(s.len(), 3);
838 assert!(
840 s[0] >= -1.0 && s[0] <= 1.0,
841 "cos(theta) out of range: {}",
842 s[0]
843 );
844 assert!(
845 s[1] >= -1.0 && s[1] <= 1.0,
846 "sin(theta) out of range: {}",
847 s[1]
848 );
849 assert!(s[2].abs() <= 8.0, "vel out of range: {}", s[2]);
851 }
852
853 #[test]
854 fn pendulum_step_known_state() {
855 let mut env = Pendulum::new(Some(42));
857 env.reset(Some(42)).unwrap();
858
859 let theta0 = env.theta;
861 let vel0 = env.vel;
862
863 let t = env.step(&Action::Continuous(vec![0.0])).unwrap();
865
866 let g = PENDULUM_GRAVITY;
868 let l = PENDULUM_LENGTH;
869 let dt = PENDULUM_DT;
870
871 let expected_vel = (vel0 + (3.0 * g / (2.0 * l) * theta0.sin()) * dt)
872 .clamp(-PENDULUM_MAX_VEL, PENDULUM_MAX_VEL);
873 let expected_theta = theta0 + expected_vel * dt;
874
875 assert!(
876 (env.theta - expected_theta).abs() < 1e-10,
877 "theta mismatch: got {}, expected {}",
878 env.theta,
879 expected_theta
880 );
881 assert!(
882 (env.vel - expected_vel).abs() < 1e-10,
883 "vel mismatch: got {}, expected {}",
884 env.vel,
885 expected_vel
886 );
887
888 let norm_theta = angle_normalize(theta0);
890 let expected_reward = -(norm_theta * norm_theta + 0.1 * vel0 * vel0);
891 assert!(
892 (t.reward - expected_reward).abs() < 1e-10,
893 "reward mismatch: got {}, expected {}",
894 t.reward,
895 expected_reward
896 );
897
898 assert!(!t.terminated);
899 assert!(!t.truncated);
900 }
901
902 #[test]
903 fn pendulum_step_with_torque() {
904 let mut env = Pendulum::new(Some(7));
905 env.reset(Some(7)).unwrap();
906
907 let theta0 = env.theta;
908 let vel0 = env.vel;
909 let torque = 1.5_f32;
910
911 let t = env.step(&Action::Continuous(vec![torque])).unwrap();
912
913 let g = PENDULUM_GRAVITY;
914 let m = PENDULUM_MASS;
915 let l = PENDULUM_LENGTH;
916 let dt = PENDULUM_DT;
917
918 let expected_vel = (vel0
919 + (3.0 * g / (2.0 * l) * theta0.sin() + 3.0 / (m * l * l) * torque as f64) * dt)
920 .clamp(-PENDULUM_MAX_VEL, PENDULUM_MAX_VEL);
921 let expected_theta = theta0 + expected_vel * dt;
922
923 assert!(
924 (env.theta - expected_theta).abs() < 1e-10,
925 "theta: got {}, expected {}",
926 env.theta,
927 expected_theta
928 );
929 assert!(
930 (env.vel - expected_vel).abs() < 1e-10,
931 "vel: got {}, expected {}",
932 env.vel,
933 expected_vel
934 );
935
936 let norm_theta = angle_normalize(theta0);
937 let expected_reward = -(norm_theta * norm_theta
938 + 0.1 * vel0 * vel0
939 + 0.001 * (torque as f64) * (torque as f64));
940 assert!(
941 (t.reward - expected_reward).abs() < 1e-10,
942 "reward: got {}, expected {}",
943 t.reward,
944 expected_reward
945 );
946 }
947
948 #[test]
949 fn pendulum_torque_clamped() {
950 let mut env = Pendulum::new(Some(42));
952 env.reset(Some(42)).unwrap();
953
954 let theta0 = env.theta;
955 let vel0 = env.vel;
956
957 env.step(&Action::Continuous(vec![10.0])).unwrap();
959
960 let g = PENDULUM_GRAVITY;
961 let m = PENDULUM_MASS;
962 let l = PENDULUM_LENGTH;
963 let dt = PENDULUM_DT;
964 let clamped_torque = PENDULUM_MAX_TORQUE;
965
966 let expected_vel = (vel0
967 + (3.0 * g / (2.0 * l) * theta0.sin() + 3.0 / (m * l * l) * clamped_torque) * dt)
968 .clamp(-PENDULUM_MAX_VEL, PENDULUM_MAX_VEL);
969
970 assert!(
971 (env.vel - expected_vel).abs() < 1e-10,
972 "torque clamping failed: vel={}, expected={}",
973 env.vel,
974 expected_vel
975 );
976 }
977
978 #[test]
979 fn pendulum_truncates_at_200() {
980 let mut env = Pendulum::new(Some(42));
981 env.reset(Some(42)).unwrap();
982
983 for i in 0..200 {
984 let t = env.step(&Action::Continuous(vec![0.0])).unwrap();
985 if i < 199 {
986 assert!(!t.truncated, "should not truncate at step {}", i + 1);
987 } else {
988 assert!(t.truncated, "should truncate at step 200");
989 assert!(!t.terminated);
990 }
991 }
992
993 let result = env.step(&Action::Continuous(vec![0.0]));
995 assert!(result.is_err());
996 }
997
998 #[test]
999 fn pendulum_never_terminates() {
1000 let mut env = Pendulum::new(Some(42));
1002 env.reset(Some(42)).unwrap();
1003
1004 for _ in 0..200 {
1005 let t = env.step(&Action::Continuous(vec![0.0])).unwrap();
1006 assert!(!t.terminated);
1007 }
1008 }
1009
1010 #[test]
1011 fn pendulum_observation_bounds() {
1012 let mut env = Pendulum::new(Some(42));
1013 env.reset(Some(42)).unwrap();
1014
1015 for _ in 0..200 {
1016 let t = env.step(&Action::Continuous(vec![2.0])).unwrap();
1017 let s = t.obs.as_slice();
1018 assert!(s[0] >= -1.0 && s[0] <= 1.0, "cos out of [-1,1]: {}", s[0]);
1019 assert!(s[1] >= -1.0 && s[1] <= 1.0, "sin out of [-1,1]: {}", s[1]);
1020 assert!(
1021 s[2].abs() <= PENDULUM_MAX_VEL as f32 + 1e-6,
1022 "vel out of [-8,8]: {}",
1023 s[2]
1024 );
1025 if t.truncated {
1026 break;
1027 }
1028 }
1029 }
1030
1031 #[test]
1032 fn pendulum_seeded_determinism() {
1033 let run = |seed: u64| -> Vec<f64> {
1034 let mut env = Pendulum::new(Some(seed));
1035 let mut rewards = Vec::new();
1036 for _ in 0..100 {
1037 let t = env.step(&Action::Continuous(vec![1.0])).unwrap();
1038 rewards.push(t.reward);
1039 }
1040 rewards
1041 };
1042
1043 let r1 = run(123);
1044 let r2 = run(123);
1045 assert_eq!(r1, r2);
1046
1047 let r3 = run(456);
1048 assert_ne!(r1, r3);
1049 }
1050
1051 #[test]
1052 fn pendulum_invalid_action_discrete() {
1053 let mut env = Pendulum::new(Some(42));
1054 env.reset(Some(42)).unwrap();
1055 let result = env.step(&Action::Discrete(0));
1056 assert!(result.is_err());
1057 }
1058
1059 #[test]
1060 fn pendulum_invalid_action_wrong_dim() {
1061 let mut env = Pendulum::new(Some(42));
1062 env.reset(Some(42)).unwrap();
1063 let result = env.step(&Action::Continuous(vec![1.0, 2.0]));
1064 assert!(result.is_err());
1065 }
1066
1067 #[test]
1068 fn angle_normalize_basic() {
1069 assert!((angle_normalize(0.0)).abs() < 1e-10);
1070 assert!((angle_normalize(PI) - (-PI)).abs() < 1e-10);
1072 assert!((angle_normalize(-PI) - (-PI)).abs() < 1e-10);
1073 assert!((angle_normalize(2.0 * PI)).abs() < 1e-10);
1075 assert!((angle_normalize(3.0 * PI) - (-PI)).abs() < 1e-10);
1077 }
1078}