//! Error types for PINN benchmark operations use thiserror::Error; /// Errors that can occur during PINN operations #[derive(Debug, Error)] pub enum PINNError { /// Network architecture error #[error("Network error: {0}")] Network(String), /// Training error #[error("Training error: {0}")] Training(String), /// Benchmark error #[error("Benchmark error: {0}")] Benchmark(String), /// Invalid configuration #[error("Invalid configuration: {0}")] InvalidConfig(String), /// Numerical computation error #[error("Numerical error: {0}")] Numerical(String), /// I/O error #[error("I/O error: {0}")] Io(#[from] std::io::Error), /// Serialization error #[error("Serialization error: {0}")] Serialization(#[from] serde_json::Error), } /// Result type for PINN operations pub type Result = std::result::Result; impl PINNError { /// Creates a network error pub fn network>(msg: S) -> Self { Self::Network(msg.into()) } /// Creates a training error pub fn training>(msg: S) -> Self { Self::Training(msg.into()) } /// Creates a benchmark error pub fn benchmark>(msg: S) -> Self { Self::Benchmark(msg.into()) } /// Creates an invalid configuration error pub fn invalid_config>(msg: S) -> Self { Self::InvalidConfig(msg.into()) } /// Creates a numerical error pub fn numerical>(msg: S) -> Self { Self::Numerical(msg.into()) } } #[cfg(test)] mod tests { use super::*; #[test] fn test_network_error_creation() { let err = PINNError::network("invalid layer size"); assert!(matches!(err, PINNError::Network(_))); assert_eq!(err.to_string(), "Network error: invalid layer size"); } #[test] fn test_training_error_creation() { let err = PINNError::training("loss diverged"); assert!(matches!(err, PINNError::Training(_))); assert_eq!(err.to_string(), "Training error: loss diverged"); } #[test] fn test_benchmark_error_creation() { let err = PINNError::benchmark("invalid problem type"); assert!(matches!(err, PINNError::Benchmark(_))); assert_eq!(err.to_string(), "Benchmark error: invalid problem type"); } #[test] fn test_invalid_config_error_creation() { let err = PINNError::invalid_config("learning rate must be positive"); assert!(matches!(err, PINNError::InvalidConfig(_))); assert!(err.to_string().contains("learning rate")); } #[test] fn test_numerical_error_creation() { let err = PINNError::numerical("NaN detected"); assert!(matches!(err, PINNError::Numerical(_))); assert_eq!(err.to_string(), "Numerical error: NaN detected"); } #[test] fn test_error_is_send_and_sync() { fn assert_send_sync() {} assert_send_sync::(); } #[test] fn test_error_display_network() { let err = PINNError::Network("test message".to_string()); let display = format!("{err}"); assert_eq!(display, "Network error: test message"); } #[test] fn test_error_display_training() { let err = PINNError::Training("gradient overflow".to_string()); let display = format!("{err}"); assert_eq!(display, "Training error: gradient overflow"); } #[test] fn test_error_from_string_slice() { let err = PINNError::network("test"); assert!(matches!(err, PINNError::Network(_))); } #[test] fn test_error_from_string() { let msg = String::from("owned string"); let err = PINNError::training(msg); assert!(matches!(err, PINNError::Training(_))); } #[test] fn test_result_type_ok() { let result: Result = Ok(42); assert!(result.is_ok()); assert_eq!(result.unwrap(), 42); } #[test] fn test_result_type_err() { let result: Result = Err(PINNError::network("failed")); assert!(result.is_err()); } #[test] fn test_error_chain_with_question_mark() { fn inner_fn() -> Result { Err(PINNError::network("inner error")) } fn outer_fn() -> Result { inner_fn()?; Ok(42) } let result = outer_fn(); assert!(result.is_err()); assert!(matches!(result.unwrap_err(), PINNError::Network(_))); } #[test] fn test_multiple_error_types() { let errors = vec![ PINNError::network("net"), PINNError::training("train"), PINNError::benchmark("bench"), PINNError::invalid_config("config"), PINNError::numerical("num"), ]; assert_eq!(errors.len(), 5); for err in errors { assert!(!err.to_string().is_empty()); } } #[test] fn test_error_debug_format() { let err = PINNError::network("debug test"); let debug = format!("{err:?}"); assert!(debug.contains("Network")); assert!(debug.contains("debug test")); } #[test] fn test_serde_json_error_conversion() { let json_err = serde_json::from_str::("invalid json"); assert!(json_err.is_err()); let pinn_err = PINNError::from(json_err.unwrap_err()); assert!(matches!(pinn_err, PINNError::Serialization(_))); } }