use indicatif::{ProgressBar, ProgressStyle};
use std::path::Path;
use tracing::{error, info, warn};
use crate::core::{
image_loader::{normalize_image, ImageCapture},
sensor::Sensor,
GlintAlgorithm, ImageLoader, PostProcessor,
};
use crate::error::{GlintError, Result};
pub struct Masker {
loader: Box<dyn ImageLoader>,
algorithm: Box<dyn GlintAlgorithm>,
postprocessor: Box<dyn PostProcessor>,
sensor: Sensor,
}
type Callback = dyn Fn(&ImageCapture) + Send + Sync;
impl Masker {
pub fn builder() -> MaskerBuilder {
MaskerBuilder::new()
}
pub fn process_directory(
&self,
input_dir: &Path,
output_dir: &Path,
callback: Option<Box<Callback>>,
) -> Result<ProcessingStats> {
self.process_directory_with_progress(input_dir, output_dir, callback, true)
}
pub fn process_directory_with_progress(
&self,
input_dir: &Path,
output_dir: &Path,
callback: Option<Box<Callback>>,
show_progress: bool,
) -> Result<ProcessingStats> {
if !show_progress {
info!("Starting glint mask processing");
info!("Input directory: {}", input_dir.display());
info!("Output directory: {}", output_dir.display());
info!("Sensor: {} ({})", self.sensor.name, self.sensor.id);
info!("Algorithm: {}", self.algorithm.name());
info!("Post-processor: {}", self.postprocessor.name());
}
let captures = self.loader.discover_captures(input_dir, output_dir)?;
if !show_progress {
info!("Found {} image captures", captures.len());
}
if captures.is_empty() {
warn!("No image captures found in input directory");
return Ok(ProcessingStats::new());
}
let progress_bar = if show_progress {
info!("Processing {} images...", captures.len());
let pb = ProgressBar::new(captures.len() as u64);
pb.set_style(
ProgressStyle::default_bar()
.template("{spinner:.green} [{elapsed_precise}] [{bar:40.cyan/blue}] {pos}/{len} ({eta})")
.unwrap()
.progress_chars("#>-")
);
pb.set_message("Processing images...");
Some(pb)
} else {
None
};
let mut results = Vec::new();
for (i, capture) in captures.iter().enumerate() {
if let Some(ref pb) = progress_bar {
pb.set_message(format!("Processing {}", capture.id));
pb.set_position(i as u64);
}
let result = self.process_single_capture(capture);
if let Some(ref cb) = callback {
cb(capture);
}
results.push(result);
}
if let Some(ref pb) = progress_bar {
pb.finish_with_message("Processing complete");
}
let mut stats = ProcessingStats::new();
stats.total_captures = captures.len();
for result in results {
match result {
Ok(()) => stats.successful_captures += 1,
Err(e) => {
error!("Processing failed: {}", e);
stats.failed_captures += 1;
stats.errors.push(e.to_string());
}
}
}
if !show_progress {
info!(
"Processing complete: {}/{} successful",
stats.successful_captures, stats.total_captures
);
if stats.failed_captures > 0 {
warn!("{} captures failed to process", stats.failed_captures);
}
}
Ok(stats)
}
fn process_single_capture(&self, capture: &ImageCapture) -> Result<()> {
self.loader.validate_capture(capture)?;
if let Some(big_tiff_loader) = self
.loader
.as_any()
.downcast_ref::<crate::loaders::BigTiffLoader>()
{
return big_tiff_loader.process_chunked_image(
capture,
&*self.algorithm,
&*self.postprocessor,
self.sensor.bit_depth,
0, );
}
let image = self.loader.load_image(capture)?;
let (height, width, bands) = image.dim();
if bands != self.sensor.band_count() {
return Err(GlintError::BandCountMismatch {
expected: self.sensor.band_count(),
actual: bands,
});
}
if !self.algorithm.supports_bands(bands) {
return Err(GlintError::validation(format!(
"Algorithm '{}' does not support {} bands",
self.algorithm.name(),
bands
)));
}
let normalized_image = normalize_image(&image, self.sensor.bit_depth)?;
let mask = self.algorithm.detect_glint(&normalized_image)?;
if mask.dim() != (height, width) {
return Err(GlintError::DimensionMismatch {
expected: (width as u32, height as u32),
actual: (mask.dim().1 as u32, mask.dim().0 as u32),
});
}
let processed_mask = self.postprocessor.process_mask(&mask)?;
self.loader.save_masks(&processed_mask, capture)?;
Ok(())
}
pub fn sensor(&self) -> &Sensor {
&self.sensor
}
pub fn algorithm(&self) -> &dyn GlintAlgorithm {
&*self.algorithm
}
pub fn postprocessor(&self) -> &dyn PostProcessor {
&*self.postprocessor
}
}
pub struct MaskerBuilder {
loader: Option<Box<dyn ImageLoader>>,
algorithm: Option<Box<dyn GlintAlgorithm>>,
postprocessor: Option<Box<dyn PostProcessor>>,
sensor: Option<Sensor>,
}
impl MaskerBuilder {
pub fn new() -> Self {
Self {
loader: None,
algorithm: None,
postprocessor: None,
sensor: None,
}
}
pub fn with_loader(mut self, loader: Box<dyn ImageLoader>) -> Self {
self.loader = Some(loader);
self
}
pub fn with_algorithm(mut self, algorithm: Box<dyn GlintAlgorithm>) -> Self {
self.algorithm = Some(algorithm);
self
}
pub fn with_postprocessor(mut self, postprocessor: Box<dyn PostProcessor>) -> Self {
self.postprocessor = Some(postprocessor);
self
}
pub fn with_sensor(mut self, sensor: Sensor) -> Self {
self.sensor = Some(sensor);
self
}
pub fn build(self) -> Result<Masker> {
let loader = self
.loader
.ok_or_else(|| GlintError::validation("Image loader must be specified"))?;
let algorithm = self
.algorithm
.ok_or_else(|| GlintError::validation("Algorithm must be specified"))?;
let postprocessor = self
.postprocessor
.ok_or_else(|| GlintError::validation("Post-processor must be specified"))?;
let sensor = self
.sensor
.ok_or_else(|| GlintError::validation("Sensor configuration must be specified"))?;
sensor.validate()?;
if loader.band_count() != sensor.band_count() {
return Err(GlintError::BandCountMismatch {
expected: sensor.band_count(),
actual: loader.band_count(),
});
}
if loader.bit_depth() != sensor.bit_depth {
return Err(GlintError::InvalidBitDepth {
bit_depth: loader.bit_depth(),
});
}
algorithm.validate_parameters()?;
postprocessor.validate_parameters()?;
Ok(Masker {
loader,
algorithm,
postprocessor,
sensor,
})
}
}
impl Default for MaskerBuilder {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct ProcessingStats {
pub total_captures: usize,
pub successful_captures: usize,
pub failed_captures: usize,
pub errors: Vec<String>,
}
impl ProcessingStats {
fn new() -> Self {
Self {
total_captures: 0,
successful_captures: 0,
failed_captures: 0,
errors: Vec::new(),
}
}
pub fn success_rate(&self) -> f64 {
if self.total_captures == 0 {
0.0
} else {
(self.successful_captures as f64 / self.total_captures as f64) * 100.0
}
}
pub fn all_successful(&self) -> bool {
self.failed_captures == 0 && self.total_captures > 0
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::sensor::{Band, Sensor};
use ndarray::Array3;
use std::path::PathBuf;
struct MockLoader;
impl ImageLoader for MockLoader {
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn discover_captures(
&self,
_input_dir: &Path,
_output_dir: &Path,
) -> Result<Vec<ImageCapture>> {
Ok(vec![ImageCapture {
id: "test".to_string(),
paths: vec![PathBuf::from("test.jpg")],
mask_paths: vec![PathBuf::from("test_mask.png")],
}])
}
fn load_image(&self, _capture: &ImageCapture) -> Result<Array3<f64>> {
Ok(Array3::from_shape_vec((10, 10, 3), vec![128.0; 300]).unwrap())
}
fn band_count(&self) -> usize {
3
}
fn bit_depth(&self) -> u8 {
8
}
fn supported_extensions(&self) -> Vec<String> {
vec!["jpg".to_string()]
}
}
struct MockAlgorithm;
impl GlintAlgorithm for MockAlgorithm {
fn detect_glint(&self, _image: &Array3<f64>) -> Result<ndarray::Array2<u8>> {
Ok(ndarray::Array2::zeros((10, 10)))
}
fn name(&self) -> &'static str {
"Mock"
}
fn description(&self) -> &'static str {
"Mock algorithm"
}
}
struct MockPostProcessor;
impl PostProcessor for MockPostProcessor {
fn process_mask(&self, mask: &ndarray::Array2<u8>) -> Result<ndarray::Array2<u8>> {
Ok(mask.clone())
}
fn name(&self) -> &'static str {
"Mock"
}
fn description(&self) -> &'static str {
"Mock post-processor"
}
}
#[test]
fn test_masker_builder() {
let sensor = Sensor::new(
"test",
"Test Sensor",
vec![
Band::new("Red", 0.9),
Band::new("Green", 0.8),
Band::new("Blue", 0.7),
],
8,
"mock",
);
let masker = Masker::builder()
.with_loader(Box::new(MockLoader))
.with_algorithm(Box::new(MockAlgorithm))
.with_postprocessor(Box::new(MockPostProcessor))
.with_sensor(sensor)
.build();
assert!(masker.is_ok());
}
#[test]
fn test_masker_builder_validation() {
let result = Masker::builder().build();
assert!(result.is_err());
let sensor = Sensor::new(
"test",
"Test Sensor",
vec![Band::new("Red", 0.9)], 8,
"mock",
);
let result = Masker::builder()
.with_loader(Box::new(MockLoader)) .with_algorithm(Box::new(MockAlgorithm))
.with_postprocessor(Box::new(MockPostProcessor))
.with_sensor(sensor)
.build();
assert!(result.is_err());
}
#[test]
fn test_processing_stats() {
let mut stats = ProcessingStats::new();
stats.total_captures = 10;
stats.successful_captures = 8;
stats.failed_captures = 2;
assert_eq!(stats.success_rate(), 80.0);
assert!(!stats.all_successful());
stats.failed_captures = 0;
stats.successful_captures = 10;
assert!(stats.all_successful());
}
}