* This file is part of the openHiTLS project.
*
* openHiTLS is licensed under the Mulan PSL v2.
* You can use this software according to the terms and conditions of the Mulan PSL v2.
* You may obtain a copy of Mulan PSL v2 at:
*
* http://license.coscl.org.cn/MulanPSL2
*
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND,
* EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,
* MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE.
* See the Mulan PSL v2 for more details.
*/
#include "hitls_build.h"
#ifdef HITLS_CRYPTO_MLDSA
#include <string.h>
#include "bsl_errno.h"
#include "bsl_sal.h"
#include "crypt_utils.h"
#include "crypt_sha3.h"
#include "crypt_errno.h"
#include "crypt_util_rand.h"
#include "bsl_err_internal.h"
#include "ml_dsa_local.h"
#include "eal_md_local.h"
#define BITS_OF_BYTE 8
#define MLDSA_SET_VECTOR_MEM(ptr, buf) {ptr = buf; buf += MLDSA_N;}
static int32_t MLDSAInitHashCtx(int32_t mdId, const EAL_MdMethod **hashMethod, void **mdCtx)
{
const EAL_MdMethod *method = EAL_MdFindDefaultMethod(mdId);
if (method == NULL) {
BSL_ERR_PUSH_ERROR(CRYPT_EAL_ALG_NOT_SUPPORT);
return CRYPT_EAL_ALG_NOT_SUPPORT;
}
void *ctx = method->newCtx(NULL, method->id);
if (ctx == NULL) {
BSL_ERR_PUSH_ERROR(CRYPT_MEM_ALLOC_FAIL);
return CRYPT_MEM_ALLOC_FAIL;
}
int32_t ret = method->init(ctx, NULL);
if (ret != CRYPT_SUCCESS) {
method->freeCtx(ctx);
BSL_ERR_PUSH_ERROR(ret);
return ret;
}
*hashMethod = method;
*mdCtx = ctx;
return CRYPT_SUCCESS;
}
static int32_t HashFuncH(const uint8_t *inPutA, uint32_t lenA, const uint8_t *inPutB, uint32_t lenB,
uint8_t *out, uint32_t outLen)
{
uint32_t len = outLen;
int32_t ret = 0;
const EAL_MdMethod *hashMethod = NULL;
void *mdCtx = NULL;
RETURN_RET_IF_ERR_EX(MLDSAInitHashCtx(CRYPT_MD_SHAKE256, &hashMethod, &mdCtx), ret);
GOTO_ERR_IF(hashMethod->update(mdCtx, inPutA, lenA), ret);
if (inPutB != NULL) {
GOTO_ERR_IF(hashMethod->update(mdCtx, inPutB, lenB), ret);
}
GOTO_ERR_IF(hashMethod->final(mdCtx, out, &len), ret);
ERR:
hashMethod->freeCtx(mdCtx);
return ret;
}
typedef struct {
int32_t *bufAddr;
uint32_t bufSize;
int32_t *matrix[MLDSA_L_MAX];
int32_t *s2[MLDSA_K_MAX];
int32_t *t0[MLDSA_K_MAX];
int32_t *t1[MLDSA_K_MAX];
int32_t *s1[MLDSA_L_MAX];
int32_t *s1Ntt[MLDSA_L_MAX];
} MLDSA_KeyGenMatrixSt;
static void MLDSASetMatrixMem(uint8_t k, uint8_t l, int32_t *matrix[MLDSA_K_MAX][MLDSA_L_MAX], int32_t *buf)
{
for (uint8_t i = 0; i < k; i++) {
for (uint8_t j = 0; j < l; j++) {
matrix[i][j] = buf;
buf += MLDSA_N;
}
}
}
static int32_t MLDSAKeyGenCreateMatrix(uint8_t k, uint8_t l, MLDSA_KeyGenMatrixSt *st)
{
st->bufSize = (3 * k + 3 * l) * MLDSA_N * sizeof(int32_t);
int32_t *buf = BSL_SAL_Malloc(st->bufSize);
if (buf == NULL) {
return BSL_MALLOC_FAIL;
}
st->bufAddr = buf;
for (uint8_t i = 0; i < k; i++) {
MLDSA_SET_VECTOR_MEM(st->t0[i], buf);
MLDSA_SET_VECTOR_MEM(st->t1[i], buf);
MLDSA_SET_VECTOR_MEM(st->s2[i], buf);
}
for (uint8_t i = 0; i < l; i++) {
MLDSA_SET_VECTOR_MEM(st->matrix[i], buf);
}
for (uint8_t i = 0; i < l; i++) {
MLDSA_SET_VECTOR_MEM(st->s1[i], buf);
}
for (uint8_t i = 0; i < l; i++) {
MLDSA_SET_VECTOR_MEM(st->s1Ntt[i], buf);
}
return CRYPT_SUCCESS;
}
typedef struct {
int32_t *bufAddr;
uint32_t bufSize;
int32_t *matrix[MLDSA_K_MAX][MLDSA_L_MAX];
int32_t *t0[MLDSA_K_MAX];
int32_t *r0[MLDSA_K_MAX];
int32_t *s2[MLDSA_K_MAX];
int32_t *w[MLDSA_K_MAX];
int32_t *w1[MLDSA_K_MAX];
int32_t *y[MLDSA_K_MAX];
int32_t *s1[MLDSA_L_MAX];
int32_t *z[MLDSA_L_MAX];
} MLDSA_SignMatrixSt;
static int32_t MLDSASignCreateMatrix(uint8_t k, uint8_t l, MLDSA_SignMatrixSt *st)
{
st->bufSize = (k * l + 6 * k + 2 * l) * MLDSA_N * sizeof(int32_t);
int32_t *buf = BSL_SAL_Malloc(st->bufSize);
if (buf == NULL) {
return BSL_MALLOC_FAIL;
}
st->bufAddr = buf;
MLDSASetMatrixMem(k, l, st->matrix, buf);
buf += k * l * MLDSA_N;
for (uint8_t i = 0; i < k; i++) {
MLDSA_SET_VECTOR_MEM(st->r0[i], buf);
MLDSA_SET_VECTOR_MEM(st->t0[i], buf);
MLDSA_SET_VECTOR_MEM(st->s2[i], buf);
MLDSA_SET_VECTOR_MEM(st->w[i], buf);
MLDSA_SET_VECTOR_MEM(st->w1[i], buf);
}
for (uint8_t i = 0; i < k; i++) {
MLDSA_SET_VECTOR_MEM(st->y[i], buf);
}
for (uint8_t i = 0; i < l; i++) {
MLDSA_SET_VECTOR_MEM(st->s1[i], buf);
}
for (uint8_t i = 0; i < l; i++) {
MLDSA_SET_VECTOR_MEM(st->z[i], buf);
}
return CRYPT_SUCCESS;
}
typedef struct {
int32_t *bufAddr;
uint32_t bufSize;
int32_t *matrix[MLDSA_L_MAX];
int32_t *t1[MLDSA_K_MAX];
int32_t *h[MLDSA_K_MAX];
int32_t *w[MLDSA_K_MAX];
int32_t *z[MLDSA_L_MAX];
} MLDSA_VerifyMatrixSt;
static int32_t MLDSAVerifyCreateMatrix(uint8_t k, uint8_t l, MLDSA_VerifyMatrixSt *st)
{
st->bufSize = (3 * k + 2 * l) * MLDSA_N * sizeof(int32_t);
int32_t *buf = BSL_SAL_Malloc(st->bufSize);
if (buf == NULL) {
return BSL_MALLOC_FAIL;
}
st->bufAddr = buf;
for (uint8_t i = 0; i < k; i++) {
MLDSA_SET_VECTOR_MEM(st->t1[i], buf);
MLDSA_SET_VECTOR_MEM(st->h[i], buf);
MLDSA_SET_VECTOR_MEM(st->w[i], buf);
}
for (uint8_t i = 0; i < l; i++) {
MLDSA_SET_VECTOR_MEM(st->matrix[i], buf);
MLDSA_SET_VECTOR_MEM(st->z[i], buf);
}
return CRYPT_SUCCESS;
}
static int32_t RejNTTPoly(int32_t a[MLDSA_N], const uint8_t seed[MLDSA_SEED_EXTEND_BYTES_LEN])
{
int32_t ret;
unsigned int outlen = CRYPT_SHAKE128_BLOCKSIZE;
const uint32_t buflen = CRYPT_SHAKE128_BLOCKSIZE / 4;
uint32_t buf[CRYPT_SHAKE128_BLOCKSIZE / 4];
const EAL_MdMethod *hashMethod = NULL;
void *mdCtx = NULL;
RETURN_RET_IF_ERR_EX(MLDSAInitHashCtx(CRYPT_MD_SHAKE128, &hashMethod, &mdCtx), ret);
GOTO_ERR_IF(hashMethod->update(mdCtx, seed, MLDSA_SEED_EXTEND_BYTES_LEN), ret);
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, (uint8_t *)buf, outlen), ret);
uint32_t j = 0;
for (uint32_t i = 0; i < MLDSA_N;) {
const uint32_t w0 = CRYPT_HTOLE32(buf[j]);
const uint32_t w1 = CRYPT_HTOLE32(buf[j + 1]);
const uint32_t w2 = CRYPT_HTOLE32(buf[j + 2]);
int32_t t0 = w0;
int32_t t1 = (w0 >> 24) | (w1 << 8);
int32_t t2 = (w1 >> 16) | (w2 << 16);
int32_t t3 = (w2 >> 8);
t0 &= 0x7FFFFFU;
t1 &= 0x7FFFFFU;
t2 &= 0x7FFFFFU;
t3 &= 0x7FFFFFU;
const int32_t m0 = (MLDSA_Q - 1 - t0) >> 31;
const int32_t m1 = (MLDSA_Q - 1 - t1) >> 31;
const int32_t m2 = (MLDSA_Q - 1 - t2) >> 31;
const int32_t m3 = (MLDSA_Q - 1 - t3) >> 31;
a[i] = t0 & ~m0;
i += 1 + m0;
if (i < MLDSA_N) {
a[i] = t1 & ~m1;
i += 1 + m1;
}
if (i < MLDSA_N) {
a[i] = t2 & ~m2;
i += 1 + m2;
}
if (i < MLDSA_N) {
a[i] = t3 & ~m3;
i += 1 + m3;
}
j += 3;
if (j >= buflen && i < MLDSA_N) {
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, (uint8_t *)buf, outlen), ret);
j = 0;
}
}
ERR:
hashMethod->freeCtx(mdCtx);
return ret;
}
static int32_t ExpandA(const CRYPT_ML_DSA_Ctx *ctx, const uint8_t *pubSeed, int32_t *matrix[MLDSA_K_MAX][MLDSA_L_MAX])
{
uint8_t k = ctx->info->k;
uint8_t l = ctx->info->l;
uint8_t seed[MLDSA_SEED_EXTEND_BYTES_LEN];
memcpy(seed, pubSeed, MLDSA_PUBLIC_SEED_LEN);
for (uint8_t i = 0; i < k; i++) {
for (uint8_t j = 0; j < l; j++) {
seed[MLDSA_PUBLIC_SEED_LEN] = j;
seed[MLDSA_PUBLIC_SEED_LEN + 1] = i;
int32_t ret = RejNTTPoly(matrix[i][j], seed);
RETURN_RET_IF(ret != CRYPT_SUCCESS, ret);
}
}
return CRYPT_SUCCESS;
}
static int32_t RejBoundedPolyEta2(int32_t *a, const uint8_t *s)
{
uint8_t buf[CRYPT_SHAKE256_BLOCKSIZE];
uint32_t bufLen = CRYPT_SHAKE256_BLOCKSIZE;
int32_t ret = CRYPT_SUCCESS;
const EAL_MdMethod *hashMethod = NULL;
void *mdCtx = NULL;
RETURN_RET_IF_ERR_EX(MLDSAInitHashCtx(CRYPT_MD_SHAKE256, &hashMethod, &mdCtx), ret);
GOTO_ERR_IF(hashMethod->update(mdCtx, s, MLDSA_PRIVATE_SEED_LEN + 2), ret);
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, buf, bufLen), ret);
for (uint32_t i = 0, j = 0; i < MLDSA_N; j++) {
if (j == CRYPT_SHAKE256_BLOCKSIZE) {
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, buf, CRYPT_SHAKE256_BLOCKSIZE), ret);
j = 0;
}
int32_t z0 = (int32_t)(buf[j] & 0x0F);
int32_t z1 = (int32_t)(buf[j] >> 4u);
int32_t mask = (0xE - z0) >> 31;
z0 = z0 - ((205 * z0) >> 10) * 5;
a[i] = (2 - z0) & ~mask;
i += 1 + mask;
if (i < MLDSA_N) {
mask = (0xE - z1) >> 31;
z1 = z1 - ((205 * z1) >> 10) * 5;
a[i] = (2 - z1) & ~mask;
i += 1 + mask;
}
}
ERR:
hashMethod->freeCtx(mdCtx);
return ret;
}
static int32_t RejBoundedPolyEta4(int32_t *a, const uint8_t *s)
{
uint8_t buf[CRYPT_SHAKE256_BLOCKSIZE];
uint32_t bufLen = CRYPT_SHAKE256_BLOCKSIZE;
int32_t ret = CRYPT_SUCCESS;
const EAL_MdMethod *hashMethod = NULL;
void *mdCtx = NULL;
RETURN_RET_IF_ERR_EX(MLDSAInitHashCtx(CRYPT_MD_SHAKE256, &hashMethod, &mdCtx), ret);
GOTO_ERR_IF(hashMethod->update(mdCtx, s, MLDSA_PRIVATE_SEED_LEN + 2), ret);
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, buf, bufLen), ret);
for (uint32_t i = 0, j = 0; i < MLDSA_N; j++) {
if (j == CRYPT_SHAKE256_BLOCKSIZE) {
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, buf, CRYPT_SHAKE256_BLOCKSIZE), ret);
j = 0;
}
int32_t z0 = (int32_t)(buf[j] & 0x0F);
int32_t z1 = (int32_t)(buf[j] >> 4u);
int32_t mask = (0x8 - z0) >> 31;
a[i] = (4 - z0) & ~mask;
i += 1 + mask;
if (i < MLDSA_N) {
mask = (0x8 - z1) >> 31;
a[i] = (4 - z1) & ~mask;
i += 1 + mask;
}
}
ERR:
hashMethod->freeCtx(mdCtx);
return ret;
}
static int32_t ExpandS(const CRYPT_ML_DSA_Ctx *ctx, const uint8_t *prvSeed,
int32_t *s1[MLDSA_L_MAX], int32_t *s2[MLDSA_K_MAX])
{
int32_t ret;
uint8_t k = ctx->info->k;
uint8_t l = ctx->info->l;
uint8_t seed[MLDSA_PRIVATE_SEED_LEN + 2];
memcpy(seed, prvSeed, MLDSA_PRIVATE_SEED_LEN);
seed[MLDSA_PRIVATE_SEED_LEN + 1] = 0;
int32_t (*rejBoundedPoly)(int32_t *a, const uint8_t *s);
if (ctx->info->eta == 2) {
rejBoundedPoly = RejBoundedPolyEta2;
} else {
rejBoundedPoly = RejBoundedPolyEta4;
}
for (uint8_t i = 0; i < l; i++) {
seed[MLDSA_PRIVATE_SEED_LEN] = i;
ret = rejBoundedPoly(s1[i], seed);
if (ret != CRYPT_SUCCESS) {
BSL_SAL_CleanseData(seed, MLDSA_PRIVATE_SEED_LEN + 2);
return ret;
}
}
for (uint8_t i = 0; i < k; i++) {
seed[MLDSA_PRIVATE_SEED_LEN] = l + i;
ret = rejBoundedPoly(s2[i], seed);
if (ret != CRYPT_SUCCESS) {
BSL_SAL_CleanseData(seed, MLDSA_PRIVATE_SEED_LEN + 2);
return ret;
}
}
BSL_SAL_CleanseData(seed, MLDSA_PRIVATE_SEED_LEN + 2);
return CRYPT_SUCCESS;
}
static void ComputesNTT(const CRYPT_ML_DSA_Ctx *ctx, int32_t *const s[MLDSA_L_MAX], int32_t *sOut[MLDSA_L_MAX])
{
for (uint8_t i = 0; i < ctx->info->l; i++) {
memcpy(sOut[i], s[i], sizeof(int32_t) * MLDSA_N);
MLDSA_ComputesNTT(sOut[i]);
}
}
static void VectorsMul(int32_t *t, const int32_t *matrix, const int32_t *s)
{
for (uint32_t i = 0; i < MLDSA_N; i++) {
t[i] = MLDSA_PlantardMulReduce((uint64_t)matrix[i] * (uint64_t)s[i] * (uint64_t)MLDSA_PLANTARD_INV);
}
}
static void MatrixMul(const CRYPT_ML_DSA_Ctx *ctx, int32_t *t, int32_t *const matrix[MLDSA_L_MAX],
int32_t *const s[MLDSA_L_MAX])
{
int64_t tmp[MLDSA_N] = { 0 };
for (uint32_t i = 0; i < ctx->info->l; i++) {
for (uint32_t j = 0; j < MLDSA_N; j++) {
tmp[j] += (int64_t)matrix[i][j] * s[i][j];
}
}
for (uint32_t j = 0; j < MLDSA_N; j++) {
t[j] = MLDSA_PlantardMulReduce((uint64_t)tmp[j] * (uint64_t)MLDSA_PLANTARD_INV);
}
}
static int32_t ComputesT(const CRYPT_ML_DSA_Ctx *ctx, int32_t *t[MLDSA_K_MAX], MLDSA_KeyGenMatrixSt *st, uint8_t*pub)
{
uint8_t seed[MLDSA_SEED_EXTEND_BYTES_LEN];
(void)memcpy(seed, pub, MLDSA_PUBLIC_SEED_LEN);
for (uint8_t i = 0; i < ctx->info->k; i++) {
for (uint8_t j = 0; j < ctx->info->l; j++) {
seed[MLDSA_PUBLIC_SEED_LEN] = j;
seed[MLDSA_PUBLIC_SEED_LEN + 1] = i;
int32_t ret = RejNTTPoly(st->matrix[j], seed);
RETURN_RET_IF(ret != CRYPT_SUCCESS, ret);
}
MatrixMul(ctx, t[i], st->matrix, st->s1Ntt);
MLDSA_ComputesINVNTT(t[i]);
for (int32_t j = 0; j < MLDSA_N; j++) {
t[i][j] = t[i][j] + st->s2[i][j];
t[i][j] = t[i][j] + (MLDSA_Q & (t[i][j] >> 31));
}
}
return CRYPT_SUCCESS;
}
static void ComputesPower2Round(const CRYPT_ML_DSA_Ctx *ctx, int32_t *t0[MLDSA_K_MAX], int32_t *t1[MLDSA_K_MAX])
{
for (uint32_t i = 0; i < ctx->info->k; i++) {
for (int32_t j = 0; j < MLDSA_N; j++) {
int32_t t = (t1[i][j] + (1 << (MLDSA_D - 1)) - 1) >> MLDSA_D;
t0[i][j] = t1[i][j] - (t << MLDSA_D);
t1[i][j] = t;
}
}
}
static void ByteEncode(uint8_t *buf, const uint32_t *t, uint32_t bits)
{
if (bits == 10u) {
for (uint32_t i = 0; i < MLDSA_N / 4; i++) {
buf[5 * i + 0] = (uint8_t)(t[4 * i + 0] >> 0);
buf[5 * i + 1u] = (uint8_t)((t[4 * i + 0] >> 8u) | (t[4 * i + 1u] << 2u));
buf[5 * i + 2u] = (uint8_t)((t[4 * i + 1u] >> 6u) | (t[4 * i + 2u] << 4u));
buf[5 * i + 3u] = (uint8_t)((t[4 * i + 2u] >> 4u) | (t[4 * i + 3u] << 6u));
buf[5 * i + 4u] = (uint8_t)(t[4 * i + 3u] >> 2u);
}
} else if (bits == 6u) {
for (uint32_t i = 0; i < MLDSA_N / 4; i++) {
buf[3 * i + 0] = (uint8_t)(t[4 * i] | (t[4 * i + 1] << 6u));
buf[3 * i + 1u] = (uint8_t)(t[4 * i + 1u] >> 2 | (t[4 * i + 2u] << 4u));
buf[3 * i + 2u] = (uint8_t)(t[4 * i + 2u] >> 4 | (t[4 * i + 3u] << 2u));
}
} else if (bits == 4u) {
for (uint32_t i = 0; i < MLDSA_N / 2; i++) {
buf[i] = (uint8_t)(t[2 * i] | (t[2 * i + 1] << 4u));
}
}
}
static void ByteDecode(const uint8_t *buf, uint32_t *t, uint32_t bits)
{
if (bits == 10u) {
for (uint32_t i = 0; i < MLDSA_N / 4; i++) {
t[4 * i + 0] = (buf[5 * i + 0] | ((uint32_t)buf[5 * i + 1] << 8)) & 0x03ff;
t[4 * i + 1u] = ((buf[5 * i + 1u] >> 2u) | ((uint32_t)buf[5 * i + 2u] << 6u)) & 0x03ff;
t[4 * i + 2u] = ((buf[5 * i + 2u] >> 4u) | ((uint32_t)buf[5 * i + 3u] << 4u)) & 0x03ff;
t[4 * i + 3u] = ((buf[5 * i + 3u] >> 6u) | ((uint32_t)buf[5 * i + 4u] << 2u)) & 0x03ff;
}
}
}
static void BitPack(uint8_t *buf, const uint32_t w[MLDSA_N], uint32_t bits, uint32_t b)
{
uint32_t t[8] = {0};
uint32_t i;
uint32_t n;
if (bits == 3u) {
for (i = 0; i < MLDSA_N / 8; i++) {
for (uint32_t j = 0; j < 8; j++) {
t[j] = b - (uint32_t)w[i * 8 + j];
}
n = bits * i;
buf[n + 0] = (uint8_t)((t[0]) | (t[1] << 3u) | (t[2] << 6u));
buf[n + 1u] = (uint8_t)((t[2] >> 2u) | (t[3] << 1u) | (t[4] << 4u) | (t[5] << 7u));
buf[n + 2u] = (uint8_t)((t[5] >> 1u) | (t[6] << 2u) | (t[7] << 5u));
}
} else if (bits == 4u) {
for (i = 0; i < MLDSA_N / 2; i++) {
t[0] = (int32_t)b - w[i * 2];
t[1] = (int32_t)b - w[i * 2 + 1];
buf[i] = (uint8_t)(t[0] | (t[1] << 4u));
}
} else if (bits == MLDSA_D) {
for (i = 0; i < MLDSA_N / 8; i++) {
for (uint32_t j = 0; j < 8; j++) {
t[j] = b - w[i * 8 + j];
}
n = bits * i;
buf[n + 0] = (uint8_t)t[0];
buf[n + 1] = (uint8_t)(t[0] >> 8u);
buf[n + 1] |= (uint8_t)(t[1] << 5u);
buf[n + 2] = (uint8_t)(t[1] >> 3u);
buf[n + 3] = (uint8_t)(t[1] >> 11u);
buf[n + 3] |= (uint8_t)(t[2] << 2u);
buf[n + 4] = (uint8_t)(t[2] >> 6u);
buf[n + 4] |= (uint8_t)(t[3] << 7u);
buf[n + 5] = (uint8_t)(t[3] >> 1u);
buf[n + 6] = (uint8_t)(t[3] >> 9u);
buf[n + 6] |= (uint8_t)(t[4] << 4u);
buf[n + 7] = (uint8_t)(t[4] >> 4u);
buf[n + 8] = (uint8_t)(t[4] >> 12u);
buf[n + 8] |= (uint8_t)(t[5] << 1u);
buf[n + 9] = (uint8_t)(t[5] >> 7u);
buf[n + 9] |= (uint8_t)(t[6] << 6u);
buf[n + 10] = (uint8_t)(t[6] >> 2u);
buf[n + 11] = (uint8_t)(t[6] >> 10u);
buf[n + 11] |= (uint8_t)(t[7] << 3u);
buf[n + 12] = (uint8_t)(t[7] >> 5u);
}
}
}
static void BitUnPake(const uint8_t *v, uint32_t w[MLDSA_N], uint32_t bits, uint32_t b)
{
uint32_t t[8] = {0};
uint32_t i;
uint32_t n;
if (bits == 3u) {
for (i = 0; i < MLDSA_N / 8; i++) {
n = bits * i;
t[0] = (v[n + 0]) & 0x07;
t[1] = (v[n + 0] >> 3u) & 0x07;
t[2] = ((v[n + 0] >> 6u) | (v[n + 1] << 2u)) & 0x07;
t[3] = (v[n + 1u] >> 1u) & 0x07;
t[4] = (v[n + 1u] >> 4u) & 0x07;
t[5] = ((v[n + 1u] >> 7u) | (v[n + 2] << 1u)) & 0x07;
t[6] = (v[n + 2u] >> 2u) & 0x07;
t[7] = (v[n + 2u] >> 5u) & 0x07;
for (uint32_t j = 0; j < 8; j++) {
w[i * 8 + j] = b - t[j];
}
}
} else if (bits == 4u) {
for (i = 0; i < MLDSA_N / 2; i++) {
t[0] = v[i] & 0x0f;
t[1] = (v[i] >> 4u) & 0x0f;
w[i * 2] = b - t[0];
w[i * 2 + 1] = b - t[1];
}
} else if (bits == MLDSA_D) {
for (i = 0; i < MLDSA_N / 8; i++) {
n = bits * i;
t[0] = (v[n + 0] | ((uint32_t)v[n + 1] << 8u)) & 0x1fff;
t[1] = (v[n + 1] >> 5u | ((uint32_t)v[n + 2u] << 3u) |
((uint32_t)v[n + 3u] << 11u)) & 0x1fff;
t[2] = (v[n + 3u] >> 2u | ((uint32_t)v[n + 4u] << 6u)) & 0x1fff;
t[3] = (v[n + 4u] >> 7u | ((uint32_t)v[n + 5u] << 1u) |
((uint32_t)v[n + 6u] << 9u)) & 0x1fff;
t[4] = (v[n + 6u] >> 4u | ((uint32_t)v[n + 7u] << 4u) |
((uint32_t)v[n + 8u] << 12u)) & 0x1fff;
t[5] = (v[n + 8u] >> 1u | ((uint32_t)v[n + 9u] << 7u)) & 0x1fff;
t[6] = (v[n + 9u] >> 6u | ((uint32_t)v[n + 10u] << 2u) |
((uint32_t)v[n + 11u] << 10u)) & 0x1fff;
t[7] = (v[n + 11u] >> 3u | ((uint32_t)v[n + 12u] << 5u)) & 0x1fff;
for (uint32_t j = 0; j < 8; j++) {
w[i * 8 + j] = b - t[j];
}
}
}
}
static void SignBitPack(uint8_t *buf, const uint32_t w[MLDSA_N], uint32_t bits, uint32_t b)
{
uint32_t t[4] = {0};
uint32_t i;
uint32_t n;
if (bits == GAMMA_BITS_OF_MLDSA_44) {
for (i = 0; i < MLDSA_N / 4; i++) {
for (uint32_t j = 0; j < 4; j++) {
t[j] = b - w[i * 4 + j];
}
n = 9 * i;
buf[n + 0] = (uint8_t)t[0];
buf[n + 1u] = (uint8_t)(t[0] >> 8u);
buf[n + 2u] = (uint8_t)(t[0] >> 16u | t[1] << 2u);
buf[n + 3u] = (uint8_t)(t[1] >> 6u);
buf[n + 4u] = (uint8_t)(t[1] >> 14u | t[2] << 4u);
buf[n + 5u] = (uint8_t)(t[2] >> 4u);
buf[n + 6u] = (uint8_t)(t[2] >> 12u | t[3] << 6u);
buf[n + 7u] = (uint8_t)(t[3] >> 2u);
buf[n + 8u] = (uint8_t)(t[3] >> 10u);
}
} else if (bits == GAMMA_BITS_OF_MLDSA_65_87) {
for (i = 0; i < MLDSA_N / 2; i++) {
t[0] = b - w[i * 2];
t[1] = b - w[i * 2 + 1u];
n = 5 * i;
buf[n + 0] = (uint8_t)t[0];
buf[n + 1u] = (uint8_t)(t[0] >> 8u);
buf[n + 2u] = (uint8_t)(t[0] >> 16u | t[1] << 4u);
buf[n + 3u] = (uint8_t)(t[1] >> 4u);
buf[n + 4u] = (uint8_t)(t[1] >> 12u);
}
}
}
static void SignBitUnPack(const uint8_t *v, uint32_t w[MLDSA_N], uint32_t bits, uint32_t b)
{
uint32_t t[4] = {0};
uint32_t i;
uint32_t n;
if (bits == GAMMA_BITS_OF_MLDSA_44) {
for (i = 0; i < MLDSA_N / 4; i++) {
n = 9 * i;
t[0] = (v[n + 0] | ((uint32_t)v[n + 1] << 8) | ((uint32_t)v[n + 2] << 16)) & 0x3ffff;
t[1] = (v[n + 2u] >> 2u | ((uint32_t)v[n + 3u] << 6u) | ((uint32_t)v[n + 4u] << 14u)) & 0x3ffff;
t[2] = (v[n + 4u] >> 4u | ((uint32_t)v[n + 5u] << 4u) | ((uint32_t)v[n + 6u] << 12u)) & 0x3ffff;
t[3] = (v[n + 6u] >> 6u | ((uint32_t)v[n + 7u] << 2u) | ((uint32_t)v[n + 8u] << 10u)) & 0x3ffff;
n = 4 * i;
w[n] = b - t[0];
w[n + 1u] = b - t[1];
w[n + 2u] = b - t[2];
w[n + 3u] = b - t[3];
}
} else if (bits == GAMMA_BITS_OF_MLDSA_65_87) {
for (i = 0; i < MLDSA_N / 2; i++) {
n = 5 * i;
t[0] = (v[n + 0] | ((uint32_t)v[n + 1] << 8u) | ((uint32_t)v[n + 2u] << 16u)) & 0xfffff;
t[1] = (v[n + 2u] >> 4u | ((uint32_t)v[n + 3u] << 4u) | ((uint32_t)v[n + 4u] << 12u)) & 0xfffff;
w[i * 2] = b - t[0];
w[i * 2 + 1u] = b - t[1];
}
}
}
static void PkEncode(const CRYPT_ML_DSA_Ctx *ctx, const uint8_t *seed, int32_t *const t[MLDSA_K_MAX])
{
memcpy(ctx->pubKey, seed, MLDSA_PUBLIC_SEED_LEN);
for (int32_t i = 0; i < ctx->info->k; i++) {
ByteEncode(ctx->pubKey + MLDSA_PUBLIC_SEED_LEN + i * MLDSA_PUBKEY_POLYT_PACKEDBYTES, (uint32_t *)t[i], 10);
}
}
static void PkDecode(const CRYPT_ML_DSA_Ctx *ctx, uint8_t *seed, int32_t *t[MLDSA_K_MAX])
{
memcpy(seed, ctx->pubKey, MLDSA_PUBLIC_SEED_LEN);
for (int32_t i = 0; i < ctx->info->k; i++) {
ByteDecode(ctx->pubKey + MLDSA_PUBLIC_SEED_LEN + i * MLDSA_PUBKEY_POLYT_PACKEDBYTES, (uint32_t *)t[i], 10);
}
}
static void SkEncode(const CRYPT_ML_DSA_Ctx *ctx, const uint8_t *pubSeed, const uint8_t *signSeed, const uint8_t *tr,
const MLDSA_KeyGenMatrixSt *st)
{
uint32_t i;
uint32_t bitLen = ctx->info->eta == 2 ? 3 : 4;
uint32_t index = MLDSA_PUBLIC_SEED_LEN;
memcpy(ctx->prvKey, pubSeed, MLDSA_PUBLIC_SEED_LEN);
memcpy(ctx->prvKey + index, signSeed, MLDSA_SIGNING_SEED_LEN);
index += MLDSA_SIGNING_SEED_LEN;
memcpy(ctx->prvKey + index, tr, MLDSA_TR_MSG_LEN);
index += MLDSA_TR_MSG_LEN;
for (i = 0; i < ctx->info->l; i++) {
BitPack(ctx->prvKey + index, (uint32_t *)st->s1[i], bitLen, ctx->info->eta);
index += MLDSA_N_BYTE * bitLen;
}
for (i = 0; i < ctx->info->k; i++) {
BitPack(ctx->prvKey + index, (uint32_t *)st->s2[i], bitLen, ctx->info->eta);
index += MLDSA_N_BYTE * bitLen;
}
for (i = 0; i < ctx->info->k; i++) {
BitPack(ctx->prvKey + index, (uint32_t *)st->t0[i], MLDSA_D, 4096);
index += MLDSA_N_BYTE * MLDSA_D;
}
}
static void SkDecode(const CRYPT_ML_DSA_Ctx *ctx, uint8_t *pubSeed, uint8_t *signSeed, uint8_t *tr,
MLDSA_SignMatrixSt *st)
{
uint32_t i;
uint32_t bitLen = ctx->info->eta == 2 ? 3 : 4;
uint32_t index = MLDSA_PUBLIC_SEED_LEN;
memcpy(pubSeed, ctx->prvKey, MLDSA_PUBLIC_SEED_LEN);
memcpy(signSeed, ctx->prvKey + index, MLDSA_SIGNING_SEED_LEN);
index += MLDSA_SIGNING_SEED_LEN;
memcpy(tr, ctx->prvKey + index, MLDSA_PRIVATE_SEED_LEN);
index += MLDSA_PRIVATE_SEED_LEN;
for (i = 0; i < ctx->info->l; i++) {
BitUnPake(ctx->prvKey + index, (uint32_t *)st->s1[i], bitLen, ctx->info->eta);
index += MLDSA_N_BYTE * bitLen;
}
for (i = 0; i < ctx->info->k; i++) {
BitUnPake(ctx->prvKey + index, (uint32_t *)st->s2[i], bitLen, ctx->info->eta);
index += MLDSA_N_BYTE * bitLen;
}
for (i = 0; i < ctx->info->k; i++) {
BitUnPake(ctx->prvKey + index, (uint32_t *)st->t0[i], MLDSA_D, 4096);
index += MLDSA_N_BYTE * MLDSA_D;
}
}
static void SignCalNtt(const CRYPT_ML_DSA_Ctx *ctx, MLDSA_SignMatrixSt *st)
{
uint32_t i;
for (i = 0; i < ctx->info->l; i++) {
MLDSA_ComputesNTT(st->s1[i]);
}
for (i = 0; i < ctx->info->k; i++) {
MLDSA_ComputesNTT(st->s2[i]);
}
for (i = 0; i < ctx->info->k; i++) {
MLDSA_ComputesNTT(st->t0[i]);
}
}
static int32_t ExpandMask(const CRYPT_ML_DSA_Ctx *ctx, int32_t *y[MLDSA_L_MAX], uint8_t *p, uint16_t u)
{
uint16_t n = 0;
uint8_t v[640];
uint32_t bits = (ctx->info->k == K_VALUE_OF_MLDSA_44) ? GAMMA_BITS_OF_MLDSA_44 : GAMMA_BITS_OF_MLDSA_65_87;
for (uint16_t i = 0; i < ctx->info->l; i++) {
n = u + i;
p[MLDSA_PRIVATE_SEED_LEN] = (uint8_t)n;
p[MLDSA_PRIVATE_SEED_LEN + 1] = (uint8_t)(n >> BITS_OF_BYTE);
int32_t ret = HashFuncH(p, MLDSA_PRIVATE_SEED_LEN + 2, NULL, 0, v, 32 * bits);
if (ret != CRYPT_SUCCESS) {
return ret;
}
SignBitUnPack(v, (uint32_t *)y[i], bits, ctx->info->gamma1);
}
return CRYPT_SUCCESS;
}
static void Decompose(const CRYPT_ML_DSA_Ctx *ctx, int32_t r, int32_t *r1, int32_t *r0)
{
int32_t t = (int32_t)(((uint32_t)r + 0x7f) >> 7u);
if (ctx->info->k == K_VALUE_OF_MLDSA_44) {
t = (t * 11275u + (1 << 23u)) >> 24u;
t ^= ((43 - t) >> 31u) & t;
} else {
t = (t * 1025u + (1 << 21u)) >> 22u;
t &= 0x0f;
}
*r0 = r - t * 2 * ctx->info->gamma2;
*r0 -= (((MLDSA_Q - 1) / 2 - *r0) >> 31u) & MLDSA_Q;
*r1 = t;
}
static void ComputesW(const CRYPT_ML_DSA_Ctx *ctx, int32_t *w[MLDSA_L_MAX], int32_t *w1[MLDSA_L_MAX],
int32_t *const matrix[MLDSA_K_MAX][MLDSA_L_MAX], int32_t *const y[MLDSA_L_MAX])
{
for (uint8_t i = 0; i < ctx->info->k; i++) {
MatrixMul(ctx, w[i], matrix[i], y);
MLDSA_ComputesINVNTT(w[i]);
for (int32_t j = 0; j < MLDSA_N; j++) {
w[i][j] = w[i][j] + (MLDSA_Q & (w[i][j] >> 31));
Decompose(ctx, w[i][j], &w1[i][j], &w[i][j]);
}
}
}
static void W1Encode(const CRYPT_ML_DSA_Ctx *ctx, uint8_t *buf, int32_t *const w[MLDSA_K_MAX])
{
uint32_t bitLen = ctx->info->k == K_VALUE_OF_MLDSA_44 ? 6 : 4;
uint32_t blockSize = ctx->info->k == K_VALUE_OF_MLDSA_44 ? 192 : 128;
for (uint32_t i = 0; i < ctx->info->k; i++) {
ByteEncode(buf + i * blockSize, (const uint32_t *)w[i], bitLen);
}
}
static int32_t SampleInBall(const CRYPT_ML_DSA_Ctx *ctx, const uint8_t *p, uint32_t pLen, int32_t c[MLDSA_N])
{
uint8_t s[CRYPT_SHAKE256_BLOCKSIZE] = {0};
uint32_t sLen = CRYPT_SHAKE256_BLOCKSIZE;
uint64_t h = 0;
uint32_t index = 0;
uint8_t j = 0;
int32_t ret;
const EAL_MdMethod *hashMethod = NULL;
void *mdCtx = NULL;
RETURN_RET_IF_ERR_EX(MLDSAInitHashCtx(CRYPT_MD_SHAKE256, &hashMethod, &mdCtx), ret);
GOTO_ERR_IF(hashMethod->update(mdCtx, p, pLen), ret);
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, s, sLen), ret);
for (index = 0; index < 8; index++) {
h = h | ((uint64_t)s[index] << (8 * index));
}
for (uint32_t i = MLDSA_N - ctx->info->tau; i < MLDSA_N; i++) {
do {
if (index == CRYPT_SHAKE256_BLOCKSIZE) {
GOTO_ERR_IF(hashMethod->squeeze(mdCtx, s, sLen), ret);
index = 0;
}
j = s[index];
index++;
} while (j > i);
c[i] = c[j];
c[j] = 1 - ((h & 1) << 1);
h >>= 1;
}
ERR:
hashMethod->freeCtx(mdCtx);
return ret;
}
static void MLDSA_VectorsAdd(int32_t *t, int32_t *a, int32_t *b)
{
for (uint32_t i = 0; i < MLDSA_N; i++) {
t[i] = a[i] + b[i];
MLDSA_MOD_Q(t[i]);
}
}
static void MLDSA_VectorsSub(int32_t *t, int32_t *a, int32_t *b)
{
for (uint32_t i = 0; i < MLDSA_N; i++) {
t[i] = a[i] - b[i];
MLDSA_MOD_Q(t[i]);
}
}
static void ComputesZ(const CRYPT_ML_DSA_Ctx *ctx, int32_t *y[MLDSA_L_MAX], const int32_t *c,
int32_t *const s[MLDSA_L_MAX], int32_t *const z[MLDSA_L_MAX])
{
for (uint8_t i = 0; i < ctx->info->l; i++) {
VectorsMul(z[i], c, s[i]);
MLDSA_ComputesINVNTT(z[i]);
MLDSA_VectorsAdd(z[i], y[i], z[i]);
}
}
static bool ValidityChecks(const int32_t *z, uint32_t t)
{
uint32_t n;
uint32_t result = 0;
for (uint32_t j = 0; j < MLDSA_N; j++) {
n = z[j] >> 31;
n = z[j] - (n & ((uint32_t)z[j] << 1));
result |= ((t - 1 - n) >> 31) & 1;
}
return (result == 0);
}
static bool ValidityChecksL(const CRYPT_ML_DSA_Ctx *ctx, int32_t *const z[MLDSA_L_MAX], uint32_t t)
{
bool valid = true;
for (uint8_t i = 0; i < ctx->info->l; i++) {
valid &= ValidityChecks(z[i], t);
}
return valid;
}
static bool ValidityChecksK(const CRYPT_ML_DSA_Ctx *ctx, int32_t *const z[MLDSA_K_MAX], uint32_t t)
{
bool valid = true;
for (uint8_t i = 0; i < ctx->info->k; i++) {
valid &= ValidityChecks(z[i], t);
}
return valid;
}
static void ComputesR(const CRYPT_ML_DSA_Ctx *ctx, const int32_t *c, MLDSA_SignMatrixSt *st)
{
for (uint8_t i = 0; i < ctx->info->k; i++) {
VectorsMul(st->y[i], c, st->s2[i]);
MLDSA_ComputesINVNTT(st->y[i]);
MLDSA_VectorsSub(st->r0[i], st->w[i], st->y[i]);
}
}
static void ComputesCT(const CRYPT_ML_DSA_Ctx *ctx, const int32_t *c,
int32_t *const t[MLDSA_K_MAX], int32_t *ct[MLDSA_K_MAX])
{
for (uint8_t i = 0; i < ctx->info->k; i++) {
VectorsMul(ct[i], c, t[i]);
MLDSA_ComputesINVNTT(ct[i]);
}
}
static uint32_t MakeHint(const CRYPT_ML_DSA_Ctx *ctx, MLDSA_SignMatrixSt *st)
{
uint32_t num = 0;
int32_t g = (int32_t)ctx->info->gamma2;
for (uint32_t i = 0; i < ctx->info->k; i++) {
for (uint32_t j = 0; j < MLDSA_N; j++) {
int32_t v = st->w[i][j] + st->r0[i][j] - st->y[i][j];
MLDSA_MOD_Q(v);
uint32_t x = (uint32_t)(v + g);
uint32_t c1 = ((uint32_t)(g - v) >> 31) & 1;
uint32_t c2 = (x >> 31) & 1;
uint32_t isZero = ((x | (0 - x)) >> 31) ^ 1;
uint32_t y = (uint32_t)st->w1[i][j];
uint32_t isNonZero = ((y | (0 - y)) >> 31) & 1;
uint32_t bit = c1 | c2 | (isZero & isNonZero);
st->w[i][j] = (int32_t)bit;
num += bit;
}
}
return num;
}
static void SigEncode(const CRYPT_ML_DSA_Ctx *ctx, uint8_t *out, uint32_t outLen, int32_t *const z[MLDSA_L_MAX],
int32_t *const h[MLDSA_K_MAX])
{
uint32_t bits = (ctx->info->k == K_VALUE_OF_MLDSA_44) ? GAMMA_BITS_OF_MLDSA_44 : GAMMA_BITS_OF_MLDSA_65_87;
uint32_t blockSize = MLDSA_N / BITS_OF_BYTE * bits;
uint8_t *ptr = out;
uint32_t index = 0;
for (uint32_t i = 0; i < ctx->info->l; i++) {
SignBitPack(ptr, (const uint32_t *)z[i], bits, ctx->info->gamma1);
ptr += blockSize;
}
memset(ptr, 0, outLen - blockSize * ctx->info->l);
for (uint32_t i = 0; i < ctx->info->k; i++) {
for (uint32_t j = 0; j < MLDSA_N; j++) {
if (h[i][j] != 0) {
ptr[index] = j;
index++;
}
}
ptr[ctx->info->omega + i] = index;
}
}
static int32_t SigDecode(const CRYPT_ML_DSA_Ctx *ctx, const uint8_t *in, int32_t *z[MLDSA_L_MAX],
int32_t *h[MLDSA_K_MAX])
{
uint32_t bits = (ctx->info->k == K_VALUE_OF_MLDSA_44) ? GAMMA_BITS_OF_MLDSA_44 : GAMMA_BITS_OF_MLDSA_65_87;
uint32_t blockSize = MLDSA_N / BITS_OF_BYTE * bits;
const uint8_t *ptr = in;
uint32_t index = 0;
for (int32_t i = 0; i < ctx->info->l; i++) {
SignBitUnPack(ptr, (uint32_t *)z[i], bits, ctx->info->gamma1);
ptr += blockSize;
}
for (int32_t i = 0; i < ctx->info->k; i++) {
if (ptr[ctx->info->omega + i] < index || ptr[ctx->info->omega + i] > ctx->info->omega) {
BSL_ERR_PUSH_ERROR(CRYPT_MLDSA_SIGN_DATA_ERROR);
return CRYPT_MLDSA_SIGN_DATA_ERROR;
}
uint32_t first = index;
memset(h[i], 0, sizeof(int32_t) * MLDSA_N);
while (index < ptr[ctx->info->omega + i]) {
if (index > first && (ptr[index - 1] >= ptr[index])) {
BSL_ERR_PUSH_ERROR(CRYPT_MLDSA_SIGN_DATA_ERROR);
return CRYPT_MLDSA_SIGN_DATA_ERROR;
}
h[i][ptr[index]] = 1;
index++;
}
}
for (int32_t i = index; i <= (ctx->info->omega - 1); i++) {
RETURN_RET_IF(ptr[i] != 0, CRYPT_MLDSA_SIGN_DATA_ERROR);
}
return CRYPT_SUCCESS;
}
static int32_t ComputesApproxW(const CRYPT_ML_DSA_Ctx *ctx, MLDSA_VerifyMatrixSt *st, const uint8_t *pubSeed,
int32_t *c, int32_t *w[MLDSA_K_MAX])
{
uint8_t seed[MLDSA_SEED_EXTEND_BYTES_LEN];
(void)memcpy(seed, pubSeed, MLDSA_PUBLIC_SEED_LEN);
MLDSA_ComputesNTT(c);
for (uint8_t i = 0; i < ctx->info->l; i++) {
MLDSA_ComputesNTT(st->z[i]);
}
for (uint8_t i = 0; i < ctx->info->k; i++) {
for (uint8_t j = 0; j < ctx->info->l; j++) {
seed[MLDSA_PUBLIC_SEED_LEN] = j;
seed[MLDSA_PUBLIC_SEED_LEN + 1] = i;
int32_t ret = RejNTTPoly(st->matrix[j], seed);
RETURN_RET_IF(ret != CRYPT_SUCCESS, ret);
}
for (int32_t j = 0; j < MLDSA_N; j++) {
st->t1[i][j] = (int32_t)((uint32_t)st->t1[i][j] << MLDSA_D);
}
MLDSA_ComputesNTT(st->t1[i]);
VectorsMul(st->t1[i], st->t1[i], c);
MatrixMul(ctx, w[i], st->matrix, st->z);
MLDSA_VectorsSub(w[i], w[i], st->t1[i]);
MLDSA_ComputesINVNTT(w[i]);
}
return CRYPT_SUCCESS;
}
static void UseHint(const CRYPT_ML_DSA_Ctx *ctx, int32_t *const h[MLDSA_K_MAX], int32_t *w[MLDSA_K_MAX])
{
int32_t r1;
int32_t r0;
for (uint8_t i = 0; i < ctx->info->k; i++) {
for (uint32_t j = 0; j < MLDSA_N; j++) {
w[i][j] = w[i][j] + (MLDSA_Q & (w[i][j] >> 31));
Decompose(ctx, w[i][j], &r1, &r0);
if (h[i][j] == 0) {
w[i][j] = r1;
continue;
}
if (ctx->info->gamma2 == 95232) {
w[i][j] = (r0 > 0) ? ((r1 == 43) ? 0 : (r1 + 1)) : ((r1 == 0) ? 43 : (r1 - 1));
continue;
}
w[i][j] = ((r0 > 0) ? (r1 + 1) : (r1 - 1)) & 0x0f;
}
}
}
int32_t MLDSA_KeyGenInternal(CRYPT_ML_DSA_Ctx *ctx, const uint8_t *d)
{
uint8_t k = ctx->info->k;
uint8_t l = ctx->info->l;
uint8_t seed[MLDSA_SEED_EXTEND_BYTES_LEN] = { 0 };
uint8_t digest[MLDSA_EXPANDED_SEED_BYTES_LEN] = { 0 };
uint8_t tr[MLDSA_TR_MSG_LEN] = { 0 };
MLDSA_KeyGenMatrixSt st = { 0 };
int32_t ret;
GOTO_ERR_IF(MLDSAKeyGenCreateMatrix(k, l, &st), ret);
memcpy(seed, d, MLDSA_SEED_BYTES_LEN);
seed[MLDSA_SEED_BYTES_LEN] = k;
seed[MLDSA_SEED_BYTES_LEN + 1] = l;
GOTO_ERR_IF(HashFuncH(seed, sizeof(seed), NULL, 0, digest, MLDSA_EXPANDED_SEED_BYTES_LEN), ret);
uint8_t *pubSeed = digest;
uint8_t *prvSeed = digest + MLDSA_PUBLIC_SEED_LEN;
uint8_t *signSeed = digest + MLDSA_PUBLIC_SEED_LEN + MLDSA_PRIVATE_SEED_LEN;
GOTO_ERR_IF(ExpandS(ctx, prvSeed, st.s1, st.s2), ret);
ComputesNTT(ctx, st.s1, st.s1Ntt);
GOTO_ERR_IF(ComputesT(ctx, st.t1, &st, pubSeed), ret);
ComputesPower2Round(ctx, st.t0, st.t1);
PkEncode(ctx, pubSeed, st.t1);
GOTO_ERR_IF(HashFuncH(ctx->pubKey, ctx->pubLen, NULL, 0, tr, MLDSA_TR_MSG_LEN), ret);
SkEncode(ctx, pubSeed, signSeed, tr, &st);
ctx->hasSeed = true;
memcpy(ctx->seed, d, MLDSA_SEED_BYTES_LEN);
ERR:
BSL_SAL_ClearFree(st.bufAddr, st.bufSize);
BSL_SAL_CleanseData(seed, sizeof(seed));
BSL_SAL_CleanseData(digest, sizeof(digest));
return ret;
}
int32_t MLDSA_SignInternal(const CRYPT_ML_DSA_Ctx *ctx, const CRYPT_Data *msg, uint8_t *out, uint32_t *outLen,
const uint8_t *rand)
{
int32_t ret = CRYPT_SUCCESS;
uint8_t pubSeed[MLDSA_PUBLIC_SEED_LEN];
uint8_t uBuf[MLDSA_XOF_MSG_LEN];
uint8_t tr[MLDSA_TR_MSG_LEN];
uint8_t signSeed[MLDSA_SIGNING_SEED_LEN + MLDSA_SEED_BYTES_LEN];
memcpy(signSeed + MLDSA_SIGNING_SEED_LEN, rand, MLDSA_SEED_BYTES_LEN);
uint32_t w1Len = (ctx->info->k == 4 || ctx->info->k == 6) ? 768 : 1024;
uint8_t *w1Buf = BSL_SAL_Malloc(w1Len);
RETURN_RET_IF(w1Buf == NULL, CRYPT_MEM_ALLOC_FAIL);
MLDSA_SignMatrixSt st = { 0 };
GOTO_ERR_IF(MLDSASignCreateMatrix(ctx->info->k, ctx->info->l, &st), ret);
SkDecode(ctx, pubSeed, signSeed, tr, &st);
SignCalNtt(ctx, &st);
GOTO_ERR_IF(ExpandA(ctx, pubSeed, st.matrix), ret);
if (ctx->isMuMsg) {
memcpy(uBuf, msg->data, msg->len);
} else {
GOTO_ERR_IF(HashFuncH(tr, MLDSA_TR_MSG_LEN, msg->data, msg->len, uBuf, MLDSA_XOF_MSG_LEN), ret);
}
uint8_t p[MLDSA_XOF_MSG_LEN + 2];
GOTO_ERR_IF(HashFuncH(signSeed, sizeof(signSeed), uBuf, MLDSA_XOF_MSG_LEN, p, MLDSA_XOF_MSG_LEN), ret);
uint16_t u = 0;
uint32_t cBufLen = ctx->info->secBits / 4;
int32_t c[MLDSA_N];
do {
GOTO_ERR_IF(ExpandMask(ctx, st.y, p, u), ret);
u = u + ctx->info->l;
ComputesNTT(ctx, st.y, st.z);
ComputesW(ctx, st.w, st.w1, st.matrix, st.z);
W1Encode(ctx, w1Buf, st.w1);
GOTO_ERR_IF(HashFuncH(uBuf, MLDSA_XOF_MSG_LEN, w1Buf, w1Len, out, cBufLen), ret);
memset(c, 0, sizeof(c));
GOTO_ERR_IF(SampleInBall(ctx, out, cBufLen, c), ret);
MLDSA_ComputesNTT(c);
ComputesZ(ctx, st.y, c, st.s1, st.z);
if (ValidityChecksL(ctx, st.z, ctx->info->gamma1 - ctx->info->beta) == false) {
continue;
}
ComputesR(ctx, c, &st);
if (ValidityChecksK(ctx, st.r0, ctx->info->gamma2 - ctx->info->beta) == false) {
continue;
}
ComputesCT(ctx, c, st.t0, st.r0);
if (ValidityChecksK(ctx, st.r0, ctx->info->gamma2) == false) {
continue;
}
if (MakeHint(ctx, &st) > ctx->info->omega) {
continue;
}
break;
} while (true);
*outLen = ctx->info->signatureLen;
SigEncode(ctx, out + cBufLen, *outLen - cBufLen, st.z, st.w);
ERR:
BSL_SAL_ClearFree(st.bufAddr, st.bufSize);
BSL_SAL_ClearFree(w1Buf, w1Len);
BSL_SAL_CleanseData(signSeed, sizeof(signSeed));
BSL_SAL_CleanseData(p, sizeof(p));
return ret;
}
int32_t MLDSA_VerifyInternal(const CRYPT_ML_DSA_Ctx *ctx, const CRYPT_Data *msg, const uint8_t *sign, uint32_t signLen)
{
(void)signLen;
uint8_t k = ctx->info->k;
uint8_t l = ctx->info->l;
uint8_t pubSeed[MLDSA_PUBLIC_SEED_LEN];
uint8_t uBuf[MLDSA_XOF_MSG_LEN];
uint8_t cBuf[MLDSA_XOF_MSG_LEN];
uint8_t tr[MLDSA_TR_MSG_LEN];
uint32_t cBufLen = ctx->info->secBits / 4;
MLDSA_VerifyMatrixSt st = { 0 };
int32_t c[MLDSA_N] = { 0 };
int32_t ret;
uint32_t w1Len = (k == 4 || k == 6) ? 768 : 1024;
uint8_t *w1Buf = BSL_SAL_Malloc(w1Len);
RETURN_RET_IF(w1Buf == NULL, CRYPT_MEM_ALLOC_FAIL);
GOTO_ERR_IF(MLDSAVerifyCreateMatrix(k, l, &st), ret);
PkDecode(ctx, pubSeed, st.t1);
GOTO_ERR_IF(SigDecode(ctx, sign + cBufLen, st.z, st.h), ret);
if (ValidityChecksL(ctx, st.z, ctx->info->gamma1 - ctx->info->beta) == false) {
ret = CRYPT_MLDSA_SIGN_DATA_ERROR;
goto ERR;
}
if (ctx->isMuMsg) {
memcpy(uBuf, msg->data, msg->len);
} else {
GOTO_ERR_IF(HashFuncH(ctx->pubKey, ctx->pubLen, NULL, 0, tr, MLDSA_TR_MSG_LEN), ret);
GOTO_ERR_IF(HashFuncH(tr, MLDSA_TR_MSG_LEN, msg->data, msg->len, uBuf, MLDSA_XOF_MSG_LEN), ret);
}
GOTO_ERR_IF(SampleInBall(ctx, sign, cBufLen, c), ret);
GOTO_ERR_IF(ComputesApproxW(ctx, &st, pubSeed, c, st.w), ret);
UseHint(ctx, st.h, st.w);
W1Encode(ctx, w1Buf, st.w);
GOTO_ERR_IF(HashFuncH(uBuf, MLDSA_XOF_MSG_LEN, w1Buf, w1Len, cBuf, cBufLen), ret);
if (memcmp(sign, cBuf, cBufLen) != 0) {
BSL_ERR_PUSH_ERROR(CRYPT_MLDSA_VERIFY_FAIL);
ret = CRYPT_MLDSA_VERIFY_FAIL;
goto ERR;
}
ERR:
BSL_SAL_Free(st.bufAddr);
BSL_SAL_Free(w1Buf);
return ret;
}
static void DecodePrvKey(const CRYPT_ML_DSA_Ctx *ctx, uint8_t *pubSeed, MLDSA_KeyGenMatrixSt *st)
{
uint32_t bitLen = ctx->info->eta == 2 ? 3 : 4;
uint32_t index = MLDSA_PUBLIC_SEED_LEN + MLDSA_SIGNING_SEED_LEN + MLDSA_PRIVATE_SEED_LEN;
(void)memcpy(pubSeed, ctx->prvKey, MLDSA_PUBLIC_SEED_LEN);
uint32_t i = 0;
for (i = 0; i < ctx->info->l; i++) {
BitUnPake(ctx->prvKey + index, (uint32_t *)st->s1[i], bitLen, ctx->info->eta);
index += MLDSA_N_BYTE * bitLen;
}
for (i = 0; i < ctx->info->k; i++) {
BitUnPake(ctx->prvKey + index, (uint32_t *)st->s2[i], bitLen, ctx->info->eta);
index += MLDSA_N_BYTE * bitLen;
}
for (i = 0; i < ctx->info->k; i++) {
BitUnPake(ctx->prvKey + index, (uint32_t *)st->t0[i], MLDSA_D, 4096);
index += MLDSA_N_BYTE * MLDSA_D;
}
}
int32_t MLDSA_CalPub(const CRYPT_ML_DSA_Ctx *ctx, uint8_t *pub, uint32_t pubLen)
{
int32_t ret;
MLDSA_KeyGenMatrixSt st = { 0 };
uint8_t pubSeed[MLDSA_PUBLIC_SEED_LEN];
GOTO_ERR_IF(MLDSAKeyGenCreateMatrix(ctx->info->k, ctx->info->l, &st), ret);
DecodePrvKey(ctx, pubSeed, &st);
ComputesNTT(ctx, st.s1, st.s1Ntt);
GOTO_ERR_IF(ComputesT(ctx, st.t1, &st, pubSeed), ret);
ComputesPower2Round(ctx, st.s2, st.t1);
for (int32_t i = 0; i < ctx->info->k; i++) {
if (memcmp(st.s2[i], st.t0[i], MLDSA_N * sizeof(int32_t)) != 0) {
BSL_ERR_PUSH_ERROR(CRYPT_MLDSA_PAIRWISE_CHECK_FAIL);
ret = CRYPT_MLDSA_PAIRWISE_CHECK_FAIL;
goto ERR;
}
}
if (MLDSA_PUBLIC_SEED_LEN > pubLen) {
BSL_ERR_PUSH_ERROR(CRYPT_MLDSA_LEN_NOT_ENOUGH);
ret = CRYPT_MLDSA_LEN_NOT_ENOUGH;
goto ERR;
}
memcpy(pub, pubSeed, MLDSA_PUBLIC_SEED_LEN);
for (int32_t i = 0; i < ctx->info->k; i++) {
ByteEncode(pub + MLDSA_PUBLIC_SEED_LEN + i * MLDSA_PUBKEY_POLYT_PACKEDBYTES, (uint32_t *)st.t1[i], 10);
}
ERR:
BSL_SAL_ClearFree(st.bufAddr, st.bufSize);
BSL_SAL_CleanseData(pubSeed, sizeof(pubSeed));
return ret;
}
int32_t MLDSA_KeyConsistenceCheck(CRYPT_ML_DSA_Ctx *ctx)
{
int32_t ret = CRYPT_SUCCESS;
uint8_t *pubKey = BSL_SAL_Malloc(ctx->info->publicKeyLen);
if (pubKey == NULL) {
BSL_ERR_PUSH_ERROR(CRYPT_MEM_ALLOC_FAIL);
return CRYPT_MEM_ALLOC_FAIL;
}
ret = MLDSA_CalPub(ctx, pubKey, ctx->info->publicKeyLen);
if (ret != CRYPT_SUCCESS) {
BSL_SAL_FREE(pubKey);
BSL_ERR_PUSH_ERROR(ret);
return ret;
}
uint8_t tr[MLDSA_TR_MSG_LEN] = {0};
ret = HashFuncH(pubKey, ctx->info->publicKeyLen, NULL, 0, tr, MLDSA_TR_MSG_LEN);
if (ret != CRYPT_SUCCESS) {
BSL_SAL_FREE(pubKey);
BSL_ERR_PUSH_ERROR(ret);
return ret;
}
if (memcmp(tr, ctx->prvKey + MLDSA_PUBLIC_SEED_LEN + MLDSA_SIGNING_SEED_LEN, MLDSA_TR_MSG_LEN) != 0) {
BSL_SAL_FREE(pubKey);
BSL_ERR_PUSH_ERROR(CRYPT_MLDSA_PAIRWISE_CHECK_FAIL);
return CRYPT_MLDSA_PAIRWISE_CHECK_FAIL;
}
if (ctx->pubKey == NULL) {
ctx->pubKey = pubKey;
ctx->pubLen = ctx->info->publicKeyLen;
} else {
if (memcmp(pubKey, ctx->pubKey, ctx->info->publicKeyLen) != 0) {
BSL_SAL_FREE(pubKey);
BSL_ERR_PUSH_ERROR(CRYPT_MLDSA_PAIRWISE_CHECK_FAIL);
return CRYPT_MLDSA_PAIRWISE_CHECK_FAIL;
}
BSL_SAL_FREE(pubKey);
}
return CRYPT_SUCCESS;
}
#endif