use crate::error::CodeError;
const NOTBIT0: u64 = 0xfefe_fefe_fefe_fefe;
const BIT7: u64 = 0x8080_8080_8080_8080;
const GF8POLY: u64 = 0x1d1d_1d1d_1d1d_1d1d;
#[inline]
const fn gf2_mul2_lanes(q: u64) -> u64 {
let m = q & BIT7;
((q << 1) & NOTBIT0) ^ (((m << 1).wrapping_sub(m >> 7)) & GF8POLY)
}
#[inline]
const fn gf2_mul2_byte(q: u8) -> u8 {
(q << 1) ^ (if q & 0x80 != 0 { 0x1d } else { 0 })
}
fn check_lens(sources: &[&[u8]], len: usize, min_sources: usize) -> Result<(), CodeError> {
if sources.len() < min_sources {
return Err(CodeError::ShardCount {
expected: min_sources,
got: sources.len(),
});
}
for (index, s) in sources.iter().enumerate() {
if s.len() != len {
return Err(CodeError::ShardLength {
index,
expected: len,
got: s.len(),
});
}
}
Ok(())
}
pub fn xor_gen(sources: &[&[u8]], parity: &mut [u8]) -> Result<(), CodeError> {
check_lens(sources, parity.len(), 2)?;
crate::kernel::SCALAR_CENSUS_BYTES.fetch_add(
(sources.len() * parity.len()) as u64,
core::sync::atomic::Ordering::Relaxed,
);
let (first, rest) = sources.split_first().expect("count checked");
parity.copy_from_slice(first);
for src in rest {
for (d, &s) in parity.iter_mut().zip(*src) {
*d ^= s;
}
}
Ok(())
}
pub fn xor_check(vects: &[&[u8]]) -> Result<bool, CodeError> {
let len = vects.first().map_or(0, |v| v.len());
check_lens(vects, len, 2)?;
let words = len / 8;
for i in 0..words {
let o = i * 8;
let mut acc = 0u64;
for v in vects {
acc ^= u64::from_ne_bytes(v[o..o + 8].try_into().expect("in range"));
}
if acc != 0 {
return Ok(false);
}
}
for i in words * 8..len {
let mut acc = 0u8;
for v in vects {
acc ^= v[i];
}
if acc != 0 {
return Ok(false);
}
}
Ok(true)
}
pub fn pq_gen(sources: &[&[u8]], p: &mut [u8], q: &mut [u8]) -> Result<(), CodeError> {
let len = p.len();
if q.len() != len {
return Err(CodeError::ShardLength {
index: 1,
expected: len,
got: q.len(),
});
}
check_lens(sources, len, 2)?;
crate::kernel::SCALAR_CENSUS_BYTES.fetch_add(
(sources.len() * len) as u64,
core::sync::atomic::Ordering::Relaxed,
);
let last = sources.len() - 1;
let blocks = len / 32;
for i in 0..blocks {
let o = i * 32;
let load = |s: &[u8], w: usize| {
u64::from_ne_bytes(s[o + w * 8..o + w * 8 + 8].try_into().expect("in range"))
};
let mut pw = [0u64; 4];
let mut qw = [0u64; 4];
for w in 0..4 {
pw[w] = load(sources[last], w);
qw[w] = pw[w];
}
for j in (0..last).rev() {
for w in 0..4 {
let s = load(sources[j], w);
pw[w] ^= s;
qw[w] = s ^ gf2_mul2_lanes(qw[w]);
}
}
for w in 0..4 {
p[o + w * 8..o + w * 8 + 8].copy_from_slice(&pw[w].to_ne_bytes());
q[o + w * 8..o + w * 8 + 8].copy_from_slice(&qw[w].to_ne_bytes());
}
}
for i in blocks * 32..len {
let last = sources.len() - 1;
let mut pb = sources[last][i];
let mut qb = pb;
for j in (0..last).rev() {
let s = sources[j][i];
pb ^= s;
qb = s ^ gf2_mul2_byte(qb);
}
p[i] = pb;
q[i] = qb;
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PqParity {
P,
Q,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PqMismatch {
pub index: usize,
pub parity: PqParity,
}
pub fn pq_check(sources: &[&[u8]], p: &[u8], q: &[u8]) -> Result<Option<PqMismatch>, CodeError> {
let len = p.len();
if q.len() != len {
return Err(CodeError::ShardLength {
index: 1,
expected: len,
got: q.len(),
});
}
check_lens(sources, len, 2)?;
let last = sources.len() - 1;
let byte_scan = |from: usize, to: usize| -> Option<PqMismatch> {
for i in from..to {
let mut pb = sources[last][i];
let mut qb = pb;
for j in (0..last).rev() {
let s = sources[j][i];
pb ^= s;
qb = s ^ gf2_mul2_byte(qb);
}
if p[i] != pb {
return Some(PqMismatch {
index: i,
parity: PqParity::P,
});
}
if q[i] != qb {
return Some(PqMismatch {
index: i,
parity: PqParity::Q,
});
}
}
None
};
let words = len / 8;
for i in 0..words {
let o = i * 8;
let load = |s: &[u8]| u64::from_ne_bytes(s[o..o + 8].try_into().expect("in range"));
let mut pw = load(sources[last]);
let mut qw = pw;
for j in (0..last).rev() {
let s = load(sources[j]);
pw ^= s;
qw = s ^ gf2_mul2_lanes(qw);
}
if pw != load(p) || qw != load(q) {
return Ok(byte_scan(o, o + 8));
}
}
Ok(byte_scan(words * 8, len))
}