use pyo3::prelude::*;
use tempoch::qtty::unit::{Day, Hour, Second};
use tempoch::{Interval, ModifiedJulianDate, Period, MJD, TT, UTC};
use crate::errors::{ensure_finite, map_conversion_error, map_invalid_interval_error};
use crate::jd::PyJulianDate;
use crate::mjd::PyModifiedJulianDate;
#[pyclass(name = "TimePeriod", module = "tempoch", from_py_object)]
#[derive(Clone, Copy)]
pub struct PyTimePeriod {
inner: Period<TT>,
}
impl PyTimePeriod {
pub fn from_inner(inner: Period<TT>) -> Self {
Self { inner }
}
}
#[allow(clippy::wrong_self_convention)]
#[pymethods]
impl PyTimePeriod {
#[new]
fn new(start_mjd: f64, end_mjd: f64) -> PyResult<Self> {
ensure_finite(start_mjd)?;
ensure_finite(end_mjd)?;
let start = ModifiedJulianDate::<TT>::new(start_mjd).to_j2000s();
let end = ModifiedJulianDate::<TT>::new(end_mjd).to_j2000s();
let inner = Interval::try_new(start, end).map_err(map_invalid_interval_error)?;
Ok(Self { inner })
}
#[staticmethod]
fn from_mjd(start: &PyModifiedJulianDate, end: &PyModifiedJulianDate) -> PyResult<Self> {
let inner = Interval::try_new(start.inner.to_j2000s(), end.inner.to_j2000s())
.map_err(map_invalid_interval_error)?;
Ok(Self { inner })
}
#[staticmethod]
fn from_jd(start: &PyJulianDate, end: &PyJulianDate) -> PyResult<Self> {
let start_mjd = start.inner.to_j2000s();
let end_mjd = end.inner.to_j2000s();
let inner = Interval::try_new(start_mjd, end_mjd).map_err(map_invalid_interval_error)?;
Ok(Self { inner })
}
#[staticmethod]
fn from_utc(start_utc: &str, end_utc: &str) -> PyResult<Self> {
use chrono::{DateTime, FixedOffset, Utc};
let start_dt: DateTime<FixedOffset> = start_utc.parse().map_err(|e| {
pyo3::exceptions::PyValueError::new_err(format!("Invalid start UTC: {e}"))
})?;
let end_dt: DateTime<FixedOffset> = end_utc.parse().map_err(|e| {
pyo3::exceptions::PyValueError::new_err(format!("Invalid end UTC: {e}"))
})?;
let start = tempoch::Time::<UTC>::try_from_chrono(start_dt.with_timezone(&Utc))
.map_err(map_conversion_error)?
.to::<TT>();
let end = tempoch::Time::<UTC>::try_from_chrono(end_dt.with_timezone(&Utc))
.map_err(map_conversion_error)?
.to::<TT>();
let inner = Interval::try_new(start, end).map_err(map_invalid_interval_error)?;
Ok(Self { inner })
}
#[getter]
fn start(&self) -> PyModifiedJulianDate {
PyModifiedJulianDate::from_inner(self.inner.start.to::<MJD>())
}
#[getter]
fn end(&self) -> PyModifiedJulianDate {
PyModifiedJulianDate::from_inner(self.inner.end.to::<MJD>())
}
#[getter]
fn start_mjd(&self) -> f64 {
self.inner.start.to::<MJD>().value()
}
#[getter]
fn end_mjd(&self) -> f64 {
self.inner.end.to::<MJD>().value()
}
fn duration_days(&self) -> f64 {
(self.inner.end - self.inner.start).to::<Day>().value()
}
fn duration_seconds(&self) -> f64 {
(self.inner.end - self.inner.start).to::<Second>().value()
}
fn duration_hours(&self) -> f64 {
(self.inner.end - self.inner.start).to::<Hour>().value()
}
fn to_utc(&self) -> PyResult<(String, String)> {
let start_utc = self
.inner
.start
.to::<UTC>()
.try_to_chrono()
.map_err(map_conversion_error)?;
let end_utc = self
.inner
.end
.to::<UTC>()
.try_to_chrono()
.map_err(map_conversion_error)?;
Ok((start_utc.to_rfc3339(), end_utc.to_rfc3339()))
}
fn intersection(&self, other: &PyTimePeriod) -> Option<Self> {
self.inner
.intersection(&other.inner)
.map(|p| Self { inner: p })
}
fn contains(&self, mjd: f64) -> bool {
if !mjd.is_finite() {
return false;
}
let t = ModifiedJulianDate::<TT>::new(mjd).to_j2000s();
t >= self.inner.start && t < self.inner.end
}
fn contains_mjd(&self, mjd: &PyModifiedJulianDate) -> bool {
let instant = mjd.inner.to_j2000s();
instant >= self.inner.start && instant < self.inner.end
}
fn __eq__(&self, other: &PyTimePeriod) -> bool {
self.inner == other.inner
}
fn __ne__(&self, other: &PyTimePeriod) -> bool {
self.inner != other.inner
}
fn __repr__(&self) -> String {
format!(
"TimePeriod({}, {})",
self.inner.start.to::<MJD>().value(),
self.inner.end.to::<MJD>().value()
)
}
fn __str__(&self) -> String {
format!(
"TimePeriod(MJD {} to {})",
self.inner.start.to::<MJD>().value(),
self.inner.end.to::<MJD>().value()
)
}
fn __reduce__(&self, py: Python<'_>) -> PyResult<(Py<PyAny>, (f64, f64))> {
let cls = py.get_type::<Self>().into_any().unbind();
Ok((
cls,
(
self.inner.start.to::<MJD>().value(),
self.inner.end.to::<MJD>().value(),
),
))
}
fn __hash__(&self) -> u64 {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
self.inner
.start
.to::<MJD>()
.value()
.to_bits()
.hash(&mut hasher);
self.inner
.end
.to::<MJD>()
.value()
.to_bits()
.hash(&mut hasher);
hasher.finish()
}
}
#[pyfunction]
pub fn intersect_periods_py(
periods: Vec<PyTimePeriod>,
bounds: &PyTimePeriod,
) -> Vec<PyTimePeriod> {
let rust_periods: Vec<Period<TT>> = periods.iter().map(|p| p.inner).collect();
let bound_period = bounds.inner;
rust_periods
.iter()
.filter_map(|p| p.intersection(&bound_period))
.map(PyTimePeriod::from_inner)
.collect()
}