xref: /freebsd/crypto/openssl/crypto/ml_dsa/ml_dsa_vector.h (revision 78e936b2d0b5e6554425009199be31e76bc67c10)
1 /*
2  * Copyright 2024-2026 The OpenSSL Project Authors. All Rights Reserved.
3  *
4  * Licensed under the Apache License 2.0 (the "License").  You may not use
5  * this file except in compliance with the License.  You can obtain a copy
6  * in the file LICENSE in the source distribution or at
7  * https://www.openssl.org/source/license.html
8  */
9 
10 #include <assert.h>
11 #include <openssl/crypto.h>
12 #include "ml_dsa_poly.h"
13 
14 struct vector_st {
15     POLY *poly;
16     size_t num_poly;
17 };
18 
19 /**
20  * @brief Initialize a Vector object.
21  *
22  * @param v The vector to initialize.
23  * @param polys Preallocated storage for an array of Polynomials blocks. |v|
24  *              does not own/free this.
25  * @param num_polys The number of |polys| blocks (k or l)
26  */
vector_init(VECTOR * v,POLY * polys,size_t num_polys)27 static ossl_inline ossl_unused void vector_init(VECTOR *v, POLY *polys, size_t num_polys)
28 {
29     v->poly = polys;
30     v->num_poly = num_polys;
31 }
32 
vector_alloc(VECTOR * v,size_t num_polys)33 static ossl_inline ossl_unused int vector_alloc(VECTOR *v, size_t num_polys)
34 {
35     v->poly = OPENSSL_malloc(num_polys * sizeof(POLY));
36     if (v->poly == NULL)
37         return 0;
38     v->num_poly = num_polys;
39     return 1;
40 }
41 
vector_free(VECTOR * v)42 static ossl_inline ossl_unused void vector_free(VECTOR *v)
43 {
44     OPENSSL_free(v->poly);
45     v->poly = NULL;
46     v->num_poly = 0;
47 }
48 
49 /* @brief zeroize a vectors polynomial coefficients */
vector_zero(VECTOR * va)50 static ossl_inline ossl_unused void vector_zero(VECTOR *va)
51 {
52     if (va->poly != NULL)
53         memset(va->poly, 0, va->num_poly * sizeof(va->poly[0]));
54 }
55 
56 /*
57  * @brief copy a vector
58  * The assumption is that |dst| has already been initialized
59  */
60 static ossl_inline ossl_unused void
vector_copy(VECTOR * dst,const VECTOR * src)61 vector_copy(VECTOR *dst, const VECTOR *src)
62 {
63     assert(dst->num_poly == src->num_poly);
64     memcpy(dst->poly, src->poly, src->num_poly * sizeof(src->poly[0]));
65 }
66 
67 /* @brief return 1 if 2 vectors are equal, or 0 otherwise */
68 static ossl_inline ossl_unused int
vector_equal(const VECTOR * a,const VECTOR * b)69 vector_equal(const VECTOR *a, const VECTOR *b)
70 {
71     size_t i;
72 
73     if (a->num_poly != b->num_poly)
74         return 0;
75     for (i = 0; i < a->num_poly; ++i) {
76         if (!poly_equal(a->poly + i, b->poly + i))
77             return 0;
78     }
79     return 1;
80 }
81 
82 /* @brief add 2 vectors */
83 static ossl_inline ossl_unused void
vector_add(const VECTOR * lhs,const VECTOR * rhs,VECTOR * out)84 vector_add(const VECTOR *lhs, const VECTOR *rhs, VECTOR *out)
85 {
86     size_t i;
87 
88     for (i = 0; i < lhs->num_poly; i++)
89         poly_add(lhs->poly + i, rhs->poly + i, out->poly + i);
90 }
91 
92 /* @brief subtract 2 vectors */
93 static ossl_inline ossl_unused void
vector_sub(const VECTOR * lhs,const VECTOR * rhs,VECTOR * out)94 vector_sub(const VECTOR *lhs, const VECTOR *rhs, VECTOR *out)
95 {
96     size_t i;
97 
98     for (i = 0; i < lhs->num_poly; i++)
99         poly_sub(lhs->poly + i, rhs->poly + i, out->poly + i);
100 }
101 
102 /* @brief convert a vector in place into NTT form */
103 static ossl_inline ossl_unused void
vector_ntt(VECTOR * va)104 vector_ntt(VECTOR *va)
105 {
106     size_t i;
107 
108     for (i = 0; i < va->num_poly; i++)
109         ossl_ml_dsa_poly_ntt(va->poly + i);
110 }
111 
112 /* @brief convert a vector in place into inverse NTT form */
113 static ossl_inline ossl_unused void
vector_ntt_inverse(VECTOR * va)114 vector_ntt_inverse(VECTOR *va)
115 {
116     size_t i;
117 
118     for (i = 0; i < va->num_poly; i++)
119         ossl_ml_dsa_poly_ntt_inverse(va->poly + i);
120 }
121 
122 /* @brief multiply a vector by a SCALAR polynomial */
123 static ossl_inline ossl_unused void
vector_mult_scalar(const VECTOR * lhs,const POLY * rhs,VECTOR * out)124 vector_mult_scalar(const VECTOR *lhs, const POLY *rhs, VECTOR *out)
125 {
126     size_t i;
127 
128     for (i = 0; i < lhs->num_poly; i++)
129         ossl_ml_dsa_poly_ntt_mult(lhs->poly + i, rhs, out->poly + i);
130 }
131 
132 static ossl_inline ossl_unused int
vector_expand_S(EVP_MD_CTX * h_ctx,const EVP_MD * md,int eta,const uint8_t * seed,VECTOR * s1,VECTOR * s2)133 vector_expand_S(EVP_MD_CTX *h_ctx, const EVP_MD *md, int eta,
134     const uint8_t *seed, VECTOR *s1, VECTOR *s2)
135 {
136     return ossl_ml_dsa_vector_expand_S(h_ctx, md, eta, seed, s1, s2);
137 }
138 
139 static ossl_inline ossl_unused void
vector_expand_mask(VECTOR * out,const uint8_t * rho_prime,size_t rho_prime_len,uint32_t kappa,uint32_t gamma1,EVP_MD_CTX * h_ctx,const EVP_MD * md)140 vector_expand_mask(VECTOR *out, const uint8_t *rho_prime, size_t rho_prime_len,
141     uint32_t kappa, uint32_t gamma1,
142     EVP_MD_CTX *h_ctx, const EVP_MD *md)
143 {
144     size_t i;
145     uint8_t derived_seed[ML_DSA_RHO_PRIME_BYTES + 2];
146 
147     memcpy(derived_seed, rho_prime, ML_DSA_RHO_PRIME_BYTES);
148 
149     for (i = 0; i < out->num_poly; i++) {
150         size_t index = kappa + i;
151 
152         derived_seed[ML_DSA_RHO_PRIME_BYTES] = index & 0xFF;
153         derived_seed[ML_DSA_RHO_PRIME_BYTES + 1] = (index >> 8) & 0xFF;
154         poly_expand_mask(out->poly + i, derived_seed, sizeof(derived_seed),
155             gamma1, h_ctx, md);
156     }
157     OPENSSL_cleanse(derived_seed, sizeof(derived_seed));
158 }
159 
160 /* Scale back previously rounded value */
161 static ossl_inline ossl_unused void
vector_scale_power2_round_ntt(const VECTOR * in,VECTOR * out)162 vector_scale_power2_round_ntt(const VECTOR *in, VECTOR *out)
163 {
164     size_t i;
165 
166     for (i = 0; i < in->num_poly; i++)
167         poly_scale_power2_round(in->poly + i, out->poly + i);
168     vector_ntt(out);
169 }
170 
171 /*
172  * @brief Decompose all polynomial coefficients of a vector into (t1, t0) such
173  * that coeff[i] == t1[i] * 2^13 + t0[i] mod q.
174  * See FIPS 204, Algorithm 35, Power2Round()
175  */
176 static ossl_inline ossl_unused void
vector_power2_round(const VECTOR * t,VECTOR * t1,VECTOR * t0)177 vector_power2_round(const VECTOR *t, VECTOR *t1, VECTOR *t0)
178 {
179     size_t i;
180 
181     for (i = 0; i < t->num_poly; i++)
182         poly_power2_round(t->poly + i, t1->poly + i, t0->poly + i);
183 }
184 
185 static ossl_inline ossl_unused void
vector_high_bits(const VECTOR * in,uint32_t gamma2,VECTOR * out)186 vector_high_bits(const VECTOR *in, uint32_t gamma2, VECTOR *out)
187 {
188     size_t i;
189 
190     for (i = 0; i < out->num_poly; i++)
191         poly_high_bits(in->poly + i, gamma2, out->poly + i);
192 }
193 
194 static ossl_inline ossl_unused void
vector_low_bits(const VECTOR * in,uint32_t gamma2,VECTOR * out)195 vector_low_bits(const VECTOR *in, uint32_t gamma2, VECTOR *out)
196 {
197     size_t i;
198 
199     for (i = 0; i < out->num_poly; i++)
200         poly_low_bits(in->poly + i, gamma2, out->poly + i);
201 }
202 
203 static ossl_inline ossl_unused uint32_t
vector_max(const VECTOR * v)204 vector_max(const VECTOR *v)
205 {
206     size_t i;
207     uint32_t mx = 0;
208 
209     for (i = 0; i < v->num_poly; i++)
210         poly_max(v->poly + i, &mx);
211     return mx;
212 }
213 
214 static ossl_inline ossl_unused uint32_t
vector_max_signed(const VECTOR * v)215 vector_max_signed(const VECTOR *v)
216 {
217     size_t i;
218     uint32_t mx = 0;
219 
220     for (i = 0; i < v->num_poly; i++)
221         poly_max_signed(v->poly + i, &mx);
222     return mx;
223 }
224 
225 static ossl_inline ossl_unused size_t
vector_count_ones(const VECTOR * v)226 vector_count_ones(const VECTOR *v)
227 {
228     int j;
229     size_t i, count = 0;
230 
231     for (i = 0; i < v->num_poly; i++)
232         for (j = 0; j < ML_DSA_NUM_POLY_COEFFICIENTS; j++)
233             count += v->poly[i].coeff[j];
234     return count;
235 }
236 
237 static ossl_inline ossl_unused void
vector_make_hint(const VECTOR * ct0,const VECTOR * cs2,const VECTOR * w,uint32_t gamma2,VECTOR * out)238 vector_make_hint(const VECTOR *ct0, const VECTOR *cs2, const VECTOR *w,
239     uint32_t gamma2, VECTOR *out)
240 {
241     size_t i;
242 
243     for (i = 0; i < out->num_poly; i++)
244         poly_make_hint(ct0->poly + i, cs2->poly + i, w->poly + i, gamma2,
245             out->poly + i);
246 }
247 
248 static ossl_inline ossl_unused void
vector_use_hint(const VECTOR * h,const VECTOR * r,uint32_t gamma2,VECTOR * out)249 vector_use_hint(const VECTOR *h, const VECTOR *r, uint32_t gamma2, VECTOR *out)
250 {
251     size_t i;
252 
253     for (i = 0; i < out->num_poly; i++)
254         poly_use_hint(h->poly + i, r->poly + i, gamma2, out->poly + i);
255 }
256