use ndarray::Array2;
use crate::{Error, Result};
pub use crate::upstream::numpy::percentile_linear;
pub use crate::upstream::skimage::{gradient3x3, morphological_gradient};
const MIN_LAKE_COMPONENT: usize = 12;
pub fn masked_gradient3x3(img: &[f64], excluded: &[bool], nrows: usize, ncols: usize) -> Vec<f32> {
let mut out = vec![0f32; nrows * ncols];
for r in 0..nrows {
let (r0, r1) = (r.saturating_sub(1), (r + 1).min(nrows - 1));
for c in 0..ncols {
let (c0, c1) = (c.saturating_sub(1), (c + 1).min(ncols - 1));
let bounds = (r0..=r1)
.flat_map(|rr| (c0..=c1).map(move |cc| rr * ncols + cc))
.filter(|&i| !excluded[i])
.map(|i| img[i])
.fold(None, |acc: Option<(f64, f64)>, v| {
Some(acc.map_or((v, v), |(lo, hi)| (lo.min(v), hi.max(v))))
});
out[r * ncols + c] = bounds.map_or(0.0, |(lo, hi)| (hi - lo) as f32);
}
}
out
}
fn keep_large_components(mask: &[bool], nrows: usize, ncols: usize, min_size: usize) -> Vec<bool> {
let mut out = vec![false; nrows * ncols];
let mut visited = vec![false; nrows * ncols];
for start in 0..nrows * ncols {
if !mask[start] || visited[start] {
continue;
}
let mut component = Vec::new();
let mut stack = vec![start];
visited[start] = true;
while let Some(i) = stack.pop() {
component.push(i);
let (r, c) = (i / ncols, i % ncols);
let neighbours = [
r.checked_sub(1).map(|rr| (rr, c)),
(r + 1 < nrows).then_some((r + 1, c)),
c.checked_sub(1).map(|cc| (r, cc)),
(c + 1 < ncols).then_some((r, c + 1)),
];
for j in neighbours
.into_iter()
.flatten()
.map(|(rr, cc)| rr * ncols + cc)
{
if mask[j] && !visited[j] {
visited[j] = true;
stack.push(j);
}
}
}
if component.len() >= min_size {
for i in component {
out[i] = true;
}
}
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn preprocess(
values: &Array2<f64>,
water_ntile: f64,
lake_flatness: i32,
vertical_ratio: f64,
flatness_scale: f64,
stats_region: Option<&[bool]>,
) -> Result<Array2<f64>> {
let (nrows, ncols) = values.dim();
let mut v = values.to_owned();
let nan_mask: Vec<bool> = v.iter().map(|x| x.is_nan()).collect();
let in_stats = |i: usize| match stats_region {
Some(region) => region[i],
None => true,
};
let finite: Vec<f64> = v
.iter()
.enumerate()
.filter(|(i, x)| !x.is_nan() && in_stats(*i))
.map(|(_, x)| *x)
.collect();
if finite.is_empty() {
return Err(Error::EmptyData);
}
let min = finite.iter().cloned().fold(f64::INFINITY, f64::min);
let max = finite.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
for (i, x) in v.iter_mut().enumerate() {
if nan_mask[i] {
*x = min;
}
}
let span = max - min;
if span > 0.0 {
for x in v.iter_mut() {
*x = (*x - min) / span;
}
} else {
for x in v.iter_mut() {
*x = 0.0;
}
}
let mut sorted: Vec<f64> = v
.iter()
.enumerate()
.filter(|(i, _)| !nan_mask[*i] && in_stats(*i))
.map(|(_, x)| *x)
.collect();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
let water_level = percentile_linear(&sorted, water_ntile.clamp(0.0, 100.0));
let is_water: Vec<bool> = v
.iter()
.enumerate()
.map(|(i, x)| !nan_mask[i] && *x < water_level)
.collect();
let excluded: Vec<bool> = (0..v.len()).map(|i| nan_mask[i] || is_water[i]).collect();
let grad = masked_gradient3x3(v.as_slice().expect("contiguous"), &excluded, nrows, ncols);
let threshold = (lake_flatness as f64 / 255.0) * flatness_scale;
let candidate: Vec<bool> = (0..v.len())
.map(|i| !excluded[i] && (grad[i] as f64) < threshold)
.collect();
let is_lake = keep_large_components(&candidate, nrows, ncols, MIN_LAKE_COMPONENT);
let mut out = Array2::<f64>::from_elem((nrows, ncols), f64::NAN);
for r in 0..nrows {
let src_row = nrows - 1 - r;
for c in 0..ncols {
let i = src_row * ncols + c;
if nan_mask[i] {
continue; }
let val = v[(src_row, c)];
if val < water_level || is_lake[i] {
continue; }
out[(r, c)] = val * vertical_ratio;
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::{array, s};
#[test]
fn masks_water_and_flat() {
let mut a = Array2::<f64>::zeros((2, 4));
a[(0, 0)] = 0.0;
a[(0, 1)] = 50.0;
a[(0, 2)] = 50.0;
a[(0, 3)] = 100.0;
a[(1, 0)] = 10.0;
a[(1, 1)] = 90.0;
a[(1, 2)] = 95.0;
a[(1, 3)] = 100.0;
let out = preprocess(&a, 25.0, 3, 40.0, 1.0, None).unwrap();
assert!(out[(1, 0)].is_nan(), "deep cell masked as water");
assert!(out[(1, 2 - 2)].is_nan() || !out[(1, 1)].is_nan());
assert_eq!(out[(0, 1)], 0.9 * 40.0);
assert_eq!(out[(0, 3)], 1.0 * 40.0);
assert_eq!(out[(1, 3)], 1.0 * 40.0);
}
#[test]
fn rows_are_flipped() {
let a = array![[100.0, 200.0], [0.0, 50.0]];
let out = preprocess(&a, 0.0, 0, 10.0, 1.0, None).unwrap();
assert!((out[(0, 0)] - 0.0).abs() < 1e-12 || out[(0, 0)].is_nan());
assert!((out[(0, 1)] - 2.5).abs() < 1e-12);
assert!((out[(1, 0)] - 5.0).abs() < 1e-12);
assert!((out[(1, 1)] - 10.0).abs() < 1e-12);
}
#[test]
fn lake_mask_needs_coherence_and_ignores_water_edges() {
let n = 40;
let mut a = Array2::from_shape_fn((n, n), |(r, _)| 100.0 + 0.6 * r as f64);
a.slice_mut(s![8..13, 12..22]).fill(105.0); a.slice_mut(s![14..24, 12..22]).fill(107.0); a.slice_mut(s![28..30, 30..32]).fill(116.0);
let out = preprocess(&a, 0.0, 3, 1.0, 1.0, None).unwrap();
let first_unmasked_in = |rows: std::ops::Range<usize>, cols: std::ops::Range<usize>| {
out.indexed_iter()
.find(|&((r, c), v)| rows.contains(&r) && cols.contains(&c) && !v.is_nan())
.map(|((r, c), _)| (r, c))
};
let unmasked_pond = first_unmasked_in(17..24, 13..20);
assert!(
unmasked_pond.is_none(),
"pond cell {:?} should be masked",
unmasked_pond
);
let unmasked_bench = first_unmasked_in(28..30, 13..20);
assert!(
unmasked_bench.is_none(),
"bench cell {:?} should be masked",
unmasked_bench
);
let unexpected_hole = out.indexed_iter().find(|&((r, c), v)| {
let in_expected_region = ((17..=24).contains(&r) && (13..=20).contains(&c))
|| ((28..=30).contains(&r) && (13..=20).contains(&c));
v.is_nan() && !in_expected_region
});
assert!(
unexpected_hole.is_none(),
"unexpected hole at {:?}",
unexpected_hole.map(|((r, c), _)| (r, c))
);
}
#[test]
fn water_percentile_ignores_nan_padding() {
let n = 10;
let a = Array2::from_shape_fn((n, n), |(r, c)| {
if c < 7 {
10.0 + 5.0 * r as f64 + c as f64 } else {
f64::NAN }
});
let out = preprocess(&a, 15.0, 0, 1.0, 1.0, None).unwrap();
assert!(
out.slice(s![n - 1, ..7]).iter().all(|v| v.is_nan()),
"river cells (out row {}) must be water-masked",
n - 1
);
assert!(
out.slice(s![4, ..7]).iter().all(|v| v.is_finite()),
"upland cells (out row 4) must survive"
);
}
#[test]
fn all_nan_errors() {
let a = Array2::<f64>::from_elem((2, 2), f64::NAN);
assert!(preprocess(&a, 10.0, 3, 40.0, 1.0, None).is_err());
}
}