xref: /linux/lib/crypto/tests/aead-test-template.h (revision fc8c78bce3335860ff4ae9fcfd3b2eb1f674efbb)
1*9396dc94SEric Biggers /* SPDX-License-Identifier: GPL-2.0-or-later */
2*9396dc94SEric Biggers /*
3*9396dc94SEric Biggers  * Shared KUnit test cases for AEAD algorithms, including a benchmark
4*9396dc94SEric Biggers  *
5*9396dc94SEric Biggers  * Copyright 2026 Google LLC
6*9396dc94SEric Biggers  */
7*9396dc94SEric Biggers 
8*9396dc94SEric Biggers /*
9*9396dc94SEric Biggers  * This file implements KUnit test cases shared by the different KUnit test
10*9396dc94SEric Biggers  * suites for Authenticated Encryption with Associated Data (AEAD) algorithms.
11*9396dc94SEric Biggers  *
12*9396dc94SEric Biggers  * Test suites including this file must #define the following:
13*9396dc94SEric Biggers  *
14*9396dc94SEric Biggers  * Data structs:
15*9396dc94SEric Biggers  * - AEAD_KEY: name of key struct
16*9396dc94SEric Biggers  * - AEAD_CTX: name of context for incremental computation
17*9396dc94SEric Biggers  *
18*9396dc94SEric Biggers  * Constants:
19*9396dc94SEric Biggers  * - AEAD_VALID_KEY_LENS: array of all valid key lengths in bytes
20*9396dc94SEric Biggers  * - AEAD_VALID_NONCE_LENS: array of all valid nonce lengths in bytes
21*9396dc94SEric Biggers  * - AEAD_VALID_TAG_LENS: array of all valid authtag lengths in bytes
22*9396dc94SEric Biggers  * - AEAD_MAX_KEY_LEN: max key length in bytes (assumed to fit on stack)
23*9396dc94SEric Biggers  * - AEAD_MAX_NONCE_LEN: max nonce length in bytes (assumed to fit on stack)
24*9396dc94SEric Biggers  * - AEAD_MAX_TAG_LEN: max authtag length in bytes (assumed to fit on stack)
25*9396dc94SEric Biggers  * - AEAD_MONTE_CARLO_CHECKSUM: checksum of a deterministically generated series
26*9396dc94SEric Biggers  *   of (ciphertext, authtag) pairs (see test_aead_monte_carlo())
27*9396dc94SEric Biggers  *
28*9396dc94SEric Biggers  * Functions:
29*9396dc94SEric Biggers  * - AEAD_PREPAREKEY: key preparation
30*9396dc94SEric Biggers  * - AEAD_ENCRYPT and AEAD_DECRYPT: one-shot encryption and decryption
31*9396dc94SEric Biggers  * - AEAD_INIT, AEAD_AUTH_UPDATE, AEAD_ENCRYPT_UPDATE, AEAD_ENCRYPT_FINAL,
32*9396dc94SEric Biggers  *   AEAD_DECRYPT_UPDATE, AEAD_DECRYPT_FINAL: functions for incremental
33*9396dc94SEric Biggers  *   encryption and decryption
34*9396dc94SEric Biggers  *
35*9396dc94SEric Biggers  * Function prototypes and their behavior must match the AES-CCM API.
36*9396dc94SEric Biggers  */
37*9396dc94SEric Biggers 
38*9396dc94SEric Biggers #include <crypto/blake2s.h>
39*9396dc94SEric Biggers #include <kunit/run-in-irq-context.h>
40*9396dc94SEric Biggers #include <kunit/test.h>
41*9396dc94SEric Biggers #include <linux/ktime.h>
42*9396dc94SEric Biggers #include <linux/preempt.h>
43*9396dc94SEric Biggers #include "test-utils.h"
44*9396dc94SEric Biggers 
45*9396dc94SEric Biggers /*
46*9396dc94SEric Biggers  * Allocate a KUnit-managed struct AEAD_KEY and prepare it with a random key,
47*9396dc94SEric Biggers  * using a random key length and random authentication tag length.
48*9396dc94SEric Biggers  */
49*9396dc94SEric Biggers static struct AEAD_KEY *aead_alloc_random_key(struct kunit *test,
50*9396dc94SEric Biggers 					      size_t *tag_len_ret)
51*9396dc94SEric Biggers {
52*9396dc94SEric Biggers 	size_t key_len =
53*9396dc94SEric Biggers 		AEAD_VALID_KEY_LENS[rand32() % ARRAY_SIZE(AEAD_VALID_KEY_LENS)];
54*9396dc94SEric Biggers 	size_t tag_len =
55*9396dc94SEric Biggers 		AEAD_VALID_TAG_LENS[rand32() % ARRAY_SIZE(AEAD_VALID_TAG_LENS)];
56*9396dc94SEric Biggers 	u8 raw_key[AEAD_MAX_KEY_LEN];
57*9396dc94SEric Biggers 	struct AEAD_KEY *key = alloc_buf(test, sizeof(*key));
58*9396dc94SEric Biggers 	int err;
59*9396dc94SEric Biggers 
60*9396dc94SEric Biggers 	rand_bytes(raw_key, key_len);
61*9396dc94SEric Biggers 	err = AEAD_PREPAREKEY(key, raw_key, key_len, tag_len);
62*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ(test, 0, err);
63*9396dc94SEric Biggers 	*tag_len_ret = tag_len;
64*9396dc94SEric Biggers 	return key;
65*9396dc94SEric Biggers }
66*9396dc94SEric Biggers 
67*9396dc94SEric Biggers /*
68*9396dc94SEric Biggers  * Allocate a KUnit-managed slab buffer of length @len bytes and initialize it
69*9396dc94SEric Biggers  * with random data.
70*9396dc94SEric Biggers  */
71*9396dc94SEric Biggers static u8 *aead_alloc_random_data(struct kunit *test, size_t len)
72*9396dc94SEric Biggers {
73*9396dc94SEric Biggers 	u8 *buf = alloc_buf(test, len);
74*9396dc94SEric Biggers 
75*9396dc94SEric Biggers 	rand_bytes(buf, len);
76*9396dc94SEric Biggers 	return buf;
77*9396dc94SEric Biggers }
78*9396dc94SEric Biggers 
79*9396dc94SEric Biggers /*
80*9396dc94SEric Biggers  * Allocate a KUnit-managed guarded buffer of length @len bytes and initialize
81*9396dc94SEric Biggers  * it with random data.
82*9396dc94SEric Biggers  */
83*9396dc94SEric Biggers static u8 *aead_alloc_random_data_guarded(struct kunit *test, size_t len)
84*9396dc94SEric Biggers {
85*9396dc94SEric Biggers 	u8 *buf = alloc_guarded_buf(test, len);
86*9396dc94SEric Biggers 
87*9396dc94SEric Biggers 	rand_bytes(buf, len);
88*9396dc94SEric Biggers 	return buf;
89*9396dc94SEric Biggers }
90*9396dc94SEric Biggers 
91*9396dc94SEric Biggers /* Process the given associated data using a random incremental strategy. */
92*9396dc94SEric Biggers static size_t aead_auth_incrementally(struct AEAD_CTX *ctx, const u8 *ad,
93*9396dc94SEric Biggers 				      size_t ad_len)
94*9396dc94SEric Biggers {
95*9396dc94SEric Biggers 	size_t num_parts = 0;
96*9396dc94SEric Biggers 	size_t pos = 0;
97*9396dc94SEric Biggers 
98*9396dc94SEric Biggers 	while (rand_bool()) {
99*9396dc94SEric Biggers 		size_t part_len = rand_length(ad_len - pos);
100*9396dc94SEric Biggers 
101*9396dc94SEric Biggers 		AEAD_AUTH_UPDATE(ctx, &ad[pos], part_len);
102*9396dc94SEric Biggers 		pos += part_len;
103*9396dc94SEric Biggers 		num_parts++;
104*9396dc94SEric Biggers 	}
105*9396dc94SEric Biggers 	if (pos < ad_len || rand_bool()) {
106*9396dc94SEric Biggers 		AEAD_AUTH_UPDATE(ctx, &ad[pos], ad_len - pos);
107*9396dc94SEric Biggers 		num_parts++;
108*9396dc94SEric Biggers 	}
109*9396dc94SEric Biggers 	return num_parts;
110*9396dc94SEric Biggers }
111*9396dc94SEric Biggers 
112*9396dc94SEric Biggers /* Process the given en/decrypted data using a random incremental strategy. */
113*9396dc94SEric Biggers static size_t aead_crypt_incrementally(struct AEAD_CTX *ctx, u8 *dst,
114*9396dc94SEric Biggers 				       const u8 *src, size_t data_len, bool enc)
115*9396dc94SEric Biggers {
116*9396dc94SEric Biggers 	size_t num_parts = 0;
117*9396dc94SEric Biggers 	size_t pos = 0;
118*9396dc94SEric Biggers 
119*9396dc94SEric Biggers 	while (rand_bool()) {
120*9396dc94SEric Biggers 		size_t part_len = rand_length(data_len - pos);
121*9396dc94SEric Biggers 
122*9396dc94SEric Biggers 		if (enc)
123*9396dc94SEric Biggers 			AEAD_ENCRYPT_UPDATE(ctx, &dst[pos], &src[pos],
124*9396dc94SEric Biggers 					    part_len);
125*9396dc94SEric Biggers 		else
126*9396dc94SEric Biggers 			AEAD_DECRYPT_UPDATE(ctx, &dst[pos], &src[pos],
127*9396dc94SEric Biggers 					    part_len);
128*9396dc94SEric Biggers 		pos += part_len;
129*9396dc94SEric Biggers 		num_parts++;
130*9396dc94SEric Biggers 	}
131*9396dc94SEric Biggers 	if (pos < data_len || rand_bool()) {
132*9396dc94SEric Biggers 		if (enc)
133*9396dc94SEric Biggers 			AEAD_ENCRYPT_UPDATE(ctx, &dst[pos], &src[pos],
134*9396dc94SEric Biggers 					    data_len - pos);
135*9396dc94SEric Biggers 		else
136*9396dc94SEric Biggers 			AEAD_DECRYPT_UPDATE(ctx, &dst[pos], &src[pos],
137*9396dc94SEric Biggers 					    data_len - pos);
138*9396dc94SEric Biggers 		num_parts++;
139*9396dc94SEric Biggers 	}
140*9396dc94SEric Biggers 	return num_parts;
141*9396dc94SEric Biggers }
142*9396dc94SEric Biggers 
143*9396dc94SEric Biggers struct aead_incremental_info {
144*9396dc94SEric Biggers 	size_t num_data_parts;
145*9396dc94SEric Biggers 	size_t num_ad_parts;
146*9396dc94SEric Biggers };
147*9396dc94SEric Biggers 
148*9396dc94SEric Biggers static const char *aead_incr_info_str(struct kunit *test,
149*9396dc94SEric Biggers 				      const struct aead_incremental_info *info)
150*9396dc94SEric Biggers {
151*9396dc94SEric Biggers 	const size_t max_str_len = 64;
152*9396dc94SEric Biggers 	char *str = alloc_buf(test, max_str_len);
153*9396dc94SEric Biggers 
154*9396dc94SEric Biggers 	snprintf(str, max_str_len, "num_data_parts=%zu num_ad_parts=%zu",
155*9396dc94SEric Biggers 		 info->num_data_parts, info->num_ad_parts);
156*9396dc94SEric Biggers 	return str;
157*9396dc94SEric Biggers }
158*9396dc94SEric Biggers 
159*9396dc94SEric Biggers /*
160*9396dc94SEric Biggers  * Encrypt data using a random incremental strategy.
161*9396dc94SEric Biggers  * Return information about the incremental strategy used.
162*9396dc94SEric Biggers  */
163*9396dc94SEric Biggers static struct aead_incremental_info
164*9396dc94SEric Biggers aead_encrypt_incrementally(struct kunit *test, struct AEAD_CTX *ctx, u8 *dst,
165*9396dc94SEric Biggers 			   const u8 *src, size_t data_len, u8 *tag,
166*9396dc94SEric Biggers 			   const u8 *ad, size_t ad_len, const u8 *nonce,
167*9396dc94SEric Biggers 			   size_t nonce_len, const struct AEAD_KEY *key)
168*9396dc94SEric Biggers {
169*9396dc94SEric Biggers 	struct aead_incremental_info info;
170*9396dc94SEric Biggers 	int err;
171*9396dc94SEric Biggers 
172*9396dc94SEric Biggers 	err = AEAD_INIT(ctx, data_len, ad_len, nonce, nonce_len, key);
173*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ(test, 0, err);
174*9396dc94SEric Biggers 	info.num_ad_parts = aead_auth_incrementally(ctx, ad, ad_len);
175*9396dc94SEric Biggers 	info.num_data_parts = aead_crypt_incrementally(ctx, dst, src, data_len,
176*9396dc94SEric Biggers 						       /* enc= */ true);
177*9396dc94SEric Biggers 	AEAD_ENCRYPT_FINAL(ctx, tag);
178*9396dc94SEric Biggers 	KUNIT_ASSERT_TRUE_MSG(test, mem_is_zero(ctx, sizeof(*ctx)),
179*9396dc94SEric Biggers 			      "encrypt_final didn't zeroize context");
180*9396dc94SEric Biggers 	return info;
181*9396dc94SEric Biggers }
182*9396dc94SEric Biggers 
183*9396dc94SEric Biggers /*
184*9396dc94SEric Biggers  * Decrypt authentic data using a random incremental strategy.
185*9396dc94SEric Biggers  * Return information about the incremental strategy used.
186*9396dc94SEric Biggers  */
187*9396dc94SEric Biggers static struct aead_incremental_info
188*9396dc94SEric Biggers aead_decrypt_incrementally(struct kunit *test, struct AEAD_CTX *ctx, u8 *dst,
189*9396dc94SEric Biggers 			   const u8 *src, size_t data_len, const u8 *tag,
190*9396dc94SEric Biggers 			   const u8 *ad, size_t ad_len, const u8 *nonce,
191*9396dc94SEric Biggers 			   size_t nonce_len, const struct AEAD_KEY *key)
192*9396dc94SEric Biggers {
193*9396dc94SEric Biggers 	struct aead_incremental_info info;
194*9396dc94SEric Biggers 	int err;
195*9396dc94SEric Biggers 
196*9396dc94SEric Biggers 	err = AEAD_INIT(ctx, data_len, ad_len, nonce, nonce_len, key);
197*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ(test, 0, err);
198*9396dc94SEric Biggers 	info.num_ad_parts = aead_auth_incrementally(ctx, ad, ad_len);
199*9396dc94SEric Biggers 	info.num_data_parts = aead_crypt_incrementally(ctx, dst, src, data_len,
200*9396dc94SEric Biggers 						       /* enc= */ false);
201*9396dc94SEric Biggers 	err = AEAD_DECRYPT_FINAL(ctx, tag);
202*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ(test, 0, err);
203*9396dc94SEric Biggers 	KUNIT_ASSERT_TRUE_MSG(test, mem_is_zero(ctx, sizeof(*ctx)),
204*9396dc94SEric Biggers 			      "decrypt_final didn't zeroize context");
205*9396dc94SEric Biggers 	return info;
206*9396dc94SEric Biggers }
207*9396dc94SEric Biggers 
208*9396dc94SEric Biggers /* Return true if key_len is declared to be a valid key length. */
209*9396dc94SEric Biggers static bool aead_is_key_len_expected_valid(size_t key_len)
210*9396dc94SEric Biggers {
211*9396dc94SEric Biggers 	for (size_t i = 0; i < ARRAY_SIZE(AEAD_VALID_KEY_LENS); i++) {
212*9396dc94SEric Biggers 		if (AEAD_VALID_KEY_LENS[i] == key_len)
213*9396dc94SEric Biggers 			return true;
214*9396dc94SEric Biggers 	}
215*9396dc94SEric Biggers 	return false;
216*9396dc94SEric Biggers }
217*9396dc94SEric Biggers 
218*9396dc94SEric Biggers /* Return true if nonce_len is declared to be a valid nonce length. */
219*9396dc94SEric Biggers static bool aead_is_nonce_len_expected_valid(size_t nonce_len)
220*9396dc94SEric Biggers {
221*9396dc94SEric Biggers 	for (size_t i = 0; i < ARRAY_SIZE(AEAD_VALID_NONCE_LENS); i++) {
222*9396dc94SEric Biggers 		if (AEAD_VALID_NONCE_LENS[i] == nonce_len)
223*9396dc94SEric Biggers 			return true;
224*9396dc94SEric Biggers 	}
225*9396dc94SEric Biggers 	return false;
226*9396dc94SEric Biggers }
227*9396dc94SEric Biggers 
228*9396dc94SEric Biggers /* Return true if tag_len is declared to be a valid tag length. */
229*9396dc94SEric Biggers static bool aead_is_tag_len_expected_valid(size_t tag_len)
230*9396dc94SEric Biggers {
231*9396dc94SEric Biggers 	for (size_t i = 0; i < ARRAY_SIZE(AEAD_VALID_TAG_LENS); i++) {
232*9396dc94SEric Biggers 		if (AEAD_VALID_TAG_LENS[i] == tag_len)
233*9396dc94SEric Biggers 			return true;
234*9396dc94SEric Biggers 	}
235*9396dc94SEric Biggers 	return false;
236*9396dc94SEric Biggers }
237*9396dc94SEric Biggers 
238*9396dc94SEric Biggers struct aead_basic_validation_test_ctx {
239*9396dc94SEric Biggers 	struct AEAD_KEY key;
240*9396dc94SEric Biggers 	struct AEAD_CTX ctx;
241*9396dc94SEric Biggers 	u8 *raw_key_buf_end;
242*9396dc94SEric Biggers 	u8 *nonce_buf_end;
243*9396dc94SEric Biggers 	u8 *tag_buf_end;
244*9396dc94SEric Biggers 	u8 pt[64]; /* plaintext */
245*9396dc94SEric Biggers 	u8 ct[64]; /* ciphertext */
246*9396dc94SEric Biggers 	u8 decrypted[64];
247*9396dc94SEric Biggers 	u8 ad[16]; /* associated data */
248*9396dc94SEric Biggers 	u8 *unused_buf;
249*9396dc94SEric Biggers 	size_t data_len;
250*9396dc94SEric Biggers 	size_t ad_len;
251*9396dc94SEric Biggers };
252*9396dc94SEric Biggers 
253*9396dc94SEric Biggers static struct aead_basic_validation_test_ctx *
254*9396dc94SEric Biggers aead_alloc_basic_validation_test_ctx(struct kunit *test)
255*9396dc94SEric Biggers {
256*9396dc94SEric Biggers 	struct aead_basic_validation_test_ctx *ctx =
257*9396dc94SEric Biggers 		alloc_buf(test, sizeof(*ctx));
258*9396dc94SEric Biggers 
259*9396dc94SEric Biggers 	memset(ctx, 0, sizeof(*ctx));
260*9396dc94SEric Biggers 	ctx->raw_key_buf_end =
261*9396dc94SEric Biggers 		aead_alloc_random_data_guarded(test, AEAD_MAX_KEY_LEN) +
262*9396dc94SEric Biggers 		AEAD_MAX_KEY_LEN;
263*9396dc94SEric Biggers 	ctx->nonce_buf_end =
264*9396dc94SEric Biggers 		aead_alloc_random_data_guarded(test, AEAD_MAX_NONCE_LEN) +
265*9396dc94SEric Biggers 		AEAD_MAX_NONCE_LEN;
266*9396dc94SEric Biggers 	ctx->tag_buf_end =
267*9396dc94SEric Biggers 		aead_alloc_random_data_guarded(test, AEAD_MAX_TAG_LEN) +
268*9396dc94SEric Biggers 		AEAD_MAX_TAG_LEN;
269*9396dc94SEric Biggers 
270*9396dc94SEric Biggers 	/*
271*9396dc94SEric Biggers 	 * A pointer to this buffer is passed when passing a length that is
272*9396dc94SEric Biggers 	 * expected to be invalid.  It should never actually be accessed.
273*9396dc94SEric Biggers 	 */
274*9396dc94SEric Biggers 	ctx->unused_buf =
275*9396dc94SEric Biggers 		alloc_buf(test, max3(AEAD_MAX_KEY_LEN, AEAD_MAX_NONCE_LEN,
276*9396dc94SEric Biggers 				     AEAD_MAX_TAG_LEN));
277*9396dc94SEric Biggers 
278*9396dc94SEric Biggers 	ctx->data_len = sizeof(ctx->pt);
279*9396dc94SEric Biggers 	ctx->ad_len = sizeof(ctx->ad);
280*9396dc94SEric Biggers 	return ctx;
281*9396dc94SEric Biggers }
282*9396dc94SEric Biggers 
283*9396dc94SEric Biggers /*
284*9396dc94SEric Biggers  * Given an expected-valid key_len, nonce_len, and tag_len, verify round-trip
285*9396dc94SEric Biggers  * encryption and decryption with them.  Use guarded buffers for each of the raw
286*9396dc94SEric Biggers  * key, nonce, and tag to detect any buffer overruns in them.  Also, verify that
287*9396dc94SEric Biggers  * every byte of the tag is actually checked.
288*9396dc94SEric Biggers  */
289*9396dc94SEric Biggers static void aead_do_basic_checks(struct kunit *test,
290*9396dc94SEric Biggers 				 struct aead_basic_validation_test_ctx *ctx,
291*9396dc94SEric Biggers 				 size_t key_len, size_t nonce_len,
292*9396dc94SEric Biggers 				 size_t tag_len)
293*9396dc94SEric Biggers {
294*9396dc94SEric Biggers 	/* Set up exact-size guarded buffers for (raw_key, nonce, tag). */
295*9396dc94SEric Biggers 	const u8 *raw_key = ctx->raw_key_buf_end - key_len;
296*9396dc94SEric Biggers 	const u8 *nonce = ctx->nonce_buf_end - nonce_len;
297*9396dc94SEric Biggers 	u8 *tag = ctx->tag_buf_end - tag_len;
298*9396dc94SEric Biggers 	int err;
299*9396dc94SEric Biggers 
300*9396dc94SEric Biggers 	/* Key preparation should succeed. */
301*9396dc94SEric Biggers 	err = AEAD_PREPAREKEY(&ctx->key, raw_key, key_len, tag_len);
302*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(test, 0, err,
303*9396dc94SEric Biggers 			    "key_len=%zu, tag_len=%zu wasn't accepted", key_len,
304*9396dc94SEric Biggers 			    tag_len);
305*9396dc94SEric Biggers 
306*9396dc94SEric Biggers 	/* Encryption should succeed. */
307*9396dc94SEric Biggers 	err = AEAD_ENCRYPT(ctx->ct, ctx->pt, ctx->data_len, tag, ctx->ad,
308*9396dc94SEric Biggers 			   ctx->ad_len, nonce, nonce_len, &ctx->key);
309*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(
310*9396dc94SEric Biggers 		test, 0, err,
311*9396dc94SEric Biggers 		"Encryption failed with key_len=%zu, nonce_len=%zu, tag_len=%zu",
312*9396dc94SEric Biggers 		key_len, nonce_len, tag_len);
313*9396dc94SEric Biggers 
314*9396dc94SEric Biggers 	/* Decryption should succeed and give the original data. */
315*9396dc94SEric Biggers 	err = AEAD_DECRYPT(ctx->decrypted, ctx->ct, ctx->data_len, tag, ctx->ad,
316*9396dc94SEric Biggers 			   ctx->ad_len, nonce, nonce_len, &ctx->key);
317*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(
318*9396dc94SEric Biggers 		test, 0, err,
319*9396dc94SEric Biggers 		"Decryption failed with key_len=%zu, nonce_len=%zu, tag_len=%zu",
320*9396dc94SEric Biggers 		key_len, nonce_len, tag_len);
321*9396dc94SEric Biggers 	KUNIT_ASSERT_MEMEQ_MSG(
322*9396dc94SEric Biggers 		test, ctx->pt, ctx->decrypted, ctx->data_len,
323*9396dc94SEric Biggers 		"Decryption gave wrong output with key_len=%zu, nonce_len=%zu, tag_len=%zu",
324*9396dc94SEric Biggers 		key_len, nonce_len, tag_len);
325*9396dc94SEric Biggers 
326*9396dc94SEric Biggers 	/*
327*9396dc94SEric Biggers 	 * Every byte of the tag should actually be checked.
328*9396dc94SEric Biggers 	 * And on authentication failure, the dst buffer should be cleared.
329*9396dc94SEric Biggers 	 */
330*9396dc94SEric Biggers 	for (size_t i = 0; i < tag_len; i++) {
331*9396dc94SEric Biggers 		memset(ctx->decrypted, 0xff, ctx->data_len);
332*9396dc94SEric Biggers 		tag[i] ^= 1;
333*9396dc94SEric Biggers 		err = AEAD_DECRYPT(ctx->decrypted, ctx->ct, ctx->data_len, tag,
334*9396dc94SEric Biggers 				   ctx->ad, ctx->ad_len, nonce, nonce_len,
335*9396dc94SEric Biggers 				   &ctx->key);
336*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ_MSG(
337*9396dc94SEric Biggers 			test, -EBADMSG, err,
338*9396dc94SEric Biggers 			"Decryption with bad auth tag with key_len=%zu, nonce_len=%zu, tag_len=%zu didn't fail with -EBADMSG",
339*9396dc94SEric Biggers 			key_len, nonce_len, tag_len);
340*9396dc94SEric Biggers 		KUNIT_ASSERT_TRUE_MSG(
341*9396dc94SEric Biggers 			test, mem_is_zero(ctx->decrypted, ctx->data_len),
342*9396dc94SEric Biggers 			"dst wasn't cleared on authentication failure");
343*9396dc94SEric Biggers 		tag[i] ^= 1;
344*9396dc94SEric Biggers 	}
345*9396dc94SEric Biggers }
346*9396dc94SEric Biggers 
347*9396dc94SEric Biggers /* Verify that the given expected-invalid key_len is actually rejected. */
348*9396dc94SEric Biggers static void
349*9396dc94SEric Biggers aead_verify_invalid_key_len(struct kunit *test,
350*9396dc94SEric Biggers 			    struct aead_basic_validation_test_ctx *ctx,
351*9396dc94SEric Biggers 			    size_t key_len)
352*9396dc94SEric Biggers {
353*9396dc94SEric Biggers 	int err;
354*9396dc94SEric Biggers 
355*9396dc94SEric Biggers 	/*
356*9396dc94SEric Biggers 	 * The preparekey function should reject the key_len.  It should do so
357*9396dc94SEric Biggers 	 * before writing to the key struct.
358*9396dc94SEric Biggers 	 */
359*9396dc94SEric Biggers 	memset(&ctx->key, 0, sizeof(ctx->key));
360*9396dc94SEric Biggers 	err = AEAD_PREPAREKEY(&ctx->key, ctx->unused_buf, key_len,
361*9396dc94SEric Biggers 			      AEAD_MAX_TAG_LEN);
362*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(test, -EINVAL, err,
363*9396dc94SEric Biggers 			    "key_len=%zu wasn't rejected with -EINVAL",
364*9396dc94SEric Biggers 			    key_len);
365*9396dc94SEric Biggers 	KUNIT_ASSERT_TRUE_MSG(
366*9396dc94SEric Biggers 		test, mem_is_zero(&ctx->key, sizeof(ctx->key)),
367*9396dc94SEric Biggers 		"Key struct was written to before length validation");
368*9396dc94SEric Biggers }
369*9396dc94SEric Biggers 
370*9396dc94SEric Biggers /*
371*9396dc94SEric Biggers  * Test that every valid key length is accepted and basic checks pass with it,
372*9396dc94SEric Biggers  * and test that invalid key lengths are rejected.
373*9396dc94SEric Biggers  */
374*9396dc94SEric Biggers static void test_aead_all_key_lens(struct kunit *test)
375*9396dc94SEric Biggers {
376*9396dc94SEric Biggers 	struct aead_basic_validation_test_ctx *ctx =
377*9396dc94SEric Biggers 		aead_alloc_basic_validation_test_ctx(test);
378*9396dc94SEric Biggers 
379*9396dc94SEric Biggers 	for (size_t key_len = 0; key_len <= AEAD_MAX_KEY_LEN; key_len++) {
380*9396dc94SEric Biggers 		if (aead_is_key_len_expected_valid(key_len))
381*9396dc94SEric Biggers 			aead_do_basic_checks(test, ctx, key_len,
382*9396dc94SEric Biggers 					     AEAD_MAX_NONCE_LEN,
383*9396dc94SEric Biggers 					     AEAD_MAX_TAG_LEN);
384*9396dc94SEric Biggers 		else
385*9396dc94SEric Biggers 			aead_verify_invalid_key_len(test, ctx, key_len);
386*9396dc94SEric Biggers 	}
387*9396dc94SEric Biggers 	aead_verify_invalid_key_len(test, ctx, AEAD_MAX_KEY_LEN + 1);
388*9396dc94SEric Biggers 	aead_verify_invalid_key_len(test, ctx, AEAD_MAX_KEY_LEN * 2);
389*9396dc94SEric Biggers 	aead_verify_invalid_key_len(test, ctx, U32_MAX);
390*9396dc94SEric Biggers 	aead_verify_invalid_key_len(test, ctx, SIZE_MAX);
391*9396dc94SEric Biggers }
392*9396dc94SEric Biggers 
393*9396dc94SEric Biggers /* Verify that the given expected-invalid nonce_len is actually rejected. */
394*9396dc94SEric Biggers static void
395*9396dc94SEric Biggers aead_verify_invalid_nonce_len(struct kunit *test,
396*9396dc94SEric Biggers 			      struct aead_basic_validation_test_ctx *ctx,
397*9396dc94SEric Biggers 			      size_t nonce_len)
398*9396dc94SEric Biggers {
399*9396dc94SEric Biggers 	static const u8 raw_key[AEAD_MAX_KEY_LEN];
400*9396dc94SEric Biggers 	int err;
401*9396dc94SEric Biggers 
402*9396dc94SEric Biggers 	/* Key preparation should succeed, as nonce_len isn't given yet. */
403*9396dc94SEric Biggers 	err = AEAD_PREPAREKEY(&ctx->key, raw_key, sizeof(raw_key),
404*9396dc94SEric Biggers 			      AEAD_MAX_TAG_LEN);
405*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ(test, 0, err);
406*9396dc94SEric Biggers 
407*9396dc94SEric Biggers 	/* The init function should reject the nonce_len. */
408*9396dc94SEric Biggers 	memset(&ctx->ctx, 0, sizeof(ctx->ctx));
409*9396dc94SEric Biggers 	err = AEAD_INIT(&ctx->ctx, ctx->data_len, ctx->ad_len, ctx->unused_buf,
410*9396dc94SEric Biggers 			nonce_len, &ctx->key);
411*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(test, -EINVAL, err,
412*9396dc94SEric Biggers 			    "nonce_len=%zu wasn't rejected with -EINVAL (init)",
413*9396dc94SEric Biggers 			    nonce_len);
414*9396dc94SEric Biggers 	KUNIT_ASSERT_TRUE_MSG(
415*9396dc94SEric Biggers 		test, mem_is_zero(&ctx->ctx, sizeof(ctx->ctx)),
416*9396dc94SEric Biggers 		"Context struct was written to before length validation");
417*9396dc94SEric Biggers 
418*9396dc94SEric Biggers 	/* The encrypt function should reject the nonce_len. */
419*9396dc94SEric Biggers 	err = AEAD_ENCRYPT(ctx->ct, ctx->pt, ctx->data_len, ctx->unused_buf,
420*9396dc94SEric Biggers 			   ctx->ad, ctx->ad_len, ctx->unused_buf, nonce_len,
421*9396dc94SEric Biggers 			   &ctx->key);
422*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(
423*9396dc94SEric Biggers 		test, -EINVAL, err,
424*9396dc94SEric Biggers 		"nonce_len=%zu wasn't rejected with -EINVAL (encrypt)",
425*9396dc94SEric Biggers 		nonce_len);
426*9396dc94SEric Biggers 
427*9396dc94SEric Biggers 	/* The decrypt function should reject the nonce_len. */
428*9396dc94SEric Biggers 	err = AEAD_DECRYPT(ctx->pt, ctx->ct, ctx->data_len, ctx->unused_buf,
429*9396dc94SEric Biggers 			   ctx->ad, ctx->ad_len, ctx->unused_buf, nonce_len,
430*9396dc94SEric Biggers 			   &ctx->key);
431*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(
432*9396dc94SEric Biggers 		test, -EINVAL, err,
433*9396dc94SEric Biggers 		"nonce_len=%zu wasn't rejected with -EINVAL (decrypt)",
434*9396dc94SEric Biggers 		nonce_len);
435*9396dc94SEric Biggers }
436*9396dc94SEric Biggers 
437*9396dc94SEric Biggers /*
438*9396dc94SEric Biggers  * Test that every valid nonce length is accepted and basic checks pass with it,
439*9396dc94SEric Biggers  * and test that invalid nonce lengths are rejected.
440*9396dc94SEric Biggers  */
441*9396dc94SEric Biggers static void test_aead_all_nonce_lens(struct kunit *test)
442*9396dc94SEric Biggers {
443*9396dc94SEric Biggers 	struct aead_basic_validation_test_ctx *ctx =
444*9396dc94SEric Biggers 		aead_alloc_basic_validation_test_ctx(test);
445*9396dc94SEric Biggers 
446*9396dc94SEric Biggers 	for (size_t nonce_len = 0; nonce_len <= AEAD_MAX_NONCE_LEN;
447*9396dc94SEric Biggers 	     nonce_len++) {
448*9396dc94SEric Biggers 		if (aead_is_nonce_len_expected_valid(nonce_len))
449*9396dc94SEric Biggers 			aead_do_basic_checks(test, ctx, AEAD_MAX_KEY_LEN,
450*9396dc94SEric Biggers 					     nonce_len, AEAD_MAX_TAG_LEN);
451*9396dc94SEric Biggers 		else
452*9396dc94SEric Biggers 			aead_verify_invalid_nonce_len(test, ctx, nonce_len);
453*9396dc94SEric Biggers 	}
454*9396dc94SEric Biggers 	aead_verify_invalid_nonce_len(test, ctx, AEAD_MAX_NONCE_LEN + 1);
455*9396dc94SEric Biggers 	aead_verify_invalid_nonce_len(test, ctx, AEAD_MAX_NONCE_LEN * 2);
456*9396dc94SEric Biggers 	aead_verify_invalid_nonce_len(test, ctx, U32_MAX);
457*9396dc94SEric Biggers 	aead_verify_invalid_nonce_len(test, ctx, SIZE_MAX);
458*9396dc94SEric Biggers }
459*9396dc94SEric Biggers 
460*9396dc94SEric Biggers /* Verify that the given expected-invalid tag_len is actually rejected. */
461*9396dc94SEric Biggers static void
462*9396dc94SEric Biggers aead_verify_invalid_tag_len(struct kunit *test,
463*9396dc94SEric Biggers 			    struct aead_basic_validation_test_ctx *ctx,
464*9396dc94SEric Biggers 			    size_t tag_len)
465*9396dc94SEric Biggers {
466*9396dc94SEric Biggers 	static const u8 raw_key[AEAD_MAX_KEY_LEN];
467*9396dc94SEric Biggers 	int err;
468*9396dc94SEric Biggers 
469*9396dc94SEric Biggers 	/*
470*9396dc94SEric Biggers 	 * The preparekey function should reject the tag_len.  It should do so
471*9396dc94SEric Biggers 	 * before writing to the key struct.
472*9396dc94SEric Biggers 	 */
473*9396dc94SEric Biggers 	memset(&ctx->key, 0, sizeof(ctx->key));
474*9396dc94SEric Biggers 	err = AEAD_PREPAREKEY(&ctx->key, raw_key, sizeof(raw_key), tag_len);
475*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ_MSG(test, -EINVAL, err,
476*9396dc94SEric Biggers 			    "tag_len=%zu wasn't rejected with -EINVAL",
477*9396dc94SEric Biggers 			    tag_len);
478*9396dc94SEric Biggers 	KUNIT_ASSERT_TRUE_MSG(
479*9396dc94SEric Biggers 		test, mem_is_zero(&ctx->key, sizeof(ctx->key)),
480*9396dc94SEric Biggers 		"Key struct was written to before length validation");
481*9396dc94SEric Biggers }
482*9396dc94SEric Biggers 
483*9396dc94SEric Biggers /*
484*9396dc94SEric Biggers  * Test that every valid authentication tag length is accepted and basic checks
485*9396dc94SEric Biggers  * pass with it, and test that invalid authentication tag lengths are rejected.
486*9396dc94SEric Biggers  */
487*9396dc94SEric Biggers static void test_aead_all_tag_lens(struct kunit *test)
488*9396dc94SEric Biggers {
489*9396dc94SEric Biggers 	struct aead_basic_validation_test_ctx *ctx =
490*9396dc94SEric Biggers 		aead_alloc_basic_validation_test_ctx(test);
491*9396dc94SEric Biggers 
492*9396dc94SEric Biggers 	for (size_t tag_len = 0; tag_len <= AEAD_MAX_TAG_LEN; tag_len++) {
493*9396dc94SEric Biggers 		if (aead_is_tag_len_expected_valid(tag_len))
494*9396dc94SEric Biggers 			aead_do_basic_checks(test, ctx, AEAD_MAX_KEY_LEN,
495*9396dc94SEric Biggers 					     AEAD_MAX_NONCE_LEN, tag_len);
496*9396dc94SEric Biggers 		else
497*9396dc94SEric Biggers 			aead_verify_invalid_tag_len(test, ctx, tag_len);
498*9396dc94SEric Biggers 	}
499*9396dc94SEric Biggers 	aead_verify_invalid_tag_len(test, ctx, AEAD_MAX_TAG_LEN + 1);
500*9396dc94SEric Biggers 	aead_verify_invalid_tag_len(test, ctx, AEAD_MAX_TAG_LEN * 2);
501*9396dc94SEric Biggers 	aead_verify_invalid_tag_len(test, ctx, U32_MAX);
502*9396dc94SEric Biggers 	aead_verify_invalid_tag_len(test, ctx, SIZE_MAX);
503*9396dc94SEric Biggers }
504*9396dc94SEric Biggers 
505*9396dc94SEric Biggers /*
506*9396dc94SEric Biggers  * Test that one-shot encryption and decryption are consistent with each other
507*9396dc94SEric Biggers  * and with incremental encryption and decryption.
508*9396dc94SEric Biggers  */
509*9396dc94SEric Biggers static void test_aead_incremental_updates(struct kunit *test)
510*9396dc94SEric Biggers {
511*9396dc94SEric Biggers 	const size_t max_data_len = 1024;
512*9396dc94SEric Biggers 	const size_t max_ad_len = 512;
513*9396dc94SEric Biggers 	const size_t nonce_len = AEAD_MAX_NONCE_LEN;
514*9396dc94SEric Biggers 	size_t tag_len;
515*9396dc94SEric Biggers 	struct AEAD_KEY *key = aead_alloc_random_key(test, &tag_len);
516*9396dc94SEric Biggers 	struct AEAD_CTX *ctx = alloc_buf(test, sizeof(*ctx));
517*9396dc94SEric Biggers 	u8 *pt = aead_alloc_random_data(test, max_data_len);
518*9396dc94SEric Biggers 	u8 *ad = aead_alloc_random_data(test, max_ad_len);
519*9396dc94SEric Biggers 	u8 *nonce = aead_alloc_random_data(test, nonce_len);
520*9396dc94SEric Biggers 	u8 *ct = alloc_buf(test, max_data_len);
521*9396dc94SEric Biggers 	u8 *ct2 = alloc_buf(test, max_data_len);
522*9396dc94SEric Biggers 	u8 *decrypted = alloc_buf(test, max_data_len);
523*9396dc94SEric Biggers 	u8 *tag = alloc_buf(test, tag_len);
524*9396dc94SEric Biggers 	u8 *tag2 = alloc_buf(test, tag_len);
525*9396dc94SEric Biggers 	int err;
526*9396dc94SEric Biggers 
527*9396dc94SEric Biggers 	for (int i = 0; i < 500; i++) {
528*9396dc94SEric Biggers 		/* Select the lengths to test. */
529*9396dc94SEric Biggers 		const size_t data_len = rand_length(max_data_len);
530*9396dc94SEric Biggers 		const size_t ad_len = rand_length(max_ad_len);
531*9396dc94SEric Biggers 		struct aead_incremental_info incr_info;
532*9396dc94SEric Biggers 
533*9396dc94SEric Biggers 		/* Try one-shot encryption and decryption. */
534*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(ct, pt, data_len, tag, ad, ad_len, nonce,
535*9396dc94SEric Biggers 				   nonce_len, key);
536*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
537*9396dc94SEric Biggers 		err = AEAD_DECRYPT(decrypted, ct, data_len, tag, ad, ad_len,
538*9396dc94SEric Biggers 				   nonce, nonce_len, key);
539*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
540*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ_MSG(
541*9396dc94SEric Biggers 			test, pt, decrypted, data_len,
542*9396dc94SEric Biggers 			"Decryption didn't invert encryption; data_len=%zu, ad_len=%zu",
543*9396dc94SEric Biggers 			data_len, ad_len);
544*9396dc94SEric Biggers 
545*9396dc94SEric Biggers 		/* Try incremental encryption and decryption. */
546*9396dc94SEric Biggers 		incr_info = aead_encrypt_incrementally(test, ctx, ct2, pt,
547*9396dc94SEric Biggers 						       data_len, tag2, ad,
548*9396dc94SEric Biggers 						       ad_len, nonce, nonce_len,
549*9396dc94SEric Biggers 						       key);
550*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ_MSG(
551*9396dc94SEric Biggers 			test, ct, ct2, data_len,
552*9396dc94SEric Biggers 			"One-shot and incremental encryption gave different ciphertexts; data_len=%zu ad_len=%zu %s",
553*9396dc94SEric Biggers 			data_len, ad_len, aead_incr_info_str(test, &incr_info));
554*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ_MSG(
555*9396dc94SEric Biggers 			test, tag, tag2, tag_len,
556*9396dc94SEric Biggers 			"One-shot and incremental encryption gave different auth tags; data_len=%zu ad_len=%zu %s",
557*9396dc94SEric Biggers 			data_len, ad_len, aead_incr_info_str(test, &incr_info));
558*9396dc94SEric Biggers 		incr_info = aead_decrypt_incrementally(test, ctx, decrypted,
559*9396dc94SEric Biggers 						       ct2, data_len, tag2, ad,
560*9396dc94SEric Biggers 						       ad_len, nonce, nonce_len,
561*9396dc94SEric Biggers 						       key);
562*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ_MSG(
563*9396dc94SEric Biggers 			test, pt, decrypted, data_len,
564*9396dc94SEric Biggers 			"One-shot and incremental decryption gave different plaintexts; data_len=%zu ad_len=%zu %s",
565*9396dc94SEric Biggers 			data_len, ad_len, aead_incr_info_str(test, &incr_info));
566*9396dc94SEric Biggers 	}
567*9396dc94SEric Biggers }
568*9396dc94SEric Biggers 
569*9396dc94SEric Biggers /*
570*9396dc94SEric Biggers  * Test using guarded buffers for the plaintext, ciphertext, and associated
571*9396dc94SEric Biggers  * data.  This detects out-of-bounds accesses, even in assembly code.
572*9396dc94SEric Biggers  *
573*9396dc94SEric Biggers  * Note: other test cases cover overrun of raw_key, nonce, and tag.
574*9396dc94SEric Biggers  */
575*9396dc94SEric Biggers static void test_aead_data_buffer_overruns(struct kunit *test)
576*9396dc94SEric Biggers {
577*9396dc94SEric Biggers 	const size_t max_data_len = 1024;
578*9396dc94SEric Biggers 	const size_t max_ad_len = 512;
579*9396dc94SEric Biggers 	const size_t nonce_len = AEAD_MAX_NONCE_LEN;
580*9396dc94SEric Biggers 	size_t tag_len;
581*9396dc94SEric Biggers 	struct AEAD_KEY *key = aead_alloc_random_key(test, &tag_len);
582*9396dc94SEric Biggers 	const u8 *nonce = aead_alloc_random_data(test, nonce_len);
583*9396dc94SEric Biggers 	const u8 *pt_end = aead_alloc_random_data_guarded(test, max_data_len) +
584*9396dc94SEric Biggers 			   max_data_len;
585*9396dc94SEric Biggers 	const u8 *ad_end =
586*9396dc94SEric Biggers 		aead_alloc_random_data_guarded(test, max_ad_len) + max_ad_len;
587*9396dc94SEric Biggers 	u8 *ct_end = alloc_guarded_buf(test, max_data_len) + max_data_len;
588*9396dc94SEric Biggers 	u8 *decrypted_end =
589*9396dc94SEric Biggers 		alloc_guarded_buf(test, max_data_len) + max_data_len;
590*9396dc94SEric Biggers 	u8 *tag = alloc_buf(test, tag_len);
591*9396dc94SEric Biggers 
592*9396dc94SEric Biggers 	for (int i = 0; i < 200; i++) {
593*9396dc94SEric Biggers 		/* Select the lengths to test. */
594*9396dc94SEric Biggers 		const size_t data_len = rand_length(max_data_len);
595*9396dc94SEric Biggers 		const size_t ad_len = rand_length(max_ad_len);
596*9396dc94SEric Biggers 		/* Set up exact-size guarded buffers. */
597*9396dc94SEric Biggers 		const u8 *pt = pt_end - data_len;
598*9396dc94SEric Biggers 		const u8 *ad = ad_end - ad_len;
599*9396dc94SEric Biggers 		u8 *ct = ct_end - data_len;
600*9396dc94SEric Biggers 		u8 *decrypted = decrypted_end - data_len;
601*9396dc94SEric Biggers 		int err;
602*9396dc94SEric Biggers 
603*9396dc94SEric Biggers 		/* Encrypt and decrypt. */
604*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(ct, pt, data_len, tag, ad, ad_len, nonce,
605*9396dc94SEric Biggers 				   nonce_len, key);
606*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
607*9396dc94SEric Biggers 		err = AEAD_DECRYPT(decrypted, ct, data_len, tag, ad, ad_len,
608*9396dc94SEric Biggers 				   nonce, nonce_len, key);
609*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
610*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ_MSG(
611*9396dc94SEric Biggers 			test, pt, decrypted, data_len,
612*9396dc94SEric Biggers 			"Decryption didn't invert encryption; data_len=%zu, ad_len=%zu",
613*9396dc94SEric Biggers 			data_len, ad_len);
614*9396dc94SEric Biggers 	}
615*9396dc94SEric Biggers }
616*9396dc94SEric Biggers 
617*9396dc94SEric Biggers /*
618*9396dc94SEric Biggers  * Test that encryption and decryption produce the same results regardless of
619*9396dc94SEric Biggers  * how the buffers are aligned in memory.
620*9396dc94SEric Biggers  */
621*9396dc94SEric Biggers static void test_aead_alignment_consistency(struct kunit *test)
622*9396dc94SEric Biggers {
623*9396dc94SEric Biggers 	const size_t max_data_len = 4096;
624*9396dc94SEric Biggers 	const size_t max_ad_len = 4096;
625*9396dc94SEric Biggers 	const size_t max_offset = 128;
626*9396dc94SEric Biggers 	const size_t nonce_len = AEAD_MAX_NONCE_LEN;
627*9396dc94SEric Biggers 	const size_t key_len = AEAD_MAX_KEY_LEN;
628*9396dc94SEric Biggers 	const size_t tag_len = AEAD_MAX_TAG_LEN;
629*9396dc94SEric Biggers 	u8 *raw_key1_buf = alloc_buf(test, key_len + max_offset);
630*9396dc94SEric Biggers 	u8 *raw_key2_buf = alloc_buf(test, key_len + max_offset);
631*9396dc94SEric Biggers 	u8 *nonce1_buf = alloc_buf(test, nonce_len + max_offset);
632*9396dc94SEric Biggers 	u8 *nonce2_buf = alloc_buf(test, nonce_len + max_offset);
633*9396dc94SEric Biggers 	u8 *pt1_buf = alloc_buf(test, max_data_len);
634*9396dc94SEric Biggers 	u8 *pt2_buf = alloc_buf(test, max_data_len);
635*9396dc94SEric Biggers 	u8 *ct1_buf = alloc_buf(test, max_data_len);
636*9396dc94SEric Biggers 	u8 *ct2_buf = alloc_buf(test, max_data_len);
637*9396dc94SEric Biggers 	u8 *ad1_buf = alloc_buf(test, max_ad_len);
638*9396dc94SEric Biggers 	u8 *ad2_buf = alloc_buf(test, max_ad_len);
639*9396dc94SEric Biggers 	u8 *tag1_buf = alloc_buf(test, tag_len + max_offset);
640*9396dc94SEric Biggers 	u8 *tag2_buf = alloc_buf(test, tag_len + max_offset);
641*9396dc94SEric Biggers 	struct AEAD_KEY *key = alloc_buf(test, sizeof(*key));
642*9396dc94SEric Biggers 	int err;
643*9396dc94SEric Biggers 
644*9396dc94SEric Biggers 	for (int i = 0; i < 100; i++) {
645*9396dc94SEric Biggers 		/* Generate lengths. */
646*9396dc94SEric Biggers 		size_t data_len = rand_length(max_data_len);
647*9396dc94SEric Biggers 		size_t ad_len = rand_length(max_ad_len);
648*9396dc94SEric Biggers 
649*9396dc94SEric Biggers 		/* Generate two sets of alignments. */
650*9396dc94SEric Biggers 		u8 *raw_key1 = raw_key1_buf + rand_offset(max_offset);
651*9396dc94SEric Biggers 		u8 *raw_key2 = raw_key2_buf + rand_offset(max_offset);
652*9396dc94SEric Biggers 		u8 *nonce1 = nonce1_buf + rand_offset(max_offset);
653*9396dc94SEric Biggers 		u8 *nonce2 = nonce2_buf + rand_offset(max_offset);
654*9396dc94SEric Biggers 		u8 *pt1 = pt1_buf + rand_offset(max_data_len - data_len);
655*9396dc94SEric Biggers 		u8 *pt2 = pt2_buf + rand_offset(max_data_len - data_len);
656*9396dc94SEric Biggers 		u8 *ct1 = ct1_buf + rand_offset(max_data_len - data_len);
657*9396dc94SEric Biggers 		u8 *ct2 = ct2_buf + rand_offset(max_data_len - data_len);
658*9396dc94SEric Biggers 		u8 *ad1 = ad1_buf + rand_offset(max_ad_len - ad_len);
659*9396dc94SEric Biggers 		u8 *ad2 = ad2_buf + rand_offset(max_ad_len - ad_len);
660*9396dc94SEric Biggers 		u8 *tag1 = tag1_buf + rand_offset(max_offset);
661*9396dc94SEric Biggers 		u8 *tag2 = tag2_buf + rand_offset(max_offset);
662*9396dc94SEric Biggers 
663*9396dc94SEric Biggers 		/*
664*9396dc94SEric Biggers 		 * Generate inputs in the first set of buffers using the first
665*9396dc94SEric Biggers 		 * set of alignments.
666*9396dc94SEric Biggers 		 */
667*9396dc94SEric Biggers 		rand_bytes(raw_key1, key_len);
668*9396dc94SEric Biggers 		rand_bytes(nonce1, nonce_len);
669*9396dc94SEric Biggers 		rand_bytes(pt1, data_len);
670*9396dc94SEric Biggers 		rand_bytes(ad1, ad_len);
671*9396dc94SEric Biggers 
672*9396dc94SEric Biggers 		/*
673*9396dc94SEric Biggers 		 * Copy the inputs to the second set of buffers using the second
674*9396dc94SEric Biggers 		 * set of alignments.
675*9396dc94SEric Biggers 		 */
676*9396dc94SEric Biggers 		memcpy(raw_key2, raw_key1, key_len);
677*9396dc94SEric Biggers 		memcpy(nonce2, nonce1, nonce_len);
678*9396dc94SEric Biggers 		memcpy(pt2, pt1, data_len);
679*9396dc94SEric Biggers 		memcpy(ad2, ad1, ad_len);
680*9396dc94SEric Biggers 
681*9396dc94SEric Biggers 		/* Verify encryption consistency. */
682*9396dc94SEric Biggers 
683*9396dc94SEric Biggers 		err = AEAD_PREPAREKEY(key, raw_key1, key_len, tag_len);
684*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
685*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(ct1, pt1, data_len, tag1, ad1, ad_len,
686*9396dc94SEric Biggers 				   nonce1, nonce_len, key);
687*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
688*9396dc94SEric Biggers 
689*9396dc94SEric Biggers 		err = AEAD_PREPAREKEY(key, raw_key2, key_len, tag_len);
690*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
691*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(ct2, pt2, data_len, tag2, ad2, ad_len,
692*9396dc94SEric Biggers 				   nonce2, nonce_len, key);
693*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
694*9396dc94SEric Biggers 
695*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ(test, ct1, ct2, data_len);
696*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ(test, tag1, tag2, tag_len);
697*9396dc94SEric Biggers 
698*9396dc94SEric Biggers 		/* Verify decryption consistency. */
699*9396dc94SEric Biggers 
700*9396dc94SEric Biggers 		err = AEAD_PREPAREKEY(key, raw_key1, key_len, tag_len);
701*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
702*9396dc94SEric Biggers 		err = AEAD_DECRYPT(pt1, ct1, data_len, tag1, ad1, ad_len,
703*9396dc94SEric Biggers 				   nonce1, nonce_len, key);
704*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
705*9396dc94SEric Biggers 
706*9396dc94SEric Biggers 		err = AEAD_PREPAREKEY(key, raw_key2, key_len, tag_len);
707*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
708*9396dc94SEric Biggers 		err = AEAD_DECRYPT(pt2, ct2, data_len, tag2, ad2, ad_len,
709*9396dc94SEric Biggers 				   nonce2, nonce_len, key);
710*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
711*9396dc94SEric Biggers 
712*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ(test, pt1, pt2, data_len);
713*9396dc94SEric Biggers 	}
714*9396dc94SEric Biggers }
715*9396dc94SEric Biggers 
716*9396dc94SEric Biggers static void test_aead_inplace(struct kunit *test)
717*9396dc94SEric Biggers {
718*9396dc94SEric Biggers 	const size_t max_data_len = 1024;
719*9396dc94SEric Biggers 	const size_t max_ad_len = 512;
720*9396dc94SEric Biggers 	const size_t nonce_len = AEAD_MAX_NONCE_LEN;
721*9396dc94SEric Biggers 	size_t tag_len;
722*9396dc94SEric Biggers 	struct AEAD_KEY *key = aead_alloc_random_key(test, &tag_len);
723*9396dc94SEric Biggers 	u8 *data = aead_alloc_random_data(test, max_data_len + tag_len);
724*9396dc94SEric Biggers 	u8 *data2 = alloc_buf(test, max_data_len + tag_len);
725*9396dc94SEric Biggers 	u8 *ad = aead_alloc_random_data(test, max_ad_len);
726*9396dc94SEric Biggers 	const u8 *nonce = aead_alloc_random_data(test, nonce_len);
727*9396dc94SEric Biggers 
728*9396dc94SEric Biggers 	for (int i = 0; i < 100; i++) {
729*9396dc94SEric Biggers 		size_t data_len = rand_length(max_data_len);
730*9396dc94SEric Biggers 		size_t ad_len = rand_length(max_ad_len);
731*9396dc94SEric Biggers 		int err;
732*9396dc94SEric Biggers 
733*9396dc94SEric Biggers 		/* Encrypt out-of-place. */
734*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(data2, data, data_len, data2 + data_len, ad,
735*9396dc94SEric Biggers 				   ad_len, nonce, nonce_len, key);
736*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
737*9396dc94SEric Biggers 
738*9396dc94SEric Biggers 		/* Encrypt in-place. */
739*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(data, data, data_len, data + data_len, ad,
740*9396dc94SEric Biggers 				   ad_len, nonce, nonce_len, key);
741*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
742*9396dc94SEric Biggers 
743*9396dc94SEric Biggers 		/* Compare the results. */
744*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ(test, data2, data, data_len + tag_len);
745*9396dc94SEric Biggers 
746*9396dc94SEric Biggers 		/* Decrypt out-of-place. */
747*9396dc94SEric Biggers 		err = AEAD_DECRYPT(data2, data, data_len, data + data_len, ad,
748*9396dc94SEric Biggers 				   ad_len, nonce, nonce_len, key);
749*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
750*9396dc94SEric Biggers 
751*9396dc94SEric Biggers 		/* Decrypt in-place. */
752*9396dc94SEric Biggers 		err = AEAD_DECRYPT(data, data, data_len, data + data_len, ad,
753*9396dc94SEric Biggers 				   ad_len, nonce, nonce_len, key);
754*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
755*9396dc94SEric Biggers 
756*9396dc94SEric Biggers 		/* Compare the results. */
757*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ(test, data2, data, data_len);
758*9396dc94SEric Biggers 	}
759*9396dc94SEric Biggers }
760*9396dc94SEric Biggers 
761*9396dc94SEric Biggers /*
762*9396dc94SEric Biggers  * Monte-Carlo test for AEAD algorithms.  This deterministically generates
763*9396dc94SEric Biggers  * random AEAD inputs, encrypts them, verifies decryption, and computes and
764*9396dc94SEric Biggers  * verifies the checksum of all computed (ciphertext, tag) pairs.
765*9396dc94SEric Biggers  */
766*9396dc94SEric Biggers static void test_aead_monte_carlo(struct kunit *test)
767*9396dc94SEric Biggers {
768*9396dc94SEric Biggers 	const size_t max_data_len = 1024;
769*9396dc94SEric Biggers 	const size_t max_ad_len = 293;
770*9396dc94SEric Biggers 	u8 raw_key[AEAD_MAX_KEY_LEN];
771*9396dc94SEric Biggers 	u8 nonce[AEAD_MAX_NONCE_LEN];
772*9396dc94SEric Biggers 	u8 tag[AEAD_MAX_TAG_LEN];
773*9396dc94SEric Biggers 	u8 *pt = alloc_buf(test, max_data_len);
774*9396dc94SEric Biggers 	u8 *ct = alloc_buf(test, max_data_len);
775*9396dc94SEric Biggers 	u8 *decrypted = alloc_buf(test, max_data_len);
776*9396dc94SEric Biggers 	u8 *ad = alloc_buf(test, max_ad_len);
777*9396dc94SEric Biggers 	struct AEAD_KEY *key = alloc_buf(test, sizeof(*key));
778*9396dc94SEric Biggers 	struct blake2s_ctx checksum_ctx;
779*9396dc94SEric Biggers 	u8 actual_checksum[BLAKE2S_HASH_SIZE];
780*9396dc94SEric Biggers 	int err;
781*9396dc94SEric Biggers 
782*9396dc94SEric Biggers 	blake2s_init(&checksum_ctx, BLAKE2S_HASH_SIZE);
783*9396dc94SEric Biggers 
784*9396dc94SEric Biggers 	for (size_t data_len = 0; data_len <= max_data_len; data_len++) {
785*9396dc94SEric Biggers 		size_t ad_len = data_len % max_ad_len;
786*9396dc94SEric Biggers 		size_t key_len =
787*9396dc94SEric Biggers 			AEAD_VALID_KEY_LENS[data_len %
788*9396dc94SEric Biggers 					    ARRAY_SIZE(AEAD_VALID_KEY_LENS)];
789*9396dc94SEric Biggers 		size_t nonce_len =
790*9396dc94SEric Biggers 			AEAD_VALID_NONCE_LENS[data_len %
791*9396dc94SEric Biggers 					      ARRAY_SIZE(AEAD_VALID_NONCE_LENS)];
792*9396dc94SEric Biggers 		size_t tag_len =
793*9396dc94SEric Biggers 			AEAD_VALID_TAG_LENS[data_len %
794*9396dc94SEric Biggers 					    ARRAY_SIZE(AEAD_VALID_TAG_LENS)];
795*9396dc94SEric Biggers 
796*9396dc94SEric Biggers 		rand_bytes_seeded_from_len(pt, data_len);
797*9396dc94SEric Biggers 		rand_bytes_seeded_from_len(ad, ad_len);
798*9396dc94SEric Biggers 		rand_bytes_seeded_from_len(raw_key, key_len);
799*9396dc94SEric Biggers 		rand_bytes_seeded_from_len(nonce, nonce_len);
800*9396dc94SEric Biggers 
801*9396dc94SEric Biggers 		err = AEAD_PREPAREKEY(key, raw_key, key_len, tag_len);
802*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
803*9396dc94SEric Biggers 
804*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(ct, pt, data_len, tag, ad, ad_len, nonce,
805*9396dc94SEric Biggers 				   nonce_len, key);
806*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ_MSG(
807*9396dc94SEric Biggers 			test, 0, err,
808*9396dc94SEric Biggers 			"Encryption failed with data_len=%zu, ad_len=%zu",
809*9396dc94SEric Biggers 			data_len, ad_len);
810*9396dc94SEric Biggers 		err = AEAD_DECRYPT(decrypted, ct, data_len, tag, ad, ad_len,
811*9396dc94SEric Biggers 				   nonce, nonce_len, key);
812*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ_MSG(
813*9396dc94SEric Biggers 			test, 0, err,
814*9396dc94SEric Biggers 			"Decryption failed with data_len=%zu, ad_len=%zu",
815*9396dc94SEric Biggers 			data_len, ad_len);
816*9396dc94SEric Biggers 		KUNIT_ASSERT_MEMEQ_MSG(
817*9396dc94SEric Biggers 			test, pt, decrypted, data_len,
818*9396dc94SEric Biggers 			"Decryption didn't invert encryption; data_len=%zu, ad_len=%zu",
819*9396dc94SEric Biggers 			data_len, ad_len);
820*9396dc94SEric Biggers 
821*9396dc94SEric Biggers 		blake2s_update(&checksum_ctx, ct, data_len);
822*9396dc94SEric Biggers 		blake2s_update(&checksum_ctx, tag, tag_len);
823*9396dc94SEric Biggers 	}
824*9396dc94SEric Biggers 
825*9396dc94SEric Biggers 	blake2s_final(&checksum_ctx, actual_checksum);
826*9396dc94SEric Biggers 	KUNIT_EXPECT_MEMEQ_MSG(test, actual_checksum, AEAD_MONTE_CARLO_CHECKSUM,
827*9396dc94SEric Biggers 			       BLAKE2S_HASH_SIZE,
828*9396dc94SEric Biggers 			       "Monte-Carlo checksum mismatch");
829*9396dc94SEric Biggers }
830*9396dc94SEric Biggers 
831*9396dc94SEric Biggers #define IRQ_TEST_DATA_LEN 256
832*9396dc94SEric Biggers #define IRQ_TEST_NUM_BUFFERS 3 /* matches max concurrency level */
833*9396dc94SEric Biggers 
834*9396dc94SEric Biggers struct aead_irq_test_slot {
835*9396dc94SEric Biggers 	/* Fields written only at test case initialization time */
836*9396dc94SEric Biggers 	u8 raw_key[AEAD_MAX_KEY_LEN];
837*9396dc94SEric Biggers 	u8 nonce[AEAD_MAX_NONCE_LEN];
838*9396dc94SEric Biggers 	u8 pt[IRQ_TEST_DATA_LEN];
839*9396dc94SEric Biggers 	u8 ct[IRQ_TEST_DATA_LEN + AEAD_MAX_TAG_LEN];
840*9396dc94SEric Biggers 	u8 ad[IRQ_TEST_DATA_LEN];
841*9396dc94SEric Biggers 
842*9396dc94SEric Biggers 	/* Fields written throughout the test case */
843*9396dc94SEric Biggers 	struct AEAD_KEY key;
844*9396dc94SEric Biggers 	u8 scratch_buf[IRQ_TEST_DATA_LEN + AEAD_MAX_TAG_LEN];
845*9396dc94SEric Biggers 	int phase;
846*9396dc94SEric Biggers 	atomic_t in_use;
847*9396dc94SEric Biggers };
848*9396dc94SEric Biggers 
849*9396dc94SEric Biggers struct aead_irq_test_state {
850*9396dc94SEric Biggers 	struct aead_irq_test_slot slots[IRQ_TEST_NUM_BUFFERS];
851*9396dc94SEric Biggers };
852*9396dc94SEric Biggers 
853*9396dc94SEric Biggers static bool aead_irq_test_func(void *state_)
854*9396dc94SEric Biggers {
855*9396dc94SEric Biggers 	struct aead_irq_test_state *state = state_;
856*9396dc94SEric Biggers 	struct aead_irq_test_slot *slot;
857*9396dc94SEric Biggers 	size_t data_len;
858*9396dc94SEric Biggers 	bool ok = true;
859*9396dc94SEric Biggers 
860*9396dc94SEric Biggers 	/*
861*9396dc94SEric Biggers 	 * Find a free slot.  This should always succeed, since the number of
862*9396dc94SEric Biggers 	 * slots is equal to the max concurrency level of kunit_run_irq_test().
863*9396dc94SEric Biggers 	 */
864*9396dc94SEric Biggers 	for (slot = &state->slots[0];
865*9396dc94SEric Biggers 	     slot < &state->slots[ARRAY_SIZE(state->slots)]; slot++) {
866*9396dc94SEric Biggers 		if (atomic_cmpxchg(&slot->in_use, 0, 1) == 0)
867*9396dc94SEric Biggers 			break;
868*9396dc94SEric Biggers 	}
869*9396dc94SEric Biggers 	if (WARN_ON_ONCE(slot == &state->slots[ARRAY_SIZE(state->slots)]))
870*9396dc94SEric Biggers 		return false;
871*9396dc94SEric Biggers 	/*
872*9396dc94SEric Biggers 	 * This execution context now has exclusive access to 'slot'.
873*9396dc94SEric Biggers 	 * Next, execute the next operation that the slot is set to perform.
874*9396dc94SEric Biggers 	 */
875*9396dc94SEric Biggers 
876*9396dc94SEric Biggers 	data_len = sizeof(slot->pt);
877*9396dc94SEric Biggers 	if (slot->phase == 0) {
878*9396dc94SEric Biggers 		/* Phase 0: Prepare slot's key in current context. */
879*9396dc94SEric Biggers 		ok = ok && AEAD_PREPAREKEY(&slot->key, slot->raw_key,
880*9396dc94SEric Biggers 					   sizeof(slot->raw_key),
881*9396dc94SEric Biggers 					   AEAD_MAX_TAG_LEN) == 0;
882*9396dc94SEric Biggers 	} else if (slot->phase == 1) {
883*9396dc94SEric Biggers 		/*
884*9396dc94SEric Biggers 		 * Phase 1: Encrypt plaintext using key that may have been
885*9396dc94SEric Biggers 		 * prepared in a different context.
886*9396dc94SEric Biggers 		 */
887*9396dc94SEric Biggers 		ok = ok && AEAD_ENCRYPT(slot->scratch_buf, slot->pt, data_len,
888*9396dc94SEric Biggers 					&slot->scratch_buf[data_len], slot->ad,
889*9396dc94SEric Biggers 					sizeof(slot->ad), slot->nonce,
890*9396dc94SEric Biggers 					sizeof(slot->nonce), &slot->key) == 0;
891*9396dc94SEric Biggers 		/* Verify the ciphertext (with concatenated auth tag) matches */
892*9396dc94SEric Biggers 		ok = ok &&
893*9396dc94SEric Biggers 		     memcmp(slot->scratch_buf, slot->ct, sizeof(slot->ct)) == 0;
894*9396dc94SEric Biggers 	} else {
895*9396dc94SEric Biggers 		/*
896*9396dc94SEric Biggers 		 * Phase 2: Decrypt ciphertext using key that may have been
897*9396dc94SEric Biggers 		 * prepared in a different context.
898*9396dc94SEric Biggers 		 */
899*9396dc94SEric Biggers 		ok = ok && AEAD_DECRYPT(slot->scratch_buf, slot->ct, data_len,
900*9396dc94SEric Biggers 					&slot->ct[data_len], slot->ad,
901*9396dc94SEric Biggers 					sizeof(slot->ad), slot->nonce,
902*9396dc94SEric Biggers 					sizeof(slot->nonce), &slot->key) == 0;
903*9396dc94SEric Biggers 		/* Verify the plaintext matches. */
904*9396dc94SEric Biggers 		ok = ok && memcmp(slot->scratch_buf, slot->pt, data_len) == 0;
905*9396dc94SEric Biggers 	}
906*9396dc94SEric Biggers 	slot->phase = (slot->phase + 1) % 3;
907*9396dc94SEric Biggers 	atomic_set_release(&slot->in_use, 0);
908*9396dc94SEric Biggers 	return ok;
909*9396dc94SEric Biggers }
910*9396dc94SEric Biggers 
911*9396dc94SEric Biggers /*
912*9396dc94SEric Biggers  * Test that encryption and decryption produce the correct results in task,
913*9396dc94SEric Biggers  * softirq, and hardirq contexts running concurrently -- including with keys
914*9396dc94SEric Biggers  * prepared in other contexts.  This is needed to cover fallback code paths that
915*9396dc94SEric Biggers  * execute in contexts where FPU or vector registers cannot be used.
916*9396dc94SEric Biggers  */
917*9396dc94SEric Biggers static void test_aead_interrupt_context(struct kunit *test)
918*9396dc94SEric Biggers {
919*9396dc94SEric Biggers 	struct aead_irq_test_state *state = alloc_buf(test, sizeof(*state));
920*9396dc94SEric Biggers 
921*9396dc94SEric Biggers 	memset(state, 0, sizeof(*state));
922*9396dc94SEric Biggers 
923*9396dc94SEric Biggers 	/*
924*9396dc94SEric Biggers 	 * For each slot, generate a set of AEAD inputs: a key, a nonce, a
925*9396dc94SEric Biggers 	 * plaintext, and some associated data.  Then generate the corresponding
926*9396dc94SEric Biggers 	 * ciphertext with concatenated auth tag.
927*9396dc94SEric Biggers 	 */
928*9396dc94SEric Biggers 	for (int i = 0; i < IRQ_TEST_NUM_BUFFERS; i++) {
929*9396dc94SEric Biggers 		struct aead_irq_test_slot *slot = &state->slots[i];
930*9396dc94SEric Biggers 		int err;
931*9396dc94SEric Biggers 
932*9396dc94SEric Biggers 		rand_bytes(slot->raw_key, sizeof(slot->raw_key));
933*9396dc94SEric Biggers 		rand_bytes(slot->nonce, sizeof(slot->nonce));
934*9396dc94SEric Biggers 		rand_bytes(slot->pt, sizeof(slot->pt));
935*9396dc94SEric Biggers 		rand_bytes(slot->ad, sizeof(slot->ad));
936*9396dc94SEric Biggers 		err = AEAD_PREPAREKEY(&slot->key, slot->raw_key,
937*9396dc94SEric Biggers 				      sizeof(slot->raw_key), AEAD_MAX_TAG_LEN);
938*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
939*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(slot->ct, slot->pt, sizeof(slot->pt),
940*9396dc94SEric Biggers 				   &slot->ct[sizeof(slot->pt)], slot->ad,
941*9396dc94SEric Biggers 				   sizeof(slot->ad), slot->nonce,
942*9396dc94SEric Biggers 				   sizeof(slot->nonce), &slot->key);
943*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
944*9396dc94SEric Biggers 	}
945*9396dc94SEric Biggers 
946*9396dc94SEric Biggers 	kunit_run_irq_test(test, aead_irq_test_func, 100000, state);
947*9396dc94SEric Biggers }
948*9396dc94SEric Biggers 
949*9396dc94SEric Biggers /* Benchmark AEAD encryption and decryption on various data lengths. */
950*9396dc94SEric Biggers static void benchmark_aead(struct kunit *test)
951*9396dc94SEric Biggers {
952*9396dc94SEric Biggers 	static const size_t data_lens_to_test[] = {
953*9396dc94SEric Biggers 		16, 64, 128, 256, 512, 1024, 1420, 4096, 16384,
954*9396dc94SEric Biggers 	};
955*9396dc94SEric Biggers 	const size_t max_data_len = 16384;
956*9396dc94SEric Biggers 	const size_t ad_len = 16;
957*9396dc94SEric Biggers 	const size_t key_len = AEAD_MAX_KEY_LEN;
958*9396dc94SEric Biggers 	const size_t nonce_len = AEAD_MAX_NONCE_LEN;
959*9396dc94SEric Biggers 	const size_t tag_len = AEAD_MAX_TAG_LEN;
960*9396dc94SEric Biggers 	const u8 *raw_key, *nonce, *ad;
961*9396dc94SEric Biggers 	u8 *pt, *ct, *tag;
962*9396dc94SEric Biggers 	struct AEAD_KEY *key;
963*9396dc94SEric Biggers 	int err;
964*9396dc94SEric Biggers 
965*9396dc94SEric Biggers 	if (!IS_ENABLED(CONFIG_CRYPTO_LIB_BENCHMARK))
966*9396dc94SEric Biggers 		kunit_skip(test, "not enabled");
967*9396dc94SEric Biggers 
968*9396dc94SEric Biggers 	raw_key = aead_alloc_random_data(test, key_len);
969*9396dc94SEric Biggers 	nonce = aead_alloc_random_data(test, nonce_len);
970*9396dc94SEric Biggers 	ad = aead_alloc_random_data(test, ad_len);
971*9396dc94SEric Biggers 	pt = aead_alloc_random_data(test, max_data_len);
972*9396dc94SEric Biggers 	ct = alloc_buf(test, max_data_len);
973*9396dc94SEric Biggers 	tag = alloc_buf(test, tag_len);
974*9396dc94SEric Biggers 
975*9396dc94SEric Biggers 	key = alloc_buf(test, sizeof(*key));
976*9396dc94SEric Biggers 	err = AEAD_PREPAREKEY(key, raw_key, key_len, tag_len);
977*9396dc94SEric Biggers 	KUNIT_ASSERT_EQ(test, 0, err);
978*9396dc94SEric Biggers 
979*9396dc94SEric Biggers 	/* Warm-up */
980*9396dc94SEric Biggers 	for (size_t i = 0; i < 10000000; i += max_data_len) {
981*9396dc94SEric Biggers 		err = AEAD_ENCRYPT(ct, pt, max_data_len, tag, ad, ad_len, nonce,
982*9396dc94SEric Biggers 				   nonce_len, key);
983*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
984*9396dc94SEric Biggers 		err = AEAD_DECRYPT(pt, ct, max_data_len, tag, ad, ad_len, nonce,
985*9396dc94SEric Biggers 				   nonce_len, key);
986*9396dc94SEric Biggers 		KUNIT_ASSERT_EQ(test, 0, err);
987*9396dc94SEric Biggers 	}
988*9396dc94SEric Biggers 
989*9396dc94SEric Biggers 	for (size_t i = 0; i < ARRAY_SIZE(data_lens_to_test); i++) {
990*9396dc94SEric Biggers 		size_t data_len = data_lens_to_test[i];
991*9396dc94SEric Biggers 		size_t num_iters = 10000000 / (data_len + 128);
992*9396dc94SEric Biggers 		u64 t_enc, t_dec;
993*9396dc94SEric Biggers 		bool ok = true;
994*9396dc94SEric Biggers 
995*9396dc94SEric Biggers 		KUNIT_ASSERT_LE(test, data_len, max_data_len);
996*9396dc94SEric Biggers 
997*9396dc94SEric Biggers 		preempt_disable();
998*9396dc94SEric Biggers 
999*9396dc94SEric Biggers 		t_enc = ktime_get_ns();
1000*9396dc94SEric Biggers 		for (size_t j = 0; j < num_iters; j++) {
1001*9396dc94SEric Biggers 			err = AEAD_ENCRYPT(ct, pt, data_len, tag, ad, ad_len,
1002*9396dc94SEric Biggers 					   nonce, nonce_len, key);
1003*9396dc94SEric Biggers 			ok &= (err == 0);
1004*9396dc94SEric Biggers 		}
1005*9396dc94SEric Biggers 		t_enc = ktime_get_ns() - t_enc;
1006*9396dc94SEric Biggers 
1007*9396dc94SEric Biggers 		t_dec = ktime_get_ns();
1008*9396dc94SEric Biggers 		for (size_t j = 0; j < num_iters; j++) {
1009*9396dc94SEric Biggers 			err = AEAD_DECRYPT(pt, ct, data_len, tag, ad, ad_len,
1010*9396dc94SEric Biggers 					   nonce, nonce_len, key);
1011*9396dc94SEric Biggers 			ok &= (err == 0);
1012*9396dc94SEric Biggers 		}
1013*9396dc94SEric Biggers 		t_dec = ktime_get_ns() - t_dec;
1014*9396dc94SEric Biggers 
1015*9396dc94SEric Biggers 		preempt_enable();
1016*9396dc94SEric Biggers 
1017*9396dc94SEric Biggers 		KUNIT_ASSERT_TRUE_MSG(test, ok, "data_len=%zu", data_len);
1018*9396dc94SEric Biggers 
1019*9396dc94SEric Biggers 		kunit_info(test, "data_len=%zu: enc %llu MB/s, dec %llu MB/s",
1020*9396dc94SEric Biggers 			   data_len,
1021*9396dc94SEric Biggers 			   div64_u64((u64)data_len * num_iters * 1000,
1022*9396dc94SEric Biggers 				     t_enc ?: 1),
1023*9396dc94SEric Biggers 			   div64_u64((u64)data_len * num_iters * 1000,
1024*9396dc94SEric Biggers 				     t_dec ?: 1));
1025*9396dc94SEric Biggers 	}
1026*9396dc94SEric Biggers }
1027*9396dc94SEric Biggers 
1028*9396dc94SEric Biggers /* clang-format off */
1029*9396dc94SEric Biggers #define AEAD_KUNIT_CASES				\
1030*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_all_key_lens),		\
1031*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_all_nonce_lens),		\
1032*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_all_tag_lens),		\
1033*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_incremental_updates),	\
1034*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_data_buffer_overruns),	\
1035*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_alignment_consistency),	\
1036*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_inplace),			\
1037*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_monte_carlo),		\
1038*9396dc94SEric Biggers 	KUNIT_CASE(test_aead_interrupt_context),	\
1039*9396dc94SEric Biggers 	KUNIT_CASE(benchmark_aead)
1040