oxiproj-transformations 0.1.1

Datum transformations and coordinate conversions for OxiProj.
Documentation
#![forbid(unsafe_code)]
//! Generic grid shift (`gridshift`) — port of PROJ 9.8.0 generic gridshift.
//! Detects the grid format from raw bytes (NTv2, GTX, GeoTIFF) and applies
//! the appropriate horizontal, vertical, or 3-D shift. Acts as a dispatcher.

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;

/// Determines how the interpolated band values are interpreted as shifts.
#[derive(Debug)]
enum ShiftKind {
    /// 2-band NTv2: band0=lat-shift (arcsec), band1=lon-shift (arcsec).
    HorizontalArcSec,
    /// 1-band GTX: band0=vertical shift (meters).
    Vertical,
    /// 2-band (degrees): band0=lon-shift, band1=lat-shift (both in degrees).
    HorizontalDeg,
    /// 3-band: band0=lon-shift (deg), band1=lat-shift (deg), band2=dz (m).
    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 => {
                // No iteration needed for pure vertical shift.
                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)
            }
            _ => {
                // Iterative inversion for horizontal components.
                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())
                            }
                            // Vertical is handled by the outer match arm above;
                            // this arm is logically unreachable, but we return
                            // the input coordinates unchanged rather than panicking.
                            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;
                        }
                    }
                    // Compute dz correction for Xyz kind.
                    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
    }
}

/// Construct a `gridshift` transform from parsed parameters.
///
/// Tries formats in order: NTv2, GTX, GeoTIFF. The shift kind is determined
/// by the detected format and band count.
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)?;

    // Try NTv2 first (most common for horizontal shifts).
    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,
            ));
        }
    }

    // Try GTX (vertical shifts).
    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,
            ));
        }
    }

    // Try GeoTIFF (multi-band).
    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> {
        // Minimal 2×2 NTv2 covering lat 0-1°, lon -1-0°
        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()); // south_lat=0
        buf.extend_from_slice(&0.0f64.to_be_bytes()); // west_lon=0
        buf.extend_from_slice(&1.0f64.to_be_bytes()); // lat_inc=1
        buf.extend_from_slice(&1.0f64.to_be_bytes()); // lon_inc=1
        buf.extend_from_slice(&2i32.to_be_bytes()); // rows=2
        buf.extend_from_slice(&2i32.to_be_bytes()); // cols=2
        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() {
        // NTv2 grid covers lat 0-1°, lon -1-0° (NTv2 positive-west convention)
        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(&reg),
        };
        let tb = new(&p).unwrap();
        // Point inside the grid: lat=0.5°, lon=-0.5°
        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();
        // NTv2 lat_shift=1800 arcsec → 0.5 deg shift, lon_shift=900 arcsec → 0.25 deg shift
        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() {
        // GTX grid covers lat 0-1°, lon 0-1°
        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(&reg),
        };
        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() {
        // NTv2 grid covers lat 0-1°, lon -1-0°
        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(&reg),
        };
        let tb = new(&p).unwrap();
        // Point inside the grid: lat=0.4°, lon=-0.7°
        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");
    }
}