pub fn trim_i8_encode(f: &[i8], nbits: u32, d: &mut [u8]) -> usize {
assert!((((f.len() as u32) * nbits) & 0x07) == 0);
let mut k = 0;
let mut acc = 0;
let mut acc_len = 0;
let mask = (1u32 << nbits) - 1;
for i in 0..f.len() {
acc |= (((f[i] as u8) as u32) & mask) << acc_len;
acc_len += nbits;
while acc_len >= 8 {
d[k] = acc as u8;
k += 1;
acc >>= 8;
acc_len -= 8;
}
}
k
}
pub fn trim_i8_decode(d: &[u8], f: &mut [i8], nbits: u32) -> Option<usize> {
let n = f.len();
let needed = n * (nbits as usize);
assert!((needed & 0x07) == 0);
let needed = needed >> 3;
if d.len() < needed {
return None;
}
let mut j = 0;
let mut acc = 0;
let mut acc_len = 0;
let mask1 = (1 << nbits) - 1;
let mask2 = 1 << (nbits - 1);
for i in 0..needed {
acc |= (d[i] as u32) << acc_len;
acc_len += 8;
while acc_len >= nbits {
let w = acc & mask1;
acc >>= nbits;
acc_len -= nbits;
let w = w | (w & mask2).wrapping_neg();
if w == mask2.wrapping_neg() {
return None;
}
if j >= n {
return None;
}
f[j] = w as i8;
j += 1;
}
}
Some(needed)
}
pub fn modq_encode(h: &[u16], d: &mut [u8]) -> usize {
assert!((h.len() & 3) == 0);
let mut j = 0;
for i in 0..(h.len() >> 2) {
let x0 = h[4 * i + 0] as u64;
let x1 = h[4 * i + 1] as u64;
let x2 = h[4 * i + 2] as u64;
let x3 = h[4 * i + 3] as u64;
let x = (x3 << 42) | (x2 << 28) | (x1 << 14) | x0;
d[j..(j + 7)].copy_from_slice(&x.to_le_bytes()[0..7]);
j += 7;
}
j
}
pub fn modq_decode(d: &[u8], h: &mut [u16]) -> Option<usize> {
let n = h.len();
if n == 0 {
return Some(0);
}
assert!((n & 3) == 0);
let needed = 7 * (n >> 2);
if d.len() != needed {
return None;
}
let mut ov = 0xFFFF;
if n >= 8 {
for i in 0..((n >> 2) - 1) {
let x = u64::from_le_bytes(
*<&[u8; 8]>::try_from(&d[(7 * i)..(7 * i + 8)]).unwrap());
let h0 = (x as u32) & 0x3FFF;
let h1 = ((x >> 14) as u32) & 0x3FFF;
let h2 = ((x >> 28) as u32) & 0x3FFF;
let h3 = ((x >> 42) as u32) & 0x3FFF;
ov &= h0.wrapping_sub(12289);
ov &= h1.wrapping_sub(12289);
ov &= h2.wrapping_sub(12289);
ov &= h3.wrapping_sub(12289);
h[4 * i + 0] = h0 as u16;
h[4 * i + 1] = h1 as u16;
h[4 * i + 2] = h2 as u16;
h[4 * i + 3] = h3 as u16;
}
}
let j = d.len() - 7;
let x = (d[j + 0] as u64)
| ((d[j + 1] as u64) << 8)
| ((d[j + 2] as u64) << 16)
| ((d[j + 3] as u64) << 24)
| ((d[j + 4] as u64) << 32)
| ((d[j + 5] as u64) << 40)
| ((d[j + 6] as u64) << 48);
let h0 = (x as u32) & 0x3FFF;
let h1 = ((x >> 14) as u32) & 0x3FFF;
let h2 = ((x >> 28) as u32) & 0x3FFF;
let h3 = ((x >> 42) as u32) & 0x3FFF;
ov &= h0.wrapping_sub(12289);
ov &= h1.wrapping_sub(12289);
ov &= h2.wrapping_sub(12289);
ov &= h3.wrapping_sub(12289);
h[n - 4] = h0 as u16;
h[n - 3] = h1 as u16;
h[n - 2] = h2 as u16;
h[n - 1] = h3 as u16;
if (ov & 0x8000) == 0 {
return None;
}
Some(needed)
}
pub const B_INF: i32 = 840;
pub fn comp_encode(s: &[i16], d: &mut [u8]) -> bool {
let mut acc = 0;
let mut acc_len = 0;
let mut j = 0;
for i in 0..s.len() {
let x = s[i] as i32;
let sw = (x >> 16) as u32;
let w = ((x as u32) ^ sw).wrapping_sub(sw);
if w > (B_INF as u32) {
return false;
}
acc |= ((sw & 1) | ((w & 0x7F) << 1)) << acc_len;
acc_len += 8;
let wh = w >> 7;
acc |= 1u32 << (acc_len + wh);
acc_len += wh + 1;
while acc_len >= 8 {
if j >= d.len() {
return false;
}
d[j] = acc as u8;
j += 1;
acc >>= 8;
acc_len -= 8;
}
}
if acc_len > 0 {
if j >= d.len() {
return false;
}
d[j] = acc as u8;
j += 1;
}
for k in j..d.len() {
d[k] = 0;
}
true
}
pub fn comp_decode(d: &[u8], v: &mut [i16]) -> bool {
let mut j = 0;
let mut acc = 0;
let mut acc_len = 0;
for i in 0..v.len() {
if j >= d.len() {
return false;
}
acc |= (d[j] as u32) << acc_len;
j += 1;
let s = acc & 1;
let m = (acc >> 1) & 0x7F;
acc >>= 8;
if acc == 0 {
if j >= d.len() {
return false;
}
acc |= (d[j] as u32) << acc_len;
j += 1;
acc_len += 8;
if acc == 0 {
return false;
}
}
let tz = acc.trailing_zeros();
let m = m + (tz << 7);
if m > (B_INF as u32) {
return false;
}
acc >>= tz + 1;
acc_len -= tz + 1;
if (s & (m.wrapping_sub(1) >> 31)) != 0 {
return false;
}
let sw = s.wrapping_neg();
let w = (m ^ sw).wrapping_sub(sw);
v[i] = w as i16;
}
if acc != 0 {
return false;
}
for k in j..d.len() {
if d[k] != 0 {
return false;
}
}
true
}
#[cfg(test)]
mod tests {
use super::*;
use crate::PRNG;
use crate::shake::SHAKE256_PRNG;
#[test]
fn modq() {
let mut tmp = [0u16; 2048];
let mut bb = [0u8; 1792];
for logn in 2..11 {
let n = 1usize << logn;
let (t1, tx) = tmp.split_at_mut(n);
let (t2, _) = tx.split_at_mut(n);
let (b1, _) = bb.split_at_mut(7 << (logn - 2));
for r in 0..64 {
let mut rng = SHAKE256_PRNG::new(&[logn as u8, r as u8]);
for i in 0..n {
t1[i] = rng.next_u16() % 12289;
}
assert!(modq_encode(t1, b1) == b1.len());
assert!(modq_decode(b1, t2).unwrap() == b1.len());
assert!(t1 == t2);
let j = (r & (n - 1)) * 14;
for k in 0..14 {
let off = (j + k) >> 3;
let m = 1u8 << ((j + k) & 7);
if ((12289u32 >> k) & 1) == 0 {
b1[off] &= !m;
} else {
b1[off] |= m;
}
}
assert!(modq_decode(b1, t2).is_none());
}
}
}
#[test]
fn compressed() {
let mut tmp = [0i16; 2048];
let mut bb = [0u8; 3850];
for logn in 2..11 {
let n = 1usize << logn;
let (t1, tx) = tmp.split_at_mut(n);
let (t2, _) = tx.split_at_mut(n);
let blen = (n << 1) - (n >> 3) + 5;
let (b1, bx) = bb.split_at_mut(blen);
let (b2, _) = bx.split_at_mut(blen);
for r in 0..64 {
let mut rng = SHAKE256_PRNG::new(&[logn as u8, r as u8]);
for i in 0..n {
let x = rng.next_u16() as i32;
t1[i] = (x % (1 + 2 * B_INF) - B_INF) as i16;
}
assert!(comp_encode(t1, b1));
assert!(comp_decode(b1, t2));
assert!(t1 == t2);
let mut k = b1.len();
while k > 0 && b1[k - 1] == 0 {
k -= 1;
}
assert!(comp_decode(&b1[..k], t2));
assert!(t1 == t2);
let mut g = 8;
while ((b1[k - 1] as u32) & (1u32 << (g - 1))) == 0 {
g -= 1;
}
for j in g..32 {
let m = 1u8 << (j & 7);
let off = j >> 3;
b1[(k - 1) + off] ^= m;
assert!(!comp_decode(&b1[..(k + 3)], t2));
b1[(k - 1) + off] ^= m;
}
let s = r & (n - 1);
t1[s] = 836;
assert!(comp_encode(t1, b1));
t1[s] = 837;
assert!(comp_encode(t1, b2));
let mut pos = 0;
let mut val;
loop {
val = b1[pos] ^ b2[pos];
if val != 0 {
break;
}
pos = pos + 1;
}
t1[s] = 841;
assert!(!comp_encode(t1, b1));
t1[s] = -841;
assert!(!comp_encode(t1, b1));
t1[s] = 840;
assert!(comp_encode(t1, b1));
assert!(comp_decode(b1, t2));
assert!(t1 == t2);
b1[pos] ^= val;
assert!(!comp_decode(b1, t2));
t1[s] = -840;
assert!(comp_encode(t1, b1));
assert!(comp_decode(b1, t2));
assert!(t1 == t2);
b1[pos] ^= val;
assert!(!comp_decode(b1, t2));
t1[s] = 1;
assert!(comp_encode(t1, b1));
t1[s] = -1;
assert!(comp_encode(t1, b2));
let mut pos = 0;
let mut val;
loop {
val = b1[pos] ^ b2[pos];
if val != 0 {
break;
}
pos = pos + 1;
}
t1[s] = 0;
assert!(comp_encode(t1, b1));
b1[pos] ^= val;
assert!(!comp_decode(b1, t2));
}
}
}
}