#![forbid(unsafe_code)]
use crate::{TransBuild, TransParams};
use oxiproj_core::{Coord, IoUnits, Operation, ProjError, ProjResult};
use oxiproj_grids::{read_geotiff, read_gtx, read_ntv2, sample_grid, GridSet};
use std::f64::consts::PI;
const ARCSEC_TO_RAD: f64 = PI / 648_000.0;
const MAX_ITER: usize = 10;
const ITER_TOL: f64 = 1e-12;
#[derive(Debug)]
enum ShiftKind {
HorizontalArcSec,
Vertical,
HorizontalDeg,
Xyz,
}
#[derive(Debug)]
struct GridShift {
grids: Vec<GridSet>,
kind: ShiftKind,
}
impl Operation for GridShift {
fn forward_4d(&self, c: Coord) -> ProjResult<Coord> {
let v = c.v();
let lon_deg = v[0].to_degrees();
let lat_deg = v[1].to_degrees();
for gs in &self.grids {
if let Some(shifts) = sample_grid(gs, lat_deg, lon_deg) {
return match self.kind {
ShiftKind::HorizontalArcSec => Ok(Coord::new(
v[0] + shifts[1] * ARCSEC_TO_RAD,
v[1] + shifts[0] * ARCSEC_TO_RAD,
v[2],
v[3],
)),
ShiftKind::Vertical => Ok(Coord::new(v[0], v[1], v[2] + shifts[0], v[3])),
ShiftKind::HorizontalDeg => Ok(Coord::new(
v[0] + shifts[0].to_radians(),
v[1] + shifts[1].to_radians(),
v[2],
v[3],
)),
ShiftKind::Xyz => {
let dz = if shifts.len() >= 3 { shifts[2] } else { 0.0 };
Ok(Coord::new(
v[0] + shifts[0].to_radians(),
v[1] + shifts[1].to_radians(),
v[2] + dz,
v[3],
))
}
};
}
}
Err(ProjError::OutsideGrid)
}
fn inverse_4d(&self, c: Coord) -> ProjResult<Coord> {
let v = c.v();
match self.kind {
ShiftKind::Vertical => {
for gs in &self.grids {
if let Some(shifts) = sample_grid(gs, v[1].to_degrees(), v[0].to_degrees()) {
return Ok(Coord::new(v[0], v[1], v[2] - shifts[0], v[3]));
}
}
Err(ProjError::OutsideGrid)
}
_ => {
let mut lon_r = v[0];
let mut lat_r = v[1];
'outer: for gs in &self.grids {
if sample_grid(gs, lat_r.to_degrees(), lon_r.to_degrees()).is_none() {
continue;
}
for _ in 0..MAX_ITER {
let shifts = match sample_grid(gs, lat_r.to_degrees(), lon_r.to_degrees()) {
Some(s) => s,
None => break 'outer,
};
let (new_lon, new_lat) = match self.kind {
ShiftKind::HorizontalArcSec => (
v[0] - shifts[1] * ARCSEC_TO_RAD,
v[1] - shifts[0] * ARCSEC_TO_RAD,
),
ShiftKind::HorizontalDeg | ShiftKind::Xyz => {
(v[0] - shifts[0].to_radians(), v[1] - shifts[1].to_radians())
}
ShiftKind::Vertical => (v[0], v[1]),
};
let dlon = (new_lon - lon_r).abs();
let dlat = (new_lat - lat_r).abs();
lon_r = new_lon;
lat_r = new_lat;
if dlon < ITER_TOL && dlat < ITER_TOL {
break;
}
}
if let ShiftKind::Xyz = self.kind {
if let Some(shifts) =
sample_grid(gs, lat_r.to_degrees(), lon_r.to_degrees())
{
let dz = if shifts.len() >= 3 { shifts[2] } else { 0.0 };
return Ok(Coord::new(lon_r, lat_r, v[2] - dz, v[3]));
}
}
return Ok(Coord::new(lon_r, lat_r, v[2], v[3]));
}
Err(ProjError::OutsideGrid)
}
}
}
fn has_inverse(&self) -> bool {
true
}
}
pub fn new(p: &TransParams) -> ProjResult<TransBuild> {
let grid_name = p.params.get_str("grids").ok_or(ProjError::MissingArg)?;
let registry = p.registry.ok_or(ProjError::FileNotFound)?;
let data = registry
.get_grid(grid_name)
.ok_or(ProjError::FileNotFound)?;
if let Ok(grids) = read_ntv2(data, grid_name) {
if !grids.is_empty() && grids[0].bands.len() >= 2 {
return Ok(TransBuild::new(
Box::new(GridShift {
grids,
kind: ShiftKind::HorizontalArcSec,
}),
IoUnits::Radians,
IoUnits::Radians,
));
}
}
if let Ok(gs) = read_gtx(data, grid_name) {
if gs.bands.len() == 1 {
return Ok(TransBuild::new(
Box::new(GridShift {
grids: vec![gs],
kind: ShiftKind::Vertical,
}),
IoUnits::Radians,
IoUnits::Radians,
));
}
}
if let Ok(gs) = read_geotiff(data, grid_name) {
let kind = match gs.bands.len() {
1 => ShiftKind::Vertical,
2 => ShiftKind::HorizontalDeg,
_ => ShiftKind::Xyz,
};
return Ok(TransBuild::new(
Box::new(GridShift {
grids: vec![gs],
kind,
}),
IoUnits::Radians,
IoUnits::Radians,
));
}
Err(ProjError::FileNotFound)
}
#[cfg(test)]
mod tests {
use super::*;
use oxiproj_core::DEG_TO_RAD;
fn build_ntv2(lat_shift: f32, lon_shift: f32) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(b"NUM_OREC");
buf.extend_from_slice(&11i32.to_le_bytes());
buf.extend_from_slice(&[0u8; 4]);
buf.extend_from_slice(b"NUM_SREC");
buf.extend_from_slice(&11i32.to_le_bytes());
buf.extend_from_slice(&[0u8; 4]);
buf.extend_from_slice(b"NUM_FILE");
buf.extend_from_slice(&1u32.to_le_bytes());
buf.extend_from_slice(&[0u8; 4]);
buf.extend_from_slice(b"GS_TYPE ");
buf.extend_from_slice(b"SECONDS ");
buf.extend_from_slice(&[0u8; 112]);
buf.extend_from_slice(b"SUB_NAME");
buf.extend_from_slice(b"TESTGRID");
buf.extend_from_slice(b"PARENT ");
buf.extend_from_slice(b"NONE ");
buf.extend_from_slice(b"CREATED ");
buf.extend_from_slice(b"20240101");
buf.extend_from_slice(b"UPDATED ");
buf.extend_from_slice(b"20240101");
buf.extend_from_slice(b"S_LAT ");
buf.extend_from_slice(&0.0f64.to_le_bytes());
buf.extend_from_slice(b"N_LAT ");
buf.extend_from_slice(&3600.0f64.to_le_bytes());
buf.extend_from_slice(b"E_LONG ");
buf.extend_from_slice(&0.0f64.to_le_bytes());
buf.extend_from_slice(b"W_LONG ");
buf.extend_from_slice(&3600.0f64.to_le_bytes());
buf.extend_from_slice(b"LAT_INC ");
buf.extend_from_slice(&3600.0f64.to_le_bytes());
buf.extend_from_slice(b"LONG_INC");
buf.extend_from_slice(&3600.0f64.to_le_bytes());
buf.extend_from_slice(b"GS_COUNT");
buf.extend_from_slice(&4i32.to_le_bytes());
buf.extend_from_slice(&[0u8; 4]);
for _ in 0..4 {
buf.extend_from_slice(&lat_shift.to_le_bytes());
buf.extend_from_slice(&lon_shift.to_le_bytes());
buf.extend_from_slice(&0.0f32.to_le_bytes());
buf.extend_from_slice(&0.0f32.to_le_bytes());
}
buf
}
fn build_gtx(shift: f32) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(&0.0f64.to_be_bytes()); buf.extend_from_slice(&0.0f64.to_be_bytes()); buf.extend_from_slice(&1.0f64.to_be_bytes()); buf.extend_from_slice(&1.0f64.to_be_bytes()); buf.extend_from_slice(&2i32.to_be_bytes()); buf.extend_from_slice(&2i32.to_be_bytes()); for _ in 0..4 {
buf.extend_from_slice(&shift.to_be_bytes());
}
buf
}
struct InMemReg(std::collections::HashMap<String, Vec<u8>>);
impl crate::GridRegistry for InMemReg {
fn get_grid(&self, name: &str) -> Option<&[u8]> {
self.0.get(name).map(|v| v.as_slice())
}
}
struct TestParams {
grids: String,
}
impl crate::TransParamLookup for TestParams {
fn get_str(&self, key: &str) -> Option<&str> {
if key == "grids" {
Some(&self.grids)
} else {
None
}
}
fn get_f64(&self, _: &str) -> Option<f64> {
None
}
fn get_dms(&self, _: &str) -> Option<f64> {
None
}
fn get_int(&self, _: &str) -> Option<i64> {
None
}
fn get_bool(&self, _: &str) -> bool {
false
}
fn exists(&self, key: &str) -> bool {
key == "grids"
}
}
#[test]
fn test_gridshift_detects_ntv2() {
let data = build_ntv2(1800.0, 900.0);
let mut map = std::collections::HashMap::new();
map.insert("grid.gsb".to_string(), data);
let reg = InMemReg(map);
let tp = TestParams {
grids: "grid.gsb".to_string(),
};
let ell = oxiproj_core::Ellipsoid::named("WGS84").unwrap();
let p = crate::TransParams {
ellipsoid: &ell,
params: &tp,
registry: Some(®),
};
let tb = new(&p).unwrap();
let input = Coord::new(-0.5 * DEG_TO_RAD, 0.5 * DEG_TO_RAD, 0.0, 0.0);
let out = tb.operation.forward_4d(input).unwrap();
let expected_lat = (0.5 + 0.5) * DEG_TO_RAD;
let expected_lon = (-0.5 + 0.25) * DEG_TO_RAD;
assert!((out.v()[1] - expected_lat).abs() < 1e-9, "lat");
assert!((out.v()[0] - expected_lon).abs() < 1e-9, "lon");
}
#[test]
fn test_gridshift_detects_gtx() {
let data = build_gtx(5.0);
let mut map = std::collections::HashMap::new();
map.insert("grid.gtx".to_string(), data);
let reg = InMemReg(map);
let tp = TestParams {
grids: "grid.gtx".to_string(),
};
let ell = oxiproj_core::Ellipsoid::named("WGS84").unwrap();
let p = crate::TransParams {
ellipsoid: &ell,
params: &tp,
registry: Some(®),
};
let tb = new(&p).unwrap();
let input = Coord::new(0.5 * DEG_TO_RAD, 0.5 * DEG_TO_RAD, 0.0, 0.0);
let out = tb.operation.forward_4d(input).unwrap();
assert!((out.v()[2] - 5.0).abs() < 1e-6, "z should be 5.0");
}
#[test]
fn test_gridshift_round_trip_ntv2() {
let data = build_ntv2(1800.0, 900.0);
let mut map = std::collections::HashMap::new();
map.insert("grid.gsb".to_string(), data);
let reg = InMemReg(map);
let tp = TestParams {
grids: "grid.gsb".to_string(),
};
let ell = oxiproj_core::Ellipsoid::named("WGS84").unwrap();
let p = crate::TransParams {
ellipsoid: &ell,
params: &tp,
registry: Some(®),
};
let tb = new(&p).unwrap();
let input = Coord::new(-0.7 * DEG_TO_RAD, 0.4 * DEG_TO_RAD, 0.0, 0.0);
let fwd = tb.operation.forward_4d(input).unwrap();
let inv = tb.operation.inverse_4d(fwd).unwrap();
let vi = inv.v();
let vi0 = input.v();
assert!((vi[0] - vi0[0]).abs() < 1e-9, "lon round-trip");
assert!((vi[1] - vi0[1]).abs() < 1e-9, "lat round-trip");
}
}