Files
rustytorch/demos/segmentation-shared/src/ipc.rs
T
2026-03-04 00:08:42 +00:00

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();
}
}
}