#ifndef MLD_ROUNDING_H
#define MLD_ROUNDING_H
#include "cbmc.h"
#include "common.h"
#include "ct.h"
#include "debug.h"
#define mld_power2round MLD_ADD_PARAM_SET(mld_power2round)
#define mld_decompose MLD_ADD_PARAM_SET(mld_decompose)
#define mld_make_hint MLD_ADD_PARAM_SET(mld_make_hint)
#define mld_use_hint MLD_ADD_PARAM_SET(mld_use_hint)
#define MLD_2_POW_D (1 << MLDSA_D)
static MLD_INLINE void mld_power2round(int32_t *a0, int32_t *a1, int32_t a)
__contract__(
requires(memory_no_alias(a0, sizeof(int32_t)))
requires(memory_no_alias(a1, sizeof(int32_t)))
requires(a >= 0 && a < MLDSA_Q)
assigns(memory_slice(a0, sizeof(int32_t)))
assigns(memory_slice(a1, sizeof(int32_t)))
ensures(*a0 > -(MLD_2_POW_D/2) && *a0 <= (MLD_2_POW_D/2))
ensures(*a1 >= 0 && *a1 <= (MLDSA_Q - 1) / MLD_2_POW_D)
ensures((*a1 * MLD_2_POW_D + *a0 - a) % MLDSA_Q == 0)
)
{
*a1 = (a + (1 << (MLDSA_D - 1)) - 1) >> MLDSA_D;
*a0 = a - (*a1 << MLDSA_D);
}
static MLD_INLINE void mld_decompose(int32_t *a0, int32_t *a1, int32_t a)
__contract__(
requires(memory_no_alias(a0, sizeof(int32_t)))
requires(memory_no_alias(a1, sizeof(int32_t)))
requires(a >= 0 && a < MLDSA_Q)
assigns(memory_slice(a0, sizeof(int32_t)))
assigns(memory_slice(a1, sizeof(int32_t)))
ensures(*a0 >= -MLDSA_GAMMA2 && *a0 <= MLDSA_GAMMA2)
ensures(*a1 >= 0 && *a1 < (MLDSA_Q-1)/(2*MLDSA_GAMMA2))
ensures((*a1 * 2 * MLDSA_GAMMA2 + *a0 - a) % MLDSA_Q == 0)
)
{
*a1 = (a + 127) >> 7;
mld_assert(*a1 >= 0 && *a1 <= 65472);
#if MLD_CONFIG_PARAMETER_SET == 44
*a1 = (*a1 * 11275 + ((int32_t)1 << 23)) >> 24;
mld_assert(*a1 >= 0 && *a1 <= 44);
*a1 = mld_ct_sel_int32(0, *a1, mld_ct_cmask_neg_i32(43 - *a1));
mld_assert(*a1 >= 0 && *a1 <= 43);
#else
*a1 = (*a1 * 1025 + ((int32_t)1 << 21)) >> 22;
mld_assert(*a1 >= 0 && *a1 <= 16);
*a1 &= 15;
mld_assert(*a1 >= 0 && *a1 <= 15);
#endif
*a0 = a - *a1 * 2 * MLDSA_GAMMA2;
*a0 = mld_ct_sel_int32(*a0 - MLDSA_Q, *a0,
mld_ct_cmask_neg_i32((MLDSA_Q - 1) / 2 - *a0));
}
MLD_MUST_CHECK_RETURN_VALUE
static MLD_INLINE unsigned int mld_make_hint(int32_t a0, int32_t a1)
__contract__(
ensures(return_value >= 0 && return_value <= 1)
)
{
if (a0 > MLDSA_GAMMA2 || a0 < -MLDSA_GAMMA2 ||
(a0 == -MLDSA_GAMMA2 && a1 != 0))
{
return 1;
}
return 0;
}
MLD_MUST_CHECK_RETURN_VALUE
static MLD_INLINE int32_t mld_use_hint(int32_t a, int32_t hint)
__contract__(
requires(hint >= 0 && hint <= 1)
requires(a >= 0 && a < MLDSA_Q)
ensures(return_value >= 0 && return_value < (MLDSA_Q-1)/(2*MLDSA_GAMMA2))
)
{
int32_t a0, a1;
mld_decompose(&a0, &a1, a);
if (hint == 0)
{
return a1;
}
#if MLD_CONFIG_PARAMETER_SET == 44
if (a0 > 0)
{
return (a1 == 43) ? 0 : a1 + 1;
}
else
{
return (a1 == 0) ? 43 : a1 - 1;
}
#else
if (a0 > 0)
{
return (a1 + 1) & 15;
}
else
{
return (a1 - 1) & 15;
}
#endif
}
#endif