#!/usr/bin/env python3
"""tapcat — shovel Ethernet frames between a tap device and a UDP socket.

    tapcat.py tap0 6000                      # listen, learn the peer from the first frame
    tapcat.py tap0 198.51.100.10 6000        # connect to a listener
    TAPCAT_ZSTD=frame tapcat.py tap0 6000    # compress each datagram on its own
    TAPCAT_ZSTD=stream tapcat.py tap0 6000   # keep the window across frames
    TAPCAT_DICT=tunnel.dict TAPCAT_ZSTD=frame tapcat.py tap0 6000

One read from the tap device returns exactly one frame, and one datagram
carries exactly one frame, so there is no framing to invent and nothing to
resynchronise after a loss. Works over IPv4 or IPv6 — whichever getaddrinfo
hands back first.

Wire format: one header byte, then the frame. 0x00 means the frame follows as
it is, 0x01 means it is a self-contained zstd frame, 0x02 means it is a block
from a continuing zstd stream. The header is always there, even with
compression off, and it names the mode rather than just saying "compressed" —
so two ends started with different settings still understand each other.

Compression modes, and the difference is not marginal:

  stream  FLUSH_BLOCK — keeps the compression window across frames. Measured
          three times smaller on real traffic. Every frame then depends on
          every frame before it, so this needs a carrier that delivers
          everything in order: TCP, or TLS over TCP. Not UDP, not DTLS.
  frame   FLUSH_FRAME — each datagram is a complete, self-contained zstd
          frame, so datagram N still decodes after 1..N-1 were lost. This is
          the only mode a datagram carrier can use. Pair it with a trained
          dictionary (TAPCAT_DICT) to win back most of what the window gave.

Each direction ramps on its own: a link starts raw for RAMP seconds and only
then begins compressing, so it comes up on the simplest path it has and gets
clever afterwards.
"""
import fcntl, os, select, socket, struct, sys, time

TUNSETIFF = 0x400454CA
IFF_TAP   = 0x0002          # IFF_TUN is 0x0001 if you want layer 3 instead
IFF_NO_PI = 0x1000          # no 4-byte packet-info header on every read

MTU    = 1600               # big enough for 1500 plus the Ethernet header
RAW    = b"\x00"          # the frame follows as it is
ZFRAME = b"\x01"          # a complete, self-contained zstd frame
ZBLOCK = b"\x02"          # a block from a continuing zstd stream
RAMP   = 1.0                # seconds of raw frames before compression starts
ZLEVEL = 1                  # 3 and 9 buy about 1% and cost most of the throughput

try:
    from compression.zstd import (ZstdCompressor, ZstdDecompressor,  # Python 3.14+
                                  ZstdDict, decompress as zstd_decompress)
except ImportError:
    ZstdCompressor = ZstdDecompressor = ZstdDict = zstd_decompress = None


def open_tap(name):
    fd = os.open("/dev/net/tun", os.O_RDWR)
    fcntl.ioctl(fd, TUNSETIFF, struct.pack("16sH", name.encode(), IFF_TAP | IFF_NO_PI))
    return fd


class Codec:
    """Header on everything; compression only after the link has settled."""

    def __init__(self, mode, dict_path=None):
        self.started = time.monotonic()
        self.mode = mode if ZstdCompressor else None
        self.zdict = None
        if dict_path and ZstdDict:
            # Load it once. Passing a dictionary per call measured 5 MB/s.
            self.zdict = ZstdDict(open(dict_path, "rb").read())
        self.streaming = mode == "stream"
        self.flush = (ZstdCompressor.FLUSH_BLOCK if self.streaming
                      else ZstdCompressor.FLUSH_FRAME) if self.mode else None
        self.tag = ZBLOCK if self.streaming else ZFRAME
        self.c = (ZstdCompressor(level=ZLEVEL, zstd_dict=self.zdict)
                  if self.mode else None)
        self.d = None           # the stream decoder, made on the first block

    def pack(self, frame):
        if self.c is None or time.monotonic() - self.started < RAMP:
            return RAW + frame
        out = self.c.compress(frame, mode=self.flush)
        # An already-encrypted payload grows, and a 64-byte ACK grows by the
        # 10-byte zstd header. Only send the compressed form if it is smaller.
        # In stream mode the window has to stay in step, so it always goes.
        if self.streaming or len(out) < len(frame):
            return self.tag + out
        return RAW + frame

    def unpack(self, datagram):
        tag, body = datagram[:1], datagram[1:]
        if tag == RAW:
            return body
        if zstd_decompress is None:
            raise RuntimeError("peer is compressing and this end has no zstd")
        if tag == ZFRAME:                       # stands on its own
            return zstd_decompress(body, zstd_dict=self.zdict)
        if self.d is None:                      # one decoder for the whole stream
            self.d = ZstdDecompressor(zstd_dict=self.zdict)
        return self.d.decompress(body)


def main(argv):
    dev = argv[1]
    listening = len(argv) == 3
    host, port = ("::", int(argv[2])) if listening else (argv[2], int(argv[3]))

    family, stype, proto, _, addr = socket.getaddrinfo(
        host, port, type=socket.SOCK_DGRAM,
        flags=socket.AI_PASSIVE if listening else 0)[0]
    sock = socket.socket(family, stype, proto)

    peer = None
    if listening:
        sock.bind(addr)
    else:
        peer = addr
        sock.sendto(RAW, peer)          # open the path so the listener can answer

    codec = Codec(os.environ.get("TAPCAT_ZSTD"), os.environ.get("TAPCAT_DICT"))
    tap = open_tap(dev)
    while True:
        ready, _, _ = select.select([tap, sock], [], [])
        if tap in ready and peer:
            sock.sendto(codec.pack(os.read(tap, MTU)), peer)
        if sock in ready:
            datagram, src = sock.recvfrom(MTU + 64)
            peer = src                  # last speaker wins; see the note below
            if len(datagram) > 1:
                os.write(tap, codec.unpack(datagram))


if __name__ == "__main__":
    if not 3 <= len(sys.argv) <= 4:
        sys.exit(__doc__)
    main(sys.argv)
