use crate::context::Context;
use crate::error::{Error, Result};
use crate::od::{DetermineResult, ODConfig, Observations};
use std::ptr::NonNull;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct SessionDiff {
pub reduced_chi2_delta: f64,
pub iterations_delta: i64,
pub n_observations_delta: i64,
pub update_norm_current: f64,
pub update_norm_prior: f64,
}
impl SessionDiff {
fn from_ffi(d: &empyrean_sys::EmpyreanSessionDiff) -> Self {
Self {
reduced_chi2_delta: d.reduced_chi2_delta,
iterations_delta: d.iterations_delta,
n_observations_delta: d.n_observations_delta,
update_norm_current: d.update_norm_current,
update_norm_prior: d.update_norm_prior,
}
}
}
pub struct Session {
raw: NonNull<empyrean_sys::EmpyreanSession>,
}
unsafe impl Send for Session {}
impl Session {
pub fn new(observations: Observations, config: ODConfig) -> Result<Self> {
let (obs_ptr, obs_len) = observations.as_ffi_slice();
let (ffi_config, _perturbers_keep) = config.to_ffi_with();
let raw = unsafe { empyrean_sys::empyrean_session_new(obs_ptr, obs_len, &ffi_config) };
let _ = observations;
NonNull::new(raw)
.map(|raw| Session { raw })
.ok_or_else(Error::from_null_ptr)
}
pub fn n_observations(&self) -> usize {
unsafe { empyrean_sys::empyrean_session_n_observations(self.raw.as_ptr()) }
}
pub fn n_masked(&self) -> usize {
unsafe { empyrean_sys::empyrean_session_n_masked(self.raw.as_ptr()) }
}
pub fn n_active(&self) -> usize {
unsafe { empyrean_sys::empyrean_session_n_active(self.raw.as_ptr()) }
}
pub fn mask(&mut self, idx: usize) -> Result<()> {
let code = unsafe { empyrean_sys::empyrean_session_mask(self.raw.as_ptr(), idx) };
if code != 0 {
return Err(Error::capture(code));
}
Ok(())
}
pub fn unmask(&mut self, idx: usize) -> Result<()> {
let code = unsafe { empyrean_sys::empyrean_session_unmask(self.raw.as_ptr(), idx) };
if code != 0 {
return Err(Error::capture(code));
}
Ok(())
}
pub fn unmask_all(&mut self) -> Result<()> {
let code = unsafe { empyrean_sys::empyrean_session_unmask_all(self.raw.as_ptr()) };
if code != 0 {
return Err(Error::capture(code));
}
Ok(())
}
pub fn is_masked(&self, idx: usize) -> bool {
unsafe { empyrean_sys::empyrean_session_is_masked(self.raw.as_ptr(), idx) == 1 }
}
pub fn refine(&mut self, ctx: &Context) -> Result<DetermineResult> {
let mut ffi_result = empyrean_sys::EmpyreanODResult::default();
let code = unsafe {
empyrean_sys::empyrean_session_refine(self.raw.as_ptr(), ctx.as_raw(), &mut ffi_result)
};
if code != 0 {
return Err(Error::capture(code));
}
let det = od_result_from_ffi(&ffi_result);
unsafe { empyrean_sys::empyrean_od_result_free(&mut ffi_result) };
det
}
pub fn history_len(&self) -> usize {
unsafe { empyrean_sys::empyrean_session_history_len(self.raw.as_ptr()) }
}
pub fn history(&self, idx: usize) -> Result<DetermineResult> {
let mut ffi_result = empyrean_sys::EmpyreanODResult::default();
let code = unsafe {
empyrean_sys::empyrean_session_get_history(self.raw.as_ptr(), idx, &mut ffi_result)
};
if code != 0 {
return Err(Error::capture(code));
}
let det = od_result_from_ffi(&ffi_result);
unsafe { empyrean_sys::empyrean_od_result_free(&mut ffi_result) };
det
}
pub fn diff(&self, prior_idx: usize) -> Result<SessionDiff> {
let mut ffi_diff = empyrean_sys::EmpyreanSessionDiff {
reduced_chi2_delta: 0.0,
iterations_delta: 0,
n_observations_delta: 0,
update_norm_current: 0.0,
update_norm_prior: 0.0,
};
let code = unsafe {
empyrean_sys::empyrean_session_diff(self.raw.as_ptr(), prior_idx, &mut ffi_diff)
};
if code != 0 {
return Err(Error::capture(code));
}
Ok(SessionDiff::from_ffi(&ffi_diff))
}
}
impl Drop for Session {
fn drop(&mut self) {
unsafe { empyrean_sys::empyrean_session_free(self.raw.as_ptr()) }
}
}
fn od_result_from_ffi(result: &empyrean_sys::EmpyreanODResult) -> Result<DetermineResult> {
crate::od::ffi_od_result_to_rust_pub(result)
}