370 lines
11 KiB
Rust
370 lines
11 KiB
Rust
// 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<u8>,
|
|
/// 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<String>,
|
|
},
|
|
/// Class information
|
|
Classes {
|
|
/// Class names for the model
|
|
class_names: Vec<String>,
|
|
},
|
|
/// 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<String>) -> Self {
|
|
Self::Models { models }
|
|
}
|
|
|
|
/// Create classes response
|
|
pub fn classes(class_names: Vec<String>) -> 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();
|
|
}
|
|
}
|
|
}
|