use core::fmt;
use crate::ids::{Lag, VariableId};
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum TemporalIndexError {
Invalid {
message: &'static str,
},
UnknownVariable {
id: VariableId,
},
}
impl fmt::Display for TemporalIndexError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Invalid { message } => write!(f, "{message}"),
Self::UnknownVariable { id } => write!(f, "unknown variable id {}", id.raw()),
}
}
}
impl std::error::Error for TemporalIndexError {}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct TemporalNodeKey {
pub variable: VariableId,
pub offset: i32,
}
impl TemporalNodeKey {
#[must_use]
pub const fn contemporaneous(variable: VariableId) -> Self {
Self { variable, offset: 0 }
}
#[must_use]
pub fn lagged(variable: VariableId, lag: Lag) -> Option<Self> {
let lag_i = i32::try_from(lag.raw()).ok()?;
Some(Self { variable, offset: -lag_i })
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct TemporalIndexer {
variable_count: u32,
history: u32,
horizon: u32,
dense_len: usize,
}
impl TemporalIndexer {
pub fn new(
variable_count: u32,
history: u32,
horizon: u32,
) -> Result<Self, TemporalIndexError> {
if variable_count == 0 || horizon == 0 {
return Err(TemporalIndexError::Invalid {
message: "variable_count and horizon must be non-zero",
});
}
let slices = history
.checked_add(horizon)
.ok_or(TemporalIndexError::Invalid { message: "history+horizon overflow" })?;
let dense_u32 = variable_count
.checked_mul(slices)
.ok_or(TemporalIndexError::Invalid { message: "dense index space overflow" })?;
let dense_len = usize::try_from(dense_u32)
.map_err(|_| TemporalIndexError::Invalid { message: "dense index space overflow" })?;
Ok(Self { variable_count, history, horizon, dense_len })
}
#[must_use]
pub const fn variable_count(&self) -> u32 {
self.variable_count
}
#[must_use]
pub const fn history(&self) -> u32 {
self.history
}
#[must_use]
pub const fn horizon(&self) -> u32 {
self.horizon
}
#[must_use]
pub const fn dense_len(&self) -> usize {
self.dense_len
}
pub fn dense_id(&self, key: TemporalNodeKey) -> Result<u32, TemporalIndexError> {
let v = key.variable.raw();
if v >= self.variable_count {
return Err(TemporalIndexError::UnknownVariable { id: key.variable });
}
let slice = i64::from(key.offset) + i64::from(self.history);
if slice < 0 || slice >= i64::from(self.history) + i64::from(self.horizon) {
return Err(TemporalIndexError::Invalid {
message: "temporal offset outside unfolding window",
});
}
let slice_u = u64::try_from(slice).map_err(|_| TemporalIndexError::Invalid {
message: "temporal offset outside unfolding window",
})?;
let dense = slice_u * u64::from(self.variable_count) + u64::from(v);
u32::try_from(dense)
.map_err(|_| TemporalIndexError::Invalid { message: "dense id exceeds u32" })
}
pub fn key_of(&self, dense: u32) -> Result<TemporalNodeKey, TemporalIndexError> {
let dense_usize = usize::try_from(dense)
.map_err(|_| TemporalIndexError::Invalid { message: "dense id out of range" })?;
if dense_usize >= self.dense_len() {
return Err(TemporalIndexError::Invalid { message: "dense id out of range" });
}
let vc = self.variable_count;
let slice = dense / vc;
let var = dense % vc;
let offset = i32::try_from(i64::from(slice) - i64::from(self.history))
.map_err(|_| TemporalIndexError::Invalid { message: "offset overflow" })?;
Ok(TemporalNodeKey { variable: VariableId::from_raw(var), offset })
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn time_major_dense_round_trip() {
let idx = TemporalIndexer::new(3, 2, 4).unwrap();
assert_eq!(idx.dense_len(), 18);
let key = TemporalNodeKey { variable: VariableId::from_raw(1), offset: -1 };
let dense = idx.dense_id(key).unwrap();
assert_eq!(dense, 4);
assert_eq!(idx.key_of(dense).unwrap(), key);
}
}