use rtx_rl::{Environment, Step}; #[derive(Clone, Debug)] struct TestEnv { state: Vec, step_count: usize, max_steps: usize, } impl Environment for TestEnv { type State = Vec; type Action = Vec; type Reward = f32; fn reset(&mut self) -> Self::State { self.state = vec![0.0, 0.0]; self.step_count = 0; self.state.clone() } fn step(&mut self, action: &Self::Action) -> Step { self.state[0] += action[0]; self.state[1] += action[1]; self.step_count += 1; let reward = -(self.state[0].powi(2) + self.state[1].powi(2)); let done = self.step_count >= self.max_steps; Step { state: self.state.clone(), reward, done, info: std::collections::HashMap::new(), } } fn action_space(&self) -> (Vec, Vec) { (vec![-1.0, -1.0], vec![1.0, 1.0]) } fn observation_space(&self) -> (Vec, Vec) { (vec![-10.0, -10.0], vec![10.0, 10.0]) } } #[tokio::test] async fn test_environment_reset() { let mut env = TestEnv { state: vec![5.0, 5.0], step_count: 10, max_steps: 100, }; let state = env.reset(); assert_eq!(state, vec![0.0, 0.0]); assert_eq!(env.step_count, 0); } #[tokio::test] async fn test_environment_step() { let mut env = TestEnv { state: vec![0.0, 0.0], step_count: 0, max_steps: 100, }; let action = vec![0.5, -0.3]; let step_result = env.step(&action); assert_eq!(step_result.state, vec![0.5, -0.3]); assert!(step_result.reward < 0.0); // Negative quadratic reward assert!(!step_result.done); assert_eq!(env.step_count, 1); } #[tokio::test] async fn test_environment_termination() { let mut env = TestEnv { state: vec![0.0, 0.0], step_count: 99, max_steps: 100, }; let action = vec![0.1, 0.1]; let step_result = env.step(&action); assert!(step_result.done); assert_eq!(env.step_count, 100); } #[tokio::test] async fn test_environment_spaces() { let env = TestEnv { state: vec![0.0, 0.0], step_count: 0, max_steps: 100, }; let (action_low, action_high) = env.action_space(); assert_eq!(action_low, vec![-1.0, -1.0]); assert_eq!(action_high, vec![1.0, 1.0]); let (obs_low, obs_high) = env.observation_space(); assert_eq!(obs_low, vec![-10.0, -10.0]); assert_eq!(obs_high, vec![10.0, 10.0]); }