xref: /linux/tools/testing/selftests/arm64/fp/fp-ptrace.c (revision 85cdaca6970028bf6f544c355c90035586836ddf)
1 // SPDX-License-Identifier: GPL-2.0-only
2 /*
3  * Copyright (C) 2023 ARM Limited.
4  * Original author: Mark Brown <broonie@kernel.org>
5  */
6 
7 #define _GNU_SOURCE
8 
9 #include <errno.h>
10 #include <stdbool.h>
11 #include <stddef.h>
12 #include <stdio.h>
13 #include <stdlib.h>
14 #include <string.h>
15 #include <unistd.h>
16 
17 #include <sys/auxv.h>
18 #include <sys/prctl.h>
19 #include <sys/ptrace.h>
20 #include <sys/types.h>
21 #include <sys/uio.h>
22 #include <sys/wait.h>
23 
24 #include <linux/kernel.h>
25 
26 #include <asm/sigcontext.h>
27 #include <asm/sve_context.h>
28 #include <asm/ptrace.h>
29 
30 #include "kselftest.h"
31 
32 #include "fp-ptrace.h"
33 
34 #include <linux/bits.h>
35 
36 #define FPMR_LSCALE2_MASK                               GENMASK(37, 32)
37 #define FPMR_NSCALE_MASK                                GENMASK(31, 24)
38 #define FPMR_LSCALE_MASK                                GENMASK(22, 16)
39 #define FPMR_OSC_MASK                                   GENMASK(15, 15)
40 #define FPMR_OSM_MASK                                   GENMASK(14, 14)
41 
42 /* <linux/elf.h> and <sys/auxv.h> don't like each other, so: */
43 #ifndef NT_ARM_SVE
44 #define NT_ARM_SVE 0x405
45 #endif
46 
47 #ifndef NT_ARM_SSVE
48 #define NT_ARM_SSVE 0x40b
49 #endif
50 
51 #ifndef NT_ARM_ZA
52 #define NT_ARM_ZA 0x40c
53 #endif
54 
55 #ifndef NT_ARM_ZT
56 #define NT_ARM_ZT 0x40d
57 #endif
58 
59 #ifndef NT_ARM_FPMR
60 #define NT_ARM_FPMR 0x40e
61 #endif
62 
63 #define ARCH_VQ_MAX 256
64 
65 /* VL 128..2048 in powers of 2 */
66 #define MAX_NUM_VLS 5
67 
68 /* Sentinel for detecting buffer bytes the kernel did not write */
69 #define REGSET_SENTINEL 0xa5
70 
71 /*
72  * FPMR bits we can set without doing feature checks to see if values
73  * are valid.
74  */
75 #define FPMR_SAFE_BITS (FPMR_LSCALE2_MASK | FPMR_NSCALE_MASK | \
76 			FPMR_LSCALE_MASK | FPMR_OSC_MASK | FPMR_OSM_MASK)
77 
78 #define NUM_FPR 32
79 __uint128_t v_in[NUM_FPR];
80 __uint128_t v_expected[NUM_FPR];
81 __uint128_t v_out[NUM_FPR];
82 
83 char z_in[__SVE_ZREGS_SIZE(ARCH_VQ_MAX)];
84 char z_expected[__SVE_ZREGS_SIZE(ARCH_VQ_MAX)];
85 char z_out[__SVE_ZREGS_SIZE(ARCH_VQ_MAX)];
86 
87 char p_in[__SVE_PREGS_SIZE(ARCH_VQ_MAX)];
88 char p_expected[__SVE_PREGS_SIZE(ARCH_VQ_MAX)];
89 char p_out[__SVE_PREGS_SIZE(ARCH_VQ_MAX)];
90 
91 char ffr_in[__SVE_PREG_SIZE(ARCH_VQ_MAX)];
92 char ffr_expected[__SVE_PREG_SIZE(ARCH_VQ_MAX)];
93 char ffr_out[__SVE_PREG_SIZE(ARCH_VQ_MAX)];
94 
95 char za_in[ZA_SIG_REGS_SIZE(ARCH_VQ_MAX)];
96 char za_expected[ZA_SIG_REGS_SIZE(ARCH_VQ_MAX)];
97 char za_out[ZA_SIG_REGS_SIZE(ARCH_VQ_MAX)];
98 
99 char zt_in[ZT_SIG_REG_BYTES];
100 char zt_expected[ZT_SIG_REG_BYTES];
101 char zt_out[ZT_SIG_REG_BYTES];
102 
103 uint64_t fpmr_in, fpmr_expected, fpmr_out;
104 
105 uint64_t sve_vl_out;
106 uint64_t sme_vl_out;
107 uint64_t svcr_in, svcr_expected, svcr_out;
108 
109 void load_and_save(int flags);
110 
111 static bool got_alarm;
112 
113 static void handle_alarm(int sig, siginfo_t *info, void *context)
114 {
115 	got_alarm = true;
116 }
117 
118 #ifdef CONFIG_CPU_BIG_ENDIAN
119 static __uint128_t arm64_cpu_to_le128(__uint128_t x)
120 {
121 	u64 a = swab64(x);
122 	u64 b = swab64(x >> 64);
123 
124 	return ((__uint128_t)a << 64) | b;
125 }
126 #else
127 static __uint128_t arm64_cpu_to_le128(__uint128_t x)
128 {
129 	return x;
130 }
131 #endif
132 
133 #define arm64_le128_to_cpu(x) arm64_cpu_to_le128(x)
134 
135 static bool sve_supported(void)
136 {
137 	return getauxval(AT_HWCAP) & HWCAP_SVE;
138 }
139 
140 static bool sme_supported(void)
141 {
142 	return getauxval(AT_HWCAP2) & HWCAP2_SME;
143 }
144 
145 static bool sme2_supported(void)
146 {
147 	return getauxval(AT_HWCAP2) & HWCAP2_SME2;
148 }
149 
150 static bool fa64_supported(void)
151 {
152 	return getauxval(AT_HWCAP2) & HWCAP2_SME_FA64;
153 }
154 
155 static bool fpmr_supported(void)
156 {
157 	return getauxval(AT_HWCAP2) & HWCAP2_FPMR;
158 }
159 
160 static bool compare_buffer(const char *name, void *out,
161 			   void *expected, size_t size)
162 {
163 	void *tmp;
164 
165 	if (memcmp(out, expected, size) == 0)
166 		return true;
167 
168 	ksft_print_msg("Mismatch in %s\n", name);
169 
170 	/* Did we just get zeros back? */
171 	tmp = malloc(size);
172 	if (!tmp) {
173 		ksft_print_msg("OOM allocating %lu bytes for %s\n",
174 			       size, name);
175 		ksft_exit_fail();
176 	}
177 	memset(tmp, 0, size);
178 
179 	if (memcmp(out, tmp, size) == 0)
180 		ksft_print_msg("%s is zero\n", name);
181 
182 	free(tmp);
183 
184 	return false;
185 }
186 
187 static bool buffer_is_filled(const void *buffer, size_t size,
188 			     unsigned char value)
189 {
190 	const unsigned char *bytes = buffer;
191 	size_t i;
192 
193 	for (i = 0; i < size; i++) {
194 		if (bytes[i] != value)
195 			return false;
196 	}
197 
198 	return true;
199 }
200 
201 struct test_config {
202 	int sve_vl_in;
203 	int sve_vl_expected;
204 	int sme_vl_in;
205 	int sme_vl_expected;
206 	int svcr_in;
207 	int svcr_expected;
208 };
209 
210 struct test_definition {
211 	const char *name;
212 	bool sve_vl_change;
213 	bool (*supported)(struct test_config *config);
214 	void (*set_expected_values)(struct test_config *config);
215 	void (*modify_values)(pid_t child, struct test_config *test_config);
216 };
217 
218 static int vl_in(struct test_config *config)
219 {
220 	int vl;
221 
222 	if (config->svcr_in & SVCR_SM)
223 		vl = config->sme_vl_in;
224 	else
225 		vl = config->sve_vl_in;
226 
227 	return vl;
228 }
229 
230 static int vl_expected(struct test_config *config)
231 {
232 	int vl;
233 
234 	if (config->svcr_expected & SVCR_SM)
235 		vl = config->sme_vl_expected;
236 	else
237 		vl = config->sve_vl_expected;
238 
239 	return vl;
240 }
241 
242 static void run_child(struct test_config *config)
243 {
244 	int ret, flags;
245 
246 	/* Let the parent attach to us */
247 	ret = ptrace(PTRACE_TRACEME, 0, 0, 0);
248 	if (ret < 0)
249 		ksft_exit_fail_msg("PTRACE_TRACEME failed: %s (%d)\n",
250 				   strerror(errno), errno);
251 
252 	/* VL setup */
253 	if (sve_supported()) {
254 		ret = prctl(PR_SVE_SET_VL, config->sve_vl_in);
255 		if (ret != config->sve_vl_in) {
256 			ksft_print_msg("Failed to set SVE VL %d: %d\n",
257 				       config->sve_vl_in, ret);
258 		}
259 	}
260 
261 	if (sme_supported()) {
262 		ret = prctl(PR_SME_SET_VL, config->sme_vl_in);
263 		if (ret != config->sme_vl_in) {
264 			ksft_print_msg("Failed to set SME VL %d: %d\n",
265 				       config->sme_vl_in, ret);
266 		}
267 	}
268 
269 	/* Load values and wait for the parent */
270 	flags = 0;
271 	if (sve_supported())
272 		flags |= HAVE_SVE;
273 	if (sme_supported())
274 		flags |= HAVE_SME;
275 	if (sme2_supported())
276 		flags |= HAVE_SME2;
277 	if (fa64_supported())
278 		flags |= HAVE_FA64;
279 	if (fpmr_supported())
280 		flags |= HAVE_FPMR;
281 
282 	load_and_save(flags);
283 
284 	exit(0);
285 }
286 
287 static void read_one_child_regs(pid_t child, char *name,
288 				struct iovec *iov_parent,
289 				struct iovec *iov_child)
290 {
291 	int len = iov_parent->iov_len;
292 	int ret;
293 
294 	ret = process_vm_readv(child, iov_parent, 1, iov_child, 1, 0);
295 	if (ret == -1)
296 		ksft_print_msg("%s read failed: %s (%d)\n",
297 			       name, strerror(errno), errno);
298 	else if (ret != len)
299 		ksft_print_msg("Short read of %s: %d\n", name, ret);
300 }
301 
302 static void read_child_regs(pid_t child)
303 {
304 	struct iovec iov_parent, iov_child;
305 
306 	/*
307 	 * Since the child fork()ed from us the buffer addresses are
308 	 * the same in parent and child.
309 	 */
310 	iov_parent.iov_base = &v_out;
311 	iov_parent.iov_len = sizeof(v_out);
312 	iov_child.iov_base = &v_out;
313 	iov_child.iov_len = sizeof(v_out);
314 	read_one_child_regs(child, "FPSIMD", &iov_parent, &iov_child);
315 
316 	if (sve_supported() || sme_supported()) {
317 		iov_parent.iov_base = &sve_vl_out;
318 		iov_parent.iov_len = sizeof(sve_vl_out);
319 		iov_child.iov_base = &sve_vl_out;
320 		iov_child.iov_len = sizeof(sve_vl_out);
321 		read_one_child_regs(child, "SVE VL", &iov_parent, &iov_child);
322 
323 		iov_parent.iov_base = &z_out;
324 		iov_parent.iov_len = sizeof(z_out);
325 		iov_child.iov_base = &z_out;
326 		iov_child.iov_len = sizeof(z_out);
327 		read_one_child_regs(child, "Z", &iov_parent, &iov_child);
328 
329 		iov_parent.iov_base = &p_out;
330 		iov_parent.iov_len = sizeof(p_out);
331 		iov_child.iov_base = &p_out;
332 		iov_child.iov_len = sizeof(p_out);
333 		read_one_child_regs(child, "P", &iov_parent, &iov_child);
334 
335 		iov_parent.iov_base = &ffr_out;
336 		iov_parent.iov_len = sizeof(ffr_out);
337 		iov_child.iov_base = &ffr_out;
338 		iov_child.iov_len = sizeof(ffr_out);
339 		read_one_child_regs(child, "FFR", &iov_parent, &iov_child);
340 	}
341 
342 	if (sme_supported()) {
343 		iov_parent.iov_base = &sme_vl_out;
344 		iov_parent.iov_len = sizeof(sme_vl_out);
345 		iov_child.iov_base = &sme_vl_out;
346 		iov_child.iov_len = sizeof(sme_vl_out);
347 		read_one_child_regs(child, "SME VL", &iov_parent, &iov_child);
348 
349 		iov_parent.iov_base = &svcr_out;
350 		iov_parent.iov_len = sizeof(svcr_out);
351 		iov_child.iov_base = &svcr_out;
352 		iov_child.iov_len = sizeof(svcr_out);
353 		read_one_child_regs(child, "SVCR", &iov_parent, &iov_child);
354 
355 		iov_parent.iov_base = &za_out;
356 		iov_parent.iov_len = sizeof(za_out);
357 		iov_child.iov_base = &za_out;
358 		iov_child.iov_len = sizeof(za_out);
359 		read_one_child_regs(child, "ZA", &iov_parent, &iov_child);
360 	}
361 
362 	if (sme2_supported()) {
363 		iov_parent.iov_base = &zt_out;
364 		iov_parent.iov_len = sizeof(zt_out);
365 		iov_child.iov_base = &zt_out;
366 		iov_child.iov_len = sizeof(zt_out);
367 		read_one_child_regs(child, "ZT", &iov_parent, &iov_child);
368 	}
369 
370 	if (fpmr_supported()) {
371 		iov_parent.iov_base = &fpmr_out;
372 		iov_parent.iov_len = sizeof(fpmr_out);
373 		iov_child.iov_base = &fpmr_out;
374 		iov_child.iov_len = sizeof(fpmr_out);
375 		read_one_child_regs(child, "FPMR", &iov_parent, &iov_child);
376 	}
377 }
378 
379 static bool continue_breakpoint(pid_t child,
380 				enum __ptrace_request restart_type)
381 {
382 	struct user_pt_regs pt_regs;
383 	struct iovec iov;
384 	int ret;
385 
386 	/* Get PC */
387 	iov.iov_base = &pt_regs;
388 	iov.iov_len = sizeof(pt_regs);
389 	ret = ptrace(PTRACE_GETREGSET, child, NT_PRSTATUS, &iov);
390 	if (ret < 0) {
391 		ksft_print_msg("Failed to get PC: %s (%d)\n",
392 			       strerror(errno), errno);
393 		return false;
394 	}
395 
396 	/* Skip over the BRK */
397 	pt_regs.pc += 4;
398 	ret = ptrace(PTRACE_SETREGSET, child, NT_PRSTATUS, &iov);
399 	if (ret < 0) {
400 		ksft_print_msg("Failed to skip BRK: %s (%d)\n",
401 			       strerror(errno), errno);
402 		return false;
403 	}
404 
405 	/* Restart */
406 	ret = ptrace(restart_type, child, 0, 0);
407 	if (ret < 0) {
408 		ksft_print_msg("Failed to restart child: %s (%d)\n",
409 			       strerror(errno), errno);
410 		return false;
411 	}
412 
413 	return true;
414 }
415 
416 static bool check_ptrace_values_sve(pid_t child, struct test_config *config)
417 {
418 	struct user_sve_header *sve;
419 	struct user_fpsimd_state *fpsimd;
420 	struct iovec iov;
421 	size_t buf_size;
422 	int ret, vq;
423 	bool pass = true;
424 
425 	if (!sve_supported())
426 		return true;
427 
428 	vq = __sve_vq_from_vl(config->sve_vl_in);
429 
430 	buf_size = SVE_PT_SVE_OFFSET + SVE_PT_SVE_SIZE(vq, SVE_PT_REGS_SVE);
431 	iov.iov_len = buf_size;
432 	iov.iov_base = malloc(buf_size);
433 	if (!iov.iov_base) {
434 		ksft_print_msg("OOM allocating %lu byte SVE buffer\n",
435 			       iov.iov_len);
436 		return false;
437 	}
438 
439 	memset(iov.iov_base, REGSET_SENTINEL, buf_size);
440 	ret = ptrace(PTRACE_GETREGSET, child, NT_ARM_SVE, &iov);
441 	if (ret != 0) {
442 		ksft_print_msg("Failed to read initial SVE: %s (%d)\n",
443 			       strerror(errno), errno);
444 		pass = false;
445 		goto out;
446 	}
447 
448 	sve = iov.iov_base;
449 
450 	if (sve->vl != config->sve_vl_in) {
451 		ksft_print_msg("Mismatch in initial SVE VL: %d != %d\n",
452 			       sve->vl, config->sve_vl_in);
453 		pass = false;
454 	}
455 
456 	/* If we are in streaming mode we should just read FPSIMD */
457 	if ((config->svcr_in & SVCR_SM) && (sve->flags & SVE_PT_REGS_SVE)) {
458 		ksft_print_msg("NT_ARM_SVE reports SVE with PSTATE.SM\n");
459 		pass = false;
460 	}
461 
462 	if (svcr_in & SVCR_SM) {
463 		if (sve->size != sizeof(*sve)) {
464 			ksft_print_msg("NT_ARM_SVE reports data with PSTATE.SM\n");
465 			pass = false;
466 		}
467 		if (!buffer_is_filled(iov.iov_base + sizeof(*sve),
468 				      buf_size - sizeof(*sve), REGSET_SENTINEL)) {
469 			ksft_print_msg("NT_ARM_SVE wrote beyond its header with PSTATE.SM\n");
470 			pass = false;
471 		}
472 		goto out;
473 	} else {
474 		if (sve->size != SVE_PT_SIZE(vq, sve->flags)) {
475 			ksft_print_msg("Mismatch in SVE header size: %d != %lu\n",
476 				       sve->size, SVE_PT_SIZE(vq, sve->flags));
477 			pass = false;
478 		}
479 	}
480 
481 	/* The registers might be in completely different formats! */
482 	if (sve->flags & SVE_PT_REGS_SVE) {
483 		if (!compare_buffer("initial SVE Z",
484 				    iov.iov_base + SVE_PT_SVE_ZREG_OFFSET(vq, 0),
485 				    z_in, SVE_PT_SVE_ZREGS_SIZE(vq)))
486 			pass = false;
487 
488 		if (!compare_buffer("initial SVE P",
489 				    iov.iov_base + SVE_PT_SVE_PREG_OFFSET(vq, 0),
490 				    p_in, SVE_PT_SVE_PREGS_SIZE(vq)))
491 			pass = false;
492 
493 		if (!compare_buffer("initial SVE FFR",
494 				    iov.iov_base + SVE_PT_SVE_FFR_OFFSET(vq),
495 				    ffr_in, SVE_PT_SVE_PREG_SIZE(vq)))
496 			pass = false;
497 	} else {
498 		fpsimd = iov.iov_base + SVE_PT_FPSIMD_OFFSET;
499 		if (!compare_buffer("initial V via SVE", &fpsimd->vregs[0],
500 				    v_in, sizeof(v_in)))
501 			pass = false;
502 	}
503 
504 out:
505 	free(iov.iov_base);
506 	return pass;
507 }
508 
509 static bool check_ptrace_values_ssve(pid_t child, struct test_config *config)
510 {
511 	struct user_sve_header *sve;
512 	struct user_fpsimd_state *fpsimd;
513 	struct iovec iov;
514 	size_t buf_size;
515 	int ret, vq;
516 	bool pass = true;
517 
518 	if (!sme_supported())
519 		return true;
520 
521 	vq = __sve_vq_from_vl(config->sme_vl_in);
522 
523 	buf_size = SVE_PT_SVE_OFFSET + SVE_PT_SVE_SIZE(vq, SVE_PT_REGS_SVE);
524 	iov.iov_len = buf_size;
525 	iov.iov_base = malloc(buf_size);
526 	if (!iov.iov_base) {
527 		ksft_print_msg("OOM allocating %lu byte SSVE buffer\n",
528 			       iov.iov_len);
529 		return false;
530 	}
531 
532 	memset(iov.iov_base, REGSET_SENTINEL, buf_size);
533 	ret = ptrace(PTRACE_GETREGSET, child, NT_ARM_SSVE, &iov);
534 	if (ret != 0) {
535 		ksft_print_msg("Failed to read initial SSVE: %s (%d)\n",
536 			       strerror(errno), errno);
537 		pass = false;
538 		goto out;
539 	}
540 
541 	sve = iov.iov_base;
542 
543 	if (sve->vl != config->sme_vl_in) {
544 		ksft_print_msg("Mismatch in initial SSVE VL: %d != %d\n",
545 			       sve->vl, config->sme_vl_in);
546 		pass = false;
547 	}
548 
549 	if ((config->svcr_in & SVCR_SM) && !(sve->flags & SVE_PT_REGS_SVE)) {
550 		ksft_print_msg("NT_ARM_SSVE reports FPSIMD with PSTATE.SM\n");
551 		pass = false;
552 	}
553 
554 	if (!(svcr_in & SVCR_SM)) {
555 		if (sve->size != sizeof(*sve)) {
556 			ksft_print_msg("NT_ARM_SSVE reports data without PSTATE.SM\n");
557 			pass = false;
558 		}
559 		if (!buffer_is_filled(iov.iov_base + sizeof(*sve),
560 				      buf_size - sizeof(*sve), REGSET_SENTINEL)) {
561 			ksft_print_msg("NT_ARM_SSVE wrote beyond its header without PSTATE.SM\n");
562 			pass = false;
563 		}
564 		goto out;
565 	} else {
566 		if (sve->size != SVE_PT_SIZE(vq, sve->flags)) {
567 			ksft_print_msg("Mismatch in SSVE header size: %d != %lu\n",
568 				       sve->size, SVE_PT_SIZE(vq, sve->flags));
569 			pass = false;
570 		}
571 	}
572 
573 	/* The registers might be in completely different formats! */
574 	if (sve->flags & SVE_PT_REGS_SVE) {
575 		if (!compare_buffer("initial SSVE Z",
576 				    iov.iov_base + SVE_PT_SVE_ZREG_OFFSET(vq, 0),
577 				    z_in, SVE_PT_SVE_ZREGS_SIZE(vq)))
578 			pass = false;
579 
580 		if (!compare_buffer("initial SSVE P",
581 				    iov.iov_base + SVE_PT_SVE_PREG_OFFSET(vq, 0),
582 				    p_in, SVE_PT_SVE_PREGS_SIZE(vq)))
583 			pass = false;
584 
585 		if (!compare_buffer("initial SSVE FFR",
586 				    iov.iov_base + SVE_PT_SVE_FFR_OFFSET(vq),
587 				    ffr_in, SVE_PT_SVE_PREG_SIZE(vq)))
588 			pass = false;
589 	} else {
590 		fpsimd = iov.iov_base + SVE_PT_FPSIMD_OFFSET;
591 		if (!compare_buffer("initial V via SSVE",
592 				    &fpsimd->vregs[0], v_in, sizeof(v_in)))
593 			pass = false;
594 	}
595 
596 out:
597 	free(iov.iov_base);
598 	return pass;
599 }
600 
601 static bool check_ptrace_values_za(pid_t child, struct test_config *config)
602 {
603 	struct user_za_header *za;
604 	struct iovec iov;
605 	int ret, vq;
606 	bool pass = true;
607 
608 	if (!sme_supported())
609 		return true;
610 
611 	vq = __sve_vq_from_vl(config->sme_vl_in);
612 
613 	iov.iov_len = ZA_SIG_CONTEXT_SIZE(vq);
614 	iov.iov_base = malloc(iov.iov_len);
615 	if (!iov.iov_base) {
616 		ksft_print_msg("OOM allocating %lu byte ZA buffer\n",
617 			       iov.iov_len);
618 		return false;
619 	}
620 
621 	ret = ptrace(PTRACE_GETREGSET, child, NT_ARM_ZA, &iov);
622 	if (ret != 0) {
623 		ksft_print_msg("Failed to read initial ZA: %s (%d)\n",
624 			       strerror(errno), errno);
625 		pass = false;
626 		goto out;
627 	}
628 
629 	za = iov.iov_base;
630 
631 	if (za->vl != config->sme_vl_in) {
632 		ksft_print_msg("Mismatch in initial SME VL: %d != %d\n",
633 			       za->vl, config->sme_vl_in);
634 		pass = false;
635 	}
636 
637 	/* If PSTATE.ZA is not set we should just read the header */
638 	if (config->svcr_in & SVCR_ZA) {
639 		if (za->size != ZA_PT_SIZE(vq)) {
640 			ksft_print_msg("Unexpected ZA ptrace read size: %d != %lu\n",
641 				       za->size, ZA_PT_SIZE(vq));
642 			pass = false;
643 		}
644 
645 		if (!compare_buffer("initial ZA",
646 				    iov.iov_base + ZA_PT_ZA_OFFSET,
647 				    za_in, ZA_PT_ZA_SIZE(vq)))
648 			pass = false;
649 	} else {
650 		if (za->size != sizeof(*za)) {
651 			ksft_print_msg("Unexpected ZA ptrace read size: %d != %lu\n",
652 				       za->size, sizeof(*za));
653 			pass = false;
654 		}
655 	}
656 
657 out:
658 	free(iov.iov_base);
659 	return pass;
660 }
661 
662 static bool check_ptrace_values_zt(pid_t child, struct test_config *config)
663 {
664 	uint8_t buf[512];
665 	struct iovec iov;
666 	int ret;
667 
668 	if (!sme2_supported())
669 		return true;
670 
671 	iov.iov_base = &buf;
672 	iov.iov_len = ZT_SIG_REG_BYTES;
673 	ret = ptrace(PTRACE_GETREGSET, child, NT_ARM_ZT, &iov);
674 	if (ret != 0) {
675 		ksft_print_msg("Failed to read initial ZT: %s (%d)\n",
676 			       strerror(errno), errno);
677 		return false;
678 	}
679 
680 	return compare_buffer("initial ZT", buf, zt_in, ZT_SIG_REG_BYTES);
681 }
682 
683 static bool check_ptrace_values_fpmr(pid_t child, struct test_config *config)
684 {
685 	uint64_t val;
686 	struct iovec iov;
687 	int ret;
688 
689 	if (!fpmr_supported())
690 		return true;
691 
692 	iov.iov_base = &val;
693 	iov.iov_len = sizeof(val);
694 	ret = ptrace(PTRACE_GETREGSET, child, NT_ARM_FPMR, &iov);
695 	if (ret != 0) {
696 		ksft_print_msg("Failed to read initial FPMR: %s (%d)\n",
697 			       strerror(errno), errno);
698 		return false;
699 	}
700 
701 	return compare_buffer("initial FPMR", &val, &fpmr_in, sizeof(val));
702 }
703 
704 static bool check_ptrace_values(pid_t child, struct test_config *config)
705 {
706 	bool pass = true;
707 	struct user_fpsimd_state fpsimd;
708 	struct iovec iov;
709 	int ret;
710 
711 	iov.iov_base = &fpsimd;
712 	iov.iov_len = sizeof(fpsimd);
713 	ret = ptrace(PTRACE_GETREGSET, child, NT_PRFPREG, &iov);
714 	if (ret == 0) {
715 		if (!compare_buffer("initial V", &fpsimd.vregs, v_in,
716 				    sizeof(v_in))) {
717 			pass = false;
718 		}
719 	} else {
720 		ksft_print_msg("Failed to read initial V: %s (%d)\n",
721 			       strerror(errno), errno);
722 		pass = false;
723 	}
724 
725 	if (!check_ptrace_values_sve(child, config))
726 		pass = false;
727 
728 	if (!check_ptrace_values_ssve(child, config))
729 		pass = false;
730 
731 	if (!check_ptrace_values_za(child, config))
732 		pass = false;
733 
734 	if (!check_ptrace_values_zt(child, config))
735 		pass = false;
736 
737 	if (!check_ptrace_values_fpmr(child, config))
738 		pass = false;
739 
740 	return pass;
741 }
742 
743 static bool run_parent(pid_t child, struct test_definition *test,
744 		       struct test_config *config)
745 {
746 	int wait_status, ret;
747 	pid_t pid;
748 	bool pass;
749 
750 	/* Initial attach */
751 	while (1) {
752 		pid = waitpid(child, &wait_status, 0);
753 		if (pid < 0) {
754 			if (errno == EINTR)
755 				continue;
756 			ksft_exit_fail_msg("waitpid() failed: %s (%d)\n",
757 					   strerror(errno), errno);
758 		}
759 
760 		if (pid == child)
761 			break;
762 	}
763 
764 	if (WIFEXITED(wait_status)) {
765 		ksft_print_msg("Child exited loading values with status %d\n",
766 			       WEXITSTATUS(wait_status));
767 		pass = false;
768 		goto out;
769 	}
770 
771 	if (WIFSIGNALED(wait_status)) {
772 		ksft_print_msg("Child died from signal %d loading values\n",
773 			       WTERMSIG(wait_status));
774 		pass = false;
775 		goto out;
776 	}
777 
778 	/* Read initial values via ptrace */
779 	pass = check_ptrace_values(child, config);
780 
781 	/* Do whatever writes we want to do */
782 	if (test->modify_values)
783 		test->modify_values(child, config);
784 
785 	if (!continue_breakpoint(child, PTRACE_CONT))
786 		goto cleanup;
787 
788 	while (1) {
789 		pid = waitpid(child, &wait_status, 0);
790 		if (pid < 0) {
791 			if (errno == EINTR)
792 				continue;
793 			ksft_exit_fail_msg("waitpid() failed: %s (%d)\n",
794 					   strerror(errno), errno);
795 		}
796 
797 		if (pid == child)
798 			break;
799 	}
800 
801 	if (WIFEXITED(wait_status)) {
802 		ksft_print_msg("Child exited saving values with status %d\n",
803 			       WEXITSTATUS(wait_status));
804 		pass = false;
805 		goto out;
806 	}
807 
808 	if (WIFSIGNALED(wait_status)) {
809 		ksft_print_msg("Child died from signal %d saving values\n",
810 			       WTERMSIG(wait_status));
811 		pass = false;
812 		goto out;
813 	}
814 
815 	/* See what happened as a result */
816 	read_child_regs(child);
817 
818 	if (!continue_breakpoint(child, PTRACE_DETACH))
819 		goto cleanup;
820 
821 	/* The child should exit cleanly */
822 	got_alarm = false;
823 	alarm(1);
824 	while (1) {
825 		if (got_alarm) {
826 			ksft_print_msg("Wait for child timed out\n");
827 			goto cleanup;
828 		}
829 
830 		pid = waitpid(child, &wait_status, 0);
831 		if (pid < 0) {
832 			if (errno == EINTR)
833 				continue;
834 			ksft_exit_fail_msg("waitpid() failed: %s (%d)\n",
835 					   strerror(errno), errno);
836 		}
837 
838 		if (pid == child)
839 			break;
840 	}
841 	alarm(0);
842 
843 	if (got_alarm) {
844 		ksft_print_msg("Timed out waiting for child\n");
845 		pass = false;
846 		goto cleanup;
847 	}
848 
849 	if (pid == child && WIFSIGNALED(wait_status)) {
850 		ksft_print_msg("Child died from signal %d cleaning up\n",
851 			       WTERMSIG(wait_status));
852 		pass = false;
853 		goto out;
854 	}
855 
856 	if (pid == child && WIFEXITED(wait_status)) {
857 		if (WEXITSTATUS(wait_status) != 0) {
858 			ksft_print_msg("Child exited with error %d\n",
859 				       WEXITSTATUS(wait_status));
860 			pass = false;
861 		}
862 	} else {
863 		ksft_print_msg("Child did not exit cleanly\n");
864 		pass = false;
865 		goto cleanup;
866 	}
867 
868 	goto out;
869 
870 cleanup:
871 	ret = kill(child, SIGKILL);
872 	if (ret != 0) {
873 		ksft_print_msg("kill() failed: %s (%d)\n",
874 			       strerror(errno), errno);
875 		return false;
876 	}
877 
878 	while (1) {
879 		pid = waitpid(child, &wait_status, 0);
880 		if (pid < 0) {
881 			if (errno == EINTR)
882 				continue;
883 			ksft_exit_fail_msg("waitpid() failed: %s (%d)\n",
884 					   strerror(errno), errno);
885 		}
886 
887 		if (pid == child)
888 			break;
889 	}
890 
891 out:
892 	return pass;
893 }
894 
895 static void fill_random(void *buf, size_t size)
896 {
897 	int i;
898 	uint32_t *lbuf = buf;
899 
900 	/* random() returns a 32 bit number regardless of the size of long */
901 	for (i = 0; i < size / sizeof(uint32_t); i++)
902 		lbuf[i] = random();
903 }
904 
905 static void fill_random_ffr(void *buf, size_t vq)
906 {
907 	uint8_t *lbuf = buf;
908 	int bits, i;
909 
910 	/*
911 	 * Only values with a continuous set of 0..n bits set are
912 	 * valid for FFR, set all bits then clear a random number of
913 	 * high bits.
914 	 */
915 	memset(buf, 0, __SVE_FFR_SIZE(vq));
916 
917 	bits = random() % (__SVE_FFR_SIZE(vq) * 8);
918 	for (i = 0; i < bits / 8; i++)
919 		lbuf[i] = 0xff;
920 	if (bits / 8 != __SVE_FFR_SIZE(vq))
921 		lbuf[i] = (1 << (bits % 8)) - 1;
922 }
923 
924 static void fpsimd_to_sve(__uint128_t *v, char *z, int vl)
925 {
926 	int vq = __sve_vq_from_vl(vl);
927 	int i;
928 	__uint128_t *p;
929 
930 	if (!vl)
931 		return;
932 
933 	for (i = 0; i < __SVE_NUM_ZREGS; i++) {
934 		p = (__uint128_t *)&z[__SVE_ZREG_OFFSET(vq, i)];
935 		*p = arm64_cpu_to_le128(v[i]);
936 	}
937 }
938 
939 static void set_initial_values(struct test_config *config)
940 {
941 	int vq = __sve_vq_from_vl(vl_in(config));
942 	int sme_vq = __sve_vq_from_vl(config->sme_vl_in);
943 
944 	svcr_in = config->svcr_in;
945 	svcr_expected = config->svcr_expected;
946 	svcr_out = 0;
947 
948 	fill_random(&v_in, sizeof(v_in));
949 	memcpy(v_expected, v_in, sizeof(v_in));
950 	memset(v_out, 0, sizeof(v_out));
951 
952 	/* Changes will be handled in the test case */
953 	if (sve_supported() || (config->svcr_in & SVCR_SM)) {
954 		/* The low 128 bits of Z are shared with the V registers */
955 		fill_random(&z_in, __SVE_ZREGS_SIZE(vq));
956 		fpsimd_to_sve(v_in, z_in, vl_in(config));
957 		memcpy(z_expected, z_in, __SVE_ZREGS_SIZE(vq));
958 		memset(z_out, 0, sizeof(z_out));
959 
960 		fill_random(&p_in, __SVE_PREGS_SIZE(vq));
961 		memcpy(p_expected, p_in, __SVE_PREGS_SIZE(vq));
962 		memset(p_out, 0, sizeof(p_out));
963 
964 		if ((config->svcr_in & SVCR_SM) && !fa64_supported())
965 			memset(ffr_in, 0, __SVE_PREG_SIZE(vq));
966 		else
967 			fill_random_ffr(&ffr_in, vq);
968 		memcpy(ffr_expected, ffr_in, __SVE_PREG_SIZE(vq));
969 		memset(ffr_out, 0, __SVE_PREG_SIZE(vq));
970 	}
971 
972 	if (config->svcr_in & SVCR_ZA)
973 		fill_random(za_in, ZA_SIG_REGS_SIZE(sme_vq));
974 	else
975 		memset(za_in, 0, ZA_SIG_REGS_SIZE(sme_vq));
976 	if (config->svcr_expected & SVCR_ZA)
977 		memcpy(za_expected, za_in, ZA_SIG_REGS_SIZE(sme_vq));
978 	else
979 		memset(za_expected, 0, ZA_SIG_REGS_SIZE(sme_vq));
980 	if (sme_supported())
981 		memset(za_out, 0, sizeof(za_out));
982 
983 	if (sme2_supported()) {
984 		if (config->svcr_in & SVCR_ZA)
985 			fill_random(zt_in, ZT_SIG_REG_BYTES);
986 		else
987 			memset(zt_in, 0, ZT_SIG_REG_BYTES);
988 		if (config->svcr_expected & SVCR_ZA)
989 			memcpy(zt_expected, zt_in, ZT_SIG_REG_BYTES);
990 		else
991 			memset(zt_expected, 0, ZT_SIG_REG_BYTES);
992 		memset(zt_out, 0, sizeof(zt_out));
993 	}
994 
995 	if (fpmr_supported()) {
996 		fill_random(&fpmr_in, sizeof(fpmr_in));
997 		fpmr_in &= FPMR_SAFE_BITS;
998 		fpmr_expected = fpmr_in;
999 	} else {
1000 		fpmr_in = 0;
1001 		fpmr_expected = 0;
1002 		fpmr_out = 0;
1003 	}
1004 }
1005 
1006 static bool check_memory_values(struct test_config *config)
1007 {
1008 	bool pass = true;
1009 	int vq, sme_vq;
1010 
1011 	if (!compare_buffer("saved V", v_out, v_expected, sizeof(v_out)))
1012 		pass = false;
1013 
1014 	vq = __sve_vq_from_vl(vl_expected(config));
1015 	sme_vq = __sve_vq_from_vl(config->sme_vl_expected);
1016 
1017 	if (svcr_out != svcr_expected) {
1018 		ksft_print_msg("Mismatch in saved SVCR %lx != %lx\n",
1019 			       svcr_out, svcr_expected);
1020 		pass = false;
1021 	}
1022 
1023 	if (sve_vl_out != config->sve_vl_expected) {
1024 		ksft_print_msg("Mismatch in SVE VL: %ld != %d\n",
1025 			       sve_vl_out, config->sve_vl_expected);
1026 		pass = false;
1027 	}
1028 
1029 	if (sme_vl_out != config->sme_vl_expected) {
1030 		ksft_print_msg("Mismatch in SME VL: %ld != %d\n",
1031 			       sme_vl_out, config->sme_vl_expected);
1032 		pass = false;
1033 	}
1034 
1035 	if (!compare_buffer("saved Z", z_out, z_expected,
1036 			    __SVE_ZREGS_SIZE(vq)))
1037 		pass = false;
1038 
1039 	if (!compare_buffer("saved P", p_out, p_expected,
1040 			    __SVE_PREGS_SIZE(vq)))
1041 		pass = false;
1042 
1043 	if (!compare_buffer("saved FFR", ffr_out, ffr_expected,
1044 			    __SVE_PREG_SIZE(vq)))
1045 		pass = false;
1046 
1047 	if (!compare_buffer("saved ZA", za_out, za_expected,
1048 			    ZA_PT_ZA_SIZE(sme_vq)))
1049 		pass = false;
1050 
1051 	if (!compare_buffer("saved ZT", zt_out, zt_expected, ZT_SIG_REG_BYTES))
1052 		pass = false;
1053 
1054 	if (fpmr_out != fpmr_expected) {
1055 		ksft_print_msg("Mismatch in saved FPMR: %lx != %lx\n",
1056 			       fpmr_out, fpmr_expected);
1057 		pass = false;
1058 	}
1059 
1060 	return pass;
1061 }
1062 
1063 static bool sve_sme_same(struct test_config *config)
1064 {
1065 	if (config->sve_vl_in != config->sve_vl_expected)
1066 		return false;
1067 
1068 	if (config->sme_vl_in != config->sme_vl_expected)
1069 		return false;
1070 
1071 	if (config->svcr_in != config->svcr_expected)
1072 		return false;
1073 
1074 	return true;
1075 }
1076 
1077 static bool sve_write_supported(struct test_config *config)
1078 {
1079 	if (!sve_supported() && !sme_supported())
1080 		return false;
1081 
1082 	if ((config->svcr_in & SVCR_ZA) != (config->svcr_expected & SVCR_ZA))
1083 		return false;
1084 
1085 	if (config->svcr_expected & SVCR_SM) {
1086 		if (config->sve_vl_in != config->sve_vl_expected) {
1087 			return false;
1088 		}
1089 
1090 		/* Changing the SME VL disables ZA */
1091 		if ((config->svcr_expected & SVCR_ZA) &&
1092 		    (config->sme_vl_in != config->sme_vl_expected)) {
1093 			return false;
1094 		}
1095 	} else {
1096 		if (config->sme_vl_in != config->sme_vl_expected) {
1097 			return false;
1098 		}
1099 
1100 		if (!sve_supported())
1101 			return false;
1102 	}
1103 
1104 	return true;
1105 }
1106 
1107 static bool sve_write_fpsimd_supported(struct test_config *config)
1108 {
1109 	if (!sve_supported() && !sme_supported())
1110 		return false;
1111 
1112 	if ((config->svcr_in & SVCR_ZA) != (config->svcr_expected & SVCR_ZA))
1113 		return false;
1114 
1115 	if (config->svcr_expected & SVCR_SM)
1116 		return false;
1117 
1118 	if (config->sme_vl_in != config->sme_vl_expected)
1119 		return false;
1120 
1121 	return true;
1122 }
1123 
1124 static void fpsimd_write_expected(struct test_config *config)
1125 {
1126 	int vl;
1127 
1128 	fill_random(&v_expected, sizeof(v_expected));
1129 
1130 	/* The SVE registers are flushed by a FPSIMD write */
1131 	vl = vl_expected(config);
1132 
1133 	memset(z_expected, 0, __SVE_ZREGS_SIZE(__sve_vq_from_vl(vl)));
1134 	memset(p_expected, 0, __SVE_PREGS_SIZE(__sve_vq_from_vl(vl)));
1135 	memset(ffr_expected, 0, __SVE_PREG_SIZE(__sve_vq_from_vl(vl)));
1136 
1137 	fpsimd_to_sve(v_expected, z_expected, vl);
1138 }
1139 
1140 static void fpsimd_write(pid_t child, struct test_config *test_config)
1141 {
1142 	struct user_fpsimd_state fpsimd;
1143 	struct iovec iov;
1144 	int ret;
1145 
1146 	memset(&fpsimd, 0, sizeof(fpsimd));
1147 	memcpy(&fpsimd.vregs, v_expected, sizeof(v_expected));
1148 
1149 	iov.iov_base = &fpsimd;
1150 	iov.iov_len = sizeof(fpsimd);
1151 	ret = ptrace(PTRACE_SETREGSET, child, NT_PRFPREG, &iov);
1152 	if (ret == -1)
1153 		ksft_print_msg("FPSIMD set failed: (%s) %d\n",
1154 			       strerror(errno), errno);
1155 }
1156 
1157 static bool fpmr_write_supported(struct test_config *config)
1158 {
1159 	if (!fpmr_supported())
1160 		return false;
1161 
1162 	if (!sve_sme_same(config))
1163 		return false;
1164 
1165 	return true;
1166 }
1167 
1168 static void fpmr_write_expected(struct test_config *config)
1169 {
1170 	fill_random(&fpmr_expected, sizeof(fpmr_expected));
1171 	fpmr_expected &= FPMR_SAFE_BITS;
1172 }
1173 
1174 static void fpmr_write(pid_t child, struct test_config *config)
1175 {
1176 	struct iovec iov;
1177 	int ret;
1178 
1179 	iov.iov_len = sizeof(fpmr_expected);
1180 	iov.iov_base = &fpmr_expected;
1181 	ret = ptrace(PTRACE_SETREGSET, child, NT_ARM_FPMR, &iov);
1182 	if (ret != 0)
1183 		ksft_print_msg("Failed to write FPMR: %s (%d)\n",
1184 			       strerror(errno), errno);
1185 }
1186 
1187 static void sve_write_expected(struct test_config *config)
1188 {
1189 	int vl = vl_expected(config);
1190 	int sme_vq = __sve_vq_from_vl(config->sme_vl_expected);
1191 
1192 	if (!vl)
1193 		return;
1194 
1195 	fill_random(z_expected, __SVE_ZREGS_SIZE(__sve_vq_from_vl(vl)));
1196 	fill_random(p_expected, __SVE_PREGS_SIZE(__sve_vq_from_vl(vl)));
1197 
1198 	if ((svcr_expected & SVCR_SM) && !fa64_supported())
1199 		memset(ffr_expected, 0, __SVE_PREG_SIZE(sme_vq));
1200 	else
1201 		fill_random_ffr(ffr_expected, __sve_vq_from_vl(vl));
1202 
1203 	/* Share the low bits of Z with V */
1204 	fill_random(&v_expected, sizeof(v_expected));
1205 	fpsimd_to_sve(v_expected, z_expected, vl);
1206 
1207 	if (config->sme_vl_in != config->sme_vl_expected) {
1208 		memset(za_expected, 0, ZA_PT_ZA_SIZE(sme_vq));
1209 		memset(zt_expected, 0, sizeof(zt_expected));
1210 	}
1211 }
1212 
1213 static void sve_write_sve(pid_t child, struct test_config *config)
1214 {
1215 	struct user_sve_header *sve;
1216 	struct iovec iov;
1217 	int ret, vl, vq, regset;
1218 
1219 	vl = vl_expected(config);
1220 	vq = __sve_vq_from_vl(vl);
1221 
1222 	if (!vl)
1223 		return;
1224 
1225 	iov.iov_len = SVE_PT_SIZE(vq, SVE_PT_REGS_SVE);
1226 	iov.iov_base = malloc(iov.iov_len);
1227 	if (!iov.iov_base) {
1228 		ksft_print_msg("Failed allocating %lu byte SVE write buffer\n",
1229 			       iov.iov_len);
1230 		return;
1231 	}
1232 	memset(iov.iov_base, 0, iov.iov_len);
1233 
1234 	sve = iov.iov_base;
1235 	sve->size = iov.iov_len;
1236 	sve->flags = SVE_PT_REGS_SVE;
1237 	sve->vl = vl;
1238 
1239 	memcpy(iov.iov_base + SVE_PT_SVE_ZREG_OFFSET(vq, 0),
1240 	       z_expected, SVE_PT_SVE_ZREGS_SIZE(vq));
1241 	memcpy(iov.iov_base + SVE_PT_SVE_PREG_OFFSET(vq, 0),
1242 	       p_expected, SVE_PT_SVE_PREGS_SIZE(vq));
1243 	memcpy(iov.iov_base + SVE_PT_SVE_FFR_OFFSET(vq),
1244 	       ffr_expected, SVE_PT_SVE_PREG_SIZE(vq));
1245 
1246 	if (svcr_expected & SVCR_SM)
1247 		regset = NT_ARM_SSVE;
1248 	else
1249 		regset = NT_ARM_SVE;
1250 
1251 	ret = ptrace(PTRACE_SETREGSET, child, regset, &iov);
1252 	if (ret != 0)
1253 		ksft_print_msg("Failed to write SVE: %s (%d)\n",
1254 			       strerror(errno), errno);
1255 
1256 	free(iov.iov_base);
1257 }
1258 
1259 static void sve_write_fpsimd(pid_t child, struct test_config *config)
1260 {
1261 	struct user_sve_header *sve;
1262 	struct user_fpsimd_state *fpsimd;
1263 	struct iovec iov;
1264 	int ret, vl, vq;
1265 
1266 	vl = vl_expected(config);
1267 	vq = __sve_vq_from_vl(vl);
1268 
1269 	iov.iov_len = SVE_PT_SIZE(vq, SVE_PT_REGS_FPSIMD);
1270 	iov.iov_base = malloc(iov.iov_len);
1271 	if (!iov.iov_base) {
1272 		ksft_print_msg("Failed allocating %lu byte SVE write buffer\n",
1273 			       iov.iov_len);
1274 		return;
1275 	}
1276 	memset(iov.iov_base, 0, iov.iov_len);
1277 
1278 	sve = iov.iov_base;
1279 	sve->size = iov.iov_len;
1280 	sve->flags = SVE_PT_REGS_FPSIMD;
1281 	sve->vl = vl;
1282 
1283 	fpsimd = iov.iov_base + SVE_PT_REGS_OFFSET;
1284 	memcpy(&fpsimd->vregs, v_expected, sizeof(v_expected));
1285 
1286 	ret = ptrace(PTRACE_SETREGSET, child, NT_ARM_SVE, &iov);
1287 	if (ret != 0)
1288 		ksft_print_msg("Failed to write SVE: %s (%d)\n",
1289 			       strerror(errno), errno);
1290 
1291 	free(iov.iov_base);
1292 }
1293 
1294 static bool za_write_supported(struct test_config *config)
1295 {
1296 	if ((config->svcr_in & SVCR_SM) != (config->svcr_expected & SVCR_SM))
1297 		return false;
1298 
1299 	return true;
1300 }
1301 
1302 static void za_write_expected(struct test_config *config)
1303 {
1304 	int sme_vq, sve_vq;
1305 
1306 	sme_vq = __sve_vq_from_vl(config->sme_vl_expected);
1307 
1308 	if (config->svcr_expected & SVCR_ZA) {
1309 		fill_random(za_expected, ZA_PT_ZA_SIZE(sme_vq));
1310 	} else {
1311 		memset(za_expected, 0, ZA_PT_ZA_SIZE(sme_vq));
1312 		memset(zt_expected, 0, sizeof(zt_expected));
1313 	}
1314 
1315 	/* Changing the SME VL flushes ZT, SVE state */
1316 	if (config->sme_vl_in != config->sme_vl_expected) {
1317 		sve_vq = __sve_vq_from_vl(vl_expected(config));
1318 		memset(z_expected, 0, __SVE_ZREGS_SIZE(sve_vq));
1319 		memset(p_expected, 0, __SVE_PREGS_SIZE(sve_vq));
1320 		memset(ffr_expected, 0, __SVE_PREG_SIZE(sve_vq));
1321 		memset(zt_expected, 0, sizeof(zt_expected));
1322 
1323 		fpsimd_to_sve(v_expected, z_expected, vl_expected(config));
1324 	}
1325 }
1326 
1327 static void za_write(pid_t child, struct test_config *config)
1328 {
1329 	struct user_za_header *za;
1330 	struct iovec iov;
1331 	int ret, vq;
1332 
1333 	vq = __sve_vq_from_vl(config->sme_vl_expected);
1334 
1335 	if (config->svcr_expected & SVCR_ZA)
1336 		iov.iov_len = ZA_PT_SIZE(vq);
1337 	else
1338 		iov.iov_len = sizeof(*za);
1339 	iov.iov_base = malloc(iov.iov_len);
1340 	if (!iov.iov_base) {
1341 		ksft_print_msg("Failed allocating %lu byte ZA write buffer\n",
1342 			       iov.iov_len);
1343 		return;
1344 	}
1345 	memset(iov.iov_base, 0, iov.iov_len);
1346 
1347 	za = iov.iov_base;
1348 	za->size = iov.iov_len;
1349 	za->vl = config->sme_vl_expected;
1350 	if (config->svcr_expected & SVCR_ZA)
1351 		memcpy(iov.iov_base + ZA_PT_ZA_OFFSET, za_expected,
1352 		       ZA_PT_ZA_SIZE(vq));
1353 
1354 	ret = ptrace(PTRACE_SETREGSET, child, NT_ARM_ZA, &iov);
1355 	if (ret != 0)
1356 		ksft_print_msg("Failed to write ZA: %s (%d)\n",
1357 			       strerror(errno), errno);
1358 
1359 	free(iov.iov_base);
1360 }
1361 
1362 static bool zt_write_supported(struct test_config *config)
1363 {
1364 	if (!sme2_supported())
1365 		return false;
1366 	if (config->sme_vl_in != config->sme_vl_expected)
1367 		return false;
1368 	if (!(config->svcr_expected & SVCR_ZA))
1369 		return false;
1370 	if ((config->svcr_in & SVCR_SM) != (config->svcr_expected & SVCR_SM))
1371 		return false;
1372 
1373 	return true;
1374 }
1375 
1376 static void zt_write_expected(struct test_config *config)
1377 {
1378 	int sme_vq;
1379 
1380 	sme_vq = __sve_vq_from_vl(config->sme_vl_expected);
1381 
1382 	if (config->svcr_expected & SVCR_ZA) {
1383 		fill_random(zt_expected, sizeof(zt_expected));
1384 	} else {
1385 		memset(za_expected, 0, ZA_PT_ZA_SIZE(sme_vq));
1386 		memset(zt_expected, 0, sizeof(zt_expected));
1387 	}
1388 }
1389 
1390 static void zt_write(pid_t child, struct test_config *config)
1391 {
1392 	struct iovec iov;
1393 	int ret;
1394 
1395 	iov.iov_len = ZT_SIG_REG_BYTES;
1396 	iov.iov_base = zt_expected;
1397 	ret = ptrace(PTRACE_SETREGSET, child, NT_ARM_ZT, &iov);
1398 	if (ret != 0)
1399 		ksft_print_msg("Failed to write ZT: %s (%d)\n",
1400 			       strerror(errno), errno);
1401 }
1402 
1403 /* Actually run a test */
1404 static void run_test(struct test_definition *test, struct test_config *config)
1405 {
1406 	pid_t child;
1407 	char name[1024];
1408 	bool pass;
1409 
1410 	if (sve_supported() && sme_supported())
1411 		snprintf(name, sizeof(name), "%s, SVE %d->%d, SME %d/%x->%d/%x",
1412 			 test->name,
1413 			 config->sve_vl_in, config->sve_vl_expected,
1414 			 config->sme_vl_in, config->svcr_in,
1415 			 config->sme_vl_expected, config->svcr_expected);
1416 	else if (sve_supported())
1417 		snprintf(name, sizeof(name), "%s, SVE %d->%d", test->name,
1418 			 config->sve_vl_in, config->sve_vl_expected);
1419 	else if (sme_supported())
1420 		snprintf(name, sizeof(name), "%s, SME %d/%x->%d/%x",
1421 			 test->name,
1422 			 config->sme_vl_in, config->svcr_in,
1423 			 config->sme_vl_expected, config->svcr_expected);
1424 	else
1425 		snprintf(name, sizeof(name), "%s", test->name);
1426 
1427 	if (test->supported && !test->supported(config)) {
1428 		ksft_test_result_skip("%s\n", name);
1429 		return;
1430 	}
1431 
1432 	set_initial_values(config);
1433 
1434 	if (test->set_expected_values)
1435 		test->set_expected_values(config);
1436 
1437 	child = fork();
1438 	if (child < 0)
1439 		ksft_exit_fail_msg("fork() failed: %s (%d)\n",
1440 				   strerror(errno), errno);
1441 	/* run_child() never returns */
1442 	if (child == 0)
1443 		run_child(config);
1444 
1445 	pass = run_parent(child, test, config);
1446 	if (!check_memory_values(config))
1447 		pass = false;
1448 
1449 	ksft_test_result(pass, "%s\n", name);
1450 }
1451 
1452 static void run_tests(struct test_definition defs[], int count,
1453 		      struct test_config *config)
1454 {
1455 	int i;
1456 
1457 	for (i = 0; i < count; i++)
1458 		run_test(&defs[i], config);
1459 }
1460 
1461 static struct test_definition base_test_defs[] = {
1462 	{
1463 		.name = "No writes",
1464 		.supported = sve_sme_same,
1465 	},
1466 	{
1467 		.name = "FPSIMD write",
1468 		.supported = sve_sme_same,
1469 		.set_expected_values = fpsimd_write_expected,
1470 		.modify_values = fpsimd_write,
1471 	},
1472 	{
1473 		.name = "FPMR write",
1474 		.supported = fpmr_write_supported,
1475 		.set_expected_values = fpmr_write_expected,
1476 		.modify_values = fpmr_write,
1477 	},
1478 };
1479 
1480 static struct test_definition sve_test_defs[] = {
1481 	{
1482 		.name = "SVE write",
1483 		.supported = sve_write_supported,
1484 		.set_expected_values = sve_write_expected,
1485 		.modify_values = sve_write_sve,
1486 	},
1487 	{
1488 		.name = "SVE write FPSIMD format",
1489 		.supported = sve_write_fpsimd_supported,
1490 		.set_expected_values = fpsimd_write_expected,
1491 		.modify_values = sve_write_fpsimd,
1492 	},
1493 };
1494 
1495 static struct test_definition za_test_defs[] = {
1496 	{
1497 		.name = "ZA write",
1498 		.supported = za_write_supported,
1499 		.set_expected_values = za_write_expected,
1500 		.modify_values = za_write,
1501 	},
1502 };
1503 
1504 static struct test_definition zt_test_defs[] = {
1505 	{
1506 		.name = "ZT write",
1507 		.supported = zt_write_supported,
1508 		.set_expected_values = zt_write_expected,
1509 		.modify_values = zt_write,
1510 	},
1511 };
1512 
1513 static int sve_vls[MAX_NUM_VLS], sme_vls[MAX_NUM_VLS];
1514 static int sve_vl_count, sme_vl_count;
1515 
1516 static void probe_vls(const char *name, int vls[], int *vl_count, int set_vl)
1517 {
1518 	unsigned int vq;
1519 	int vl;
1520 
1521 	*vl_count = 0;
1522 
1523 	for (vq = ARCH_VQ_MAX; vq > 0; vq /= 2) {
1524 		vl = prctl(set_vl, vq * 16);
1525 		if (vl == -1)
1526 			ksft_exit_fail_msg("SET_VL failed: %s (%d)\n",
1527 					   strerror(errno), errno);
1528 
1529 		vl &= PR_SVE_VL_LEN_MASK;
1530 
1531 		if (*vl_count && (vl == vls[*vl_count - 1]))
1532 			break;
1533 
1534 		vq = sve_vq_from_vl(vl);
1535 
1536 		vls[*vl_count] = vl;
1537 		*vl_count += 1;
1538 	}
1539 
1540 	if (*vl_count > 2) {
1541 		/* Just use the minimum and maximum */
1542 		vls[1] = vls[*vl_count - 1];
1543 		ksft_print_msg("%d %s VLs, using %d and %d\n",
1544 			       *vl_count, name, vls[0], vls[1]);
1545 		*vl_count = 2;
1546 	} else {
1547 		ksft_print_msg("%d %s VLs\n", *vl_count, name);
1548 	}
1549 }
1550 
1551 static struct {
1552 	int svcr_in, svcr_expected;
1553 } svcr_combinations[] = {
1554 	{ .svcr_in = 0, .svcr_expected = 0, },
1555 	{ .svcr_in = 0, .svcr_expected = SVCR_SM, },
1556 	{ .svcr_in = 0, .svcr_expected = SVCR_ZA, },
1557 	/* Can't enable both SM and ZA with a single ptrace write */
1558 
1559 	{ .svcr_in = SVCR_SM, .svcr_expected = 0, },
1560 	{ .svcr_in = SVCR_SM, .svcr_expected = SVCR_SM, },
1561 	{ .svcr_in = SVCR_SM, .svcr_expected = SVCR_ZA, },
1562 	{ .svcr_in = SVCR_SM, .svcr_expected = SVCR_SM | SVCR_ZA, },
1563 
1564 	{ .svcr_in = SVCR_ZA, .svcr_expected = 0, },
1565 	{ .svcr_in = SVCR_ZA, .svcr_expected = SVCR_SM, },
1566 	{ .svcr_in = SVCR_ZA, .svcr_expected = SVCR_ZA, },
1567 	{ .svcr_in = SVCR_ZA, .svcr_expected = SVCR_SM | SVCR_ZA, },
1568 
1569 	{ .svcr_in = SVCR_SM | SVCR_ZA, .svcr_expected = 0, },
1570 	{ .svcr_in = SVCR_SM | SVCR_ZA, .svcr_expected = SVCR_SM, },
1571 	{ .svcr_in = SVCR_SM | SVCR_ZA, .svcr_expected = SVCR_ZA, },
1572 	{ .svcr_in = SVCR_SM | SVCR_ZA, .svcr_expected = SVCR_SM | SVCR_ZA, },
1573 };
1574 
1575 static void run_sve_tests(void)
1576 {
1577 	struct test_config test_config;
1578 	int i, j;
1579 
1580 	if (!sve_supported())
1581 		return;
1582 
1583 	test_config.sme_vl_in = sme_vls[0];
1584 	test_config.sme_vl_expected = sme_vls[0];
1585 	test_config.svcr_in = 0;
1586 	test_config.svcr_expected = 0;
1587 
1588 	for (i = 0; i < sve_vl_count; i++) {
1589 		test_config.sve_vl_in = sve_vls[i];
1590 
1591 		for (j = 0; j < sve_vl_count; j++) {
1592 			test_config.sve_vl_expected = sve_vls[j];
1593 
1594 			run_tests(base_test_defs,
1595 				  ARRAY_SIZE(base_test_defs),
1596 				  &test_config);
1597 			if (sve_supported())
1598 				run_tests(sve_test_defs,
1599 					  ARRAY_SIZE(sve_test_defs),
1600 					  &test_config);
1601 		}
1602 	}
1603 }
1604 
1605 static void run_sme_tests(void)
1606 {
1607 	struct test_config test_config;
1608 	int i, j, k;
1609 
1610 	if (!sme_supported())
1611 		return;
1612 
1613 	test_config.sve_vl_in = sve_vls[0];
1614 	test_config.sve_vl_expected = sve_vls[0];
1615 
1616 	/*
1617 	 * Every SME VL/SVCR combination
1618 	 */
1619 	for (i = 0; i < sme_vl_count; i++) {
1620 		test_config.sme_vl_in = sme_vls[i];
1621 
1622 		for (j = 0; j < sme_vl_count; j++) {
1623 			test_config.sme_vl_expected = sme_vls[j];
1624 
1625 			for (k = 0; k < ARRAY_SIZE(svcr_combinations); k++) {
1626 				test_config.svcr_in = svcr_combinations[k].svcr_in;
1627 				test_config.svcr_expected = svcr_combinations[k].svcr_expected;
1628 
1629 				run_tests(base_test_defs,
1630 					  ARRAY_SIZE(base_test_defs),
1631 					  &test_config);
1632 				run_tests(sve_test_defs,
1633 					  ARRAY_SIZE(sve_test_defs),
1634 					  &test_config);
1635 				run_tests(za_test_defs,
1636 					  ARRAY_SIZE(za_test_defs),
1637 					  &test_config);
1638 
1639 				if (sme2_supported())
1640 					run_tests(zt_test_defs,
1641 						  ARRAY_SIZE(zt_test_defs),
1642 						  &test_config);
1643 			}
1644 		}
1645 	}
1646 }
1647 
1648 int main(void)
1649 {
1650 	struct test_config test_config;
1651 	struct sigaction sa;
1652 	int tests, ret, tmp;
1653 
1654 	srandom(getpid());
1655 
1656 	ksft_print_header();
1657 
1658 	if (sve_supported()) {
1659 		probe_vls("SVE", sve_vls, &sve_vl_count, PR_SVE_SET_VL);
1660 
1661 		tests = ARRAY_SIZE(base_test_defs) +
1662 			ARRAY_SIZE(sve_test_defs);
1663 		tests *= sve_vl_count * sve_vl_count;
1664 	} else {
1665 		/* Only run the FPSIMD tests */
1666 		sve_vl_count = 1;
1667 		tests = ARRAY_SIZE(base_test_defs);
1668 	}
1669 
1670 	if (sme_supported()) {
1671 		probe_vls("SME", sme_vls, &sme_vl_count, PR_SME_SET_VL);
1672 
1673 		tmp = ARRAY_SIZE(base_test_defs) + ARRAY_SIZE(sve_test_defs)
1674 			+ ARRAY_SIZE(za_test_defs);
1675 
1676 		if (sme2_supported())
1677 			tmp += ARRAY_SIZE(zt_test_defs);
1678 
1679 		tmp *= sme_vl_count * sme_vl_count;
1680 		tmp *= ARRAY_SIZE(svcr_combinations);
1681 		tests += tmp;
1682 	} else {
1683 		sme_vl_count = 1;
1684 	}
1685 
1686 	if (sme2_supported())
1687 		ksft_print_msg("SME2 supported\n");
1688 
1689 	if (fa64_supported())
1690 		ksft_print_msg("FA64 supported\n");
1691 
1692 	if (fpmr_supported())
1693 		ksft_print_msg("FPMR supported\n");
1694 
1695 	ksft_set_plan(tests);
1696 
1697 	/* Get signal handers ready before we start any children */
1698 	memset(&sa, 0, sizeof(sa));
1699 	sa.sa_sigaction = handle_alarm;
1700 	sa.sa_flags = SA_RESTART | SA_SIGINFO;
1701 	sigemptyset(&sa.sa_mask);
1702 	ret = sigaction(SIGALRM, &sa, NULL);
1703 	if (ret < 0)
1704 		ksft_print_msg("Failed to install SIGALRM handler: %s (%d)\n",
1705 			       strerror(errno), errno);
1706 
1707 	/*
1708 	 * Run the test set if there is no SVE or SME, with those we
1709 	 * have to pick a VL for each run.
1710 	 */
1711 	if (!sve_supported() && !sme_supported()) {
1712 		test_config.sve_vl_in = 0;
1713 		test_config.sve_vl_expected = 0;
1714 		test_config.sme_vl_in = 0;
1715 		test_config.sme_vl_expected = 0;
1716 		test_config.svcr_in = 0;
1717 		test_config.svcr_expected = 0;
1718 
1719 		run_tests(base_test_defs, ARRAY_SIZE(base_test_defs),
1720 			  &test_config);
1721 	}
1722 
1723 	run_sve_tests();
1724 	run_sme_tests();
1725 
1726 	ksft_finished();
1727 }
1728