#![cfg(feature = "fftpack-complex")]
use alloc::vec::Vec;
use core::convert::TryFrom;
use core::mem::{align_of, size_of};
use core::slice;
use num_complex::Complex32;
use slatec_sys::FortranInteger;
use slatec_sys::fftpack as raw;
pub use super::FftError;
use crate::runtime::lock_native;
const FACTOR_WORDS: usize = 15;
const _: () = assert!(size_of::<Complex32>() == 2 * size_of::<f32>());
const _: () = assert!(align_of::<Complex32>() == align_of::<f32>());
#[derive(Debug)]
pub struct ComplexFftPlan32 {
length: usize,
native_length: FortranInteger,
scratch: Vec<f32>,
twiddles: Vec<f32>,
factors: Vec<FortranInteger>,
}
impl ComplexFftPlan32 {
pub fn new(length: usize) -> Result<Self, FftError> {
if length < 2 {
return Err(FftError::InvalidLength { length, minimum: 2 });
}
let native_length =
FortranInteger::try_from(length).map_err(|_| FftError::DimensionOverflow)?;
let scalar_words = complex_words(length)?;
let scratch = zeroed::<f32>(scalar_words)?;
let mut twiddles = zeroed::<f32>(scalar_words)?;
let mut factors = zeroed::<FortranInteger>(FACTOR_WORDS)?;
let mut native_length_for_call = native_length;
let _native = lock_native();
unsafe {
raw::cffti1(
&mut native_length_for_call,
twiddles.as_mut_ptr(),
factors.as_mut_ptr(),
)
};
Ok(Self {
length,
native_length,
scratch,
twiddles,
factors,
})
}
#[must_use]
pub fn len(&self) -> usize {
self.length
}
#[must_use]
pub fn is_empty(&self) -> bool {
false
}
pub fn forward(&mut self, values: &mut [Complex32]) -> Result<(), FftError> {
self.apply(values, raw::cfftf1)
}
pub fn backward(&mut self, values: &mut [Complex32]) -> Result<(), FftError> {
self.apply(values, raw::cfftb1)
}
fn apply(
&mut self,
values: &mut [Complex32],
transform: unsafe extern "C" fn(
*mut FortranInteger,
*mut f32,
*mut f32,
*mut f32,
*mut FortranInteger,
),
) -> Result<(), FftError> {
self.check_values(values)?;
let words = interleaved_words(values)?;
let mut native_length = self.native_length;
let _native = lock_native();
unsafe {
transform(
&mut native_length,
words.as_mut_ptr(),
self.scratch.as_mut_ptr(),
self.twiddles.as_mut_ptr(),
self.factors.as_mut_ptr(),
)
};
Ok(())
}
fn check_values(&self, values: &[Complex32]) -> Result<(), FftError> {
if values.len() == self.length {
Ok(())
} else {
Err(FftError::LengthMismatch {
expected: self.length,
actual: values.len(),
})
}
}
}
fn complex_words(length: usize) -> Result<usize, FftError> {
length.checked_mul(2).ok_or(FftError::DimensionOverflow)
}
fn zeroed<T: Copy + Default>(length: usize) -> Result<Vec<T>, FftError> {
let mut values = Vec::new();
values
.try_reserve_exact(length)
.map_err(|_| FftError::AllocationFailed)?;
values.resize(length, T::default());
Ok(values)
}
fn interleaved_words(values: &mut [Complex32]) -> Result<&mut [f32], FftError> {
let words = complex_words(values.len())?;
if words > isize::MAX as usize / size_of::<f32>() {
return Err(FftError::DimensionOverflow);
}
Ok(unsafe { slice::from_raw_parts_mut(values.as_mut_ptr().cast::<f32>(), words) })
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn num_complex_layout_matches_the_reviewed_interleaved_contract() {
assert_eq!(size_of::<Complex32>(), 2 * size_of::<f32>());
assert_eq!(align_of::<Complex32>(), align_of::<f32>());
let mut values = [Complex32::new(1.25, -2.5), Complex32::new(3.0, 4.5)];
assert_eq!(
interleaved_words(&mut values).unwrap(),
&[1.25, -2.5, 3.0, 4.5]
);
}
#[test]
fn validates_lengths_without_native_entry() {
assert!(matches!(
ComplexFftPlan32::new(0),
Err(FftError::InvalidLength {
length: 0,
minimum: 2
})
));
assert!(matches!(
complex_words(usize::MAX),
Err(FftError::DimensionOverflow)
));
}
}