1 // SPDX-License-Identifier: GPL-2.0 2 3 #define _GNU_SOURCE 4 5 #include <arpa/inet.h> 6 #include <errno.h> 7 #include <error.h> 8 #include <fcntl.h> 9 #include <limits.h> 10 #include <linux/filter.h> 11 #include <linux/bpf.h> 12 #include <linux/if_packet.h> 13 #include <linux/if_vlan.h> 14 #include <linux/virtio_net.h> 15 #include <net/if.h> 16 #include <net/ethernet.h> 17 #include <netinet/ip.h> 18 #include <netinet/udp.h> 19 #include <poll.h> 20 #include <sched.h> 21 #include <stdbool.h> 22 #include <stdint.h> 23 #include <stdio.h> 24 #include <stdlib.h> 25 #include <string.h> 26 #include <sys/mman.h> 27 #include <sys/socket.h> 28 #include <sys/stat.h> 29 #include <sys/types.h> 30 #include <unistd.h> 31 32 #include "psock_lib.h" 33 34 static bool cfg_use_bind; 35 static bool cfg_use_csum_off; 36 static bool cfg_use_csum_off_bad; 37 static bool cfg_use_dgram; 38 static bool cfg_use_gso; 39 static bool cfg_use_qdisc_bypass; 40 static bool cfg_use_vlan; 41 static bool cfg_use_vnet; 42 static bool cfg_drop; 43 static bool cfg_aux_data; 44 static bool cfg_ignore_outgoing; 45 46 static char *cfg_ifname = "lo"; 47 static int cfg_mtu = 1500; 48 static int cfg_payload_len = DATA_LEN; 49 static int cfg_truncate_len = INT_MAX; 50 static uint16_t cfg_port = 8000; 51 52 /* test sending up to max mtu + 1 */ 53 #define TEST_SZ (sizeof(struct virtio_net_hdr) + ETH_HLEN + ETH_MAX_MTU + 1) 54 55 #define BURST_CNT (1000) 56 57 static char tbuf[TEST_SZ], rbuf[TEST_SZ]; 58 59 static unsigned long add_csum_hword(const uint16_t *start, int num_u16) 60 { 61 unsigned long sum = 0; 62 int i; 63 64 for (i = 0; i < num_u16; i++) 65 sum += start[i]; 66 67 return sum; 68 } 69 70 static uint16_t build_ip_csum(const uint16_t *start, int num_u16, 71 unsigned long sum) 72 { 73 sum += add_csum_hword(start, num_u16); 74 75 while (sum >> 16) 76 sum = (sum & 0xffff) + (sum >> 16); 77 78 return ~sum; 79 } 80 81 static int build_vnet_header(void *header) 82 { 83 struct virtio_net_hdr *vh = header; 84 85 vh->hdr_len = ETH_HLEN + sizeof(struct iphdr) + sizeof(struct udphdr); 86 87 if (cfg_use_csum_off) { 88 vh->flags |= VIRTIO_NET_HDR_F_NEEDS_CSUM; 89 vh->csum_start = ETH_HLEN + sizeof(struct iphdr); 90 vh->csum_offset = __builtin_offsetof(struct udphdr, check); 91 92 /* position check field exactly one byte beyond end of packet */ 93 if (cfg_use_csum_off_bad) 94 vh->csum_start += sizeof(struct udphdr) + cfg_payload_len - 95 vh->csum_offset - 1; 96 } 97 98 if (cfg_use_gso) { 99 vh->gso_type = VIRTIO_NET_HDR_GSO_UDP; 100 vh->gso_size = cfg_mtu - sizeof(struct iphdr); 101 } 102 103 return sizeof(*vh); 104 } 105 106 static int build_eth_header(void *header) 107 { 108 struct ethhdr *eth = header; 109 110 if (cfg_use_vlan) { 111 uint16_t *tag = header + ETH_HLEN; 112 113 eth->h_proto = htons(ETH_P_8021Q); 114 tag[1] = htons(ETH_P_IP); 115 return ETH_HLEN + 4; 116 } 117 118 eth->h_proto = htons(ETH_P_IP); 119 return ETH_HLEN; 120 } 121 122 static int build_ipv4_header(void *header, int payload_len) 123 { 124 struct iphdr *iph = header; 125 126 iph->ihl = 5; 127 iph->version = 4; 128 iph->ttl = 8; 129 iph->tot_len = htons(sizeof(*iph) + sizeof(struct udphdr) + payload_len); 130 iph->id = htons(1337); 131 iph->protocol = IPPROTO_UDP; 132 iph->saddr = htonl((172 << 24) | (17 << 16) | 2); 133 iph->daddr = htonl((172 << 24) | (17 << 16) | 1); 134 iph->check = build_ip_csum((void *) iph, iph->ihl << 1, 0); 135 136 return iph->ihl << 2; 137 } 138 139 static int build_udp_header(void *header, int payload_len) 140 { 141 const int alen = sizeof(uint32_t); 142 struct udphdr *udph = header; 143 int len = sizeof(*udph) + payload_len; 144 145 udph->source = htons(9); 146 udph->dest = htons(cfg_port); 147 udph->len = htons(len); 148 149 if (cfg_use_csum_off) 150 udph->check = build_ip_csum(header - (2 * alen), alen, 151 htons(IPPROTO_UDP) + udph->len); 152 else 153 udph->check = 0; 154 155 return sizeof(*udph); 156 } 157 158 static int build_packet(int payload_len) 159 { 160 int off = 0; 161 162 off += build_vnet_header(tbuf); 163 off += build_eth_header(tbuf + off); 164 off += build_ipv4_header(tbuf + off, payload_len); 165 off += build_udp_header(tbuf + off, payload_len); 166 167 if (off + payload_len > sizeof(tbuf)) 168 error(1, 0, "payload length exceeds max"); 169 170 memset(tbuf + off, DATA_CHAR, payload_len); 171 172 return off + payload_len; 173 } 174 175 static void do_bind_proto(int fd, uint16_t proto) 176 { 177 struct sockaddr_ll laddr = {0}; 178 179 laddr.sll_family = AF_PACKET; 180 laddr.sll_protocol = htons(proto); 181 laddr.sll_ifindex = if_nametoindex(cfg_ifname); 182 if (!laddr.sll_ifindex) 183 error(1, errno, "if_nametoindex"); 184 185 if (bind(fd, (void *)&laddr, sizeof(laddr))) 186 error(1, errno, "bind"); 187 } 188 189 static void do_bind(int fd) 190 { 191 do_bind_proto(fd, ETH_P_IP); 192 } 193 194 static void do_send(int fd, char *buf, int len) 195 { 196 int ret; 197 198 if (!cfg_use_vnet) { 199 buf += sizeof(struct virtio_net_hdr); 200 len -= sizeof(struct virtio_net_hdr); 201 } 202 if (cfg_use_dgram) { 203 buf += ETH_HLEN; 204 len -= ETH_HLEN; 205 } 206 207 if (cfg_use_bind) { 208 ret = write(fd, buf, len); 209 } else { 210 struct sockaddr_ll laddr = {0}; 211 212 laddr.sll_protocol = htons(ETH_P_IP); 213 laddr.sll_ifindex = if_nametoindex(cfg_ifname); 214 if (!laddr.sll_ifindex) 215 error(1, errno, "if_nametoindex"); 216 217 ret = sendto(fd, buf, len, 0, (void *)&laddr, sizeof(laddr)); 218 } 219 220 if (ret == -1) 221 error(1, errno, "write"); 222 if (ret != len) 223 error(1, 0, "write: %u %u", ret, len); 224 225 if (!cfg_drop) 226 fprintf(stderr, "tx: %u\n", ret); 227 } 228 229 static int do_tx(void) 230 { 231 const int one = 1; 232 int i, fd, len; 233 234 fd = socket(PF_PACKET, cfg_use_dgram ? SOCK_DGRAM : SOCK_RAW, 0); 235 if (fd == -1) 236 error(1, errno, "socket t"); 237 238 if (cfg_use_bind) 239 do_bind(fd); 240 241 if (cfg_use_qdisc_bypass && 242 setsockopt(fd, SOL_PACKET, PACKET_QDISC_BYPASS, &one, sizeof(one))) 243 error(1, errno, "setsockopt qdisc bypass"); 244 245 if (cfg_use_vnet && 246 setsockopt(fd, SOL_PACKET, PACKET_VNET_HDR, &one, sizeof(one))) 247 error(1, errno, "setsockopt vnet"); 248 249 len = build_packet(cfg_payload_len); 250 251 if (cfg_truncate_len < len) 252 len = cfg_truncate_len; 253 254 do_send(fd, tbuf, len); 255 256 if (cfg_drop) 257 for (i = 0; i < BURST_CNT; i++) 258 do_send(fd, tbuf, len); 259 260 if (close(fd)) 261 error(1, errno, "close t"); 262 263 return len; 264 } 265 266 static int setup_rx(void) 267 { 268 struct timeval tv = { .tv_usec = 100 * 1000 }; 269 struct sockaddr_in raddr = {0}; 270 int fd; 271 272 fd = socket(PF_INET, SOCK_DGRAM, 0); 273 if (fd == -1) 274 error(1, errno, "socket r"); 275 276 if (setsockopt(fd, SOL_SOCKET, SO_RCVTIMEO, &tv, sizeof(tv))) 277 error(1, errno, "setsockopt rcv timeout"); 278 279 raddr.sin_family = AF_INET; 280 raddr.sin_port = htons(cfg_port); 281 raddr.sin_addr.s_addr = htonl(INADDR_ANY); 282 283 if (bind(fd, (void *)&raddr, sizeof(raddr))) 284 error(1, errno, "bind r"); 285 286 return fd; 287 } 288 289 static void check_aux_data(struct cmsghdr *cmsg, int expected_len) 290 { 291 struct tpacket_auxdata *adata; 292 293 if (!cmsg) 294 error(1, 0, "auxdata null"); 295 296 if (cmsg->cmsg_level != SOL_PACKET) 297 error(1, 0, "cmsg_level != SOL_PACKET"); 298 299 if (cmsg->cmsg_type != PACKET_AUXDATA) 300 error(1, 0, "cmsg_type != PACKET_AUXDATA"); 301 302 adata = (struct tpacket_auxdata *)CMSG_DATA(cmsg); 303 304 if (adata->tp_net != ETH_HLEN) 305 error(1, 0, "cmsg tp_net != ETH_HLEN"); 306 307 if (adata->tp_len != expected_len) 308 error(1, 0, "cmsg tp_len != %u", expected_len); 309 310 if (adata->tp_snaplen != expected_len) 311 error(1, 0, "cmsg tp_snaplen != %u", expected_len); 312 } 313 314 /* expected_pkttype < 0 skips the sll_pkttype check. */ 315 static void do_rx(int fd, int expected_len, char *expected, bool is_psock, 316 int expected_pkttype) 317 { 318 char cmsg_buf[1024] __attribute__((aligned(8))) = {}; 319 bool aux = is_psock && cfg_aux_data; 320 struct sockaddr_ll saddr = {}; 321 struct iovec iov = { 322 .iov_base = rbuf, 323 .iov_len = sizeof(rbuf), 324 }; 325 struct msghdr msg = { 326 .msg_iov = &iov, 327 .msg_iovlen = 1, 328 }; 329 int ret; 330 331 if (aux) { 332 msg.msg_control = cmsg_buf; 333 msg.msg_controllen = sizeof(cmsg_buf); 334 } 335 if (is_psock) { 336 msg.msg_name = &saddr; 337 msg.msg_namelen = sizeof(saddr); 338 } 339 340 ret = recvmsg(fd, &msg, 0); 341 if (ret == -1) 342 error(1, errno, "recv"); 343 if (ret != expected_len) 344 error(1, 0, "recv: %u != %u", ret, expected_len); 345 346 if (memcmp(rbuf, expected, ret)) 347 error(1, 0, "recv: data mismatch"); 348 349 if (aux) 350 check_aux_data(CMSG_FIRSTHDR(&msg), expected_len); 351 352 if (expected_pkttype >= 0 && saddr.sll_pkttype != expected_pkttype) 353 error(1, 0, "recv: sll_pkttype %d != %d", 354 saddr.sll_pkttype, expected_pkttype); 355 356 fprintf(stderr, "rx: %u\n", ret); 357 } 358 359 static int setup_sniffer(void) 360 { 361 struct timeval tv = { .tv_usec = 100 * 1000 }; 362 const int one = 1; 363 int fd; 364 365 fd = socket(PF_PACKET, SOCK_RAW, 0); 366 if (fd == -1) 367 error(1, errno, "socket p"); 368 369 if (setsockopt(fd, SOL_SOCKET, SO_RCVTIMEO, &tv, sizeof(tv))) 370 error(1, errno, "setsockopt rcv timeout"); 371 372 if (cfg_drop) 373 if (setsockopt(fd, SOL_SOCKET, SO_RCVBUF, &one, sizeof(one))) 374 error(1, errno, "setsockopt SO_RCVBUF"); 375 376 if (cfg_aux_data) 377 if (setsockopt(fd, SOL_PACKET, PACKET_AUXDATA, &one, sizeof(one))) 378 error(1, errno, "setsockopt PACKET_AUXDATA"); 379 380 pair_udp_setfilter(fd); 381 382 /* binding to ETH_P_ALL adds the sniffer to ptype_all, which will see 383 * the dev_queue_xmit_nit copy. ignore_outgoing should suppress this. 384 */ 385 if (cfg_ignore_outgoing) 386 do_bind_proto(fd, ETH_P_ALL); 387 else 388 do_bind(fd); 389 390 return fd; 391 } 392 393 static void parse_opts(int argc, char **argv) 394 { 395 int c; 396 397 while ((c = getopt(argc, argv, "abcCdDgil:qt:vV")) != -1) { 398 switch (c) { 399 case 'a': 400 cfg_aux_data = true; 401 break; 402 case 'b': 403 cfg_use_bind = true; 404 break; 405 case 'c': 406 cfg_use_csum_off = true; 407 break; 408 case 'C': 409 cfg_use_csum_off_bad = true; 410 break; 411 case 'd': 412 cfg_use_dgram = true; 413 break; 414 case 'D': 415 cfg_drop = true; 416 break; 417 case 'g': 418 cfg_use_gso = true; 419 break; 420 case 'i': 421 cfg_ignore_outgoing = true; 422 break; 423 case 'l': 424 cfg_payload_len = strtoul(optarg, NULL, 0); 425 break; 426 case 'q': 427 cfg_use_qdisc_bypass = true; 428 break; 429 case 't': 430 cfg_truncate_len = strtoul(optarg, NULL, 0); 431 break; 432 case 'v': 433 cfg_use_vnet = true; 434 break; 435 case 'V': 436 cfg_use_vlan = true; 437 break; 438 default: 439 error(1, 0, "%s: parse error", argv[0]); 440 } 441 } 442 443 if (cfg_use_vlan && cfg_use_dgram) 444 error(1, 0, "option vlan (-V) conflicts with dgram (-d)"); 445 446 if (cfg_use_csum_off && !cfg_use_vnet) 447 error(1, 0, "option csum offload (-c) requires vnet (-v)"); 448 449 if (cfg_use_csum_off_bad && !cfg_use_csum_off) 450 error(1, 0, "option csum bad (-C) requires csum offload (-c)"); 451 452 if (cfg_use_gso && !cfg_use_csum_off) 453 error(1, 0, "option gso (-g) requires csum offload (-c)"); 454 455 if (cfg_aux_data && cfg_drop) 456 error(1, 0, "option aux data (-a) conflicts with drop (-D)"); 457 458 if (cfg_ignore_outgoing && (cfg_drop || cfg_aux_data)) 459 error(1, 0, 460 "option ignore outgoing (-i) conflicts with -D and -a"); 461 } 462 463 static void check_packet_stats(int fd, unsigned int expected_packets) 464 { 465 struct tpacket_stats st = {}; 466 socklen_t len = sizeof(st); 467 468 if (getsockopt(fd, SOL_PACKET, PACKET_STATISTICS, &st, &len)) 469 error(1, errno, "getsockopt packet statistics"); 470 471 if (cfg_drop) { 472 /* PACKET_STATISTICS reports all packets seen (including 473 * drops) in tp_packets 474 */ 475 if (st.tp_packets < st.tp_drops) 476 error(1, 0, "stats: tp_packets %u < tp_drops %u", 477 st.tp_packets, st.tp_drops); 478 479 if (st.tp_drops == 0) 480 error(1, 0, "stats: expected drops but tp_drops == 0"); 481 } else { 482 if (st.tp_packets != expected_packets) 483 error(1, 0, "stats: tp_packets %u != %u", 484 st.tp_packets, expected_packets); 485 486 if (st.tp_drops != 0) 487 error(1, 0, "stats: tp_drops %u != 0", st.tp_drops); 488 } 489 490 /* verify clear on read */ 491 memset(&st, 0xff, sizeof(st)); 492 len = sizeof(st); 493 494 if (getsockopt(fd, SOL_PACKET, PACKET_STATISTICS, &st, &len)) 495 error(1, errno, "getsockopt packet statistics"); 496 497 if (st.tp_packets != 0) 498 error(1, 0, "stats: tp_packets %u != 0 after clear", st.tp_packets); 499 500 if (st.tp_drops != 0) 501 error(1, 0, "stats: tp_drops %u != 0 after clear", st.tp_drops); 502 } 503 504 static void set_ignore_outgoing(int fd, int val) 505 { 506 socklen_t len = sizeof(int); 507 int got = -1; 508 509 if (setsockopt(fd, SOL_PACKET, PACKET_IGNORE_OUTGOING, 510 &val, sizeof(val))) 511 error(1, errno, "setsockopt PACKET_IGNORE_OUTGOING %d", val); 512 513 if (getsockopt(fd, SOL_PACKET, PACKET_IGNORE_OUTGOING, &got, &len)) 514 error(1, errno, "getsockopt PACKET_IGNORE_OUTGOING"); 515 if (got != val) 516 error(1, 0, "getsockopt: expected %d got %d", val, got); 517 } 518 519 static void check_ignore_outgoing_range(int fd) 520 { 521 int val; 522 523 /* Values outside [0, 1] must be rejected with -EINVAL. */ 524 val = 2; 525 if (setsockopt(fd, SOL_PACKET, PACKET_IGNORE_OUTGOING, 526 &val, sizeof(val)) != -1 || errno != EINVAL) 527 error(1, errno, 528 "setsockopt PACKET_IGNORE_OUTGOING val=2: expected EINVAL"); 529 530 val = -1; 531 if (setsockopt(fd, SOL_PACKET, PACKET_IGNORE_OUTGOING, 532 &val, sizeof(val)) != -1 || errno != EINVAL) 533 error(1, errno, 534 "setsockopt PACKET_IGNORE_OUTGOING val=-1: expected EINVAL"); 535 } 536 537 static void test_ignore_outgoing(int fds) 538 { 539 char *expected = tbuf + sizeof(struct virtio_net_hdr); 540 int expected_len; 541 542 /* ptype_all sniffer on loopback should produce two copies per packet 543 * (RX and TX). 544 */ 545 expected_len = do_tx(); 546 expected_len -= sizeof(struct virtio_net_hdr); 547 do_rx(fds, expected_len, expected, true, PACKET_OUTGOING); 548 do_rx(fds, expected_len, expected, true, PACKET_HOST); 549 check_packet_stats(fds, 2); 550 551 /* 0 and 1 accepted; anything else rejected. */ 552 set_ignore_outgoing(fds, 0); 553 set_ignore_outgoing(fds, 1); 554 check_ignore_outgoing_range(fds); 555 556 /* With PACKET_IGNORE_OUTGOING set, only the rx copy survives. */ 557 do_tx(); 558 do_rx(fds, expected_len, expected, true, PACKET_HOST); 559 if (recv(fds, rbuf, sizeof(rbuf), 0) != -1 || errno != EAGAIN) 560 error(1, errno, "expected EAGAIN, got extra packet"); 561 check_packet_stats(fds, 1); 562 } 563 564 static void run_test(void) 565 { 566 int fdr, fds, total_len; 567 568 fdr = setup_rx(); 569 fds = setup_sniffer(); 570 571 if (cfg_ignore_outgoing) { 572 test_ignore_outgoing(fds); 573 goto out; 574 } 575 576 total_len = do_tx(); 577 578 if (cfg_drop) { 579 check_packet_stats(fds, 0); 580 goto out; 581 } 582 583 /* BPF filter accepts only this length, vlan changes MAC */ 584 if (cfg_payload_len == DATA_LEN && !cfg_use_vlan) { 585 do_rx(fds, total_len - sizeof(struct virtio_net_hdr), 586 tbuf + sizeof(struct virtio_net_hdr), true, -1); 587 check_packet_stats(fds, 1); 588 } 589 590 do_rx(fdr, cfg_payload_len, tbuf + total_len - cfg_payload_len, false, -1); 591 592 out: 593 if (close(fds)) 594 error(1, errno, "close s"); 595 if (close(fdr)) 596 error(1, errno, "close r"); 597 } 598 599 int main(int argc, char **argv) 600 { 601 parse_opts(argc, argv); 602 603 if (system("ip link set dev lo mtu 1500")) 604 error(1, errno, "ip link set mtu"); 605 if (system("ip addr add dev lo 172.17.0.1/24")) 606 error(1, errno, "ip addr add"); 607 if (system("sysctl -w net.ipv4.conf.lo.accept_local=1")) 608 error(1, errno, "sysctl lo.accept_local"); 609 610 run_test(); 611 612 fprintf(stderr, "OK\n\n"); 613 return 0; 614 } 615