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;
__test_listen_socket(int backlog,void * addr,size_t addr_sz)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
__test_wait_fd(int sk,struct timeval * tv,bool write)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
test_wait_fd(int sk,time_t sec,bool write)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
__skpair_poll_should_stop(int sk,struct tcp_counters * c,test_cnt condition)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
__test_skpair_poll(int sk,bool write,uint64_t timeout,struct tcp_counters * c,test_cnt cond,volatile int * err)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
__test_connect_socket(int sk,const char * device,void * addr,size_t addr_sz,bool async)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
test_skpair_wait_poll(int sk,bool write,test_cnt cond,volatile int * err)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
_test_skpair_connect_poll(int sk,const char * device,void * addr,size_t addr_sz,test_cnt condition,volatile int * err)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
__test_set_md5(int sk,void * addr,size_t addr_sz,uint8_t prefix,int vrf,const char * password)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
test_prepare_key_sockaddr(struct tcp_ao_add * ao,const char * alg,void * addr,size_t addr_sz,bool set_current,bool set_rnext,uint8_t prefix,uint8_t vrf,uint8_t sndid,uint8_t rcvid,uint8_t maclen,uint8_t keyflags,uint8_t keylen,const char * key)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
test_get_ao_keys_nr(int sk)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
test_get_one_ao(int sk,struct tcp_ao_getsockopt * out,void * addr,size_t addr_sz,uint8_t prefix,uint8_t sndid,uint8_t rcvid,uint8_t keyflags,int ifindex)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
test_get_ao_info(int sk,struct tcp_ao_info_opt * out)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
test_set_ao_info(int sk,struct tcp_ao_info_opt * in)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
test_cmp_getsockopt_setsockopt(const struct tcp_ao_add * a,const struct tcp_ao_getsockopt * b)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
test_cmp_getsockopt_setsockopt_ao(const struct tcp_ao_info_opt * a,const struct tcp_ao_info_opt * b)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
test_get_tcp_counters(int sk,struct tcp_counters * out)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
test_cmp_counters(struct tcp_counters * before,struct tcp_counters * after)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
test_assert_counters_sk(const char * tst_name,struct tcp_counters * before,struct tcp_counters * after,test_cnt expected)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
test_assert_counters_key(const char * tst_name,struct tcp_ao_counters * before,struct tcp_ao_counters * after,test_cnt expected,int sndid,int rcvid)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
test_tcp_counters_free(struct tcp_counters * cnts)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
_test_server_run(int sk,ssize_t quota,struct tcp_counters * c,test_cnt cond,volatile int * err,time_t timeout_sec)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
test_server_run(int sk,ssize_t quota,time_t timeout_sec)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
test_skpair_server(int sk,ssize_t quota,test_cnt cond,volatile int * err)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
test_client_loop(int sk,size_t buf_sz,const size_t msg_len,struct tcp_counters * c,test_cnt cond,volatile int * err)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
test_client_verify(int sk,const size_t msg_len,const size_t nr)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
test_skpair_client(int sk,const size_t msg_len,const size_t nr,test_cnt cond,volatile int * err)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