use crate::{LongEccDecodeError, ReedSolomon};
use super::checksums::xxh64;
use super::{LongEccHeader, HEADER_SIZE};
#[cfg(test)]
thread_local! {
static DECODE_CALLS: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
const fn intersects(range: Option<(usize, usize)>, lo: usize, hi: usize) -> bool {
match range {
Some((start, end)) => start < hi && lo < end,
None => false,
}
}
fn cover(range: &mut Option<(usize, usize)>, lo: usize, hi: usize) {
*range = match *range {
Some((start, end)) => Some((start.min(lo), end.max(hi))),
None => Some((lo, hi)),
};
}
fn diff_bounds(new: &[u8], old: &[u8]) -> Option<(usize, usize)> {
let first = new.iter().zip(old).position(|(new, old)| new != old)?;
let last = new.iter().zip(old).rposition(|(new, old)| new != old)?;
Some((first, last + 1))
}
pub fn correct_in_place(codeword: &mut [u8]) -> Result<LongEccHeader, LongEccDecodeError> {
use LongEccDecodeError::{
IntegrityCheckFailed, InvalidCodeword, ReadDataError, ReadParityError,
};
const MAX_SWEEPS: u32 = 8;
let header = LongEccHeader::from_byte_slice(codeword)?;
let full_length = usize::try_from(header.full_length())?;
let codeword = codeword.get_mut(..full_length).ok_or(InvalidCodeword)?;
let parity_bytes = usize::from(header.parity_bytes());
let segment_length = usize::from(header.segment_length());
let segment_distance = usize::from(header.segment_distance());
if parity_bytes >= segment_distance {
return Err(InvalidCodeword);
}
let last_segment_length = usize::from(header.last_segment_length());
if codeword.len() < HEADER_SIZE + last_segment_length + parity_bytes {
return Err(InvalidCodeword);
}
let max_sweeps = header.segment_count().min(MAX_SWEEPS);
let mut previous_dirty: Option<(usize, usize)> = None;
let mut full_sweep = true;
let mut failure = None;
for _ in 0..max_sweeps {
let mut current_dirty: Option<(usize, usize)> = None;
failure = None;
let mut parity_index = codeword.len() - parity_bytes;
let mut data_index = parity_index - last_segment_length;
let mut data_length = last_segment_length;
loop {
let data_end = data_index + data_length;
let parity_end = parity_index + parity_bytes;
let dirty = intersects(previous_dirty, data_index, data_end)
|| intersects(previous_dirty, parity_index, parity_end)
|| intersects(current_dirty, data_index, data_end)
|| intersects(current_dirty, parity_index, parity_end);
if full_sweep || dirty {
#[cfg(test)]
DECODE_CALLS.with(|calls| calls.set(calls.get() + 1));
let (head, tail) = codeword.split_at_mut(parity_index);
let parity = tail.get_mut(..parity_bytes).ok_or(ReadParityError)?;
let data = head.get_mut(data_index..data_end).ok_or(ReadDataError)?;
let mut snapshot = [0u8; 255];
let split = data.len();
let total = split + parity.len();
snapshot[..split].copy_from_slice(data);
snapshot[split..total].copy_from_slice(parity);
match ReedSolomon::correct_detached_in_place(parity, data) {
Ok(()) => {
if let Some((lo, hi)) = diff_bounds(data, &snapshot[..split]) {
cover(&mut current_dirty, data_index + lo, data_index + hi);
}
if let Some((lo, hi)) = diff_bounds(parity, &snapshot[split..total]) {
cover(&mut current_dirty, parity_index + lo, parity_index + hi);
}
}
Err(error) => {
data.copy_from_slice(&snapshot[..split]);
parity.copy_from_slice(&snapshot[split..total]);
failure.get_or_insert(error);
}
}
}
if data_index <= HEADER_SIZE {
break;
}
data_index = data_index.saturating_sub(segment_distance);
parity_index = parity_index.saturating_sub(parity_bytes);
data_length = segment_length;
}
full_sweep = failure.is_some();
previous_dirty = current_dirty;
if previous_dirty.is_none() {
break;
}
}
if let Some(error) = failure {
return Err(error.into());
}
let payload = codeword
.get(HEADER_SIZE..full_length)
.ok_or(InvalidCodeword)?;
if xxh64(payload) != header.checksum() {
return Err(IntegrityCheckFailed);
}
codeword[..HEADER_SIZE].copy_from_slice(&header.to_bytes());
Ok(header)
}
#[cfg(test)]
mod tests {
use ps_buffer::ToBuffer;
use crate::{LongEccDecodeError, MAX_PARITY};
use super::super::{decode, encode, OverlapFactor, HEADER_SIZE};
use super::{correct_in_place, cover, diff_bounds, intersects, DECODE_CALLS};
type TestError = Box<dyn std::error::Error>;
#[test]
fn test_diff_bounds() {
assert_eq!(diff_bounds(b"", b""), None);
assert_eq!(diff_bounds(b"same", b"same"), None);
assert_eq!(diff_bounds(b"axc", b"abc"), Some((1, 2)));
assert_eq!(diff_bounds(b"xbcx", b"abca"), Some((0, 4)));
}
#[test]
fn test_intersects() {
assert!(!intersects(None, 0, 10));
assert!(intersects(Some((3, 7)), 6, 10));
assert!(!intersects(Some((3, 7)), 7, 10));
assert!(!intersects(Some((3, 7)), 0, 3));
}
#[test]
fn test_cover() {
let mut range = None;
cover(&mut range, 5, 8);
assert_eq!(range, Some((5, 8)));
cover(&mut range, 2, 6);
cover(&mut range, 7, 11);
assert_eq!(range, Some((2, 11)));
}
#[test]
fn test_long_ecc_correct_in_place_no_errors() -> Result<(), TestError> {
let message = b"Correct No Errors".to_buffer()?;
let parity: u8 = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_long_ecc_correct_in_place_one_error() -> Result<(), TestError> {
let message = b"Correct One Error".to_buffer()?;
let parity: u8 = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 5] ^= 0b0000_0001;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_long_ecc_correct_in_place_error_in_parity() -> Result<(), TestError> {
let message = b"Error In Parity".to_buffer()?;
let parity: u8 = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
let parity_start = HEADER_SIZE + message.len();
encoded[parity_start + 1] ^= 0b0000_0010;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
let decoded = decode(&encoded)?;
assert_eq!(&decoded[..], &message[..]);
Ok(())
}
#[test]
fn test_long_ecc_correct_in_place_multiple_errors_recoverable() -> Result<(), TestError> {
let message = b"Multiple Recoverable".to_buffer()?;
let parity: u8 = 3;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 1] ^= 0b0000_0001;
encoded[HEADER_SIZE + 7] ^= 0b0000_0010;
encoded[HEADER_SIZE + message.len() + 3] ^= 0b0000_0100;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_long_ecc_correct_in_place_too_many_errors() -> Result<(), TestError> {
let message = b"Too Many Errors".to_buffer()?;
let parity: u8 = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 1] ^= 0b0000_0001;
encoded[HEADER_SIZE + 3] ^= 0b0000_0010;
encoded[HEADER_SIZE + 5] ^= 0b0000_0100;
let result = correct_in_place(&mut encoded);
assert!(matches!(
result,
Err(LongEccDecodeError::RSDecodeError(
crate::RSDecodeError::RSComputeErrorsError(
crate::RSComputeErrorsError::TooManyErrors
)
))
));
Ok(())
}
#[test]
fn test_long_ecc_correct_in_place_zero_parity() -> Result<(), TestError> {
let message = b"Zero Parity Correct".to_buffer()?;
let parity: u8 = 0;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.parity(), 0);
assert_eq!(&encoded[HEADER_SIZE..], &message[..]);
Ok(())
}
#[test]
fn test_long_ecc_correct_in_place_zero_parity_detects_corruption() -> Result<(), TestError> {
let message = b"Zero Parity Corruption".to_buffer()?;
let parity: u8 = 0;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 3] ^= 0b0000_0001;
let result = correct_in_place(&mut encoded);
assert!(matches!(
result,
Err(LongEccDecodeError::IntegrityCheckFailed)
));
Ok(())
}
#[test]
fn test_correct_in_place_single_segment() -> Result<(), TestError> {
let message = b"Single segment test".to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_multiple_segments() -> Result<(), TestError> {
let message = b"This is a longer message that will span multiple segments for testing"
.repeat(10)
.to_buffer()?;
let parity = 3;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 100] ^= 0b0000_0001;
encoded[HEADER_SIZE + 400] ^= 0b0000_0010;
encoded[HEADER_SIZE + 700] ^= 0b0000_0100;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_error_correction_in_middle_segment() -> Result<(), TestError> {
let message = b"Error correction in middle segment test with sufficient length"
.repeat(8)
.to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 300] ^= 0b0000_0001;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_edge_case_two_segments() -> Result<(), TestError> {
let message = b"Two segment edge case".repeat(13).to_buffer()?;
let parity = 1;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_max_parity() -> Result<(), TestError> {
let message = b"Maximum parity correction test".to_buffer()?;
let mut encoded = encode(&message, MAX_PARITY, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 10] ^= 0b0000_0001;
encoded[HEADER_SIZE + 20] ^= 0b0000_0010;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.parity(), MAX_PARITY);
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_rejects_hash_mismatch() -> Result<(), TestError> {
let message = b"Integrity guard after correction".to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE..].fill(0);
encoded[HEADER_SIZE] = 1;
let result = correct_in_place(&mut encoded);
assert!(matches!(
result,
Err(LongEccDecodeError::IntegrityCheckFailed)
));
Ok(())
}
#[test]
fn test_correct_in_place_fixed_point_recovers_overloaded_segment() -> Result<(), TestError> {
let message = (0..=255u8).cycle().take(400).collect::<Vec<u8>>();
let message = message.to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Double)?;
encoded[HEADER_SIZE + 300] ^= 0b0100_0000;
encoded[HEADER_SIZE + 310] ^= 0b1000_0000;
encoded[HEADER_SIZE + 390] ^= 0b0000_0001;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_fixed_point_recovers_via_parity_of_parity() -> Result<(), TestError> {
let message = (0..=255u8).cycle().take(600).collect::<Vec<u8>>();
let message = message.to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 550] ^= 0b0000_0001;
encoded[HEADER_SIZE + 600] ^= 0b0000_0010;
encoded[HEADER_SIZE + 601] ^= 0b0000_0100;
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_rolls_back_uncorrectable_segment() -> Result<(), TestError> {
let message = (0..=255u8).cycle().take(600).collect::<Vec<u8>>();
let message = message.to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 510] ^= 0b0100_0000;
encoded[HEADER_SIZE + 530] ^= 0b1000_0000;
encoded[HEADER_SIZE + 550] ^= 0b0000_0001;
let corrupted = encoded.clone()?;
let result = correct_in_place(&mut encoded);
assert!(matches!(result, Err(LongEccDecodeError::RSDecodeError(_))));
assert_eq!(&encoded[..], &corrupted[..]);
Ok(())
}
#[test]
fn test_correct_in_place_skips_unchanged_segments_on_later_sweeps() -> Result<(), TestError> {
let message = (0..=255u8).cycle().take(800).collect::<Vec<u8>>();
let message = message.to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
encoded[HEADER_SIZE + 1] ^= 0b0000_0001;
DECODE_CALLS.with(|calls| calls.set(0));
let header = correct_in_place(&mut encoded)?;
let decode_calls = DECODE_CALLS.with(std::cell::Cell::get);
let segment_count = header.segment_count() as usize;
assert!(segment_count >= 3);
assert!(decode_calls > segment_count);
assert!(decode_calls < 2 * segment_count);
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
#[test]
fn test_correct_in_place_with_parity_errors() -> Result<(), TestError> {
let message = b"Parity error correction test".to_buffer()?;
let parity = 2;
let mut encoded = encode(&message, parity, OverlapFactor::Simple)?;
let parity_start = HEADER_SIZE + message.len();
if parity_start + 1 < encoded.len() {
encoded[parity_start] ^= 0b0000_0001;
encoded[parity_start + 1] ^= 0b0000_0010;
}
let header = correct_in_place(&mut encoded)?;
assert_eq!(header.message_length() as usize, message.len());
assert_eq!(
&encoded[HEADER_SIZE..HEADER_SIZE + message.len()],
&message[..]
);
Ok(())
}
}