use crate::error::{ForecastError, Result};
#[derive(Debug, Clone, Copy)]
pub struct Signal<'a> {
data: &'a [f64],
n: usize,
d: usize,
}
impl<'a> Signal<'a> {
pub fn univariate(values: &'a [f64]) -> Self {
Self {
data: values,
n: values.len(),
d: 1,
}
}
pub fn from_row_major(data: &'a [f64], n: usize, d: usize) -> Result<Self> {
if d == 0 {
return Err(ForecastError::InvalidParameter(
"Signal: dimension d must be ≥ 1".into(),
));
}
if data.len() != n * d {
return Err(ForecastError::DimensionMismatch {
expected: n * d,
got: data.len(),
});
}
Ok(Self { data, n, d })
}
pub fn n(&self) -> usize {
self.n
}
pub fn d(&self) -> usize {
self.d
}
pub fn is_univariate(&self) -> bool {
self.d == 1
}
pub fn row(&self, i: usize) -> &[f64] {
let start = i * self.d;
let end = start + self.d;
&self.data[start..end]
}
pub fn as_slice(&self) -> &[f64] {
self.data
}
}
impl<'a> From<&'a [f64]> for Signal<'a> {
fn from(values: &'a [f64]) -> Self {
Signal::univariate(values)
}
}
impl<'a, const N: usize> From<&'a [f64; N]> for Signal<'a> {
fn from(values: &'a [f64; N]) -> Self {
Signal::univariate(values.as_slice())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn univariate_from_slice_round_trip() {
let values = [1.0, 2.0, 3.0, 4.0];
let s: Signal = (&values[..]).into();
assert_eq!(s.n(), 4);
assert_eq!(s.d(), 1);
assert!(s.is_univariate());
assert_eq!(s.row(2), &[3.0]);
}
#[test]
fn multivariate_row_major() {
let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let s = Signal::from_row_major(&data, 3, 2).unwrap();
assert_eq!(s.n(), 3);
assert_eq!(s.d(), 2);
assert!(!s.is_univariate());
assert_eq!(s.row(0), &[1.0, 2.0]);
assert_eq!(s.row(1), &[3.0, 4.0]);
assert_eq!(s.row(2), &[5.0, 6.0]);
}
#[test]
fn dimension_mismatch_errors() {
let data = vec![1.0, 2.0, 3.0];
let err = Signal::from_row_major(&data, 2, 2).unwrap_err();
assert!(matches!(err, ForecastError::DimensionMismatch { .. }));
}
#[test]
fn zero_dimension_rejected() {
let data = vec![1.0, 2.0];
let err = Signal::from_row_major(&data, 2, 0).unwrap_err();
assert!(matches!(err, ForecastError::InvalidParameter(_)));
}
}