use anyhow::{Context, Result};
use indicatif::{ProgressBar, ProgressStyle};
use std::path::Path;
use std::time::Instant;
use surtgis_core::Raster;
use surtgis_core::io::{read_geotiff, write_geotiff};
use surtgis_flow::{EntrainmentParams, SimGrid, Simulation, SolverConfig, VoellmyParams};
use crate::commands::FlowCommands;
use crate::helpers::write_opts;
pub fn handle(command: FlowCommands, compress: bool) -> Result<()> {
match command {
FlowCommands::Run {
dem,
release,
outdir,
mu,
xi,
duration,
output_interval,
dump_velocity,
arrival,
erodible,
entrainment_k,
dump_erosion,
until_rest,
rest_velocity,
} => run(RunArgs {
dem_path: &dem,
release_path: &release,
outdir: &outdir,
mu,
xi,
duration,
output_interval,
dump_velocity,
arrival: arrival.as_deref(),
erodible: erodible.as_deref(),
entrainment_k,
dump_erosion,
until_rest,
rest_velocity,
compress,
}),
}
}
struct RunArgs<'a> {
dem_path: &'a Path,
release_path: &'a Path,
outdir: &'a Path,
mu: f32,
xi: f32,
duration: f64,
output_interval: f64,
dump_velocity: bool,
arrival: Option<&'a Path>,
erodible: Option<&'a Path>,
entrainment_k: f32,
dump_erosion: bool,
until_rest: Option<f64>,
rest_velocity: f32,
compress: bool,
}
fn run(args: RunArgs<'_>) -> Result<()> {
let RunArgs {
dem_path,
release_path,
outdir,
mu,
xi,
duration,
output_interval,
dump_velocity,
arrival,
erodible,
entrainment_k,
dump_erosion,
until_rest,
rest_velocity,
compress,
} = args;
anyhow::ensure!(
duration.is_finite() && output_interval.is_finite(),
"duration and output-interval must be finite (got duration={duration}, \
output-interval={output_interval})"
);
anyhow::ensure!(
duration > 0.0 && output_interval > 0.0,
"duration and output-interval must be positive"
);
let start = Instant::now();
let dem: Raster<f32> = read_geotiff(dem_path, None).context("Failed to read DEM")?;
let release: Raster<f32> =
read_geotiff(release_path, None).context("Failed to read release raster")?;
println!(
"DEM: {} x {} cells, cellsize {:.2} m",
dem.cols(),
dem.rows(),
dem.cell_size()
);
let params = VoellmyParams {
mu,
xi,
..VoellmyParams::default()
};
let config = SolverConfig {
max_substeps: 1_000_000,
..SolverConfig::default()
};
let mut sim = Simulation::new(&dem, &release, params, config)?;
let mut ent_params: Option<EntrainmentParams> = None;
if let Some(erodible_path) = erodible {
let emax: Raster<f32> =
read_geotiff(erodible_path, None).context("Failed to read erodible raster")?;
let mut p = EntrainmentParams::default();
p.k = entrainment_k;
sim.set_erodible(&emax, p)?;
ent_params = Some(p);
println!(
"Entrainment: K = {entrainment_k} /m, erodible raster {}",
erodible_path.display()
);
}
let mass0 = sim.total_mass();
println!("Release volume: {mass0:.0} m³");
std::fs::create_dir_all(outdir).context("Failed to create output directory")?;
let n_outputs = (duration / output_interval).ceil() as usize;
anyhow::ensure!(
n_outputs <= 1_000_000,
"duration / output-interval implies {n_outputs} frames (cap: 1,000,000); \
raise --output-interval or lower --duration"
);
let n_frames = n_outputs + 1;
let grid_duration = n_outputs as f64 * output_interval;
if grid_duration > duration * (1.0 + 1e-9) {
eprintln!(
"warning: --duration {duration} is not a multiple of --output-interval \
{output_interval}; simulating to {grid_duration} s so the frame grid \
stays uniform (manifest contract)"
);
}
let crs = dem.crs().cloned();
write_frames(
&sim,
crs.as_ref(),
outdir,
0,
dump_velocity,
dump_erosion,
compress,
)?;
let pb = ProgressBar::new(n_outputs as u64);
pb.set_style(
ProgressStyle::default_bar()
.template("{bar:30.green} {pos}/{len} frames t={msg}")
.unwrap(),
);
let mut total_substeps: u64 = 0;
let mut frames_written = n_outputs;
let mut rest_reached: Option<f64> = None;
for frame in 1..=n_outputs {
let elapsed_target = (frame as f64) * output_interval;
let dt = (elapsed_target - sim.time()).max(0.0);
total_substeps += u64::from(sim.step(dt as f32)?);
write_frames(
&sim,
crs.as_ref(),
outdir,
frame,
dump_velocity,
dump_erosion,
compress,
)?;
pb.set_message(format!("{:.1} s", sim.time()));
pb.inc(1);
if let Some(min_fraction) = until_rest {
let at_rest = sim.mass_fraction_at_rest(rest_velocity);
if at_rest >= min_fraction {
rest_reached = Some(sim.time());
frames_written = frame;
println!(
"Flow reached rest at t = {:.1} s ({:.1}% of the volume \
below {rest_velocity} m/s); stopping early.",
sim.time(),
at_rest * 100.0
);
break;
}
}
}
pb.finish_and_clear();
if until_rest.is_some() && rest_reached.is_none() {
eprintln!(
"warning: the flow had NOT reached rest when --duration {duration} s \
was exhausted ({:.1}% of the volume below {rest_velocity} m/s, \
target {:.1}%) — the runout of this run is a snapshot, not a \
final extent. Raise --duration.",
sim.mass_fraction_at_rest(rest_velocity) * 100.0,
until_rest.unwrap_or(0.0) * 100.0
);
}
if let Some(arrival_path) = arrival {
let raster = masked_raster(sim.grid(), dem.crs(), sim.arrival_times().to_vec());
write_geotiff(&raster, arrival_path, Some(write_opts(compress)))
.context("Failed to write arrival raster")?;
println!("Arrival times saved to: {}", arrival_path.display());
}
write_manifest(
outdir,
&dem,
output_interval,
frames_written + 1, mu,
xi,
ent_params.map(|p| (p, sim.total_eroded())),
)?;
let mass_end = sim.total_mass();
if erodible.is_some() {
println!("Eroded volume: {:.0} m³", sim.total_eroded());
}
println!(
"Simulated {:.1} s in {} frames ({total_substeps} substeps) — wall time {:.2?}",
sim.time(),
n_frames,
start.elapsed()
);
println!(
"Mass: {mass0:.0} -> {mass_end:.0} m³ ({:+.2}% through open borders)",
sim.boundary_volume() / mass0 * 100.0
);
println!("Frames saved to: {}", outdir.display());
Ok(())
}
#[allow(clippy::fn_params_excessive_bools)]
fn write_frames(
sim: &Simulation,
crs: Option<&surtgis_core::CRS>,
outdir: &Path,
frame: usize,
dump_velocity: bool,
dump_erosion: bool,
compress: bool,
) -> Result<()> {
let grid = sim.grid();
let state = sim.state();
let h = masked_raster(grid, crs, state.h.clone());
write_geotiff(
&h,
outdir.join(format!("h_t{frame:04}.tif")),
Some(write_opts(compress)),
)
.with_context(|| format!("Failed to write h frame {frame}"))?;
if dump_velocity {
let n = state.h.len();
let mut u = vec![0.0f32; n];
let mut v = vec![0.0f32; n];
for i in 0..n {
let hh = state.h[i];
if hh >= 1e-3 {
u[i] = state.hu[i] / hh;
v[i] = state.hv[i] / hh;
}
}
let u = masked_raster(grid, crs, u);
let v = masked_raster(grid, crs, v);
write_geotiff(
&u,
outdir.join(format!("u_t{frame:04}.tif")),
Some(write_opts(compress)),
)?;
write_geotiff(
&v,
outdir.join(format!("v_t{frame:04}.tif")),
Some(write_opts(compress)),
)?;
}
if dump_erosion {
let e = sim.eroded_depth();
if !e.is_empty() {
let e = masked_raster(grid, crs, e.to_vec());
write_geotiff(
&e,
outdir.join(format!("e_t{frame:04}.tif")),
Some(write_opts(compress)),
)?;
}
}
Ok(())
}
fn masked_raster(
grid: &SimGrid,
crs: Option<&surtgis_core::CRS>,
mut values: Vec<f32>,
) -> Raster<f32> {
let (rows, cols) = (grid.rows(), grid.cols());
for r in 0..rows {
for c in 0..cols {
if grid.is_solid(r, c) {
values[r * cols + c] = f32::NAN;
}
}
}
let mut raster = Raster::from_vec(values, rows, cols).expect("state length matches grid");
raster.set_transform(*grid.transform());
raster.set_crs(crs.cloned());
raster.set_nodata(Some(f32::NAN));
raster
}
fn write_manifest(
outdir: &Path,
dem: &Raster<f32>,
dt_output: f64,
n_frames: usize,
mu: f32,
xi: f32,
entrainment: Option<(EntrainmentParams, f64)>,
) -> Result<()> {
let t = dem.transform();
let crs = dem.crs().map(std::string::ToString::to_string);
let clean = |v: f32| (f64::from(v) * 1e6).round() / 1e6;
let mut manifest = serde_json::json!({
"crs": crs,
"origin": [t.origin_x, t.origin_y],
"cellsize": t.pixel_width,
"dt_output": dt_output,
"n_frames": n_frames,
"duration": dt_output * (n_frames as f64 - 1.0),
"row0": "north",
"mu": clean(mu),
"xi": clean(xi),
"units": { "h": "m", "u": "m/s", "v": "m/s", "arrival": "s" },
});
if let Some((p, total_eroded)) = entrainment {
manifest["manifest_version"] = serde_json::json!(2);
manifest["units"]["e"] = serde_json::json!("m");
manifest["entrainment"] = serde_json::json!({
"k": f64::from(p.k),
"rate_max": f64::from(p.rate_max),
"v_entr_min": f64::from(p.v_entr_min),
"f_max": f64::from(p.f_max),
"total_eroded_m3": total_eroded.round(),
});
}
let path = outdir.join("manifest.json");
std::fs::write(&path, serde_json::to_string_pretty(&manifest)?)
.context("Failed to write manifest.json")?;
println!("Manifest saved to: {}", path.display());
Ok(())
}