#include "rivide/internal/dsa_sampling.h"
#include "rivide/crypto/sha3.h"
#include "rivide/internal/dsa_ntt.h"
#include "rivide/internal/dsa_packing.h"
#include "rivide/internal/dsa_poly.h"
#include "rivide/internal/dsa_reduce.h"
#include "rivide/utils/mem.h"
void dsa_poly_uniform_eta(dsa_poly_t *p, const uint8_t seed[], size_t seedlen, uint16_t nonce,
int eta) {
uint8_t buf[136 * 2];
uint8_t extseed[66];
rivide_keccak_state_t state;
unsigned int ctr, pos, i;
size_t total = seedlen + 2;
for (i = 0; i < (unsigned int)seedlen; i++) {
extseed[i] = seed[i];
}
extseed[seedlen] = (uint8_t)(nonce & 0xFF);
extseed[seedlen + 1] = (uint8_t)(nonce >> 8);
rivide_shake256_init(&state);
rivide_shake_absorb(&state, extseed, total);
rivide_shake_squeeze(&state, buf, sizeof(buf));
ctr = 0;
pos = 0;
while (ctr < DSA_N) {
uint32_t t;
if (pos >= sizeof(buf)) {
rivide_shake_squeeze(&state, buf, sizeof(buf));
pos = 0;
}
if (eta == 2) {
t = (uint32_t)buf[pos++];
uint32_t d1 = t & 0x0F;
uint32_t d2 = t >> 4;
if (d1 < 15 && ctr < DSA_N) {
d1 = d1 - (5 * (d1 / 5));
p->coeffs[ctr++] = (int32_t)(2 - d1);
}
if (d2 < 15 && ctr < DSA_N) {
d2 = d2 - (5 * (d2 / 5));
p->coeffs[ctr++] = (int32_t)(2 - d2);
}
} else {
t = (uint32_t)buf[pos++];
uint32_t d1 = t & 0x0F;
uint32_t d2 = t >> 4;
if (d1 < 9 && ctr < DSA_N) {
p->coeffs[ctr++] = (int32_t)(4 - d1);
}
if (d2 < 9 && ctr < DSA_N) {
p->coeffs[ctr++] = (int32_t)(4 - d2);
}
}
}
rivide_cleanse(buf, sizeof(buf));
}
void dsa_poly_uniform(dsa_poly_t *p, const uint8_t seed[32], uint16_t nonce) {
rivide_keccak_state_t state;
uint8_t buf[168 * 2];
uint8_t extseed[34];
unsigned int ctr, pos, i;
for (i = 0; i < 32; i++) {
extseed[i] = seed[i];
}
extseed[32] = (uint8_t)(nonce & 0xFF);
extseed[33] = (uint8_t)(nonce >> 8);
rivide_shake128_init(&state);
rivide_shake_absorb(&state, extseed, 34);
rivide_shake_squeeze(&state, buf, sizeof(buf));
ctr = 0;
pos = 0;
while (ctr < DSA_N) {
if (pos + 3 > sizeof(buf)) {
rivide_shake_squeeze(&state, buf, sizeof(buf));
pos = 0;
}
uint32_t t =
((uint32_t)buf[pos] | ((uint32_t)buf[pos + 1] << 8) | ((uint32_t)buf[pos + 2] << 16)) &
0x7FFFFF;
pos += 3;
if (t < (uint32_t)DSA_Q) {
p->coeffs[ctr++] = (int32_t)t;
}
}
}
void dsa_poly_uniform_gamma1(dsa_poly_t *p, const uint8_t seed[], size_t seedlen, uint16_t nonce,
int32_t gamma1) {
uint8_t buf[640];
uint8_t extseed[66];
unsigned int i;
size_t total = seedlen + 2;
size_t buflen;
for (i = 0; i < (unsigned int)seedlen; i++) {
extseed[i] = seed[i];
}
extseed[seedlen] = (uint8_t)(nonce & 0xFF);
extseed[seedlen + 1] = (uint8_t)(nonce >> 8);
if (gamma1 == (1 << 17)) {
buflen = 576;
} else {
buflen = 640;
}
rivide_shake256(buf, buflen, extseed, total);
dsa_poly_unpack_z(p, buf, gamma1);
rivide_cleanse(buf, sizeof(buf));
}
void dsa_poly_challenge(dsa_poly_t *c, const uint8_t *seed, size_t len, unsigned int tau) {
rivide_keccak_state_t state;
uint8_t buf[136];
unsigned int i, pos;
uint64_t signs;
rivide_shake256_init(&state);
rivide_shake_absorb(&state, seed, len);
rivide_shake_squeeze(&state, buf, 136);
signs = 0;
for (i = 0; i < 8; i++) {
signs |= (uint64_t)buf[i] << (8 * i);
}
for (i = 0; i < DSA_N; i++) {
c->coeffs[i] = 0;
}
pos = 8;
for (i = DSA_N - tau; i < DSA_N; i++) {
uint8_t byte;
unsigned int j;
do {
if (pos >= 136) {
rivide_shake_squeeze(&state, buf, 136);
pos = 0;
}
byte = buf[pos++];
j = (unsigned int)byte;
} while (j > i);
c->coeffs[i] = c->coeffs[j];
c->coeffs[j] = (int32_t)(1 - 2 * (int32_t)(signs & 1));
signs >>= 1;
}
}
void dsa_expand_matrix_mul(dsa_polyveck_t *t, const uint8_t rho[32], const dsa_polyvecl_t *s, int k,
int l) {
dsa_poly_t a_ij, tmp;
int i, j;
for (i = 0; i < k; i++) {
for (j = 0; j < l; j++) {
uint16_t nonce = (uint16_t)((i << 8) | j);
dsa_poly_uniform(&a_ij, rho, nonce);
dsa_poly_pointwise(&tmp, &a_ij, &s->vec[j]);
if (j == 0) {
unsigned int idx;
for (idx = 0; idx < DSA_N; idx++) {
t->vec[i].coeffs[idx] = tmp.coeffs[idx];
}
} else {
dsa_poly_add(&t->vec[i], &t->vec[i], &tmp);
}
}
dsa_poly_reduce(&t->vec[i]);
}
}