use ndarray::Array3;
use regex::Regex;
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use crate::core::{
image_loader::{list_image_files, ImageCapture},
ImageLoader,
};
use crate::error::{GlintError, Result};
#[derive(Debug, Clone, PartialEq)]
pub enum BandScheme {
Numbered { count: usize },
Named { names: Vec<String> },
}
#[derive(Debug, Clone)]
pub struct MultiFileLoaderConfig {
pub base_file_pattern: String,
pub filename_template: String,
pub band_scheme: BandScheme,
pub reference_band: String,
pub extensions: Vec<String>,
pub bit_depth: u8,
}
impl MultiFileLoaderConfig {
pub fn new(
base_file_pattern: String,
filename_template: String,
band_scheme: BandScheme,
reference_band: String,
extensions: Vec<String>,
bit_depth: u8,
) -> Self {
Self {
base_file_pattern,
filename_template,
band_scheme,
reference_band,
extensions,
bit_depth,
}
}
pub fn from_config(config: &HashMap<String, String>) -> Result<Self> {
let base_file_pattern = config
.get("base_file_pattern")
.ok_or_else(|| GlintError::config("Missing base_file_pattern in loader_config"))?
.clone();
let filename_template = config
.get("filename_template")
.ok_or_else(|| GlintError::config("Missing filename_template in loader_config"))?
.clone();
let reference_band = config
.get("reference_band")
.ok_or_else(|| GlintError::config("Missing reference_band in loader_config"))?
.clone();
let extensions_str = config
.get("extensions")
.ok_or_else(|| GlintError::config("Missing extensions in loader_config"))?;
let extensions: Vec<String> = extensions_str
.split(',')
.map(|s| s.trim().to_string())
.collect();
let bit_depth = config
.get("bit_depth")
.and_then(|s| s.parse::<u8>().ok())
.unwrap_or(16);
let band_scheme = if let Some(scheme_str) = config.get("band_scheme") {
match scheme_str.as_str() {
"numbered" => {
let count = config
.get("band_count")
.and_then(|s| s.parse::<usize>().ok())
.ok_or_else(|| {
GlintError::config("Missing band_count for numbered band scheme")
})?;
BandScheme::Numbered { count }
}
"named" => {
let names_str = config.get("band_names").ok_or_else(|| {
GlintError::config("Missing band_names for named band scheme")
})?;
let names: Vec<String> =
names_str.split(',').map(|s| s.trim().to_string()).collect();
BandScheme::Named { names }
}
_ => {
return Err(GlintError::config(
"Invalid band_scheme. Must be 'numbered' or 'named'",
))
}
}
} else {
return Err(GlintError::config("Missing band_scheme in loader_config"));
};
Ok(Self::new(
base_file_pattern,
filename_template,
band_scheme,
reference_band,
extensions,
bit_depth,
))
}
pub fn band_count(&self) -> usize {
match &self.band_scheme {
BandScheme::Numbered { count } => *count,
BandScheme::Named { names } => names.len(),
}
}
pub fn band_identifiers(&self) -> Vec<String> {
match &self.band_scheme {
BandScheme::Numbered { count } => (1..=*count).map(|i| i.to_string()).collect(),
BandScheme::Named { names } => names.clone(),
}
}
}
#[derive(Debug, Clone)]
pub struct ConfigurableMultiFileLoader {
base_pattern: Regex,
config: MultiFileLoaderConfig,
}
impl ConfigurableMultiFileLoader {
pub fn new(config: MultiFileLoaderConfig) -> Result<Self> {
let base_pattern = Regex::new(&config.base_file_pattern)
.map_err(|e| GlintError::validation(format!("Invalid base file pattern: {}", e)))?;
Ok(Self {
base_pattern,
config,
})
}
pub fn from_config(config: &HashMap<String, String>) -> Result<Self> {
let loader_config = MultiFileLoaderConfig::from_config(config)?;
Self::new(loader_config)
}
fn extract_filename_parts(&self, path: &Path) -> Option<HashMap<String, String>> {
let filename = path.file_name()?.to_str()?;
if let Some(captures) = self.base_pattern.captures(filename) {
let mut parts = HashMap::new();
for name in self.base_pattern.capture_names().flatten() {
if let Some(matched) = captures.name(name) {
parts.insert(name.to_string(), matched.as_str().to_string());
}
}
Some(parts)
} else {
None
}
}
fn find_band_files(&self, base_file: &Path) -> Result<Vec<PathBuf>> {
let filename_parts = self.extract_filename_parts(base_file).ok_or_else(|| {
GlintError::processing("Could not extract filename parts from base file")
})?;
let parent_dir = base_file
.parent()
.ok_or_else(|| GlintError::processing("File has no parent directory"))?;
let mut band_files = Vec::new();
let band_identifiers = self.config.band_identifiers();
for band_id in &band_identifiers {
let mut parts = filename_parts.clone();
parts.insert("band".to_string(), band_id.clone());
let filename = self.generate_filename_from_template(&parts)?;
let band_path = parent_dir.join(filename);
if !band_path.exists() {
return Err(GlintError::MissingFiles {
files: vec![band_path],
});
}
band_files.push(band_path);
}
Ok(band_files)
}
fn generate_filename_from_template(&self, parts: &HashMap<String, String>) -> Result<String> {
let mut filename = self.config.filename_template.clone();
for (key, value) in parts {
let placeholder = format!("{{{}}}", key);
filename = filename.replace(&placeholder, value);
}
if filename.contains('{') && filename.contains('}') {
return Err(GlintError::processing(format!(
"Template contains unreplaced placeholders: {}",
filename
)));
}
Ok(filename)
}
fn generate_capture_id(&self, parts: &HashMap<String, String>) -> String {
if let Some(base) = parts.get("base") {
return base.clone();
}
if let Some(prefix) = parts.get("prefix") {
return prefix.clone();
}
let mut id_parts = Vec::new();
for key in ["prefix", "sequence", "number", "id", "capture"] {
if let Some(value) = parts.get(key) {
id_parts.push(value.clone());
}
}
if !id_parts.is_empty() {
id_parts.join("_")
} else {
parts
.iter()
.filter(|(k, _)| !matches!(k.as_str(), "band" | "extension" | "ext"))
.map(|(_, v)| v.clone())
.next()
.unwrap_or_else(|| "unknown".to_string())
}
}
fn load_band_files(&self, band_files: &[PathBuf]) -> Result<Array3<f64>> {
let expected_count = self.config.band_count();
if band_files.len() != expected_count {
return Err(GlintError::BandCountMismatch {
expected: expected_count,
actual: band_files.len(),
});
}
let mut band_arrays = Vec::new();
let mut height = 0;
let mut width = 0;
for (i, band_path) in band_files.iter().enumerate() {
let img = image::open(band_path)?;
let gray_img = img.to_luma16();
let (w, h) = gray_img.dimensions();
if i == 0 {
height = h as usize;
width = w as usize;
} else if h as usize != height || w as usize != width {
return Err(GlintError::DimensionMismatch {
expected: (width as u32, height as u32),
actual: (w, h),
});
}
let data: Vec<f64> = gray_img.into_raw().into_iter().map(|x| x as f64).collect();
let band_array = ndarray::Array2::from_shape_vec((height, width), data)
.map_err(|_| GlintError::processing("Failed to reshape band data"))?;
band_arrays.push(band_array);
}
let mut stacked_data = Vec::with_capacity(height * width * expected_count);
for y in 0..height {
for x in 0..width {
for band in &band_arrays {
stacked_data.push(band[[y, x]]);
}
}
}
let result = Array3::from_shape_vec((height, width, expected_count), stacked_data)
.map_err(|_| GlintError::processing("Failed to create stacked array"))?;
Ok(result)
}
}
impl ImageLoader for ConfigurableMultiFileLoader {
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn discover_captures(&self, input_dir: &Path, output_dir: &Path) -> Result<Vec<ImageCapture>> {
let all_files = list_image_files(input_dir, &self.config.extensions)?;
let mut captures = Vec::new();
for file_path in all_files {
let filename = file_path.file_name().and_then(|n| n.to_str()).unwrap_or("");
if self.base_pattern.is_match(filename) {
match self.find_band_files(&file_path) {
Ok(band_files) => {
let filename_parts =
self.extract_filename_parts(&file_path).unwrap_or_else(|| {
let mut parts = HashMap::new();
parts.insert("unknown".to_string(), "capture".to_string());
parts
});
let id = self.generate_capture_id(&filename_parts);
let mask_paths = self.generate_mask_paths(
&ImageCapture {
id: id.clone(),
paths: band_files.clone(),
mask_paths: Vec::new(),
},
output_dir,
);
captures.push(ImageCapture {
id,
paths: band_files,
mask_paths,
});
}
Err(e) => {
eprintln!(
"Warning: Could not load band files for {:?}: {}",
file_path, e
);
}
}
}
}
captures.sort_by(|a, b| a.id.cmp(&b.id));
Ok(captures)
}
fn load_image(&self, capture: &ImageCapture) -> Result<Array3<f64>> {
if capture.paths.len() != self.expected_file_count() {
return Err(GlintError::validation(format!(
"Configurable multifile loader expects exactly {} files, got {}",
self.expected_file_count(),
capture.paths.len()
)));
}
for path in &capture.paths {
if !path.exists() {
return Err(GlintError::MissingFiles {
files: vec![path.clone()],
});
}
}
self.load_band_files(&capture.paths)
}
fn band_count(&self) -> usize {
self.config.band_count()
}
fn bit_depth(&self) -> u8 {
self.config.bit_depth
}
fn supported_extensions(&self) -> Vec<String> {
self.config.extensions.clone()
}
fn expected_file_count(&self) -> usize {
self.config.band_count()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
use tempfile::tempdir;
#[test]
fn test_multifile_loader_config_from_hashmap() {
let mut config = HashMap::new();
config.insert(
"base_file_pattern".to_string(),
"^(?P<base>IMG_\\d{4})_(?P<band>1)(?P<extension>\\.tif)$".to_string(),
);
config.insert(
"filename_template".to_string(),
"{base}_{band}{extension}".to_string(),
);
config.insert("band_scheme".to_string(), "numbered".to_string());
config.insert("band_count".to_string(), "5".to_string());
config.insert("reference_band".to_string(), "1".to_string());
config.insert("extensions".to_string(), "tif,tiff".to_string());
config.insert("bit_depth".to_string(), "16".to_string());
let loader_config = MultiFileLoaderConfig::from_config(&config).unwrap();
assert_eq!(loader_config.band_count(), 5);
assert_eq!(loader_config.bit_depth, 16);
assert_eq!(loader_config.extensions, vec!["tif", "tiff"]);
match loader_config.band_scheme {
BandScheme::Numbered { count } => assert_eq!(count, 5),
_ => panic!("Expected numbered band scheme"),
}
}
#[test]
fn test_multifile_loader_config_named_bands() {
let mut config = HashMap::new();
config.insert(
"base_file_pattern".to_string(),
"^(?P<base>DJI_\\d+_\\d+_MS)_(?P<band>G)(?P<extension>\\.TIF)$".to_string(),
);
config.insert(
"filename_template".to_string(),
"{base}_{band}{extension}".to_string(),
);
config.insert("band_scheme".to_string(), "named".to_string());
config.insert("band_names".to_string(), "G,R,RE,NIR".to_string());
config.insert("reference_band".to_string(), "G".to_string());
config.insert("extensions".to_string(), "TIF,tif".to_string());
let loader_config = MultiFileLoaderConfig::from_config(&config).unwrap();
assert_eq!(loader_config.band_count(), 4);
assert_eq!(loader_config.bit_depth, 16);
match loader_config.band_scheme {
BandScheme::Named { names } => {
assert_eq!(names, vec!["G", "R", "RE", "NIR"]);
}
_ => panic!("Expected named band scheme"),
}
}
#[test]
fn test_configurable_multifile_loader_creation() {
let config = MultiFileLoaderConfig::new(
"^(?P<base>IMG_\\d{4})_(?P<band>1)(?P<extension>\\.tif)$".to_string(),
"{base}_{band}{extension}".to_string(),
BandScheme::Numbered { count: 5 },
"1".to_string(),
vec!["tif".to_string(), "tiff".to_string()],
16,
);
let loader = ConfigurableMultiFileLoader::new(config).unwrap();
assert_eq!(loader.band_count(), 5);
assert_eq!(loader.bit_depth(), 16);
assert_eq!(loader.expected_file_count(), 5);
assert!(loader.supported_extensions().contains(&"tif".to_string()));
}
#[test]
fn test_micasense_pattern_capture_groups() {
let config = MultiFileLoaderConfig::new(
"^(?P<base>IMG_\\d{4})_(?P<band>1)(?P<extension>\\.tif)$".to_string(),
"{base}_{band}{extension}".to_string(),
BandScheme::Numbered { count: 5 },
"1".to_string(),
vec!["tif".to_string()],
16,
);
let loader = ConfigurableMultiFileLoader::new(config).unwrap();
let path = Path::new("IMG_0001_1.tif");
let parts = loader.extract_filename_parts(path).unwrap();
assert_eq!(parts.get("base"), Some(&"IMG_0001".to_string()));
assert_eq!(parts.get("band"), Some(&"1".to_string()));
assert_eq!(parts.get("extension"), Some(&".tif".to_string()));
}
#[test]
fn test_dji_p4ms_pattern_capture_groups() {
let config = MultiFileLoaderConfig::new(
"^(?P<base>DJI_\\d{3})(?P<band>1)(?P<extension>\\.TIF)$".to_string(),
"{base}{band}{extension}".to_string(),
BandScheme::Numbered { count: 5 },
"1".to_string(),
vec!["TIF".to_string()],
16,
);
let loader = ConfigurableMultiFileLoader::new(config).unwrap();
let path = Path::new("DJI_0001.TIF");
let parts = loader.extract_filename_parts(path).unwrap();
assert_eq!(parts.get("base"), Some(&"DJI_000".to_string()));
assert_eq!(parts.get("band"), Some(&"1".to_string()));
assert_eq!(parts.get("extension"), Some(&".TIF".to_string()));
}
#[test]
fn test_dji_m3m_pattern_capture_groups() {
let config = MultiFileLoaderConfig::new(
"^(?P<base>DJI_\\d+_\\d+_MS)_(?P<band>G)(?P<extension>\\.TIF)$".to_string(),
"{base}_{band}{extension}".to_string(),
BandScheme::Named {
names: vec![
"G".to_string(),
"R".to_string(),
"RE".to_string(),
"NIR".to_string(),
],
},
"G".to_string(),
vec!["TIF".to_string()],
16,
);
let loader = ConfigurableMultiFileLoader::new(config).unwrap();
let path = Path::new("DJI_20221208115250_0001_MS_G.TIF");
let parts = loader.extract_filename_parts(path).unwrap();
assert_eq!(
parts.get("base"),
Some(&"DJI_20221208115250_0001_MS".to_string())
);
assert_eq!(parts.get("band"), Some(&"G".to_string()));
assert_eq!(parts.get("extension"), Some(&".TIF".to_string()));
}
#[test]
fn test_discover_captures_micasense_pattern() {
let temp_dir = tempdir().unwrap();
let input_dir = temp_dir.path().join("input");
let output_dir = temp_dir.path().join("output");
fs::create_dir_all(&input_dir).unwrap();
fs::create_dir_all(&output_dir).unwrap();
let base_files = ["IMG_0001", "IMG_0002"];
for base in &base_files {
for band in 1..=5 {
let filename = format!("{}_{}.tif", base, band);
let filepath = input_dir.join(filename);
fs::write(&filepath, b"fake image data").unwrap();
}
}
let config = MultiFileLoaderConfig::new(
"^(?P<base>IMG_\\d{4})_(?P<band>1)(?P<extension>\\.tif)$".to_string(),
"{base}_{band}{extension}".to_string(),
BandScheme::Numbered { count: 5 },
"1".to_string(),
vec!["tif".to_string()],
16,
);
let loader = ConfigurableMultiFileLoader::new(config).unwrap();
let captures = loader.discover_captures(&input_dir, &output_dir).unwrap();
assert_eq!(captures.len(), 2);
assert_eq!(captures[0].id, "IMG_0001");
assert_eq!(captures[1].id, "IMG_0002");
for capture in &captures {
assert_eq!(capture.paths.len(), 5);
assert_eq!(capture.mask_paths.len(), 5);
for mask_path in &capture.mask_paths {
assert!(mask_path.to_string_lossy().contains("_mask.png"));
}
}
}
#[test]
fn test_template_filename_generation() {
let config = MultiFileLoaderConfig::new(
"^(?P<base>DJI_\\d{3})(?P<band>1)(?P<extension>\\.TIF)$".to_string(),
"{base}{band}{extension}".to_string(),
BandScheme::Numbered { count: 5 },
"1".to_string(),
vec!["TIF".to_string()],
16,
);
let loader = ConfigurableMultiFileLoader::new(config).unwrap();
let mut parts = std::collections::HashMap::new();
parts.insert("base".to_string(), "DJI_000".to_string());
parts.insert("band".to_string(), "2".to_string());
parts.insert("extension".to_string(), ".TIF".to_string());
let filename = loader.generate_filename_from_template(&parts).unwrap();
assert_eq!(filename, "DJI_0002.TIF");
let config = MultiFileLoaderConfig::new(
"^(?P<base>IMG_\\d{4})_(?P<band>1)(?P<extension>\\.tif)$".to_string(),
"{base}_{band}{extension}".to_string(),
BandScheme::Numbered { count: 5 },
"1".to_string(),
vec!["tif".to_string()],
16,
);
let loader = ConfigurableMultiFileLoader::new(config).unwrap();
let mut parts = std::collections::HashMap::new();
parts.insert("base".to_string(), "IMG_0001".to_string());
parts.insert("band".to_string(), "3".to_string());
parts.insert("extension".to_string(), ".tif".to_string());
let filename = loader.generate_filename_from_template(&parts).unwrap();
assert_eq!(filename, "IMG_0001_3.tif");
}
}