use core::fmt;
pub const BLOCK_MAX: usize = 255;
const PARITY_MAX: usize = 64;
const FIELD_POLY: u16 = 0x11D;
const FCR: usize = 1;
const EXP: [u8; 512] = gf_tables().0;
const LOG: [u8; 256] = gf_tables().1;
const fn gf_tables() -> ([u8; 512], [u8; 256]) {
let mut exp = [0u8; 512];
let mut log = [0u8; 256];
let mut value: u16 = 1;
let mut i = 0;
while i < 255 {
exp[i] = value as u8;
log[value as usize] = i as u8;
value <<= 1;
if value & 0x100 != 0 {
value ^= FIELD_POLY;
}
i += 1;
}
while i < 512 {
exp[i] = exp[i - 255];
i += 1;
}
(exp, log)
}
#[inline]
const fn gf_mul(a: u8, b: u8) -> u8 {
if a == 0 || b == 0 {
0
} else {
EXP[LOG[a as usize] as usize + LOG[b as usize] as usize]
}
}
#[inline]
const fn gf_inv(a: u8) -> u8 {
if a == 0 {
0
} else {
EXP[255 - LOG[a as usize] as usize]
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum RsParity {
Two,
Four,
Six,
Eight,
Sixteen,
ThirtyTwo,
SixtyFour,
}
impl RsParity {
#[must_use]
pub const fn len(self) -> usize {
match self {
RsParity::Two => 2,
RsParity::Four => 4,
RsParity::Six => 6,
RsParity::Eight => 8,
RsParity::Sixteen => 16,
RsParity::ThirtyTwo => 32,
RsParity::SixtyFour => 64,
}
}
#[must_use]
pub const fn is_empty(self) -> bool {
false
}
#[must_use]
pub const fn correctable(self) -> usize {
self.len() / 2
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RsError {
DataTooLong {
got: usize,
max: usize,
},
ParityLengthMismatch {
got: usize,
expected: usize,
},
BlockLengthInvalid {
got: usize,
min: usize,
max: usize,
},
Uncorrectable,
}
impl fmt::Display for RsError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match *self {
RsError::DataTooLong { got, max } => {
write!(f, "data length {got} exceeds RS capacity {max}")
}
RsError::ParityLengthMismatch { got, expected } => {
write!(f, "parity slice length {got}, codec requires {expected}")
}
RsError::BlockLengthInvalid { got, min, max } => {
write!(f, "block length {got} outside valid range {min}..={max}")
}
RsError::Uncorrectable => {
write!(f, "too many symbol errors: block is uncorrectable")
}
}
}
}
impl core::error::Error for RsError {}
#[derive(Debug, Clone)]
pub struct RsCodec {
parity: RsParity,
fcr: u8,
generator: [u8; PARITY_MAX + 1],
}
impl RsCodec {
#[must_use]
pub fn new(parity: RsParity) -> Self {
#[allow(clippy::cast_possible_truncation)] Self::with_fcr(parity, FCR as u8)
}
pub(crate) fn with_fcr(parity: RsParity, fcr: u8) -> Self {
let fcr = if fcr > 1 { 1 } else { fcr };
let p = parity.len();
let mut generator = [0u8; PARITY_MAX + 1];
generator[0] = 1;
let mut degree = 0;
while degree < p {
let root = EXP[fcr as usize + degree];
let mut j = degree + 1;
while j > 0 {
generator[j] = generator[j - 1] ^ gf_mul(root, generator[j]);
j -= 1;
}
generator[0] = gf_mul(root, generator[0]);
degree += 1;
}
Self {
parity,
fcr,
generator,
}
}
#[must_use]
pub const fn parity_len(&self) -> usize {
self.parity.len()
}
#[must_use]
pub const fn data_capacity(&self) -> usize {
BLOCK_MAX - self.parity.len()
}
#[must_use]
pub const fn correctable(&self) -> usize {
self.parity.correctable()
}
pub fn encode(&self, data: &[u8], parity: &mut [u8]) -> Result<(), RsError> {
let p = self.parity.len();
if data.len() > self.data_capacity() {
return Err(RsError::DataTooLong {
got: data.len(),
max: self.data_capacity(),
});
}
if parity.len() != p {
return Err(RsError::ParityLengthMismatch {
got: parity.len(),
expected: p,
});
}
let mut reg = [0u8; PARITY_MAX];
for &byte in data {
let feedback = byte ^ reg[0];
let mut i = 0;
while i + 1 < p {
reg[i] = reg[i + 1] ^ gf_mul(feedback, self.generator[p - 1 - i]);
i += 1;
}
reg[p - 1] = gf_mul(feedback, self.generator[0]);
}
parity.copy_from_slice(®[..p]);
Ok(())
}
pub fn decode(&self, block: &mut [u8]) -> Result<usize, RsError> {
let p = self.parity.len();
let n = block.len();
if n <= p || n > BLOCK_MAX {
return Err(RsError::BlockLengthInvalid {
got: n,
min: p + 1,
max: BLOCK_MAX,
});
}
let mut syn = [0u8; PARITY_MAX];
let mut clean = true;
for j in 0..p {
let x = EXP[self.fcr as usize + j];
let mut acc = 0u8;
for &byte in block.iter() {
acc = gf_mul(acc, x) ^ byte;
}
syn[j] = acc;
clean &= acc == 0;
}
if clean {
return Ok(0);
}
let mut lambda = [0u8; PARITY_MAX + 1];
let mut prev = [0u8; PARITY_MAX + 1];
lambda[0] = 1;
prev[0] = 1;
let mut errors = 0usize; let mut shift = 1usize; let mut prev_disc = 1u8; for step in 0..p {
let mut disc = syn[step];
let mut i = 1;
while i <= errors && i <= step {
disc ^= gf_mul(lambda[i], syn[step - i]);
i += 1;
}
if disc == 0 {
shift += 1;
} else {
let coef = gf_mul(disc, gf_inv(prev_disc));
if 2 * errors <= step {
let snapshot = lambda;
let mut i = 0;
while i + shift <= PARITY_MAX {
lambda[i + shift] ^= gf_mul(coef, prev[i]);
i += 1;
}
prev = snapshot;
prev_disc = disc;
errors = step + 1 - errors;
shift = 1;
} else {
let mut i = 0;
while i + shift <= PARITY_MAX {
lambda[i + shift] ^= gf_mul(coef, prev[i]);
i += 1;
}
shift += 1;
}
}
}
if errors > self.parity.correctable() {
return Err(RsError::Uncorrectable);
}
let mut positions = [0usize; PARITY_MAX / 2];
let mut found = 0usize;
for j in 0..n {
let x_inv = if j == 0 { 1 } else { EXP[255 - j] };
let mut acc = 0u8;
let mut i = errors + 1;
while i > 0 {
i -= 1;
acc = gf_mul(acc, x_inv) ^ lambda[i];
}
if acc == 0 {
if found == positions.len() {
return Err(RsError::Uncorrectable);
}
positions[found] = j;
found += 1;
}
}
if found != errors || found == 0 {
return Err(RsError::Uncorrectable);
}
let mut omega = [0u8; PARITY_MAX];
for (i, slot) in omega.iter_mut().enumerate().take(p) {
let mut acc = 0u8;
let mut k = 0;
while k <= i && k <= errors {
acc ^= gf_mul(lambda[k], syn[i - k]);
k += 1;
}
*slot = acc;
}
for &j in positions.iter().take(found) {
let x_inv = if j == 0 { 1 } else { EXP[255 - j] };
let mut num = 0u8;
let mut i = p;
while i > 0 {
i -= 1;
num = gf_mul(num, x_inv) ^ omega[i];
}
let x_inv_sq = gf_mul(x_inv, x_inv);
let mut den = 0u8;
let mut power = 1u8; let mut i = 1;
while i <= errors {
den ^= gf_mul(lambda[i], power);
power = gf_mul(power, x_inv_sq);
i += 2;
}
if den == 0 {
return Err(RsError::Uncorrectable);
}
if self.fcr == 0 {
num = gf_mul(num, EXP[j]);
}
let magnitude = gf_mul(num, gf_inv(den));
block[n - 1 - j] ^= magnitude;
}
for j in 0..p {
let x = EXP[self.fcr as usize + j];
let mut acc = 0u8;
for &byte in block.iter() {
acc = gf_mul(acc, x) ^ byte;
}
if acc != 0 {
return Err(RsError::Uncorrectable);
}
}
Ok(found)
}
}