use super::{AdditiveFft, FftError};
use crate::{BinaryFieldExtras, Flat, HardwareField, PackedFlat};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum RsError {
BadRate { log_k: u32, log_n: u32 },
FieldTooSmall { log_n: u32, max_log_n: u32 },
BadLength { expected: usize, got: usize },
}
impl core::fmt::Display for RsError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
RsError::BadRate { log_k, log_n } => {
write!(
f,
"ReedSolomon rate: log_k {log_k} must satisfy 1 <= log_k < log_n {log_n}"
)
}
RsError::FieldTooSmall { log_n, max_log_n } => {
write!(
f,
"ReedSolomon log_n {log_n} exceeds the maximum {max_log_n} \
= min(field degree, usize::BITS - 1)"
)
}
RsError::BadLength { expected, got } => {
write!(f, "ReedSolomon buffer length {got}, expected {expected}")
}
}
}
}
impl core::error::Error for RsError {}
impl From<FftError> for RsError {
fn from(e: FftError) -> Self {
match e {
FftError::BadLength { expected, got } => RsError::BadLength { expected, got },
}
}
}
pub struct ReedSolomon<F> {
fft_k: AdditiveFft<F>,
fft_n: AdditiveFft<F>,
k: usize,
n: usize,
}
impl<F: BinaryFieldExtras + HardwareField> ReedSolomon<F> {
pub fn new(log_k: u32, log_n: u32) -> Result<Self, RsError> {
if log_k < 1 || log_k >= log_n {
return Err(RsError::BadRate { log_k, log_n });
}
let max_log_n = F::BITS.min(usize::BITS as usize - 1) as u32;
if log_n > max_log_n {
return Err(RsError::FieldTooSmall { log_n, max_log_n });
}
Ok(Self {
fft_k: AdditiveFft::new(log_k),
fft_n: AdditiveFft::new(log_n),
k: 1usize << log_k,
n: 1usize << log_n,
})
}
pub fn message_len(&self) -> usize {
self.k
}
pub fn codeword_len(&self) -> usize {
self.n
}
pub fn encode_scalar(&self, msg: &[Flat<F>], out: &mut [Flat<F>]) -> Result<(), RsError> {
self.check(msg.len(), out.len())?;
out[..self.k].copy_from_slice(msg);
out[self.k..].fill(Flat::from_raw(F::ZERO));
self.fft_k.inverse_scalar(&mut out[..self.k])?;
self.fft_n.forward_scalar(out)?;
Ok(())
}
pub fn encode(&self, msg: &[PackedFlat<F>], out: &mut [PackedFlat<F>]) -> Result<(), RsError> {
self.check(msg.len(), out.len())?;
out[..self.k].copy_from_slice(msg);
out[self.k..].fill(PackedFlat::default());
self.fft_k.inverse(&mut out[..self.k])?;
self.fft_n.forward(out)?;
Ok(())
}
fn check(&self, msg_len: usize, out_len: usize) -> Result<(), RsError> {
if msg_len != self.k {
return Err(RsError::BadLength {
expected: self.k,
got: msg_len,
});
}
if out_len != self.n {
return Err(RsError::BadLength {
expected: self.n,
got: out_len,
});
}
Ok(())
}
}