#define _DEFAULT_SOURCE
#define _POSIX_C_SOURCE 200809L
#include <sys/socket.h>
#include <netpacket/packet.h>
#include <net/if.h>
#include <linux/if_ether.h>
#include <netinet/ip.h>
#include <netinet/tcp.h>
#include <netinet/udp.h>
#include <netinet/ip_icmp.h>
#include <arpa/inet.h>
#include <unistd.h>
#include <errno.h>
#include <stdio.h>
#include <string.h>

static void
analyze_packet(const unsigned char *packet, size_t size)
{
    struct ethhdr ethernet;
    struct iphdr ip;
    char source[INET_ADDRSTRLEN], destination[INET_ADDRSTRLEN];
    const unsigned char *payload;
    size_t header_len, total_len, payload_len;
    unsigned int fragment;

    if (size < sizeof(ethernet)) return;
    memcpy(&ethernet, packet, sizeof(ethernet));
    if (ntohs(ethernet.h_proto) != ETH_P_IP) return;
    packet += sizeof(ethernet);
    size -= sizeof(ethernet);
    if (size < sizeof(ip)) return;
    /* Ethernet headers can leave the IP header unaligned. */
    memcpy(&ip, packet, sizeof(ip));
    if (ip.version != 4 || ip.ihl < 5) return;
    header_len = (size_t)ip.ihl * 4;
    total_len = ntohs(ip.tot_len);
    if (total_len < header_len || total_len > size) return;
    fragment = ntohs(ip.frag_off);
    if (fragment & IP_OFFMASK) return;
    payload = packet + header_len;
    payload_len = total_len - header_len;
    if (inet_ntop(AF_INET, &ip.saddr, source, sizeof(source)) == NULL ||
        inet_ntop(AF_INET, &ip.daddr, destination, sizeof(destination)) == NULL)
        return;

    switch (ip.protocol) {
    case IPPROTO_TCP: {
        struct tcphdr tcp;
        size_t tcp_len;
        if (payload_len < sizeof(tcp)) return;
        memcpy(&tcp, payload, sizeof(tcp));
        tcp_len = (size_t)tcp.doff * 4;
        if (tcp_len < sizeof(tcp) || tcp_len > payload_len) return;
        printf("TCP %s:%u -> %s:%u\n", source, (unsigned int)ntohs(tcp.source),
               destination, (unsigned int)ntohs(tcp.dest));
        break;
    }
    case IPPROTO_UDP: {
        struct udphdr udp;
        size_t udp_len;
        if (payload_len < sizeof(udp)) return;
        memcpy(&udp, payload, sizeof(udp));
        udp_len = ntohs(udp.len);
        if (udp_len < sizeof(udp) || (!(fragment & IP_MF) && udp_len > payload_len))
            return;
        printf("UDP %s:%u -> %s:%u\n", source, (unsigned int)ntohs(udp.source),
               destination, (unsigned int)ntohs(udp.dest));
        break;
    }
    case IPPROTO_ICMP: {
        struct icmphdr icmp;
        if (payload_len < sizeof(icmp)) return;
        memcpy(&icmp, payload, sizeof(icmp));
        printf("ICMP %s -> %s type %u code %u\n", source, destination,
               (unsigned int)icmp.type, (unsigned int)icmp.code);
        break;
    }
    }
}

int
main(int argc, char **argv)
{
    unsigned char packet[65536 + ETH_HLEN];
    struct sockaddr_ll address = {0};
    unsigned int index;
    int fd, status = 0;

    if (argc != 2) {
        fprintf(stderr, "usage: %s interface\n", argv[0]);
        return 1;
    }
    index = if_nametoindex(argv[1]);
    if (index == 0) { perror(argv[1]); return 1; }
    fd = socket(AF_PACKET, SOCK_RAW, htons(ETH_P_ALL));
    if (fd < 0) { perror("socket"); return 1; }
    address.sll_family = AF_PACKET;
    address.sll_protocol = htons(ETH_P_ALL);
    address.sll_ifindex = (int)index;
    if (bind(fd, (struct sockaddr *)&address, sizeof(address)) < 0) {
        perror("bind");
        close(fd);
        return 1;
    }
    for (;;) {
        ssize_t n = recvfrom(fd, packet, sizeof(packet), 0, NULL, NULL);
        if (n < 0 && errno == EINTR) continue;
        if (n < 0) { perror("recvfrom"); status = 1; break; }
        analyze_packet(packet, (size_t)n);
        if (fflush(stdout) == EOF) { perror("stdout"); status = 1; break; }
    }
    close(fd);
    return status;
}
