Skip to main content

siderust_py/
interop.rs

1//! Cross-extension-safe interoperability with the installed `siderust` package.
2//!
3//! PyO3 classes are local to the extension module that registers them. A
4//! downstream extension must therefore not compile its own copy of
5//! `PyObserver` or `PyDirection` and expect Python type identity to match. This
6//! module imports the installed canonical extension and asks it to extract or
7//! construct its own classes. Only primitive scalar values (`f64` payload
8//! fields and the `u32` protocol version) cross that boundary. The installed
9//! extension must expose the same [`BRIDGE_PROTOCOL_VERSION`].
10
11use pyo3::exceptions::PyImportError;
12use pyo3::prelude::*;
13use pyo3::types::PyModule;
14use siderust::coordinates::centers::Geodetic;
15use siderust::coordinates::frames::ECEF;
16use siderust::coordinates::spherical::direction;
17use siderust::qtty::{Degrees, Meters};
18
19const EXTENSION_MODULE: &str = "siderust._siderust";
20const BRIDGE_PROTOCOL_ATTRIBUTE: &str = "_bridge_protocol_version";
21const OBSERVER_FROM_PARTS: &str = "_bridge_observer_from_parts";
22const OBSERVER_TO_PARTS: &str = "_bridge_observer_to_parts";
23const DIRECTION_FROM_PARTS: &str = "_bridge_direction_from_parts";
24const DIRECTION_TO_PARTS: &str = "_bridge_direction_to_parts";
25
26/// Protocol implemented by the canonical Python extension bridge hooks.
27///
28/// Downstream extensions and the installed `siderust` Python package must use
29/// compatible siderust-py releases which expose this same protocol version.
30pub const BRIDGE_PROTOCOL_VERSION: u32 = 1;
31
32fn bridge_module<'py>(py: Python<'py>) -> PyResult<Bound<'py, PyModule>> {
33    let module = PyModule::import(py, EXTENSION_MODULE)?;
34    let actual = module
35        .getattr(BRIDGE_PROTOCOL_ATTRIBUTE)
36        .and_then(|value| value.extract::<u32>())
37        .map_err(|_| {
38            PyImportError::new_err(format!(
39                "installed {EXTENSION_MODULE} does not expose a valid \
40                 {BRIDGE_PROTOCOL_ATTRIBUTE}; expected bridge protocol \
41                 {BRIDGE_PROTOCOL_VERSION}. Install a compatible siderust Python package"
42            ))
43        })?;
44
45    if actual != BRIDGE_PROTOCOL_VERSION {
46        return Err(PyImportError::new_err(format!(
47            "incompatible siderust bridge protocol: expected \
48             {BRIDGE_PROTOCOL_VERSION}, found {actual}. Install matching \
49             siderust-py Rust and Python package versions"
50        )));
51    }
52
53    Ok(module)
54}
55
56/// Verify that the installed canonical extension supports this bridge API.
57pub fn ensure_bridge_protocol(py: Python<'_>) -> PyResult<()> {
58    bridge_module(py).map(|_| ())
59}
60
61/// Stable primitive representation of a WGS84 `Geodetic<ECEF>` observer.
62#[derive(Debug, Clone, Copy, PartialEq)]
63pub struct ObserverParts {
64    /// East-positive geodetic longitude in degrees.
65    pub longitude_degrees: f64,
66    /// North-positive geodetic latitude in degrees.
67    pub latitude_degrees: f64,
68    /// Ellipsoidal height above WGS84 in metres.
69    pub height_metres: f64,
70}
71
72impl ObserverParts {
73    /// Convert these WGS84 parts into Siderust's current observer type.
74    pub fn into_observer(self) -> Geodetic<ECEF> {
75        Geodetic::<ECEF>::new(
76            Degrees::new(self.longitude_degrees),
77            Degrees::new(self.latitude_degrees),
78            Meters::new(self.height_metres),
79        )
80    }
81}
82
83impl From<&Geodetic<ECEF>> for ObserverParts {
84    fn from(observer: &Geodetic<ECEF>) -> Self {
85        Self {
86            longitude_degrees: observer.lon.value(),
87            latitude_degrees: observer.lat.value(),
88            height_metres: observer.height.value(),
89        }
90    }
91}
92
93/// Stable primitive representation of an ICRS spherical direction.
94#[derive(Debug, Clone, Copy, PartialEq)]
95pub struct DirectionParts {
96    /// Right ascension in degrees in the ICRS frame.
97    pub right_ascension_degrees: f64,
98    /// Declination in degrees in the ICRS frame.
99    pub declination_degrees: f64,
100}
101
102impl DirectionParts {
103    /// Convert these ICRS parts into Siderust's current direction type.
104    pub fn into_direction(self) -> direction::ICRS {
105        direction::ICRS::new(
106            Degrees::new(self.right_ascension_degrees),
107            Degrees::new(self.declination_degrees),
108        )
109    }
110}
111
112impl From<&direction::ICRS> for DirectionParts {
113    fn from(value: &direction::ICRS) -> Self {
114        Self {
115            right_ascension_degrees: value.azimuth.value(),
116            declination_degrees: value.polar.value(),
117        }
118    }
119}
120
121/// Extract a canonical Python `siderust.Observer` into primitive parts.
122pub fn observer_parts_from_python(value: &Bound<'_, PyAny>) -> PyResult<ObserverParts> {
123    let (longitude_degrees, latitude_degrees, height_metres): (f64, f64, f64) =
124        bridge_module(value.py())?
125            .getattr(OBSERVER_TO_PARTS)?
126            .call1((value,))?
127            .extract()?;
128    Ok(ObserverParts {
129        longitude_degrees,
130        latitude_degrees,
131        height_metres,
132    })
133}
134
135/// Extract a canonical Python `siderust.Observer` as `Geodetic<ECEF>`.
136pub fn observer_from_python(value: &Bound<'_, PyAny>) -> PyResult<Geodetic<ECEF>> {
137    Ok(observer_parts_from_python(value)?.into_observer())
138}
139
140/// Construct the actual `siderust.Observer` class owned by the installed package.
141pub fn observer_to_python(py: Python<'_>, observer: &Geodetic<ECEF>) -> PyResult<Py<PyAny>> {
142    let parts = ObserverParts::from(observer);
143    bridge_module(py)?
144        .getattr(OBSERVER_FROM_PARTS)?
145        .call1((
146            parts.longitude_degrees,
147            parts.latitude_degrees,
148            parts.height_metres,
149        ))
150        .map(Bound::unbind)
151}
152
153/// Extract a canonical Python `siderust.Direction` into primitive parts.
154pub fn direction_parts_from_python(value: &Bound<'_, PyAny>) -> PyResult<DirectionParts> {
155    let (right_ascension_degrees, declination_degrees): (f64, f64) = bridge_module(value.py())?
156        .getattr(DIRECTION_TO_PARTS)?
157        .call1((value,))?
158        .extract()?;
159    Ok(DirectionParts {
160        right_ascension_degrees,
161        declination_degrees,
162    })
163}
164
165/// Extract a canonical Python `siderust.Direction` as the current ICRS type.
166pub fn direction_from_python(value: &Bound<'_, PyAny>) -> PyResult<direction::ICRS> {
167    Ok(direction_parts_from_python(value)?.into_direction())
168}
169
170/// Construct the actual `siderust.Direction` class owned by the installed package.
171pub fn direction_to_python(py: Python<'_>, value: &direction::ICRS) -> PyResult<Py<PyAny>> {
172    let parts = DirectionParts::from(value);
173    bridge_module(py)?
174        .getattr(DIRECTION_FROM_PARTS)?
175        .call1((parts.right_ascension_degrees, parts.declination_degrees))
176        .map(Bound::unbind)
177}
178
179#[cfg(test)]
180mod tests {
181    use super::*;
182
183    #[test]
184    fn observer_parts_round_trip() {
185        let parts = ObserverParts {
186            longitude_degrees: -17.8925,
187            latitude_degrees: 28.7543,
188            height_metres: 2396.0,
189        };
190        let round_trip = ObserverParts::from(&parts.into_observer());
191        assert!((round_trip.longitude_degrees - parts.longitude_degrees).abs() < 1e-12);
192        assert!((round_trip.latitude_degrees - parts.latitude_degrees).abs() < 1e-12);
193        assert_eq!(round_trip.height_metres, parts.height_metres);
194    }
195
196    #[test]
197    fn direction_parts_round_trip() {
198        let parts = DirectionParts {
199            right_ascension_degrees: 83.633,
200            declination_degrees: 22.014,
201        };
202        let round_trip = DirectionParts::from(&parts.into_direction());
203        assert!((round_trip.right_ascension_degrees - parts.right_ascension_degrees).abs() < 1e-12);
204        assert!((round_trip.declination_degrees - parts.declination_degrees).abs() < 1e-12);
205    }
206}