// IPC types for segmentation requests and responses use crate::{SegmentationModel, SegmentationResult}; use serde::{Deserialize, Serialize}; /// Segmentation request from frontend #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[serde(tag = "type")] pub enum SegmentRequest { /// Perform segmentation on an image Segment { /// Raw image data (RGB bytes) image_data: Vec, /// Image width width: usize, /// Image height height: usize, /// Model to use for segmentation model: SegmentationModel, }, /// Get available models GetModels, /// Get class information for a model GetClasses { /// Model to get classes for model: SegmentationModel, }, } /// Segmentation response to frontend #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[serde(tag = "type")] pub enum SegmentResponse { /// Segmentation completed successfully Success { /// Segmentation result result: SegmentationResult, }, /// Available models list Models { /// List of available model names models: Vec, }, /// Class information Classes { /// Class names for the model class_names: Vec, }, /// Error occurred Error { /// Error message message: String, }, } impl SegmentRequest { /// Validate request parameters pub fn validate(&self) -> Result<(), String> { match self { Self::Segment { image_data, width, height, .. } => { if *width == 0 || *height == 0 { return Err("width and height must be greater than 0".to_string()); } let expected_len = width * height * 3; if image_data.len() != expected_len { return Err(format!( "image_data length {} does not match width * height * 3 = {}", image_data.len(), expected_len )); } Ok(()) } Self::GetModels | Self::GetClasses { .. } => Ok(()), } } } impl SegmentResponse { /// Create success response pub fn success(result: SegmentationResult) -> Self { Self::Success { result } } /// Create models response pub fn models(models: Vec) -> Self { Self::Models { models } } /// Create classes response pub fn classes(class_names: Vec) -> Self { Self::Classes { class_names } } /// Create error response pub fn error(message: String) -> Self { Self::Error { message } } /// Check if response is an error pub fn is_error(&self) -> bool { matches!(self, Self::Error { .. }) } /// Check if response is success pub fn is_success(&self) -> bool { matches!(self, Self::Success { .. }) } } #[cfg(test)] mod tests { use super::*; use crate::SegmentationMask; // SegmentRequest tests #[test] fn test_segment_request_segment_variant() { let request = SegmentRequest::Segment { image_data: vec![0; 300], width: 10, height: 10, model: SegmentationModel::DeepLabV3, }; match request { SegmentRequest::Segment { width, height, model, .. } => { assert_eq!(width, 10); assert_eq!(height, 10); assert_eq!(model, SegmentationModel::DeepLabV3); } _ => panic!("Wrong variant"), } } #[test] fn test_segment_request_get_models_variant() { let request = SegmentRequest::GetModels; assert!(matches!(request, SegmentRequest::GetModels)); } #[test] fn test_segment_request_get_classes_variant() { let request = SegmentRequest::GetClasses { model: SegmentationModel::UNet, }; match request { SegmentRequest::GetClasses { model } => { assert_eq!(model, SegmentationModel::UNet); } _ => panic!("Wrong variant"), } } #[test] fn test_segment_request_validate_success() { let request = SegmentRequest::Segment { image_data: vec![0; 300], width: 10, height: 10, model: SegmentationModel::FCN, }; assert!(request.validate().is_ok()); } #[test] fn test_segment_request_validate_zero_width() { let request = SegmentRequest::Segment { image_data: vec![], width: 0, height: 10, model: SegmentationModel::DeepLabV3, }; assert!(request.validate().is_err()); assert!( request .validate() .unwrap_err() .contains("must be greater than 0") ); } #[test] fn test_segment_request_validate_zero_height() { let request = SegmentRequest::Segment { image_data: vec![], width: 10, height: 0, model: SegmentationModel::SegFormer, }; assert!(request.validate().is_err()); assert!( request .validate() .unwrap_err() .contains("must be greater than 0") ); } #[test] fn test_segment_request_validate_mismatched_data() { let request = SegmentRequest::Segment { image_data: vec![0; 100], width: 10, height: 10, model: SegmentationModel::UNet, }; assert!(request.validate().is_err()); let err = request.validate().unwrap_err(); assert!(err.contains("does not match")); } #[test] fn test_segment_request_validate_get_models() { let request = SegmentRequest::GetModels; assert!(request.validate().is_ok()); } #[test] fn test_segment_request_validate_get_classes() { let request = SegmentRequest::GetClasses { model: SegmentationModel::FCN, }; assert!(request.validate().is_ok()); } #[test] fn test_segment_request_serialization() { let request = SegmentRequest::Segment { image_data: vec![1, 2, 3], width: 1, height: 1, model: SegmentationModel::DeepLabV3, }; let json = serde_json::to_string(&request).unwrap(); let deserialized: SegmentRequest = serde_json::from_str(&json).unwrap(); assert_eq!(request, deserialized); } // SegmentResponse tests #[test] fn test_segment_response_success_constructor() { let mask = SegmentationMask::new(10, 10, vec![0; 100]); let result = SegmentationResult::new(mask, 100.0); let response = SegmentResponse::success(result.clone()); match response { SegmentResponse::Success { result: r } => { assert_eq!(r, result); } _ => panic!("Wrong variant"), } } #[test] fn test_segment_response_models_constructor() { let models = vec!["DeepLabV3".to_string(), "UNet".to_string()]; let response = SegmentResponse::models(models.clone()); match response { SegmentResponse::Models { models: m } => { assert_eq!(m, models); } _ => panic!("Wrong variant"), } } #[test] fn test_segment_response_classes_constructor() { let classes = vec!["background".to_string(), "person".to_string()]; let response = SegmentResponse::classes(classes.clone()); match response { SegmentResponse::Classes { class_names } => { assert_eq!(class_names, classes); } _ => panic!("Wrong variant"), } } #[test] fn test_segment_response_error_constructor() { let msg = "Test error".to_string(); let response = SegmentResponse::error(msg.clone()); match response { SegmentResponse::Error { message } => { assert_eq!(message, msg); } _ => panic!("Wrong variant"), } } #[test] fn test_segment_response_is_error() { let error_response = SegmentResponse::error("error".to_string()); assert!(error_response.is_error()); let success_response = SegmentResponse::success(SegmentationResult::new( SegmentationMask::new(10, 10, vec![0; 100]), 100.0, )); assert!(!success_response.is_error()); let models_response = SegmentResponse::models(vec!["DeepLabV3".to_string()]); assert!(!models_response.is_error()); let classes_response = SegmentResponse::classes(vec!["background".to_string()]); assert!(!classes_response.is_error()); } #[test] fn test_segment_response_is_success() { let success_response = SegmentResponse::success(SegmentationResult::new( SegmentationMask::new(10, 10, vec![0; 100]), 100.0, )); assert!(success_response.is_success()); let error_response = SegmentResponse::error("error".to_string()); assert!(!error_response.is_success()); let models_response = SegmentResponse::models(vec!["DeepLabV3".to_string()]); assert!(!models_response.is_success()); let classes_response = SegmentResponse::classes(vec!["background".to_string()]); assert!(!classes_response.is_success()); } #[test] fn test_segment_response_serialization() { let mask = SegmentationMask::new(5, 5, vec![0; 25]); let result = SegmentationResult::new(mask, 50.0); let response = SegmentResponse::success(result); let json = serde_json::to_string(&response).unwrap(); let deserialized: SegmentResponse = serde_json::from_str(&json).unwrap(); assert_eq!(response, deserialized); } #[test] fn test_all_response_variants_serialize() { let responses = vec![ SegmentResponse::success(SegmentationResult::new( SegmentationMask::new(1, 1, vec![0]), 10.0, )), SegmentResponse::models(vec!["DeepLabV3".to_string()]), SegmentResponse::classes(vec!["background".to_string()]), SegmentResponse::error("test".to_string()), ]; for response in responses { let json = serde_json::to_string(&response).unwrap(); let _deserialized: SegmentResponse = serde_json::from_str(&json).unwrap(); } } }