use crate::{Float, Result, TalibError};
#[cfg(not(feature = "std"))]
use alloc::{format, vec::Vec};
#[cfg(feature = "std")]
use std::{format, vec::Vec};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct OutputRange {
pub beg_idx: usize,
pub nb_element: usize,
}
impl OutputRange {
#[inline]
pub const fn new(beg_idx: usize, nb_element: usize) -> Self {
Self {
beg_idx,
nb_element,
}
}
#[inline]
pub const fn empty() -> Self {
Self {
beg_idx: 0,
nb_element: 0,
}
}
#[inline]
pub const fn end_idx(&self) -> usize {
self.beg_idx + self.nb_element
}
#[inline]
pub const fn is_empty(&self) -> bool {
self.nb_element == 0
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct CompactOutput<O> {
source_len: usize,
range: OutputRange,
values: O,
}
pub(crate) trait CompactPayloadLen {
fn compact_payload_len(&self) -> Result<usize>;
}
impl<T> CompactPayloadLen for Vec<T> {
#[inline]
fn compact_payload_len(&self) -> Result<usize> {
Ok(self.len())
}
}
impl<O> CompactOutput<O> {
pub(crate) fn new(source_len: usize, range: OutputRange, values: O) -> Result<Self>
where
O: CompactPayloadLen,
{
let end = range
.beg_idx
.checked_add(range.nb_element)
.ok_or_else(|| TalibError::invalid_input("Compact Output range overflow"))?;
if end > source_len {
return Err(TalibError::invalid_input(format!(
"Compact Output range {}..{} exceeds source length {}",
range.beg_idx, end, source_len
)));
}
let payload_len = values.compact_payload_len()?;
if payload_len != range.nb_element {
return Err(TalibError::invalid_input(format!(
"Compact Output payload length mismatch: range has {}, payload has {}",
range.nb_element, payload_len
)));
}
Ok(Self {
source_len,
range,
values,
})
}
#[inline]
pub const fn source_len(&self) -> usize {
self.source_len
}
#[inline]
pub const fn range(&self) -> OutputRange {
self.range
}
#[inline]
pub const fn values(&self) -> &O {
&self.values
}
#[inline]
pub fn into_values(self) -> O {
self.values
}
}
#[inline]
pub fn output_count(input_len: usize, lookback: usize) -> usize {
input_len.saturating_sub(lookback)
}
#[inline]
pub fn validate_period(name: &str, period: usize) -> Result<()> {
if period == 0 {
return Err(TalibError::invalid_period(
period,
format!("{} must be greater than zero", name),
));
}
Ok(())
}
#[inline]
pub fn period_lookback(name: &str, period: usize) -> Result<usize> {
validate_period(name, period)?;
Ok(period - 1)
}
#[inline]
pub fn validate_input_len(input_len: usize, lookback: usize) -> Result<usize> {
let count = output_count(input_len, lookback);
if count == 0 && input_len > 0 {
return Err(TalibError::insufficient_data(lookback + 1, input_len));
}
Ok(count)
}
#[inline]
pub fn validate_output_len(name: &str, output_len: usize, required: usize) -> Result<()> {
if output_len < required {
return Err(TalibError::invalid_input(format!(
"{} output buffer too small: need {}, got {}",
name, required, output_len
)));
}
Ok(())
}
#[inline]
pub fn validate_same_len(
left_name: &str,
left_len: usize,
right_name: &str,
right_len: usize,
) -> Result<()> {
if left_len != right_len {
return Err(TalibError::invalid_input(format!(
"{} and {} must have the same length: got {} and {}",
left_name, right_name, left_len, right_len
)));
}
Ok(())
}
pub fn validate_all_same_len(lengths: &[(&str, usize)]) -> Result<usize> {
let Some((first_name, first_len)) = lengths.first().copied() else {
return Ok(0);
};
for &(name, len) in &lengths[1..] {
validate_same_len(first_name, first_len, name, len)?;
}
Ok(first_len)
}
#[cold]
#[inline(never)]
fn non_finite_value_error(name: &str, idx: usize, value: Float) -> TalibError {
TalibError::invalid_input(format!("{name}[{idx}] must be finite, got {value}"))
}
#[inline(always)]
pub(crate) fn validate_finite_value(name: &str, idx: usize, value: Float) -> Result<()> {
if value.is_finite() {
Ok(())
} else {
Err(non_finite_value_error(name, idx, value))
}
}
pub fn validate_finite_slice(name: &str, values: &[Float]) -> Result<()> {
if let Some(idx) = crate::simd::dispatch::first_non_finite(values) {
Err(non_finite_value_error(name, idx, values[idx]))
} else {
Ok(())
}
}
pub fn validate_finite_slices(slices: &[(&str, &[Float])]) -> Result<()> {
for &(name, values) in slices {
validate_finite_slice(name, values)?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug)]
struct NamedColumns {
left: Vec<Float>,
right: Vec<Float>,
}
impl CompactPayloadLen for NamedColumns {
fn compact_payload_len(&self) -> Result<usize> {
if self.left.len() != self.right.len() {
return Err(TalibError::invalid_input(
"Compact Output columns must have equal lengths",
));
}
Ok(self.left.len())
}
}
#[test]
fn compact_output_accepts_valid_source_range_and_payload() {
let output = CompactOutput::new(5, OutputRange::new(2, 3), vec![2_i32, 3, 4]).unwrap();
assert_eq!(output.source_len(), 5);
assert_eq!(output.range(), OutputRange::new(2, 3));
assert_eq!(output.values(), &vec![2, 3, 4]);
}
#[test]
fn compact_output_rejects_range_overflow() {
let error = CompactOutput::new(usize::MAX, OutputRange::new(usize::MAX, 1), vec![1_i32])
.unwrap_err();
assert!(matches!(error, TalibError::InvalidInput { .. }));
}
#[test]
fn compact_output_rejects_out_of_source_range() {
let error = CompactOutput::new(4, OutputRange::new(2, 3), vec![1_i32, 2, 3]).unwrap_err();
assert!(matches!(error, TalibError::InvalidInput { .. }));
}
#[test]
fn compact_output_rejects_payload_length_mismatch() {
let error = CompactOutput::new(5, OutputRange::new(2, 3), vec![1_i32, 2]).unwrap_err();
assert!(matches!(error, TalibError::InvalidInput { .. }));
}
#[test]
fn compact_output_payload_length_machinery_supports_named_columns() {
let error = CompactOutput::new(
3,
OutputRange::new(1, 2),
NamedColumns {
left: vec![1.0 as Float, 2.0],
right: vec![3.0 as Float],
},
)
.unwrap_err();
assert!(matches!(error, TalibError::InvalidInput { .. }));
}
}