use core::sync::atomic::{AtomicU64, Ordering};
use crate::error::CodeError;
use crate::tables::{TABLE_BYTES, table_mul};
pub static SCALAR_CENSUS_BYTES: AtomicU64 = AtomicU64::new(0);
pub type EncodeFn = fn(gftbls: &[u8], data: &[&[u8]], out: &mut [&mut [u8]]);
pub type UpdateFn = fn(gftbls: &[u8], k: usize, vec_i: usize, src: &[u8], outs: &mut [&mut [u8]]);
#[derive(Clone, Copy, Debug)]
pub struct Kernels {
pub init: fn(coeffs: &[u8]) -> alloc::vec::Vec<u8>,
pub table_bytes: usize,
pub encode: EncodeFn,
pub mad: fn(tbl: &[u8], src: &[u8], dest: &mut [u8]),
pub update: UpdateFn,
pub name: &'static str,
pub census: &'static AtomicU64,
}
impl Kernels {
pub const fn scalar() -> Self {
Self {
init: crate::tables::init_tables,
table_bytes: TABLE_BYTES,
encode: scalar_encode,
mad: scalar_mad,
update: scalar_update,
name: "scalar",
census: &SCALAR_CENSUS_BYTES,
}
}
}
fn scalar_update(gftbls: &[u8], k: usize, vec_i: usize, src: &[u8], outs: &mut [&mut [u8]]) {
SCALAR_CENSUS_BYTES.fetch_add(src.len() as u64, Ordering::Relaxed);
let tbls: alloc::vec::Vec<&[u8; TABLE_BYTES]> = (0..outs.len())
.map(|l| {
let start = (l * k + vec_i) * TABLE_BYTES;
gftbls[start..start + TABLE_BYTES]
.try_into()
.expect("caller-validated tables")
})
.collect();
for (i, &s) in src.iter().enumerate() {
for (out, tbl) in outs.iter_mut().zip(&tbls) {
out[i] ^= table_mul(tbl, s);
}
}
}
fn scalar_encode(gftbls: &[u8], data: &[&[u8]], out: &mut [&mut [u8]]) {
let k = data.len();
let len = out.first().map_or(0, |b| b.len());
SCALAR_CENSUS_BYTES.fetch_add((k * len) as u64, Ordering::Relaxed);
for (l, dest) in out.iter_mut().enumerate() {
dest.fill(0);
for (j, src) in data.iter().enumerate() {
let start = (l * k + j) * TABLE_BYTES;
let tbl: &[u8; TABLE_BYTES] = gftbls[start..start + TABLE_BYTES]
.try_into()
.expect("caller-validated tables");
for (d, &s) in dest.iter_mut().zip(*src) {
*d ^= table_mul(tbl, s);
}
}
}
}
fn scalar_mad(tbl: &[u8], src: &[u8], dest: &mut [u8]) {
let tbl: &[u8; TABLE_BYTES] = tbl.try_into().expect("scalar mad takes a 32-byte table");
SCALAR_CENSUS_BYTES.fetch_add(src.len() as u64, Ordering::Relaxed);
for (d, &s) in dest.iter_mut().zip(src) {
*d ^= table_mul(tbl, s);
}
}
fn tbl32(gftbls: &[u8], index: usize) -> Result<&[u8; TABLE_BYTES], CodeError> {
let start = index * TABLE_BYTES;
let slice = gftbls
.get(start..start + TABLE_BYTES)
.ok_or(CodeError::ShardCount {
expected: index + 1,
got: gftbls.len() / TABLE_BYTES,
})?;
Ok(slice.try_into().expect("length checked above"))
}
pub fn vect_mul(dest: &mut [u8], tbl: &[u8; TABLE_BYTES], src: &[u8]) -> Result<(), CodeError> {
if dest.len() != src.len() {
return Err(CodeError::ShardLength {
index: 0,
expected: dest.len(),
got: src.len(),
});
}
for (d, &s) in dest.iter_mut().zip(src) {
*d = table_mul(tbl, s);
}
Ok(())
}
pub fn vect_mad(dest: &mut [u8], tbl: &[u8; TABLE_BYTES], src: &[u8]) -> Result<(), CodeError> {
if dest.len() != src.len() {
return Err(CodeError::ShardLength {
index: 0,
expected: dest.len(),
got: src.len(),
});
}
for (d, &s) in dest.iter_mut().zip(src) {
*d ^= table_mul(tbl, s);
}
Ok(())
}
pub fn vect_dot_prod(dest: &mut [u8], gftbls: &[u8], srcs: &[&[u8]]) -> Result<(), CodeError> {
if gftbls.len() != srcs.len() * TABLE_BYTES {
return Err(CodeError::ShardCount {
expected: srcs.len(),
got: gftbls.len() / TABLE_BYTES,
});
}
for (index, src) in srcs.iter().enumerate() {
if src.len() != dest.len() {
return Err(CodeError::ShardLength {
index,
expected: dest.len(),
got: src.len(),
});
}
}
dest.fill(0);
for (j, src) in srcs.iter().enumerate() {
let tbl = tbl32(gftbls, j)?;
for (d, &s) in dest.iter_mut().zip(*src) {
*d ^= table_mul(tbl, s);
}
}
Ok(())
}