// ml/tests/action_loader_test.rs // Unit tests for DQN action loader use ml::backtesting::{load_actions_from_csv, DQNActionRecord}; use std::io::Write; #[test] fn test_load_valid_csv() { // Test 1: Load valid CSV file with 5 actions let mut tmpfile = tempfile::NamedTempFile::new().unwrap(); writeln!( tmpfile, "timestamp,action,q_buy,q_sell,q_hold,open,high,low,close,volume" ) .unwrap(); writeln!(tmpfile, "2024-10-20T23:31:00.000000000Z,2,-658.8440,355.0268,538.5875,5914.50,5914.75,5914.25,5914.25,27").unwrap(); writeln!(tmpfile, "2024-10-20T23:32:00.000000000Z,2,-654.3466,355.2580,546.9955,5914.25,5914.25,5914.00,5914.00,41").unwrap(); writeln!(tmpfile, "2024-10-20T23:33:00.000000000Z,0,-657.4612,356.3587,541.4503,5914.00,5914.25,5914.00,5914.25,44").unwrap(); writeln!(tmpfile, "2024-10-20T23:34:00.000000000Z,1,-657.6093,356.5374,541.2177,5914.25,5914.75,5914.25,5914.50,36").unwrap(); writeln!(tmpfile, "2024-10-20T23:35:00.000000000Z,2,-659.1310,356.1713,539.7350,5914.50,5914.50,5914.50,5914.50,9").unwrap(); tmpfile.flush().unwrap(); let actions = load_actions_from_csv(tmpfile.path()).unwrap(); // Verify count assert_eq!(actions.len(), 5, "Expected 5 actions"); // Verify first action assert_eq!(actions[0].action, 2, "First action should be 2 (Hold)"); assert_eq!(actions[0].q_buy, -658.8440); assert_eq!(actions[0].q_sell, 355.0268); assert_eq!(actions[0].q_hold, 538.5875); assert_eq!(actions[0].volume, 27); // Verify action variety (Buy, Sell, Hold) assert_eq!(actions[2].action, 0, "Third action should be 0 (Buy)"); assert_eq!(actions[3].action, 1, "Fourth action should be 1 (Sell)"); // Verify all Q-values are finite for (i, action) in actions.iter().enumerate() { assert!( action.q_buy.is_finite(), "q_buy at index {} must be finite", i ); assert!( action.q_sell.is_finite(), "q_sell at index {} must be finite", i ); assert!( action.q_hold.is_finite(), "q_hold at index {} must be finite", i ); } } #[test] fn test_invalid_action_bounds() { // Test 2: Reject invalid action (action > 2) let mut tmpfile = tempfile::NamedTempFile::new().unwrap(); writeln!( tmpfile, "timestamp,action,q_buy,q_sell,q_hold,open,high,low,close,volume" ) .unwrap(); writeln!(tmpfile, "2024-10-20T23:31:00.000000000Z,3,-658.8440,355.0268,538.5875,5914.50,5914.75,5914.25,5914.25,27").unwrap(); tmpfile.flush().unwrap(); let result = load_actions_from_csv(tmpfile.path()); assert!(result.is_err(), "Should reject action > 2"); let err = result.unwrap_err(); assert!( err.contains("Invalid action 3"), "Error should mention invalid action 3" ); assert!( err.contains("action must be 0 (Buy), 1 (Sell), or 2 (Hold)"), "Error should explain valid actions" ); assert!(err.contains("row 2"), "Error should mention row number"); } #[test] fn test_validation_errors() { // Test 3: Comprehensive validation tests (NaN Q-values, timestamp ordering) // Test 3a: NaN Q-value let mut tmpfile = tempfile::NamedTempFile::new().unwrap(); writeln!( tmpfile, "timestamp,action,q_buy,q_sell,q_hold,open,high,low,close,volume" ) .unwrap(); writeln!( tmpfile, "2024-10-20T23:31:00.000000000Z,2,NaN,355.0268,538.5875,5914.50,5914.75,5914.25,5914.25,27" ) .unwrap(); tmpfile.flush().unwrap(); let result = load_actions_from_csv(tmpfile.path()); assert!(result.is_err(), "Should reject NaN Q-value"); let err = result.unwrap_err(); assert!(err.contains("Invalid q_buy"), "Error should mention q_buy"); assert!( err.contains("Q-value must be finite"), "Error should mention finite requirement" ); // Test 3b: Inf Q-value let mut tmpfile = tempfile::NamedTempFile::new().unwrap(); writeln!( tmpfile, "timestamp,action,q_buy,q_sell,q_hold,open,high,low,close,volume" ) .unwrap(); writeln!(tmpfile, "2024-10-20T23:31:00.000000000Z,2,-658.8440,inf,538.5875,5914.50,5914.75,5914.25,5914.25,27").unwrap(); tmpfile.flush().unwrap(); let result = load_actions_from_csv(tmpfile.path()); assert!(result.is_err(), "Should reject Inf Q-value"); let err = result.unwrap_err(); assert!( err.contains("Invalid q_sell"), "Error should mention q_sell" ); assert!( err.contains("Q-value must be finite"), "Error should mention finite requirement" ); // Test 3c: Timestamp ordering violation let mut tmpfile = tempfile::NamedTempFile::new().unwrap(); writeln!( tmpfile, "timestamp,action,q_buy,q_sell,q_hold,open,high,low,close,volume" ) .unwrap(); writeln!(tmpfile, "2024-10-20T23:35:00.000000000Z,2,-659.1310,356.1713,539.7350,5914.50,5914.50,5914.50,5914.50,9").unwrap(); writeln!(tmpfile, "2024-10-20T23:31:00.000000000Z,2,-658.8440,355.0268,538.5875,5914.50,5914.75,5914.25,5914.25,27").unwrap(); tmpfile.flush().unwrap(); let result = load_actions_from_csv(tmpfile.path()); assert!( result.is_err(), "Should reject timestamp ordering violation" ); let err = result.unwrap_err(); assert!( err.contains("Timestamp ordering violation"), "Error should mention timestamp ordering" ); assert!(err.contains("row 3"), "Error should mention row number"); // Test 3d: Empty CSV (no data rows) let mut tmpfile = tempfile::NamedTempFile::new().unwrap(); writeln!( tmpfile, "timestamp,action,q_buy,q_sell,q_hold,open,high,low,close,volume" ) .unwrap(); tmpfile.flush().unwrap(); let result = load_actions_from_csv(tmpfile.path()); assert!(result.is_err(), "Should reject empty CSV"); let err = result.unwrap_err(); assert!( err.contains("contains no data rows"), "Error should mention no data rows" ); }