#include "rivide/internal/kem_compress.h"
#include "rivide/internal/kem_reduce.h"
uint16_t compress_coeff(int16_t x, int d) {
int16_t c = barrett_reduce(x);
c = cond_sub_q(c);
c = (int16_t)(c + ((c >> 15) & KEM_Q));
uint32_t t = (uint32_t)(uint16_t)c;
t = (t << d) + KEM_Q / 2;
t = t / KEM_Q;
return (uint16_t)(t & ((1u << d) - 1));
}
int16_t decompress_coeff(uint16_t x, int d) {
uint32_t t = ((uint32_t)x * KEM_Q + (1u << (d - 1))) >> d;
return (int16_t)t;
}
void poly_compress(uint8_t *buf, const poly_t *p, int d) {
unsigned int i, j;
if (d == 4) {
for (i = 0; i < KEM_N / 2; i++) {
uint8_t t0 = (uint8_t)compress_coeff(p->coeffs[2 * i], d);
uint8_t t1 = (uint8_t)compress_coeff(p->coeffs[2 * i + 1], d);
buf[i] = (uint8_t)(t0 | (t1 << 4));
}
} else if (d == 5) {
for (i = 0; i < KEM_N / 8; i++) {
uint8_t t[8];
for (j = 0; j < 8; j++) {
t[j] = (uint8_t)compress_coeff(p->coeffs[8 * i + j], d);
}
buf[5 * i] = (uint8_t)(t[0] | (t[1] << 5));
buf[5 * i + 1] = (uint8_t)((t[1] >> 3) | (t[2] << 2) | (t[3] << 7));
buf[5 * i + 2] = (uint8_t)((t[3] >> 1) | (t[4] << 4));
buf[5 * i + 3] = (uint8_t)((t[4] >> 4) | (t[5] << 1) | (t[6] << 6));
buf[5 * i + 4] = (uint8_t)((t[6] >> 2) | (t[7] << 3));
}
} else if (d == 10) {
for (i = 0; i < KEM_N / 4; i++) {
uint16_t t[4];
for (j = 0; j < 4; j++) {
t[j] = compress_coeff(p->coeffs[4 * i + j], d);
}
buf[5 * i] = (uint8_t)(t[0] & 0xFF);
buf[5 * i + 1] = (uint8_t)((t[0] >> 8) | ((t[1] & 0x3F) << 2));
buf[5 * i + 2] = (uint8_t)((t[1] >> 6) | ((t[2] & 0x0F) << 4));
buf[5 * i + 3] = (uint8_t)((t[2] >> 4) | ((t[3] & 0x03) << 6));
buf[5 * i + 4] = (uint8_t)(t[3] >> 2);
}
} else if (d == 11) {
for (i = 0; i < KEM_N / 8; i++) {
uint16_t t[8];
for (j = 0; j < 8; j++) {
t[j] = compress_coeff(p->coeffs[8 * i + j], d);
}
buf[11 * i] = (uint8_t)(t[0] & 0xFF);
buf[11 * i + 1] = (uint8_t)((t[0] >> 8) | ((t[1] & 0x1F) << 3));
buf[11 * i + 2] = (uint8_t)((t[1] >> 5) | ((t[2] & 0x03) << 6));
buf[11 * i + 3] = (uint8_t)((t[2] >> 2) & 0xFF);
buf[11 * i + 4] = (uint8_t)((t[2] >> 10) | ((t[3] & 0x7F) << 1));
buf[11 * i + 5] = (uint8_t)((t[3] >> 7) | ((t[4] & 0x0F) << 4));
buf[11 * i + 6] = (uint8_t)((t[4] >> 4) | ((t[5] & 0x01) << 7));
buf[11 * i + 7] = (uint8_t)((t[5] >> 1) & 0xFF);
buf[11 * i + 8] = (uint8_t)((t[5] >> 9) | ((t[6] & 0x3F) << 2));
buf[11 * i + 9] = (uint8_t)((t[6] >> 6) | ((t[7] & 0x07) << 5));
buf[11 * i + 10] = (uint8_t)(t[7] >> 3);
}
}
}
void poly_decompress(poly_t *p, const uint8_t *buf, int d) {
unsigned int i, j;
if (d == 4) {
for (i = 0; i < KEM_N / 2; i++) {
p->coeffs[2 * i] = decompress_coeff((uint16_t)(buf[i] & 0x0F), d);
p->coeffs[2 * i + 1] = decompress_coeff((uint16_t)(buf[i] >> 4), d);
}
} else if (d == 5) {
for (i = 0; i < KEM_N / 8; i++) {
uint8_t t[8];
t[0] = (uint8_t)(buf[5 * i] & 0x1F);
t[1] = (uint8_t)((buf[5 * i] >> 5) | ((buf[5 * i + 1] & 0x03) << 3));
t[2] = (uint8_t)((buf[5 * i + 1] >> 2) & 0x1F);
t[3] = (uint8_t)((buf[5 * i + 1] >> 7) | ((buf[5 * i + 2] & 0x0F) << 1));
t[4] = (uint8_t)((buf[5 * i + 2] >> 4) | ((buf[5 * i + 3] & 0x01) << 4));
t[5] = (uint8_t)((buf[5 * i + 3] >> 1) & 0x1F);
t[6] = (uint8_t)((buf[5 * i + 3] >> 6) | ((buf[5 * i + 4] & 0x07) << 2));
t[7] = (uint8_t)((buf[5 * i + 4] >> 3) & 0x1F);
for (j = 0; j < 8; j++) {
p->coeffs[8 * i + j] = decompress_coeff((uint16_t)t[j], d);
}
}
} else if (d == 10) {
for (i = 0; i < KEM_N / 4; i++) {
uint16_t t[4];
t[0] = (uint16_t)(((uint16_t)buf[5 * i] | ((uint16_t)buf[5 * i + 1] << 8)) & 0x3FF);
t[1] = (uint16_t)((((uint16_t)buf[5 * i + 1] >> 2) | ((uint16_t)buf[5 * i + 2] << 6)) &
0x3FF);
t[2] = (uint16_t)((((uint16_t)buf[5 * i + 2] >> 4) | ((uint16_t)buf[5 * i + 3] << 4)) &
0x3FF);
t[3] = (uint16_t)((((uint16_t)buf[5 * i + 3] >> 6) | ((uint16_t)buf[5 * i + 4] << 2)) &
0x3FF);
for (j = 0; j < 4; j++) {
p->coeffs[4 * i + j] = decompress_coeff(t[j], d);
}
}
} else if (d == 11) {
for (i = 0; i < KEM_N / 8; i++) {
uint16_t t[8];
t[0] = (uint16_t)(((uint16_t)buf[11 * i] | ((uint16_t)buf[11 * i + 1] << 8)) & 0x7FF);
t[1] =
(uint16_t)((((uint16_t)buf[11 * i + 1] >> 3) | ((uint16_t)buf[11 * i + 2] << 5)) &
0x7FF);
t[2] = (uint16_t)((((uint16_t)buf[11 * i + 2] >> 6) | ((uint16_t)buf[11 * i + 3] << 2) |
((uint16_t)buf[11 * i + 4] << 10)) &
0x7FF);
t[3] =
(uint16_t)((((uint16_t)buf[11 * i + 4] >> 1) | ((uint16_t)buf[11 * i + 5] << 7)) &
0x7FF);
t[4] =
(uint16_t)((((uint16_t)buf[11 * i + 5] >> 4) | ((uint16_t)buf[11 * i + 6] << 4)) &
0x7FF);
t[5] = (uint16_t)((((uint16_t)buf[11 * i + 6] >> 7) | ((uint16_t)buf[11 * i + 7] << 1) |
((uint16_t)buf[11 * i + 8] << 9)) &
0x7FF);
t[6] =
(uint16_t)((((uint16_t)buf[11 * i + 8] >> 2) | ((uint16_t)buf[11 * i + 9] << 6)) &
0x7FF);
t[7] =
(uint16_t)((((uint16_t)buf[11 * i + 9] >> 5) | ((uint16_t)buf[11 * i + 10] << 3)) &
0x7FF);
for (j = 0; j < 8; j++) {
p->coeffs[8 * i + j] = decompress_coeff(t[j], d);
}
}
}
}
void polyvec_compress(uint8_t *buf, const polyvec_t *v, int k, int d) {
int i;
size_t poly_bytes;
if (d == 10) {
poly_bytes = 320;
} else {
poly_bytes = 352;
}
for (i = 0; i < k; i++) {
poly_compress(buf + poly_bytes * (size_t)i, &v->vec[i], d);
}
}
void polyvec_decompress(polyvec_t *v, const uint8_t *buf, int k, int d) {
int i;
size_t poly_bytes;
if (d == 10) {
poly_bytes = 320;
} else {
poly_bytes = 352;
}
for (i = 0; i < k; i++) {
poly_decompress(&v->vec[i], buf + poly_bytes * (size_t)i, d);
}
}