Remove unused

This commit is contained in:
世界 2025-09-15 18:12:04 +08:00
parent e4aedc6f6e
commit 177dea9806
No known key found for this signature in database
GPG key ID: CD109927C34A63C4
38 changed files with 0 additions and 7246 deletions

View file

@ -1,67 +0,0 @@
/* SPDX-License-Identifier: MIT
*
* Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
*/
package tun
import (
"reflect"
"testing"
"unsafe"
)
func checkAlignment(t *testing.T, name string, offset uintptr) {
t.Helper()
if offset%8 != 0 {
t.Errorf("offset of %q within struct is %d bytes, which does not align to 64-bit word boundaries (missing %d bytes). Atomic operations will crash on 32-bit systems.", name, offset, 8-(offset%8))
}
}
// TestRateJugglerAlignment checks that atomically-accessed fields are
// aligned to 64-bit boundaries, as required by the atomic package.
//
// Unfortunately, violating this rule on 32-bit platforms results in a
// hard segfault at runtime.
func TestRateJugglerAlignment(t *testing.T) {
var r rateJuggler
typ := reflect.TypeOf(&r).Elem()
t.Logf("Peer type size: %d, with fields:", typ.Size())
for i := 0; i < typ.NumField(); i++ {
field := typ.Field(i)
t.Logf("\t%30s\toffset=%3v\t(type size=%3d, align=%d)",
field.Name,
field.Offset,
field.Type.Size(),
field.Type.Align(),
)
}
checkAlignment(t, "rateJuggler.current", unsafe.Offsetof(r.current))
checkAlignment(t, "rateJuggler.nextByteCount", unsafe.Offsetof(r.nextByteCount))
checkAlignment(t, "rateJuggler.nextStartTime", unsafe.Offsetof(r.nextStartTime))
}
// TestNativeTunAlignment checks that atomically-accessed fields are
// aligned to 64-bit boundaries, as required by the atomic package.
//
// Unfortunately, violating this rule on 32-bit platforms results in a
// hard segfault at runtime.
func TestNativeTunAlignment(t *testing.T) {
var tun NativeTun
typ := reflect.TypeOf(&tun).Elem()
t.Logf("Peer type size: %d, with fields:", typ.Size())
for i := 0; i < typ.NumField(); i++ {
field := typ.Field(i)
t.Logf("\t%30s\toffset=%3v\t(type size=%3d, align=%d)",
field.Name,
field.Offset,
field.Type.Size(),
field.Type.Align(),
)
}
checkAlignment(t, "NativeTun.rate", unsafe.Offsetof(tun.rate))
}

View file

@ -1,45 +0,0 @@
//go:build amd64
package tun
import (
"golang.org/x/sys/cpu"
)
var archChecksumFuncs = []archChecksumDetails{
{
name: "generic32",
available: true,
f: checksumGeneric32,
},
{
name: "generic64",
available: true,
f: checksumGeneric64,
},
{
name: "generic32Alternate",
available: true,
f: checksumGeneric32Alternate,
},
{
name: "generic64Alternate",
available: true,
f: checksumGeneric64Alternate,
},
{
name: "AMD64",
available: true,
f: checksumAMD64,
},
{
name: "SSE2",
available: cpu.X86.HasSSE2,
f: checksumSSE2,
},
{
name: "AVX2",
available: cpu.X86.HasAVX && cpu.X86.HasAVX2 && cpu.X86.HasBMI2,
f: checksumAVX2,
},
}

View file

@ -1,26 +0,0 @@
//go:build !amd64
package tun
var archChecksumFuncs = []archChecksumDetails{
{
name: "generic32",
available: true,
f: checksumGeneric32,
},
{
name: "generic32Alternate",
available: true,
f: checksumGeneric32Alternate,
},
{
name: "generic64",
available: true,
f: checksumGeneric64,
},
{
name: "generic64Alternate",
available: true,
f: checksumGeneric64Alternate,
},
}

View file

@ -1,619 +0,0 @@
package tun
import (
"fmt"
"math"
"math/rand"
"net/netip"
"sort"
"syscall"
"testing"
"unsafe"
"gvisor.dev/gvisor/pkg/tcpip"
gvisorChecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
"gvisor.dev/gvisor/pkg/tcpip/header"
)
type archChecksumDetails struct {
name string
available bool
f func([]byte, uint16) uint16
}
func fillRandomBuffer(seed int64, buf []byte) {
rng := rand.New(rand.NewSource(seed))
n, err := rng.Read(buf)
if err != nil {
panic(err)
}
if n != len(buf) {
panic("incomplete random buffer")
}
}
func deterministicRandomBytes(seed int64, length int) []byte {
buf := make([]byte, length)
fillRandomBuffer(seed, buf)
return buf
}
func getPageAlignedRandomBytes(seed int64, length int) []byte {
alignment := syscall.Getpagesize()
buf := make([]byte, length+(alignment-1))
bufPtr := uintptr(unsafe.Pointer(&buf[0]))
alignedBufPtr := (bufPtr + uintptr(alignment-1)) & ^uintptr(alignment-1)
alignedStart := int(alignedBufPtr - bufPtr)
buf = buf[alignedStart : alignedStart+length]
fillRandomBuffer(seed, buf)
return buf
}
func TestChecksum(t *testing.T) {
alignedBuf := getPageAlignedRandomBytes(10, 8192)
allOnes := make([]byte, 65535)
for i := range allOnes {
allOnes[i] = 0xff
}
allFE := make([]byte, 65535)
for i := range allFE {
allFE[i] = 0xfe
}
tests := []struct {
name string
data []byte
initial uint16
want uint16
}{
{
name: "empty",
data: []byte{},
initial: 0,
want: 0,
},
{
name: "max initial",
data: []byte{},
initial: math.MaxUint16,
want: 0xffff,
},
{
name: "odd length",
data: []byte{0x01, 0x02, 0x01},
initial: 0,
want: 0x0202,
},
{
name: "tiny",
data: []byte{0x01, 0x02, 0x01, 0x02, 0x01, 0x02},
initial: 0,
want: 0x0306,
},
{
name: "initial",
data: []byte{0x01, 0x02, 0x01, 0x02, 0x01, 0x02},
initial: 0x1000,
want: 0x1306,
},
// cleanup0 through cleanup15 is 1024 (handled by large SIMD loops) +
// 32 (handled by small SIMD loops) + n, where n ranges from 0 to 15
// to cover all of the leftover byte sizes that are possible after small
// SIMD loops that handle 16 bytes.
{
name: "cleanup0",
data: deterministicRandomBytes(1, 1056),
initial: 0,
want: 0x11ec,
},
{
name: "cleanup1",
data: deterministicRandomBytes(1, 1057),
initial: 0,
want: 0xc5ec,
},
{
name: "cleanup2",
data: deterministicRandomBytes(1, 1058),
initial: 0,
want: 0xc6ad,
},
{
name: "cleanup3",
data: deterministicRandomBytes(1, 1059),
initial: 0,
want: 0x86ae,
},
{
name: "cleanup4",
data: deterministicRandomBytes(1, 1060),
initial: 0,
want: 0x878e,
},
{
name: "cleanup5",
data: deterministicRandomBytes(1, 1061),
initial: 0,
want: 0xdb8e,
},
{
name: "cleanup6",
data: deterministicRandomBytes(1, 1062),
initial: 0,
want: 0xdbd5,
},
{
name: "cleanup7",
data: deterministicRandomBytes(1, 1063),
initial: 0,
want: 0xcfd6,
},
{
name: "cleanup8",
data: deterministicRandomBytes(1, 1064),
initial: 0,
want: 0xd090,
},
{
name: "cleanup9",
data: deterministicRandomBytes(1, 1065),
initial: 0,
want: 0x0791,
},
{
name: "cleanup10",
data: deterministicRandomBytes(1, 1066),
initial: 0,
want: 0x079f,
},
{
name: "cleanup11",
data: deterministicRandomBytes(1, 1067),
initial: 0,
want: 0xba9f,
},
{
name: "cleanup12",
data: deterministicRandomBytes(1, 1068),
initial: 0,
want: 0xbb0c,
},
{
name: "cleanup13",
data: deterministicRandomBytes(1, 1069),
initial: 0,
want: 0x770d,
},
{
name: "cleanup14",
data: deterministicRandomBytes(1, 1070),
initial: 0,
want: 0x780a,
},
{
name: "cleanup15",
data: deterministicRandomBytes(1, 1071),
initial: 0,
want: 0x640b,
},
// small1 through small15 covers small sizes that are not large enough
// to do overlapped reads.
{
name: "small1",
data: deterministicRandomBytes(2, 1),
initial: 0x1122,
want: 0x4022,
},
{
name: "small2",
data: deterministicRandomBytes(2, 2),
initial: 0x1122,
want: 0x40a4,
},
{
name: "small3",
data: deterministicRandomBytes(2, 3),
initial: 0x1122,
want: 0xc2a4,
},
{
name: "small4",
data: deterministicRandomBytes(2, 4),
initial: 0x1122,
want: 0xc36f,
},
{
name: "small5",
data: deterministicRandomBytes(2, 5),
initial: 0x1122,
want: 0xa570,
},
{
name: "small6",
data: deterministicRandomBytes(2, 6),
initial: 0x1122,
want: 0xa669,
},
{
name: "small7",
data: deterministicRandomBytes(2, 7),
initial: 0x1122,
want: 0x0f6a,
},
{
name: "small8",
data: deterministicRandomBytes(2, 8),
initial: 0x1122,
want: 0x0fd9,
},
{
name: "small9",
data: deterministicRandomBytes(2, 9),
initial: 0x1122,
want: 0x40d9,
},
{
name: "small10",
data: deterministicRandomBytes(2, 10),
initial: 0x1122,
want: 0x411d,
},
{
name: "small11",
data: deterministicRandomBytes(2, 11),
initial: 0x1122,
want: 0x011e,
},
{
name: "small12",
data: deterministicRandomBytes(2, 12),
initial: 0x1122,
want: 0x01c8,
},
{
name: "small13",
data: deterministicRandomBytes(2, 13),
initial: 0x1122,
want: 0x4dc8,
},
{
name: "small14",
data: deterministicRandomBytes(2, 14),
initial: 0x1122,
want: 0x4eb5,
},
{
name: "small15",
data: deterministicRandomBytes(2, 15),
initial: 0x1122,
want: 0xa4b5,
},
// other small-ish sizes
{
name: "small16",
data: deterministicRandomBytes(1, 16),
initial: 0,
want: 0x02fa,
},
{
name: "small32",
data: deterministicRandomBytes(1, 32),
initial: 0,
want: 0x03ee,
},
{
name: "small64",
data: deterministicRandomBytes(1, 64),
initial: 0,
want: 0x3f85,
},
{
name: "medium",
data: deterministicRandomBytes(1, 1400),
initial: 0,
want: 0xbea5,
},
{
name: "big",
data: deterministicRandomBytes(2, 65000),
initial: 0,
want: 0x3ba7,
},
{
name: "big-initial",
data: deterministicRandomBytes(2, 65000),
initial: 0x1234,
want: 0x4ddb,
},
{
// big-small-loop is intended to exercise a few iterations of a big
// initial loop of 128 bytes or larger + a smaller loop of 16 bytes
// + some leftover
name: "big-small-loop",
data: deterministicRandomBytes(3, 1094),
initial: 0x9999,
want: 0xe65b,
},
{
name: "page-aligned",
data: alignedBuf[:4096],
initial: 0,
want: 0x963b,
},
{
name: "32-aligned",
data: alignedBuf[32:4128],
initial: 0,
want: 0x30c4,
},
{
name: "16-aligned",
data: alignedBuf[16:4112],
initial: 0,
want: 0xaeff,
},
{
name: "8-aligned",
data: alignedBuf[8:4104],
initial: 0,
want: 0x6c3b,
},
{
name: "4-aligned",
data: alignedBuf[4:4100],
initial: 0,
want: 0x2e4a,
},
{
name: "2-aligned",
data: alignedBuf[2:4098],
initial: 0,
want: 0xc702,
},
{
name: "unaligned",
data: alignedBuf[1:4097],
initial: 0,
want: 0x3bc7,
},
{
name: "unalignedAndOdd",
data: alignedBuf[1:4096],
initial: 0,
want: 0x3b13,
},
{
name: "fe1282",
data: allFE[:1282],
initial: 0,
want: 0x7c7c,
},
{
name: "fe",
data: allFE,
initial: 0,
want: 0x7e81,
},
{
name: "maximum",
data: allOnes,
initial: 0,
want: 0xff00,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
for _, fd := range archChecksumFuncs {
t.Run(fd.name, func(t *testing.T) {
if !fd.available {
t.Skip("can not run on this system")
}
if got := fd.f(tt.data, tt.initial); got != tt.want {
t.Errorf("%s checksum = %04x, want %04x", fd.name, got, tt.want)
}
})
}
t.Run("reference", func(t *testing.T) {
if got := gvisorChecksum.Checksum(tt.data, tt.initial); got != tt.want {
t.Errorf("reference checksum = %04x, want %04x", got, tt.want)
}
})
})
}
}
func TestPseudoHeaderChecksumNoFold(t *testing.T) {
tests := []struct {
name string
protocol uint8
srcAddr []byte
dstAddr []byte
totalLen uint16
want uint16
}{
{
name: "ipv4",
protocol: syscall.IPPROTO_TCP,
srcAddr: netip.MustParseAddr("192.168.1.1").AsSlice(),
dstAddr: netip.MustParseAddr("192.168.1.2").AsSlice(),
totalLen: 1492,
want: 0x892e,
},
{
name: "ipv6",
protocol: syscall.IPPROTO_TCP,
srcAddr: netip.MustParseAddr("2001:db8:3333:4444:5555:6666:7777:8888").AsSlice(),
dstAddr: netip.MustParseAddr("2001:db8:aaaa:bbbb:cccc:dddd:eeee:ffff").AsSlice(),
totalLen: 1492,
want: 0x947f,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Run("pseudoHeaderChecksum32", func(t *testing.T) {
got := pseudoHeaderChecksum32(tt.protocol, tt.srcAddr, tt.dstAddr, tt.totalLen)
if got != tt.want {
t.Errorf("got %04x, want %04x", got, tt.want)
}
})
t.Run("pseudoHeaderChecksum64", func(t *testing.T) {
got := pseudoHeaderChecksum64(tt.protocol, tt.srcAddr, tt.dstAddr, tt.totalLen)
if got != tt.want {
t.Errorf("got %04x, want %04x", got, tt.want)
}
})
t.Run("reference", func(t *testing.T) {
got := header.PseudoHeaderChecksum(
tcpip.TransportProtocolNumber(tt.protocol),
tcpip.AddrFromSlice(tt.srcAddr),
tcpip.AddrFromSlice(tt.dstAddr),
tt.totalLen)
if got != tt.want {
t.Errorf("got %04x, want %04x", got, tt.want)
}
})
})
}
}
func FuzzChecksum(f *testing.F) {
buf := getPageAlignedRandomBytes(1234, 65536)
f.Add([]byte{}, uint16(0))
f.Add([]byte{}, uint16(0x1234))
f.Add([]byte{}, uint16(0))
f.Add(buf[:15], uint16(0x1234))
f.Add(buf[:256], uint16(0x1234))
f.Add(buf[:1280], uint16(0x1234))
f.Add(buf[:1288], uint16(0x1234))
f.Add(buf[1:1050], uint16(0x1234))
f.Fuzz(func(t *testing.T, data []byte, initial uint16) {
want := gvisorChecksum.Checksum(data, initial)
for _, fd := range archChecksumFuncs {
t.Run(fd.name, func(t *testing.T) {
if !fd.available {
t.Skip("can not run on this system")
}
if got := fd.f(data, initial); got != want {
t.Errorf("%s checksum = %04x, want %04x", fd.name, got, want)
}
})
}
})
}
var result uint16
func BenchmarkChecksum(b *testing.B) {
offsets := []int{ // offsets from page alignment
0,
1,
2,
4,
8,
16,
}
lengths := []int{
0,
7,
15,
16,
31,
64,
90,
95,
128,
256,
512,
1024,
1240,
1500,
2048,
4096,
8192,
9000,
9001,
16384,
65536,
}
if !sort.IntsAreSorted(offsets) {
b.Fatal("offsets are not sorted")
}
largestLength := lengths[len(lengths)-1]
if !sort.IntsAreSorted(lengths) {
b.Fatal("lengths are not sorted")
}
largestOffset := lengths[len(offsets)-1]
alignedBuf := getPageAlignedRandomBytes(1, largestOffset+largestLength)
var r uint16
for _, offset := range offsets {
name := fmt.Sprintf("%vAligned", offset)
if offset == 0 {
name = "pageAligned"
}
offsetBuf := alignedBuf[offset:]
b.Run(name, func(b *testing.B) {
for _, length := range lengths {
b.Run(fmt.Sprintf("%d", length), func(b *testing.B) {
for _, fd := range archChecksumFuncs {
b.Run(fd.name, func(b *testing.B) {
if !fd.available {
b.Skip("can not run on this system")
}
b.SetBytes(int64(length))
for i := 0; i < b.N; i++ {
r += fd.f(offsetBuf[:length], 0)
}
})
}
})
}
})
}
result = r
}
func BenchmarkPseudoHeaderChecksum(b *testing.B) {
tests := []struct {
name string
protocol uint8
srcAddr []byte
dstAddr []byte
totalLen uint16
want uint16
}{
{
name: "ipv4",
protocol: syscall.IPPROTO_TCP,
srcAddr: []byte{192, 168, 1, 1},
dstAddr: []byte{192, 168, 1, 2},
totalLen: 1492,
want: 0x892e,
},
{
name: "ipv6",
protocol: syscall.IPPROTO_TCP,
srcAddr: netip.MustParseAddr("2001:db8:3333:4444:5555:6666:7777:8888").AsSlice(),
dstAddr: netip.MustParseAddr("2001:db8:aaaa:bbbb:cccc:dddd:eeee:ffff").AsSlice(),
totalLen: 1492,
want: 0x892e,
},
}
for _, tt := range tests {
b.Run(tt.name, func(b *testing.B) {
b.Run("pseudoHeaderChecksum32", func(b *testing.B) {
for i := 0; i < b.N; i++ {
result += pseudoHeaderChecksum32(tt.protocol, tt.srcAddr, tt.dstAddr, tt.totalLen)
}
})
b.Run("pseudoHeaderChecksum64", func(b *testing.B) {
for i := 0; i < b.N; i++ {
result += pseudoHeaderChecksum64(tt.protocol, tt.srcAddr, tt.dstAddr, tt.totalLen)
}
})
})
}
}

View file

@ -1,54 +0,0 @@
//go:build ignore
/* SPDX-License-Identifier: MIT
*
* Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
*/
package main
import (
"io"
"log"
"net/http"
"net/netip"
"github.com/tailscale/wireguard-go/conn"
"github.com/tailscale/wireguard-go/device"
"github.com/tailscale/wireguard-go/tun/netstack"
)
func main() {
tun, tnet, err := netstack.CreateNetTUN(
[]netip.Addr{netip.MustParseAddr("192.168.4.28")},
[]netip.Addr{netip.MustParseAddr("8.8.8.8")},
1420)
if err != nil {
log.Panic(err)
}
dev := device.NewDevice(tun, conn.NewDefaultBind(), device.NewLogger(device.LogLevelVerbose, ""))
err = dev.IpcSet(`private_key=087ec6e14bbed210e7215cdc73468dfa23f080a1bfb8665b2fd809bd99d28379
public_key=c4c8e984c5322c8184c72265b92b250fdb63688705f504ba003c88f03393cf28
allowed_ip=0.0.0.0/0
endpoint=127.0.0.1:58120
`)
err = dev.Up()
if err != nil {
log.Panic(err)
}
client := http.Client{
Transport: &http.Transport{
DialContext: tnet.DialContext,
},
}
resp, err := client.Get("http://192.168.4.29/")
if err != nil {
log.Panic(err)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
log.Panic(err)
}
log.Println(string(body))
}

View file

@ -1,51 +0,0 @@
//go:build ignore
/* SPDX-License-Identifier: MIT
*
* Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
*/
package main
import (
"io"
"log"
"net"
"net/http"
"net/netip"
"github.com/tailscale/wireguard-go/conn"
"github.com/tailscale/wireguard-go/device"
"github.com/tailscale/wireguard-go/tun/netstack"
)
func main() {
tun, tnet, err := netstack.CreateNetTUN(
[]netip.Addr{netip.MustParseAddr("192.168.4.29")},
[]netip.Addr{netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("8.8.4.4")},
1420,
)
if err != nil {
log.Panic(err)
}
dev := device.NewDevice(tun, conn.NewDefaultBind(), device.NewLogger(device.LogLevelVerbose, ""))
dev.IpcSet(`private_key=003ed5d73b55806c30de3f8a7bdab38af13539220533055e635690b8b87ad641
listen_port=58120
public_key=f928d4f6c1b86c12f2562c10b07c555c5c57fd00f59e90c8d8d88767271cbf7c
allowed_ip=192.168.4.28/32
persistent_keepalive_interval=25
`)
dev.Up()
listener, err := tnet.ListenTCP(&net.TCPAddr{Port: 80})
if err != nil {
log.Panicln(err)
}
http.HandleFunc("/", func(writer http.ResponseWriter, request *http.Request) {
log.Printf("> %s - %s - %s", request.RemoteAddr, request.URL.String(), request.UserAgent())
io.WriteString(writer, "Hello from userspace TCP!")
})
err = http.Serve(listener, nil)
if err != nil {
log.Panicln(err)
}
}

View file

@ -1,75 +0,0 @@
//go:build ignore
/* SPDX-License-Identifier: MIT
*
* Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
*/
package main
import (
"bytes"
"log"
"math/rand"
"net/netip"
"time"
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
"github.com/tailscale/wireguard-go/conn"
"github.com/tailscale/wireguard-go/device"
"github.com/tailscale/wireguard-go/tun/netstack"
)
func main() {
tun, tnet, err := netstack.CreateNetTUN(
[]netip.Addr{netip.MustParseAddr("192.168.4.29")},
[]netip.Addr{netip.MustParseAddr("8.8.8.8")},
1420)
if err != nil {
log.Panic(err)
}
dev := device.NewDevice(tun, conn.NewDefaultBind(), device.NewLogger(device.LogLevelVerbose, ""))
dev.IpcSet(`private_key=a8dac1d8a70a751f0f699fb14ba1cff7b79cf4fbd8f09f44c6e6a90d0369604f
public_key=25123c5dcd3328ff645e4f2a3fce0d754400d3887a0cb7c56f0267e20fbf3c5b
endpoint=163.172.161.0:12912
allowed_ip=0.0.0.0/0
`)
err = dev.Up()
if err != nil {
log.Panic(err)
}
socket, err := tnet.Dial("ping4", "zx2c4.com")
if err != nil {
log.Panic(err)
}
requestPing := icmp.Echo{
Seq: rand.Intn(1 << 16),
Data: []byte("gopher burrow"),
}
icmpBytes, _ := (&icmp.Message{Type: ipv4.ICMPTypeEcho, Code: 0, Body: &requestPing}).Marshal(nil)
socket.SetReadDeadline(time.Now().Add(time.Second * 10))
start := time.Now()
_, err = socket.Write(icmpBytes)
if err != nil {
log.Panic(err)
}
n, err := socket.Read(icmpBytes[:])
if err != nil {
log.Panic(err)
}
replyPacket, err := icmp.ParseMessage(1, icmpBytes[:n])
if err != nil {
log.Panic(err)
}
replyPing, ok := replyPacket.Body.(*icmp.Echo)
if !ok {
log.Panicf("invalid reply type: %v", replyPacket)
}
if !bytes.Equal(replyPing.Data, requestPing.Data) || replyPing.Seq != requestPing.Seq {
log.Panicf("invalid ping reply: %v", replyPing)
}
log.Printf("Ping latency: %v", time.Since(start))
}

File diff suppressed because it is too large Load diff

View file

@ -1,764 +0,0 @@
/* SPDX-License-Identifier: MIT
*
* Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
*/
package tun
import (
"net/netip"
"testing"
"github.com/tailscale/wireguard-go/conn"
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header"
)
const (
offset = virtioNetHdrLen
)
var (
ip4PortA = netip.MustParseAddrPort("192.0.2.1:1")
ip4PortB = netip.MustParseAddrPort("192.0.2.2:1")
ip4PortC = netip.MustParseAddrPort("192.0.2.3:1")
ip6PortA = netip.MustParseAddrPort("[2001:db8::1]:1")
ip6PortB = netip.MustParseAddrPort("[2001:db8::2]:1")
ip6PortC = netip.MustParseAddrPort("[2001:db8::3]:1")
)
func udp4PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, payloadLen int, ipFn func(*header.IPv4Fields)) []byte {
totalLen := 28 + payloadLen
b := make([]byte, offset+int(totalLen), 65535)
ipv4H := header.IPv4(b[offset:])
srcAs4 := srcIPPort.Addr().As4()
dstAs4 := dstIPPort.Addr().As4()
ipFields := &header.IPv4Fields{
SrcAddr: tcpip.AddrFromSlice(srcAs4[:]),
DstAddr: tcpip.AddrFromSlice(dstAs4[:]),
Protocol: unix.IPPROTO_UDP,
TTL: 64,
TotalLength: uint16(totalLen),
}
if ipFn != nil {
ipFn(ipFields)
}
ipv4H.Encode(ipFields)
udpH := header.UDP(b[offset+20:])
udpH.Encode(&header.UDPFields{
SrcPort: srcIPPort.Port(),
DstPort: dstIPPort.Port(),
Length: uint16(payloadLen + udphLen),
})
ipv4H.SetChecksum(^ipv4H.CalculateChecksum())
pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_UDP, ipv4H.SourceAddress(), ipv4H.DestinationAddress(), uint16(udphLen+payloadLen))
udpH.SetChecksum(^udpH.CalculateChecksum(pseudoCsum))
return b
}
func udp6Packet(srcIPPort, dstIPPort netip.AddrPort, payloadLen int) []byte {
return udp6PacketMutateIPFields(srcIPPort, dstIPPort, payloadLen, nil)
}
func udp6PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, payloadLen int, ipFn func(*header.IPv6Fields)) []byte {
totalLen := 48 + payloadLen
b := make([]byte, offset+int(totalLen), 65535)
ipv6H := header.IPv6(b[offset:])
srcAs16 := srcIPPort.Addr().As16()
dstAs16 := dstIPPort.Addr().As16()
ipFields := &header.IPv6Fields{
SrcAddr: tcpip.AddrFromSlice(srcAs16[:]),
DstAddr: tcpip.AddrFromSlice(dstAs16[:]),
TransportProtocol: unix.IPPROTO_UDP,
HopLimit: 64,
PayloadLength: uint16(payloadLen + udphLen),
}
if ipFn != nil {
ipFn(ipFields)
}
ipv6H.Encode(ipFields)
udpH := header.UDP(b[offset+40:])
udpH.Encode(&header.UDPFields{
SrcPort: srcIPPort.Port(),
DstPort: dstIPPort.Port(),
Length: uint16(payloadLen + udphLen),
})
pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_UDP, ipv6H.SourceAddress(), ipv6H.DestinationAddress(), uint16(udphLen+payloadLen))
udpH.SetChecksum(^udpH.CalculateChecksum(pseudoCsum))
return b
}
func udp4Packet(srcIPPort, dstIPPort netip.AddrPort, payloadLen int) []byte {
return udp4PacketMutateIPFields(srcIPPort, dstIPPort, payloadLen, nil)
}
func tcp4PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32, ipFn func(*header.IPv4Fields)) []byte {
totalLen := 40 + segmentSize
b := make([]byte, offset+int(totalLen), 65535)
ipv4H := header.IPv4(b[offset:])
srcAs4 := srcIPPort.Addr().As4()
dstAs4 := dstIPPort.Addr().As4()
ipFields := &header.IPv4Fields{
SrcAddr: tcpip.AddrFromSlice(srcAs4[:]),
DstAddr: tcpip.AddrFromSlice(dstAs4[:]),
Protocol: unix.IPPROTO_TCP,
TTL: 64,
TotalLength: uint16(totalLen),
}
if ipFn != nil {
ipFn(ipFields)
}
ipv4H.Encode(ipFields)
tcpH := header.TCP(b[offset+20:])
tcpH.Encode(&header.TCPFields{
SrcPort: srcIPPort.Port(),
DstPort: dstIPPort.Port(),
SeqNum: seq,
AckNum: 1,
DataOffset: 20,
Flags: flags,
WindowSize: 3000,
})
ipv4H.SetChecksum(^ipv4H.CalculateChecksum())
pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_TCP, ipv4H.SourceAddress(), ipv4H.DestinationAddress(), uint16(20+segmentSize))
tcpH.SetChecksum(^tcpH.CalculateChecksum(pseudoCsum))
return b
}
func tcp4Packet(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32) []byte {
return tcp4PacketMutateIPFields(srcIPPort, dstIPPort, flags, segmentSize, seq, nil)
}
func tcp6PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32, ipFn func(*header.IPv6Fields)) []byte {
totalLen := 60 + segmentSize
b := make([]byte, offset+int(totalLen), 65535)
ipv6H := header.IPv6(b[offset:])
srcAs16 := srcIPPort.Addr().As16()
dstAs16 := dstIPPort.Addr().As16()
ipFields := &header.IPv6Fields{
SrcAddr: tcpip.AddrFromSlice(srcAs16[:]),
DstAddr: tcpip.AddrFromSlice(dstAs16[:]),
TransportProtocol: unix.IPPROTO_TCP,
HopLimit: 64,
PayloadLength: uint16(segmentSize + 20),
}
if ipFn != nil {
ipFn(ipFields)
}
ipv6H.Encode(ipFields)
tcpH := header.TCP(b[offset+40:])
tcpH.Encode(&header.TCPFields{
SrcPort: srcIPPort.Port(),
DstPort: dstIPPort.Port(),
SeqNum: seq,
AckNum: 1,
DataOffset: 20,
Flags: flags,
WindowSize: 3000,
})
pseudoCsum := header.PseudoHeaderChecksum(unix.IPPROTO_TCP, ipv6H.SourceAddress(), ipv6H.DestinationAddress(), uint16(20+segmentSize))
tcpH.SetChecksum(^tcpH.CalculateChecksum(pseudoCsum))
return b
}
func tcp6Packet(srcIPPort, dstIPPort netip.AddrPort, flags header.TCPFlags, segmentSize, seq uint32) []byte {
return tcp6PacketMutateIPFields(srcIPPort, dstIPPort, flags, segmentSize, seq, nil)
}
func Test_handleVirtioRead(t *testing.T) {
tests := []struct {
name string
hdr virtioNetHdr
pktIn []byte
wantLens []int
wantErr bool
}{
{
"tcp4",
virtioNetHdr{
flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
gsoType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
gsoSize: 100,
hdrLen: 40,
csumStart: 20,
csumOffset: 16,
},
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck|header.TCPFlagPsh, 200, 1),
[]int{140, 140},
false,
},
{
"tcp6",
virtioNetHdr{
flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
gsoType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
gsoSize: 100,
hdrLen: 60,
csumStart: 40,
csumOffset: 16,
},
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck|header.TCPFlagPsh, 200, 1),
[]int{160, 160},
false,
},
{
"udp4",
virtioNetHdr{
flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
gsoType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
gsoSize: 100,
hdrLen: 28,
csumStart: 20,
csumOffset: 6,
},
udp4Packet(ip4PortA, ip4PortB, 200),
[]int{128, 128},
false,
},
{
"udp6",
virtioNetHdr{
flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
gsoType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
gsoSize: 100,
hdrLen: 48,
csumStart: 40,
csumOffset: 6,
},
udp6Packet(ip6PortA, ip6PortB, 200),
[]int{148, 148},
false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
out := make([][]byte, conn.IdealBatchSize)
sizes := make([]int, conn.IdealBatchSize)
for i := range out {
out[i] = make([]byte, 65535)
}
tt.hdr.encode(tt.pktIn)
n, err := handleVirtioRead(tt.pktIn, out, sizes, offset)
if err != nil {
if tt.wantErr {
return
}
t.Fatalf("got err: %v", err)
}
if n != len(tt.wantLens) {
t.Fatalf("got %d packets, wanted %d", n, len(tt.wantLens))
}
for i := range tt.wantLens {
if tt.wantLens[i] != sizes[i] {
t.Fatalf("wantLens[%d]: %d != outSizes: %d", i, tt.wantLens[i], sizes[i])
}
}
})
}
}
func flipTCP4Checksum(b []byte) []byte {
at := virtioNetHdrLen + 20 + 16 // 20 byte ipv4 header; tcp csum offset is 16
b[at] ^= 0xFF
b[at+1] ^= 0xFF
return b
}
func flipUDP4Checksum(b []byte) []byte {
at := virtioNetHdrLen + 20 + 6 // 20 byte ipv4 header; udp csum offset is 6
b[at] ^= 0xFF
b[at+1] ^= 0xFF
return b
}
func Fuzz_handleGRO(f *testing.F) {
pkt0 := tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1)
pkt1 := tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101)
pkt2 := tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201)
pkt3 := tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1)
pkt4 := tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101)
pkt5 := tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201)
pkt6 := udp4Packet(ip4PortA, ip4PortB, 100)
pkt7 := udp4Packet(ip4PortA, ip4PortB, 100)
pkt8 := udp4Packet(ip4PortA, ip4PortC, 100)
pkt9 := udp6Packet(ip6PortA, ip6PortB, 100)
pkt10 := udp6Packet(ip6PortA, ip6PortB, 100)
pkt11 := udp6Packet(ip6PortA, ip6PortC, 100)
f.Add(pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11, 0, offset)
f.Fuzz(func(t *testing.T, pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11 []byte, gro int, offset int) {
pkts := [][]byte{pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11}
toWrite := make([]int, 0, len(pkts))
handleGRO(pkts, offset, newTCPGROTable(), newUDPGROTable(), groDisablementFlags(gro), &toWrite)
if len(toWrite) > len(pkts) {
t.Errorf("len(toWrite): %d > len(pkts): %d", len(toWrite), len(pkts))
}
seenWriteI := make(map[int]bool)
for _, writeI := range toWrite {
if writeI < 0 || writeI > len(pkts)-1 {
t.Errorf("toWrite value (%d) outside bounds of len(pkts): %d", writeI, len(pkts))
}
if seenWriteI[writeI] {
t.Errorf("duplicate toWrite value: %d", writeI)
}
seenWriteI[writeI] = true
}
})
}
func Test_handleGRO(t *testing.T) {
tests := []struct {
name string
pktsIn [][]byte
gro groDisablementFlags
wantToWrite []int
wantLens []int
wantErr bool
}{
{
"multiple protocols and flows",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // tcp4 flow 1
udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1
udp4Packet(ip4PortA, ip4PortC, 100), // udp4 flow 2
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // tcp4 flow 1
tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201), // tcp4 flow 2
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // tcp6 flow 1
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101), // tcp6 flow 1
tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201), // tcp6 flow 2
udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1
udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1
udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1
},
0,
[]int{0, 1, 2, 4, 5, 7, 9},
[]int{240, 228, 128, 140, 260, 160, 248},
false,
},
{
"multiple protocols and flows no UDP GRO",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // tcp4 flow 1
udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1
udp4Packet(ip4PortA, ip4PortC, 100), // udp4 flow 2
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // tcp4 flow 1
tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201), // tcp4 flow 2
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // tcp6 flow 1
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101), // tcp6 flow 1
tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201), // tcp6 flow 2
udp4Packet(ip4PortA, ip4PortB, 100), // udp4 flow 1
udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1
udp6Packet(ip6PortA, ip6PortB, 100), // udp6 flow 1
},
udpGRODisabled,
[]int{0, 1, 2, 4, 5, 7, 8, 9, 10},
[]int{240, 128, 128, 140, 260, 160, 128, 148, 148},
false,
},
{
"PSH interleaved",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // v4 flow 1
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck|header.TCPFlagPsh, 100, 101), // v4 flow 1
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 201), // v4 flow 1
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 301), // v4 flow 1
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // v6 flow 1
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck|header.TCPFlagPsh, 100, 101), // v6 flow 1
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 201), // v6 flow 1
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 301), // v6 flow 1
},
0,
[]int{0, 2, 4, 6},
[]int{240, 240, 260, 260},
false,
},
{
"coalesceItemInvalidCSum",
[][]byte{
flipTCP4Checksum(tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1)), // v4 flow 1 seq 1 len 100
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // v4 flow 1 seq 101 len 100
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 201), // v4 flow 1 seq 201 len 100
flipUDP4Checksum(udp4Packet(ip4PortA, ip4PortB, 100)),
udp4Packet(ip4PortA, ip4PortB, 100),
udp4Packet(ip4PortA, ip4PortB, 100),
},
0,
[]int{0, 1, 3, 4},
[]int{140, 240, 128, 228},
false,
},
{
"out of order",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // v4 flow 1 seq 101 len 100
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // v4 flow 1 seq 1 len 100
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 201), // v4 flow 1 seq 201 len 100
},
0,
[]int{0},
[]int{340},
false,
},
{
"unequal TTL",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1),
tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) {
fields.TTL++
}),
udp4Packet(ip4PortA, ip4PortB, 100),
udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) {
fields.TTL++
}),
},
0,
[]int{0, 1, 2, 3},
[]int{140, 140, 128, 128},
false,
},
{
"unequal ToS",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1),
tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) {
fields.TOS++
}),
udp4Packet(ip4PortA, ip4PortB, 100),
udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) {
fields.TOS++
}),
},
0,
[]int{0, 1, 2, 3},
[]int{140, 140, 128, 128},
false,
},
{
"unequal flags more fragments set",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1),
tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) {
fields.Flags = 1
}),
udp4Packet(ip4PortA, ip4PortB, 100),
udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) {
fields.Flags = 1
}),
},
0,
[]int{0, 1, 2, 3},
[]int{140, 140, 128, 128},
false,
},
{
"unequal flags DF set",
[][]byte{
tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1),
tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) {
fields.Flags = 2
}),
udp4Packet(ip4PortA, ip4PortB, 100),
udp4PacketMutateIPFields(ip4PortA, ip4PortB, 100, func(fields *header.IPv4Fields) {
fields.Flags = 2
}),
},
0,
[]int{0, 1, 2, 3},
[]int{140, 140, 128, 128},
false,
},
{
"ipv6 unequal hop limit",
[][]byte{
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1),
tcp6PacketMutateIPFields(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv6Fields) {
fields.HopLimit++
}),
udp6Packet(ip6PortA, ip6PortB, 100),
udp6PacketMutateIPFields(ip6PortA, ip6PortB, 100, func(fields *header.IPv6Fields) {
fields.HopLimit++
}),
},
0,
[]int{0, 1, 2, 3},
[]int{160, 160, 148, 148},
false,
},
{
"ipv6 unequal traffic class",
[][]byte{
tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1),
tcp6PacketMutateIPFields(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv6Fields) {
fields.TrafficClass++
}),
udp6Packet(ip6PortA, ip6PortB, 100),
udp6PacketMutateIPFields(ip6PortA, ip6PortB, 100, func(fields *header.IPv6Fields) {
fields.TrafficClass++
}),
},
0,
[]int{0, 1, 2, 3},
[]int{160, 160, 148, 148},
false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
toWrite := make([]int, 0, len(tt.pktsIn))
err := handleGRO(tt.pktsIn, offset, newTCPGROTable(), newUDPGROTable(), tt.gro, &toWrite)
if err != nil {
if tt.wantErr {
return
}
t.Fatalf("got err: %v", err)
}
if len(toWrite) != len(tt.wantToWrite) {
t.Fatalf("got %d packets, wanted %d", len(toWrite), len(tt.wantToWrite))
}
for i, pktI := range tt.wantToWrite {
if tt.wantToWrite[i] != toWrite[i] {
t.Fatalf("wantToWrite[%d]: %d != toWrite: %d", i, tt.wantToWrite[i], toWrite[i])
}
if tt.wantLens[i] != len(tt.pktsIn[pktI][offset:]) {
t.Errorf("wanted len %d packet at %d, got: %d", tt.wantLens[i], i, len(tt.pktsIn[pktI][offset:]))
}
}
})
}
}
func Test_packetIsGROCandidate(t *testing.T) {
tcp4 := tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1)[virtioNetHdrLen:]
tcp4TooShort := tcp4[:39]
ip4InvalidHeaderLen := make([]byte, len(tcp4))
copy(ip4InvalidHeaderLen, tcp4)
ip4InvalidHeaderLen[0] = 0x46
ip4InvalidProtocol := make([]byte, len(tcp4))
copy(ip4InvalidProtocol, tcp4)
ip4InvalidProtocol[9] = unix.IPPROTO_GRE
tcp6 := tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1)[virtioNetHdrLen:]
tcp6TooShort := tcp6[:59]
ip6InvalidProtocol := make([]byte, len(tcp6))
copy(ip6InvalidProtocol, tcp6)
ip6InvalidProtocol[6] = unix.IPPROTO_GRE
udp4 := udp4Packet(ip4PortA, ip4PortB, 100)[virtioNetHdrLen:]
udp4TooShort := udp4[:27]
udp6 := udp6Packet(ip6PortA, ip6PortB, 100)[virtioNetHdrLen:]
udp6TooShort := udp6[:47]
tests := []struct {
name string
b []byte
gro groDisablementFlags
want groCandidateType
}{
{
"tcp4",
tcp4,
0,
tcp4GROCandidate,
},
{
"tcp4 no support",
tcp4,
tcpGRODisabled,
notGROCandidate,
},
{
"tcp6",
tcp6,
0,
tcp6GROCandidate,
},
{
"tcp6 no support",
tcp6,
tcpGRODisabled,
notGROCandidate,
},
{
"udp4",
udp4,
0,
udp4GROCandidate,
},
{
"udp4 no support",
udp4,
udpGRODisabled,
notGROCandidate,
},
{
"udp6",
udp6,
0,
udp6GROCandidate,
},
{
"udp6 no support",
udp6,
udpGRODisabled,
notGROCandidate,
},
{
"udp4 too short",
udp4TooShort,
0,
notGROCandidate,
},
{
"udp6 too short",
udp6TooShort,
0,
notGROCandidate,
},
{
"tcp4 too short",
tcp4TooShort,
0,
notGROCandidate,
},
{
"tcp6 too short",
tcp6TooShort,
0,
notGROCandidate,
},
{
"invalid IP version",
[]byte{0x00},
0,
notGROCandidate,
},
{
"invalid IP header len",
ip4InvalidHeaderLen,
0,
notGROCandidate,
},
{
"ip4 invalid protocol",
ip4InvalidProtocol,
0,
notGROCandidate,
},
{
"ip6 invalid protocol",
ip6InvalidProtocol,
0,
notGROCandidate,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := packetIsGROCandidate(tt.b, tt.gro); got != tt.want {
t.Errorf("packetIsGROCandidate() = %v, want %v", got, tt.want)
}
})
}
}
func Test_udpPacketsCanCoalesce(t *testing.T) {
udp4a := udp4Packet(ip4PortA, ip4PortB, 100)
udp4b := udp4Packet(ip4PortA, ip4PortB, 100)
udp4c := udp4Packet(ip4PortA, ip4PortB, 110)
type args struct {
pkt []byte
iphLen uint8
gsoSize uint16
item udpGROItem
bufs [][]byte
bufsOffset int
}
tests := []struct {
name string
args args
want canCoalesce
}{
{
"coalesceAppend equal gso",
args{
pkt: udp4a[offset:],
iphLen: 20,
gsoSize: 100,
item: udpGROItem{
gsoSize: 100,
iphLen: 20,
},
bufs: [][]byte{
udp4a,
udp4b,
},
bufsOffset: offset,
},
coalesceAppend,
},
{
"coalesceAppend smaller gso",
args{
pkt: udp4a[offset : len(udp4a)-90],
iphLen: 20,
gsoSize: 10,
item: udpGROItem{
gsoSize: 100,
iphLen: 20,
},
bufs: [][]byte{
udp4a,
udp4b,
},
bufsOffset: offset,
},
coalesceAppend,
},
{
"coalesceUnavailable smaller gso previously appended",
args{
pkt: udp4a[offset:],
iphLen: 20,
gsoSize: 100,
item: udpGROItem{
gsoSize: 100,
iphLen: 20,
},
bufs: [][]byte{
udp4c,
udp4b,
},
bufsOffset: offset,
},
coalesceUnavailable,
},
{
"coalesceUnavailable larger following smaller",
args{
pkt: udp4c[offset:],
iphLen: 20,
gsoSize: 110,
item: udpGROItem{
gsoSize: 100,
iphLen: 20,
},
bufs: [][]byte{
udp4a,
udp4c,
},
bufsOffset: offset,
},
coalesceUnavailable,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := udpPacketsCanCoalesce(tt.args.pkt, tt.args.iphLen, tt.args.gsoSize, tt.args.item, tt.args.bufs, tt.args.bufsOffset); got != tt.want {
t.Errorf("udpPacketsCanCoalesce() = %v, want %v", got, tt.want)
}
})
}
}

View file

@ -1,95 +0,0 @@
package tun
import (
"net/netip"
"testing"
"github.com/tailscale/wireguard-go/conn"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header"
)
func Fuzz_GSOSplit(f *testing.F) {
const segmentSize = 100
tcpFields := &header.TCPFields{
SrcPort: 1,
DstPort: 1,
SeqNum: 1,
AckNum: 1,
DataOffset: 20,
Flags: header.TCPFlagAck | header.TCPFlagPsh,
WindowSize: 3000,
}
udpFields := &header.UDPFields{
SrcPort: 1,
DstPort: 1,
Length: 8 + segmentSize,
}
gsoTCPv4 := make([]byte, 20+20+segmentSize)
header.IPv4(gsoTCPv4).Encode(&header.IPv4Fields{
SrcAddr: tcpip.AddrFromSlice(netip.MustParseAddr("192.0.2.1").AsSlice()),
DstAddr: tcpip.AddrFromSlice(netip.MustParseAddr("192.0.2.2").AsSlice()),
Protocol: ipProtoTCP,
TTL: 64,
TotalLength: uint16(len(gsoTCPv4)),
})
header.TCP(gsoTCPv4[20:]).Encode(tcpFields)
gsoUDPv4 := make([]byte, 20+8+segmentSize)
header.IPv4(gsoUDPv4).Encode(&header.IPv4Fields{
SrcAddr: tcpip.AddrFromSlice(netip.MustParseAddr("192.0.2.1").AsSlice()),
DstAddr: tcpip.AddrFromSlice(netip.MustParseAddr("192.0.2.2").AsSlice()),
Protocol: ipProtoUDP,
TTL: 64,
TotalLength: uint16(len(gsoUDPv4)),
})
header.UDP(gsoTCPv4[20:]).Encode(udpFields)
gsoTCPv6 := make([]byte, 40+20+segmentSize)
header.IPv6(gsoTCPv6).Encode(&header.IPv6Fields{
SrcAddr: tcpip.AddrFromSlice(netip.MustParseAddr("2001:db8::1").AsSlice()),
DstAddr: tcpip.AddrFromSlice(netip.MustParseAddr("2001:db8::2").AsSlice()),
TransportProtocol: ipProtoTCP,
HopLimit: 64,
PayloadLength: uint16(20 + segmentSize),
})
header.TCP(gsoTCPv6[40:]).Encode(tcpFields)
gsoUDPv6 := make([]byte, 40+8+segmentSize)
header.IPv6(gsoUDPv6).Encode(&header.IPv6Fields{
SrcAddr: tcpip.AddrFromSlice(netip.MustParseAddr("2001:db8::1").AsSlice()),
DstAddr: tcpip.AddrFromSlice(netip.MustParseAddr("2001:db8::2").AsSlice()),
TransportProtocol: ipProtoUDP,
HopLimit: 64,
PayloadLength: uint16(8 + segmentSize),
})
header.UDP(gsoUDPv6[20:]).Encode(udpFields)
out := make([][]byte, conn.IdealBatchSize)
for i := range out {
out[i] = make([]byte, 65535)
}
sizes := make([]int, conn.IdealBatchSize)
f.Add(gsoTCPv4, int(GSOTCPv4), uint16(40), uint16(20), uint16(16), uint16(100), false)
f.Add(gsoUDPv4, int(GSOUDPL4), uint16(28), uint16(20), uint16(6), uint16(100), false)
f.Add(gsoTCPv6, int(GSOTCPv6), uint16(60), uint16(40), uint16(16), uint16(100), false)
f.Add(gsoUDPv6, int(GSOUDPL4), uint16(48), uint16(40), uint16(6), uint16(100), false)
f.Fuzz(func(t *testing.T, pkt []byte, gsoType int, hdrLen, csumStart, csumOffset, gsoSize uint16, needsCsum bool) {
options := GSOOptions{
GSOType: GSOType(gsoType),
HdrLen: hdrLen,
CsumStart: csumStart,
CsumOffset: csumOffset,
GSOSize: gsoSize,
NeedsCsum: needsCsum,
}
n, _ := GSOSplit(pkt, options, out, sizes, 0)
if n > len(sizes) {
t.Errorf("n (%d) > len(sizes): %d", n, len(sizes))
}
})
}

View file

@ -1,155 +0,0 @@
/* SPDX-License-Identifier: MIT
*
* Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
*/
package tuntest
import (
"encoding/binary"
"io"
"net/netip"
"os"
"github.com/tailscale/wireguard-go/tun"
)
func Ping(dst, src netip.Addr) []byte {
localPort := uint16(1337)
seq := uint16(0)
payload := make([]byte, 4)
binary.BigEndian.PutUint16(payload[0:], localPort)
binary.BigEndian.PutUint16(payload[2:], seq)
return genICMPv4(payload, dst, src)
}
// Checksum is the "internet checksum" from https://tools.ietf.org/html/rfc1071.
func checksum(buf []byte, initial uint16) uint16 {
v := uint32(initial)
for i := 0; i < len(buf)-1; i += 2 {
v += uint32(binary.BigEndian.Uint16(buf[i:]))
}
if len(buf)%2 == 1 {
v += uint32(buf[len(buf)-1]) << 8
}
for v > 0xffff {
v = (v >> 16) + (v & 0xffff)
}
return ^uint16(v)
}
func genICMPv4(payload []byte, dst, src netip.Addr) []byte {
const (
icmpv4ProtocolNumber = 1
icmpv4Echo = 8
icmpv4ChecksumOffset = 2
icmpv4Size = 8
ipv4Size = 20
ipv4TotalLenOffset = 2
ipv4ChecksumOffset = 10
ttl = 65
headerSize = ipv4Size + icmpv4Size
)
pkt := make([]byte, headerSize+len(payload))
ip := pkt[0:ipv4Size]
icmpv4 := pkt[ipv4Size : ipv4Size+icmpv4Size]
// https://tools.ietf.org/html/rfc792
icmpv4[0] = icmpv4Echo // type
icmpv4[1] = 0 // code
chksum := ^checksum(icmpv4, checksum(payload, 0))
binary.BigEndian.PutUint16(icmpv4[icmpv4ChecksumOffset:], chksum)
// https://tools.ietf.org/html/rfc760 section 3.1
length := uint16(len(pkt))
ip[0] = (4 << 4) | (ipv4Size / 4)
binary.BigEndian.PutUint16(ip[ipv4TotalLenOffset:], length)
ip[8] = ttl
ip[9] = icmpv4ProtocolNumber
copy(ip[12:], src.AsSlice())
copy(ip[16:], dst.AsSlice())
chksum = ^checksum(ip[:], 0)
binary.BigEndian.PutUint16(ip[ipv4ChecksumOffset:], chksum)
copy(pkt[headerSize:], payload)
return pkt
}
type ChannelTUN struct {
Inbound chan []byte // incoming packets, closed on TUN close
Outbound chan []byte // outbound packets, blocks forever on TUN close
closed chan struct{}
events chan tun.Event
tun chTun
}
func NewChannelTUN() *ChannelTUN {
c := &ChannelTUN{
Inbound: make(chan []byte),
Outbound: make(chan []byte),
closed: make(chan struct{}),
events: make(chan tun.Event, 1),
}
c.tun.c = c
c.events <- tun.EventUp
return c
}
func (c *ChannelTUN) TUN() tun.Device {
return &c.tun
}
type chTun struct {
c *ChannelTUN
}
func (t *chTun) File() *os.File { return nil }
func (t *chTun) Read(packets [][]byte, sizes []int, offset int) (int, error) {
select {
case <-t.c.closed:
return 0, os.ErrClosed
case msg := <-t.c.Outbound:
n := copy(packets[0][offset:], msg)
sizes[0] = n
return 1, nil
}
}
// Write is called by the wireguard device to deliver a packet for routing.
func (t *chTun) Write(packets [][]byte, offset int) (int, error) {
if offset == -1 {
close(t.c.closed)
close(t.c.events)
return 0, io.EOF
}
for i, data := range packets {
msg := make([]byte, len(data)-offset)
copy(msg, data[offset:])
select {
case <-t.c.closed:
return i, os.ErrClosed
case t.c.Inbound <- msg:
}
}
return len(packets), nil
}
func (t *chTun) BatchSize() int {
return 1
}
const DefaultMTU = 1420
func (t *chTun) MTU() (int, error) { return DefaultMTU, nil }
func (t *chTun) Name() (string, error) { return "loopbackTun1", nil }
func (t *chTun) Events() <-chan tun.Event { return t.c.events }
func (t *chTun) Close() error {
t.Write(nil, -1)
return nil
}