use mem_core::{Trajectory, CorpusStats}; /// M5.3 Integration Tests — Training Corpus Export /// /// Verifies export to verl format: /// - Trajectories group turns by run /// - Rewards computed correctly /// - Prompts are exact byte recordings /// - Format strict (any unparsed turn = r_format 0) #[test] fn a1_trajectory_grouping() { let mut traj = Trajectory::new("run_001".to_string()); traj.add_turn(1, "Q1".to_string(), "A1".to_string(), 1, true); traj.add_turn(2, "Q2".to_string(), "A2".to_string(), -1, true); traj.add_turn(3, "Q3".to_string(), "A3".to_string(), 1, true); assert_eq!(traj.turns.len(), 3); assert_eq!(traj.turns[0].t, 1); assert_eq!(traj.turns[1].t, 2); assert_eq!(traj.turns[2].t, 3); // All turns in ascending order for i in 1..traj.turns.len() { assert!(traj.turns[i].t > traj.turns[i - 1].t); } } #[test] fn a2_r_update_signs() { let mut traj = Trajectory::new("run_001".to_string()); // Correct prediction traj.add_turn(1, "Q".to_string(), "A".to_string(), 1, true); assert_eq!(traj.turns[0].r_update, 1); // Incorrect prediction traj.add_turn(2, "Q".to_string(), "B".to_string(), -1, true); assert_eq!(traj.turns[1].r_update, -1); } #[test] fn a3_r_format_strict() { let mut traj = Trajectory::new("run_001".to_string()); // First turn OK traj.add_turn(1, "Q".to_string(), "A".to_string(), 1, true); assert_eq!(traj.r_format, 1.0); // Second turn unparsed traj.add_turn(2, "Q".to_string(), "malformed_response".to_string(), -1, false); // ENTIRE trajectory marked as unparsed assert_eq!(traj.r_format, 0.0, "Any unparsed turn makes whole trajectory unparsed"); } #[test] fn a4_r_exit_distribution() { let mut perfect = Trajectory::new("run_001".to_string()); perfect.set_exit_reward(5, 5); assert_eq!(perfect.r_exit, 0.0); let mut early = Trajectory::new("run_002".to_string()); early.set_exit_reward(3, 5); assert_eq!(early.r_exit, -0.75); let mut late = Trajectory::new("run_003".to_string()); late.set_exit_reward(7, 5); assert_eq!(late.r_exit, -0.5); } #[test] fn a5_prompt_exact_bytes() { let prompt_bytes = "Question: what is X?\n\nContext: Y is Z".to_string(); let response_bytes = "Answer: X is Y".to_string(); let mut traj = Trajectory::new("run_001".to_string()); traj.add_turn(1, prompt_bytes.clone(), response_bytes.clone(), 1, true); // Prompts should be exact recordings, not re-assembled assert_eq!(traj.turns[0].prompt, prompt_bytes); assert_eq!(traj.turns[0].response, response_bytes); } #[test] fn a6_corpus_stats_aggregation() { let mut traj1 = Trajectory::new("run_001".to_string()); traj1.add_turn(1, "Q".to_string(), "A".to_string(), 1, true); traj1.add_turn(2, "Q".to_string(), "A".to_string(), -1, true); let mut traj2 = Trajectory::new("run_002".to_string()); traj2.add_turn(1, "Q".to_string(), "A".to_string(), 1, true); let mut traj3 = Trajectory::new("run_003".to_string()); traj3.add_turn(1, "Q".to_string(), "A".to_string(), 1, false); // Unparsed let stats = CorpusStats::from_trajectories(&[traj1, traj2, traj3]); assert_eq!(stats.total_trajectories, 3); assert_eq!(stats.total_turns, 4); assert_eq!(stats.positive_r_update, 3); assert_eq!(stats.negative_r_update, 1); // 2 out of 3 have all parsed turns assert!((stats.r_format_pass_rate - 2.0/3.0).abs() < 0.01); } #[test] fn a7_r_outcome_null() { let traj = Trajectory::new("run_001".to_string()); // Should have no answer correctness signal assert!(traj.r_outcome.is_none(), "r_outcome should be null (no answer-correctness signal)"); } #[test] fn a8_trajectory_ordering() { let mut traj = Trajectory::new("run_001".to_string()); // Add in order for t in 1..=10 { traj.add_turn(t, format!("Q{}", t), format!("A{}", t), if t % 2 == 0 { 1 } else { -1 }, true); } // Check order preserved for (i, turn) in traj.turns.iter().enumerate() { assert_eq!(turn.t, i + 1); } } #[test] fn a9_multiple_trajectories() { let mut trajs = Vec::new(); for run_id in 0..5 { let mut traj = Trajectory::new(format!("run_{:03}", run_id)); for t in 1..=3 { traj.add_turn(t, format!("Q{}", t), format!("A{}", t), 1, true); } trajs.push(traj); } assert_eq!(trajs.len(), 5); assert_eq!(trajs[0].turns.len(), 3); assert_eq!(trajs[0].trajectory_id, "run_000"); assert_eq!(trajs[4].trajectory_id, "run_004"); } #[test] fn a10_trajectory_serde_roundtrip() { let mut traj = Trajectory::new("run_001".to_string()); traj.add_turn(1, "prompt1".to_string(), "response1".to_string(), 1, true); traj.add_turn(2, "prompt2".to_string(), "response2".to_string(), -1, true); traj.set_exit_reward(2, 1); let json = serde_json::to_string(&traj).expect("Should serialize"); let restored: Trajectory = serde_json::from_str(&json) .expect("Should deserialize"); assert_eq!(restored.trajectory_id, "run_001"); assert_eq!(restored.turns.len(), 2); assert_eq!(restored.r_exit, -0.5); } #[test] fn a11_corpus_stats_structure() { let traj = Trajectory::new("run_001".to_string()); let stats = CorpusStats::from_trajectories(&[traj]); // All fields present assert!(stats.total_trajectories >= 0); assert!(stats.total_turns >= 0); assert!(stats.positive_r_update >= 0); assert!(stats.negative_r_update >= 0); assert!(stats.r_format_pass_rate >= 0.0 && stats.r_format_pass_rate <= 1.0); assert!(!stats.r_exit_distribution.is_empty()); } #[test] fn a12_mixed_exit_rewards() { let mut trajs = Vec::new(); let mut perfect = Trajectory::new("perfect".to_string()); perfect.set_exit_reward(3, 3); trajs.push(perfect); let mut early = Trajectory::new("early".to_string()); early.set_exit_reward(1, 3); trajs.push(early); let mut late = Trajectory::new("late".to_string()); late.set_exit_reward(5, 3); trajs.push(late); let stats = CorpusStats::from_trajectories(&trajs); assert_eq!(stats.total_trajectories, 3); assert_eq!(stats.r_exit_distribution.len(), 3); // Should have all 3 types }