package clientwg import ( "net/netip" "os" "testing" "golang.zx2c4.com/wireguard/tun" ) type collectingSink struct{ packets [][]byte } func (s *collectingSink) Enqueue(packet []byte) bool { s.packets = append(s.packets, packet) return true } type rejectingSink struct{} func (rejectingSink) Enqueue([]byte) bool { return false } func TestPacketMuxClassifiesIPv4Destination(t *testing.T) { mux := NewPacketMux( netip.MustParsePrefix("10.88.0.0/16"), []netip.Prefix{netip.MustParsePrefix("192.168.13.0/24")}, nil, ) for _, test := range []struct { packet []byte want PacketClass }{ {ipv4Packet("10.88.0.3"), PacketOverlay}, {ipv4Packet("192.168.13.10"), PacketRemote}, {ipv4Packet("8.8.8.8"), PacketDrop}, {[]byte{0x60, 0, 0, 20}, PacketDrop}, {[]byte{0x45, 0, 0, 40}, PacketDrop}, } { if got := mux.Classify(test.packet); got != test.want { t.Errorf("Classify(%v) = %v, want %v", test.packet, got, test.want) } } } func TestMuxTunConsumesRemoteOnlyBatchUntilOverlay(t *testing.T) { base := &sequenceTUN{ events: make(chan tun.Event), batches: [][][]byte{ {ipv4Packet("192.168.13.10")}, {ipv4Packet("10.88.0.3")}, }, } sink := &collectingSink{} router := NewPacketMux(netip.MustParsePrefix("10.88.0.0/16"), []netip.Prefix{netip.MustParsePrefix("192.168.13.0/24")}, sink) mux := NewMuxTun(base) mux.SetPacketMux(router) buffer := make([]byte, 128) sizes := make([]int, 1) count, err := mux.Read([][]byte{buffer}, sizes, 4) if err != nil { t.Fatal(err) } if count != 1 || base.reads != 2 || len(sink.packets) != 1 { t.Fatalf("count=%d reads=%d remote=%d", count, base.reads, len(sink.packets)) } if destination(buffer[4:4+sizes[0]]) != netip.MustParseAddr("10.88.0.3") { t.Fatalf("returned packet destination = %s", destination(buffer[4:4+sizes[0]])) } if destination(sink.packets[0]) != netip.MustParseAddr("192.168.13.10") { t.Fatalf("queued packet destination = %s", destination(sink.packets[0])) } base.batches[0][0][16] = 1 if destination(sink.packets[0]) != netip.MustParseAddr("192.168.13.10") { t.Fatal("RemoteSink packet aliases the base TUN buffer") } } func TestMuxTunCompactsMixedBatch(t *testing.T) { base := &sequenceTUN{ events: make(chan tun.Event), batches: [][][]byte{{ ipv4Packet("192.168.13.10"), ipv4Packet("10.88.0.4"), ipv4Packet("8.8.8.8"), ipv4Packet("10.88.0.5"), }}, } sink := &collectingSink{} router := NewPacketMux(netip.MustParsePrefix("10.88.0.0/16"), []netip.Prefix{netip.MustParsePrefix("192.168.13.0/24")}, sink) mux := NewMuxTun(base) mux.SetPacketMux(router) bufs := [][]byte{make([]byte, 64), make([]byte, 64), make([]byte, 64), make([]byte, 64)} sizes := make([]int, len(bufs)) count, err := mux.Read(bufs, sizes, 0) if err != nil { t.Fatal(err) } if count != 2 || destination(bufs[0][:sizes[0]]).String() != "10.88.0.4" || destination(bufs[1][:sizes[1]]).String() != "10.88.0.5" { t.Fatalf("compacted count=%d destinations=%s,%s", count, destination(bufs[0][:sizes[0]]), destination(bufs[1][:sizes[1]])) } overlay, remote, dropped := router.Counters() if overlay != 2 || remote != 1 || dropped != 1 { t.Fatalf("counters overlay=%d remote=%d dropped=%d", overlay, remote, dropped) } } func TestPacketMuxDropCallbackContainsMetadataOnly(t *testing.T) { mux := NewPacketMux(netip.MustParsePrefix("10.88.0.0/16"), nil, nil) var events []DropEvent mux.SetDropHandler(func(event DropEvent) { events = append(events, event) }) mux.route(ipv4Packet("8.8.8.8")) mux.route([]byte{0x60}) if len(events) != 2 || events[0].Reason != DropUnmanagedDestination || events[0].Destination.String() != "8.8.8.8" || events[1].Reason != DropInvalidIPv4 { t.Fatalf("drop events = %+v", events) } } func TestPacketMuxCountsInterceptedUploadAndReportsQueueDrop(t *testing.T) { base := &sequenceTUN{ events: make(chan tun.Event), batches: [][][]byte{{ipv4Packet("192.168.13.10")}, {ipv4Packet("10.88.0.3")}}, } router := NewPacketMux(netip.MustParsePrefix("10.88.0.0/16"), []netip.Prefix{netip.MustParsePrefix("192.168.13.0/24")}, rejectingSink{}) var drops []DropEvent router.SetDropHandler(func(event DropEvent) { drops = append(drops, event) }) mux := NewMuxTun(base) mux.SetPacketMux(router) buffer := make([]byte, 64) sizes := make([]int, 1) if _, err := mux.Read([][]byte{buffer}, sizes, 0); err != nil { t.Fatal(err) } bytes, packets := router.RemoteCounters() _, remote, dropped := router.Counters() if bytes != 20 || packets != 1 || remote != 1 || dropped != 1 { t.Fatalf("Upload/drop counters bytes=%d packets=%d remote=%d dropped=%d", bytes, packets, remote, dropped) } if len(drops) != 1 || drops[0].Reason != DropRemoteQueueFull { t.Fatalf("queue drop events = %+v", drops) } } func TestPacketMuxRejectsTrailingBytesBeyondIPv4TotalLength(t *testing.T) { packet := append(ipv4Packet("192.168.13.10"), 0) mux := NewPacketMux(netip.MustParsePrefix("10.88.0.0/16"), []netip.Prefix{netip.MustParsePrefix("192.168.13.0/24")}, nil) if got := mux.Classify(packet); got != PacketDrop { t.Fatalf("Classify packet with trailing bytes = %v, want PacketDrop", got) } } func ipv4Packet(destinationText string) []byte { packet := make([]byte, 20) packet[0] = 0x45 packet[2] = 0 packet[3] = 20 packet[12] = 10 packet[13] = 88 packet[14] = 0 packet[15] = 2 destination := netip.MustParseAddr(destinationText).As4() copy(packet[16:20], destination[:]) return packet } func destination(packet []byte) netip.Addr { return netip.AddrFrom4([4]byte{packet[16], packet[17], packet[18], packet[19]}) } type sequenceTUN struct { events chan tun.Event batches [][][]byte reads int } func (t *sequenceTUN) File() *os.File { return nil } func (t *sequenceTUN) Read(bufs [][]byte, sizes []int, offset int) (int, error) { batch := t.batches[t.reads] t.reads++ for index, packet := range batch { sizes[index] = copy(bufs[index][offset:], packet) } return len(batch), nil } func (t *sequenceTUN) Write(bufs [][]byte, offset int) (int, error) { return len(bufs), nil } func (t *sequenceTUN) MTU() (int, error) { return 1280, nil } func (t *sequenceTUN) Name() (string, error) { return "RemLink", nil } func (t *sequenceTUN) Events() <-chan tun.Event { return t.events } func (t *sequenceTUN) Close() error { return nil } func (t *sequenceTUN) BatchSize() int { return 4 }