xref: /linux/tools/testing/selftests/net/tcp_ao/lib/sock.c (revision 1b78070aaef63512688aebfbc82365ef9d6660f1)
1 // SPDX-License-Identifier: GPL-2.0
2 #include <alloca.h>
3 #include <fcntl.h>
4 #include <inttypes.h>
5 #include <string.h>
6 #include "../../../../../include/linux/kernel.h"
7 #include "../../../../../include/linux/stringify.h"
8 #include "aolib.h"
9 
10 const unsigned int test_server_port = 7010;
11 int __test_listen_socket(int backlog, void *addr, size_t addr_sz)
12 {
13 	int err, sk = socket(test_family, SOCK_STREAM, IPPROTO_TCP);
14 	long flags;
15 
16 	if (sk < 0)
17 		test_error("socket()");
18 
19 	err = setsockopt(sk, SOL_SOCKET, SO_BINDTODEVICE, veth_name,
20 			 strlen(veth_name) + 1);
21 	if (err < 0)
22 		test_error("setsockopt(SO_BINDTODEVICE)");
23 
24 	if (bind(sk, (struct sockaddr *)addr, addr_sz) < 0)
25 		test_error("bind()");
26 
27 	flags = fcntl(sk, F_GETFL);
28 	if ((flags < 0) || (fcntl(sk, F_SETFL, flags | O_NONBLOCK) < 0))
29 		test_error("fcntl()");
30 
31 	if (listen(sk, backlog))
32 		test_error("listen()");
33 
34 	return sk;
35 }
36 
37 static int __test_wait_fd(int sk, struct timeval *tv, bool write)
38 {
39 	fd_set fds, efds;
40 	int ret;
41 	socklen_t slen = sizeof(ret);
42 
43 	FD_ZERO(&fds);
44 	FD_SET(sk, &fds);
45 	FD_ZERO(&efds);
46 	FD_SET(sk, &efds);
47 
48 	errno = 0;
49 	if (write)
50 		ret = select(sk + 1, NULL, &fds, &efds, tv);
51 	else
52 		ret = select(sk + 1, &fds, NULL, &efds, tv);
53 	if (ret < 0)
54 		return -errno;
55 	if (ret == 0) {
56 		errno = ETIMEDOUT;
57 		return -ETIMEDOUT;
58 	}
59 
60 	if (getsockopt(sk, SOL_SOCKET, SO_ERROR, &ret, &slen))
61 		return -errno;
62 	if (ret)
63 		return -ret;
64 	return 0;
65 }
66 
67 int test_wait_fd(int sk, time_t sec, bool write)
68 {
69 	struct timeval tv = { .tv_sec = sec, };
70 
71 	return __test_wait_fd(sk, sec ? &tv : NULL, write);
72 }
73 
74 static bool __skpair_poll_should_stop(int sk, struct tcp_counters *c,
75 				      test_cnt condition)
76 {
77 	struct tcp_counters c2;
78 	test_cnt diff;
79 
80 	if (test_get_tcp_counters(sk, &c2))
81 		test_error("test_get_tcp_counters()");
82 
83 	diff = test_cmp_counters(c, &c2);
84 	test_tcp_counters_free(&c2);
85 	return (diff & condition) == condition;
86 }
87 
88 /* How often wake up and check netns counters & paired (*err) */
89 #define POLL_USEC	150
90 static int __test_skpair_poll(int sk, bool write, uint64_t timeout,
91 			      struct tcp_counters *c, test_cnt cond,
92 			      volatile int *err)
93 {
94 	uint64_t t;
95 
96 	for (t = 0; t <= timeout * 1000000; t += POLL_USEC) {
97 		struct timeval tv = { .tv_usec = POLL_USEC, };
98 		int ret;
99 
100 		ret = __test_wait_fd(sk, &tv, write);
101 		if (ret != -ETIMEDOUT)
102 			return ret;
103 		if (c && cond && __skpair_poll_should_stop(sk, c, cond))
104 			break;
105 		if (err && *err)
106 			return *err;
107 	}
108 	if (err)
109 		*err = -ETIMEDOUT;
110 	return -ETIMEDOUT;
111 }
112 
113 int __test_connect_socket(int sk, const char *device,
114 			  void *addr, size_t addr_sz, bool async)
115 {
116 	long flags;
117 	int err;
118 
119 	if (device != NULL) {
120 		err = setsockopt(sk, SOL_SOCKET, SO_BINDTODEVICE, device,
121 				 strlen(device) + 1);
122 		if (err < 0)
123 			test_error("setsockopt(SO_BINDTODEVICE, %s)", device);
124 	}
125 
126 	flags = fcntl(sk, F_GETFL);
127 	if ((flags < 0) || (fcntl(sk, F_SETFL, flags | O_NONBLOCK) < 0))
128 		test_error("fcntl()");
129 
130 	if (connect(sk, addr, addr_sz) < 0) {
131 		if (errno != EINPROGRESS) {
132 			err = -errno;
133 			goto out;
134 		}
135 		if (async)
136 			return sk;
137 		err = test_wait_fd(sk, TEST_TIMEOUT_SEC, 1);
138 		if (err)
139 			goto out;
140 	}
141 	return sk;
142 
143 out:
144 	close(sk);
145 	return err;
146 }
147 
148 int test_skpair_wait_poll(int sk, bool write,
149 			  test_cnt cond, volatile int *err)
150 {
151 	struct tcp_counters c;
152 	int ret;
153 
154 	*err = 0;
155 	if (test_get_tcp_counters(sk, &c))
156 		test_error("test_get_tcp_counters()");
157 	synchronize_threads(); /* 1: init skpair & read nscounters */
158 
159 	ret = __test_skpair_poll(sk, write, TEST_TIMEOUT_SEC, &c, cond, err);
160 	test_tcp_counters_free(&c);
161 	return ret;
162 }
163 
164 int _test_skpair_connect_poll(int sk, const char *device,
165 			      void *addr, size_t addr_sz,
166 			      test_cnt condition, volatile int *err)
167 {
168 	struct tcp_counters c;
169 	int ret;
170 
171 	*err = 0;
172 	if (test_get_tcp_counters(sk, &c))
173 		test_error("test_get_tcp_counters()");
174 	synchronize_threads(); /* 1: init skpair & read nscounters */
175 	ret = __test_connect_socket(sk, device, addr, addr_sz, true);
176 	if (ret < 0) {
177 		test_tcp_counters_free(&c);
178 		return (*err = ret);
179 	}
180 	ret = __test_skpair_poll(sk, 1, TEST_TIMEOUT_SEC, &c, condition, err);
181 	if (ret < 0)
182 		close(sk);
183 	test_tcp_counters_free(&c);
184 	return ret;
185 }
186 
187 int __test_set_md5(int sk, void *addr, size_t addr_sz, uint8_t prefix,
188 		   int vrf, const char *password)
189 {
190 	size_t pwd_len = strlen(password);
191 	struct tcp_md5sig md5sig = {};
192 
193 	md5sig.tcpm_keylen = pwd_len;
194 	memcpy(md5sig.tcpm_key, password, pwd_len);
195 	md5sig.tcpm_flags = TCP_MD5SIG_FLAG_PREFIX;
196 	md5sig.tcpm_prefixlen = prefix;
197 	if (vrf >= 0) {
198 		md5sig.tcpm_flags |= TCP_MD5SIG_FLAG_IFINDEX;
199 		md5sig.tcpm_ifindex = (uint8_t)vrf;
200 	}
201 	memcpy(&md5sig.tcpm_addr, addr, addr_sz);
202 
203 	errno = 0;
204 	return setsockopt(sk, IPPROTO_TCP, TCP_MD5SIG_EXT,
205 			&md5sig, sizeof(md5sig));
206 }
207 
208 
209 int test_prepare_key_sockaddr(struct tcp_ao_add *ao, const char *alg,
210 		void *addr, size_t addr_sz, bool set_current, bool set_rnext,
211 		uint8_t prefix, uint8_t vrf, uint8_t sndid, uint8_t rcvid,
212 		uint8_t maclen, uint8_t keyflags,
213 		uint8_t keylen, const char *key)
214 {
215 	memset(ao, 0, sizeof(struct tcp_ao_add));
216 
217 	ao->set_current	= !!set_current;
218 	ao->set_rnext	= !!set_rnext;
219 	ao->prefix	= prefix;
220 	ao->sndid	= sndid;
221 	ao->rcvid	= rcvid;
222 	ao->maclen	= maclen;
223 	ao->keyflags	= keyflags;
224 	ao->keylen	= keylen;
225 	ao->ifindex	= vrf;
226 
227 	memcpy(&ao->addr, addr, addr_sz);
228 
229 	if (strlen(alg) > 64)
230 		return -ENOBUFS;
231 	strncpy(ao->alg_name, alg, 64);
232 
233 	memcpy(ao->key, key,
234 	       (keylen > TCP_AO_MAXKEYLEN) ? TCP_AO_MAXKEYLEN : keylen);
235 	return 0;
236 }
237 
238 static int test_get_ao_keys_nr(int sk)
239 {
240 	struct tcp_ao_getsockopt tmp = {};
241 	socklen_t tmp_sz = sizeof(tmp);
242 	int ret;
243 
244 	tmp.nkeys  = 1;
245 	tmp.get_all = 1;
246 
247 	ret = getsockopt(sk, IPPROTO_TCP, TCP_AO_GET_KEYS, &tmp, &tmp_sz);
248 	if (ret)
249 		return -errno;
250 	return (int)tmp.nkeys;
251 }
252 
253 int test_get_one_ao(int sk, struct tcp_ao_getsockopt *out,
254 		void *addr, size_t addr_sz, uint8_t prefix,
255 		uint8_t sndid, uint8_t rcvid, uint8_t keyflags, int ifindex)
256 {
257 	struct tcp_ao_getsockopt tmp = {};
258 	socklen_t tmp_sz = sizeof(tmp);
259 	int ret;
260 
261 	memcpy(&tmp.addr, addr, addr_sz);
262 	tmp.prefix = prefix;
263 	tmp.sndid  = sndid;
264 	tmp.rcvid  = rcvid;
265 	tmp.keyflags = keyflags;
266 	tmp.ifindex = ifindex;
267 	tmp.nkeys  = 1;
268 
269 	ret = getsockopt(sk, IPPROTO_TCP, TCP_AO_GET_KEYS, &tmp, &tmp_sz);
270 	if (ret)
271 		return ret;
272 	if (tmp.nkeys != 1)
273 		return -E2BIG;
274 	*out = tmp;
275 	return 0;
276 }
277 
278 int test_get_ao_info(int sk, struct tcp_ao_info_opt *out)
279 {
280 	socklen_t sz = sizeof(*out);
281 
282 	out->reserved = 0;
283 	out->reserved2 = 0;
284 	if (getsockopt(sk, IPPROTO_TCP, TCP_AO_INFO, out, &sz))
285 		return -errno;
286 	if (sz != sizeof(*out))
287 		return -EMSGSIZE;
288 	return 0;
289 }
290 
291 int test_set_ao_info(int sk, struct tcp_ao_info_opt *in)
292 {
293 	socklen_t sz = sizeof(*in);
294 
295 	in->reserved = 0;
296 	in->reserved2 = 0;
297 	if (setsockopt(sk, IPPROTO_TCP, TCP_AO_INFO, in, sz))
298 		return -errno;
299 	return 0;
300 }
301 
302 int test_cmp_getsockopt_setsockopt(const struct tcp_ao_add *a,
303 				   const struct tcp_ao_getsockopt *b)
304 {
305 	bool is_kdf_aes_128_cmac = false;
306 	bool is_cmac_aes = false;
307 
308 	if (!strcmp("cmac(aes128)", a->alg_name)) {
309 		is_kdf_aes_128_cmac = (a->keylen != 16);
310 		is_cmac_aes = true;
311 	}
312 
313 #define __cmp_ao(member)						\
314 do {									\
315 	if (b->member != a->member) {					\
316 		test_fail("getsockopt(): " __stringify(member) " %u != %u",	\
317 				b->member, a->member);			\
318 		return -1;						\
319 	}								\
320 } while(0)
321 	__cmp_ao(sndid);
322 	__cmp_ao(rcvid);
323 	__cmp_ao(prefix);
324 	__cmp_ao(keyflags);
325 	__cmp_ao(ifindex);
326 	if (a->maclen) {
327 		__cmp_ao(maclen);
328 	} else if (b->maclen != 12) {
329 		test_fail("getsockopt(): expected default maclen 12, but it's %u",
330 				b->maclen);
331 		return -1;
332 	}
333 	if (!is_kdf_aes_128_cmac) {
334 		__cmp_ao(keylen);
335 	} else if (b->keylen != 16) {
336 		test_fail("getsockopt(): expected keylen 16 for cmac(aes128), but it's %u",
337 				b->keylen);
338 		return -1;
339 	}
340 #undef __cmp_ao
341 	if (!is_kdf_aes_128_cmac && memcmp(b->key, a->key, a->keylen)) {
342 		test_fail("getsockopt(): returned key is different `%s' != `%s'",
343 				b->key, a->key);
344 		return -1;
345 	}
346 	if (memcmp(&b->addr, &a->addr, sizeof(b->addr))) {
347 		test_fail("getsockopt(): returned address is different");
348 		return -1;
349 	}
350 	if (!is_cmac_aes && strcmp(b->alg_name, a->alg_name)) {
351 		test_fail("getsockopt(): returned algorithm %s is different than %s", b->alg_name, a->alg_name);
352 		return -1;
353 	}
354 	if (is_cmac_aes && strcmp(b->alg_name, "cmac(aes)")) {
355 		test_fail("getsockopt(): returned algorithm %s is different than cmac(aes)", b->alg_name);
356 		return -1;
357 	}
358 	/* For a established key rotation test don't add a key with
359 	 * set_current = 1, as it's likely to change by peer's request;
360 	 * rather use setsockopt(TCP_AO_INFO)
361 	 */
362 	if (a->set_current != b->is_current) {
363 		test_fail("getsockopt(): returned key is not Current_key");
364 		return -1;
365 	}
366 	if (a->set_rnext != b->is_rnext) {
367 		test_fail("getsockopt(): returned key is not RNext_key");
368 		return -1;
369 	}
370 
371 	return 0;
372 }
373 
374 int test_cmp_getsockopt_setsockopt_ao(const struct tcp_ao_info_opt *a,
375 				      const struct tcp_ao_info_opt *b)
376 {
377 	/* No check for ::current_key, as it may change by the peer */
378 	if (a->ao_required != b->ao_required) {
379 		test_fail("getsockopt(): returned ao doesn't have ao_required");
380 		return -1;
381 	}
382 	if (a->accept_icmps != b->accept_icmps) {
383 		test_fail("getsockopt(): returned ao doesn't accept ICMPs");
384 		return -1;
385 	}
386 	if (a->set_rnext && a->rnext != b->rnext) {
387 		test_fail("getsockopt(): RNext KeyID has changed");
388 		return -1;
389 	}
390 #define __cmp_cnt(member)						\
391 do {									\
392 	if (b->member != a->member) {					\
393 		test_fail("getsockopt(): " __stringify(member) " %llu != %llu",	\
394 				b->member, a->member);			\
395 		return -1;						\
396 	}								\
397 } while(0)
398 	if (a->set_counters) {
399 		__cmp_cnt(pkt_good);
400 		__cmp_cnt(pkt_bad);
401 		__cmp_cnt(pkt_key_not_found);
402 		__cmp_cnt(pkt_ao_required);
403 		__cmp_cnt(pkt_dropped_icmp);
404 	}
405 #undef __cmp_cnt
406 	return 0;
407 }
408 
409 int test_get_tcp_counters(int sk, struct tcp_counters *out)
410 {
411 	struct tcp_ao_getsockopt *key_dump;
412 	socklen_t key_dump_sz = sizeof(*key_dump);
413 	struct tcp_ao_info_opt info = {};
414 	bool c1, c2, c3, c4, c5, c6, c7, c8;
415 	struct netstat *ns;
416 	int err, nr_keys;
417 
418 	memset(out, 0, sizeof(*out));
419 
420 	/* per-netns */
421 	ns = netstat_read();
422 	out->ao.netns_ao_good = netstat_get(ns, "TCPAOGood", &c1);
423 	out->ao.netns_ao_bad = netstat_get(ns, "TCPAOBad", &c2);
424 	out->ao.netns_ao_key_not_found = netstat_get(ns, "TCPAOKeyNotFound", &c3);
425 	out->ao.netns_ao_required = netstat_get(ns, "TCPAORequired", &c4);
426 	out->ao.netns_ao_dropped_icmp = netstat_get(ns, "TCPAODroppedIcmps", &c5);
427 	out->netns_md5_notfound = netstat_get(ns, "TCPMD5NotFound", &c6);
428 	out->netns_md5_unexpected = netstat_get(ns, "TCPMD5Unexpected", &c7);
429 	out->netns_md5_failure = netstat_get(ns, "TCPMD5Failure", &c8);
430 	netstat_free(ns);
431 	if (c1 || c2 || c3 || c4 || c5 || c6 || c7 || c8)
432 		return -EOPNOTSUPP;
433 
434 	err = test_get_ao_info(sk, &info);
435 	if (err == -ENOENT)
436 		return 0;
437 	if (err)
438 		return err;
439 
440 	/* per-socket */
441 	out->ao.ao_info_pkt_good = info.pkt_good;
442 	out->ao.ao_info_pkt_bad = info.pkt_bad;
443 	out->ao.ao_info_pkt_key_not_found = info.pkt_key_not_found;
444 	out->ao.ao_info_pkt_ao_required = info.pkt_ao_required;
445 	out->ao.ao_info_pkt_dropped_icmp = info.pkt_dropped_icmp;
446 
447 	/* per-key */
448 	nr_keys = test_get_ao_keys_nr(sk);
449 	if (nr_keys < 0)
450 		return nr_keys;
451 	if (nr_keys == 0)
452 		test_error("test_get_ao_keys_nr() == 0");
453 	out->ao.nr_keys = (size_t)nr_keys;
454 	key_dump = calloc(nr_keys, key_dump_sz);
455 	if (!key_dump)
456 		return -errno;
457 
458 	key_dump[0].nkeys = nr_keys;
459 	key_dump[0].get_all = 1;
460 	err = getsockopt(sk, IPPROTO_TCP, TCP_AO_GET_KEYS,
461 			 key_dump, &key_dump_sz);
462 	if (err) {
463 		free(key_dump);
464 		return -errno;
465 	}
466 
467 	out->ao.key_cnts = calloc(nr_keys, sizeof(out->ao.key_cnts[0]));
468 	if (!out->ao.key_cnts) {
469 		free(key_dump);
470 		return -errno;
471 	}
472 
473 	while (nr_keys--) {
474 		out->ao.key_cnts[nr_keys].sndid = key_dump[nr_keys].sndid;
475 		out->ao.key_cnts[nr_keys].rcvid = key_dump[nr_keys].rcvid;
476 		out->ao.key_cnts[nr_keys].pkt_good = key_dump[nr_keys].pkt_good;
477 		out->ao.key_cnts[nr_keys].pkt_bad = key_dump[nr_keys].pkt_bad;
478 	}
479 	free(key_dump);
480 
481 	return 0;
482 }
483 
484 test_cnt test_cmp_counters(struct tcp_counters *before,
485 			   struct tcp_counters *after)
486 {
487 #define __cmp(cnt, e_cnt)						\
488 do {									\
489 	if (before->cnt > after->cnt)					\
490 		test_error("counter " __stringify(cnt) " decreased");	\
491 	if (before->cnt != after->cnt)					\
492 		ret |= e_cnt;						\
493 } while (0)
494 
495 	test_cnt ret = 0;
496 	size_t i;
497 
498 	if (before->ao.nr_keys != after->ao.nr_keys)
499 		test_error("the number of keys has changed");
500 
501 	_for_each_counter(__cmp);
502 
503 	i = before->ao.nr_keys;
504 	while (i--) {
505 		__cmp(ao.key_cnts[i].pkt_good, TEST_CNT_KEY_GOOD);
506 		__cmp(ao.key_cnts[i].pkt_bad, TEST_CNT_KEY_BAD);
507 	}
508 #undef __cmp
509 	return ret;
510 }
511 
512 int test_assert_counters_sk(const char *tst_name,
513 			    struct tcp_counters *before,
514 			    struct tcp_counters *after,
515 			    test_cnt expected)
516 {
517 #define __cmp_ao(cnt, e_cnt)						\
518 do {									\
519 	if (before->cnt > after->cnt) {					\
520 		test_fail("%s: Decreased counter " __stringify(cnt) " %" PRIu64 " > %" PRIu64, \
521 			  tst_name ?: "", before->cnt, after->cnt);	\
522 		return -1;						\
523 	}								\
524 	if ((before->cnt != after->cnt) != !!(expected & e_cnt)) {	\
525 		test_fail("%s: Counter " __stringify(cnt) " was %sexpected to increase %" PRIu64 " => %" PRIu64, \
526 			  tst_name ?: "", (expected & e_cnt) ? "" : "not ",	\
527 			  before->cnt, after->cnt);			\
528 		return -1;						\
529 	}								\
530 } while (0)
531 
532 	errno = 0;
533 	_for_each_counter(__cmp_ao);
534 	return 0;
535 #undef __cmp_ao
536 }
537 
538 int test_assert_counters_key(const char *tst_name,
539 			     struct tcp_ao_counters *before,
540 			     struct tcp_ao_counters *after,
541 			     test_cnt expected, int sndid, int rcvid)
542 {
543 	size_t i;
544 #define __cmp_ao(i, cnt, e_cnt)					\
545 do {									\
546 	if (before->key_cnts[i].cnt > after->key_cnts[i].cnt) {		\
547 		test_fail("%s: Decreased counter " __stringify(cnt) " %" PRIu64 " > %" PRIu64 " for key %u:%u", \
548 			  tst_name ?: "", before->key_cnts[i].cnt,	\
549 			  after->key_cnts[i].cnt,			\
550 			  before->key_cnts[i].sndid,			\
551 			  before->key_cnts[i].rcvid);			\
552 		return -1;						\
553 	}								\
554 	if ((before->key_cnts[i].cnt != after->key_cnts[i].cnt) != !!(expected & e_cnt)) {		\
555 		test_fail("%s: Counter " __stringify(cnt) " was %sexpected to increase %" PRIu64 " => %" PRIu64 " for key %u:%u", \
556 			  tst_name ?: "", (expected & e_cnt) ? "" : "not ",\
557 			  before->key_cnts[i].cnt,			\
558 			  after->key_cnts[i].cnt,			\
559 			  before->key_cnts[i].sndid,			\
560 			  before->key_cnts[i].rcvid);			\
561 		return -1;						\
562 	}								\
563 } while (0)
564 
565 	if (before->nr_keys != after->nr_keys) {
566 		test_fail("%s: Keys changed on the socket %zu != %zu",
567 			  tst_name, before->nr_keys, after->nr_keys);
568 		return -1;
569 	}
570 
571 	/* per-key */
572 	i = before->nr_keys;
573 	while (i--) {
574 		if (sndid >= 0 && before->key_cnts[i].sndid != sndid)
575 			continue;
576 		if (rcvid >= 0 && before->key_cnts[i].rcvid != rcvid)
577 			continue;
578 		__cmp_ao(i, pkt_good, TEST_CNT_KEY_GOOD);
579 		__cmp_ao(i, pkt_bad, TEST_CNT_KEY_BAD);
580 	}
581 	return 0;
582 #undef __cmp_ao
583 }
584 
585 void test_tcp_counters_free(struct tcp_counters *cnts)
586 {
587 	free(cnts->ao.key_cnts);
588 }
589 
590 #define TEST_BUF_SIZE 4096
591 static ssize_t _test_server_run(int sk, ssize_t quota, struct tcp_counters *c,
592 				test_cnt cond, volatile int *err,
593 				time_t timeout_sec)
594 {
595 	ssize_t total = 0;
596 
597 	do {
598 		char buf[TEST_BUF_SIZE];
599 		ssize_t bytes, sent;
600 		int ret;
601 
602 		ret = __test_skpair_poll(sk, 0, timeout_sec, c, cond, err);
603 		if (ret)
604 			return ret;
605 
606 		bytes = recv(sk, buf, sizeof(buf), 0);
607 
608 		if (bytes < 0)
609 			test_error("recv(): %zd", bytes);
610 		if (bytes == 0)
611 			break;
612 
613 		ret = __test_skpair_poll(sk, 1, timeout_sec, c, cond, err);
614 		if (ret)
615 			return ret;
616 
617 		sent = send(sk, buf, bytes, 0);
618 		if (sent == 0)
619 			break;
620 		if (sent != bytes)
621 			test_error("send()");
622 		total += bytes;
623 	} while (!quota || total < quota);
624 
625 	return total;
626 }
627 
628 ssize_t test_server_run(int sk, ssize_t quota, time_t timeout_sec)
629 {
630 	return _test_server_run(sk, quota, NULL, 0, NULL,
631 				timeout_sec ?: TEST_TIMEOUT_SEC);
632 }
633 
634 int test_skpair_server(int sk, ssize_t quota, test_cnt cond, volatile int *err)
635 {
636 	struct tcp_counters c;
637 	ssize_t ret;
638 
639 	*err = 0;
640 	if (test_get_tcp_counters(sk, &c))
641 		test_error("test_get_tcp_counters()");
642 	synchronize_threads(); /* 1: init skpair & read nscounters */
643 
644 	ret = _test_server_run(sk, quota, &c, cond, err, TEST_TIMEOUT_SEC);
645 	test_tcp_counters_free(&c);
646 	return ret;
647 }
648 
649 static ssize_t test_client_loop(int sk, size_t buf_sz, const size_t msg_len,
650 				struct tcp_counters *c, test_cnt cond,
651 				volatile int *err)
652 {
653 	char msg[msg_len];
654 	int nodelay = 1;
655 	char *buf;
656 	size_t i;
657 
658 	buf = alloca(buf_sz);
659 	if (!buf)
660 		return -ENOMEM;
661 	randomize_buffer(buf, buf_sz);
662 
663 	if (setsockopt(sk, IPPROTO_TCP, TCP_NODELAY, &nodelay, sizeof(nodelay)))
664 		test_error("setsockopt(TCP_NODELAY)");
665 
666 	for (i = 0; i < buf_sz; i += min(msg_len, buf_sz - i)) {
667 		size_t sent, bytes = min(msg_len, buf_sz - i);
668 		int ret;
669 
670 		ret = __test_skpair_poll(sk, 1, TEST_TIMEOUT_SEC, c, cond, err);
671 		if (ret)
672 			return ret;
673 
674 		sent = send(sk, buf + i, bytes, 0);
675 		if (sent == 0)
676 			break;
677 		if (sent != bytes)
678 			test_error("send()");
679 
680 		bytes = 0;
681 		do {
682 			ssize_t got;
683 
684 			ret = __test_skpair_poll(sk, 0, TEST_TIMEOUT_SEC,
685 						 c, cond, err);
686 			if (ret)
687 				return ret;
688 
689 			got = recv(sk, msg + bytes, sizeof(msg) - bytes, 0);
690 			if (got <= 0)
691 				return i;
692 			bytes += got;
693 		} while (bytes < sent);
694 		if (bytes > sent)
695 			test_error("recv(): %zd > %zd", bytes, sent);
696 		if (memcmp(buf + i, msg, bytes) != 0) {
697 			test_fail("received message differs");
698 			return -1;
699 		}
700 	}
701 	return i;
702 }
703 
704 int test_client_verify(int sk, const size_t msg_len, const size_t nr)
705 {
706 	size_t buf_sz = msg_len * nr;
707 	ssize_t ret;
708 
709 	ret = test_client_loop(sk, buf_sz, msg_len, NULL, 0, NULL);
710 	if (ret < 0)
711 		return (int)ret;
712 	return ret != buf_sz ? -1 : 0;
713 }
714 
715 int test_skpair_client(int sk, const size_t msg_len, const size_t nr,
716 		       test_cnt cond, volatile int *err)
717 {
718 	struct tcp_counters c;
719 	size_t buf_sz = msg_len * nr;
720 	ssize_t ret;
721 
722 	*err = 0;
723 	if (test_get_tcp_counters(sk, &c))
724 		test_error("test_get_tcp_counters()");
725 	synchronize_threads(); /* 1: init skpair & read nscounters */
726 
727 	ret = test_client_loop(sk, buf_sz, msg_len, &c, cond, err);
728 	test_tcp_counters_free(&c);
729 	if (ret < 0)
730 		return (int)ret;
731 	return ret != buf_sz ? -1 : 0;
732 }
733