From 8f1a6a10b2a74ae86920a74e7f0989e9f57df990 Mon Sep 17 00:00:00 2001
From: Mark Puha
Date: Fri, 6 Oct 2023 02:11:27 +0530
Subject: [PATCH 01/75] Advanced security (#2)
* Advanced security header layer & config
---
.gitignore | 2 +-
conn/bind_windows.go | 2 +-
conn/bindtest/bindtest.go | 2 +-
device/bind_test.go | 2 +-
device/device.go | 278 ++++++++++++++++++++++++++-
device/device_test.go | 150 ++++++++++++---
device/keypair.go | 2 +-
device/noise-protocol.go | 26 ++-
device/noise_test.go | 4 +-
device/peer.go | 2 +-
device/queueconstants_android.go | 2 +-
device/queueconstants_default.go | 2 +-
device/receive.go | 59 ++++--
device/send.go | 104 ++++++++--
device/sticky_default.go | 4 +-
device/sticky_linux.go | 4 +-
device/tun.go | 2 +-
device/uapi.go | 146 ++++++++++++--
device/util.go | 25 +++
device/util_test.go | 27 +++
go.mod | 3 +-
go.sum | 2 +
ipc/namedpipe/namedpipe_test.go | 2 +-
ipc/uapi_linux.go | 2 +-
ipc/uapi_windows.go | 2 +-
main.go | 8 +-
main_windows.go | 8 +-
tun/netstack/examples/http_client.go | 6 +-
tun/netstack/examples/http_server.go | 6 +-
tun/netstack/examples/ping_client.go | 6 +-
tun/netstack/tun.go | 2 +-
tun/tcp_offload_linux.go | 2 +-
tun/tcp_offload_linux_test.go | 2 +-
tun/tun_linux.go | 4 +-
tun/tuntest/tuntest.go | 2 +-
35 files changed, 781 insertions(+), 121 deletions(-)
create mode 100644 device/util.go
create mode 100644 device/util_test.go
diff --git a/.gitignore b/.gitignore
index e460293..71549f4 100644
--- a/.gitignore
+++ b/.gitignore
@@ -1 +1 @@
-wireguard-go
+wireguard-go
\ No newline at end of file
diff --git a/conn/bind_windows.go b/conn/bind_windows.go
index d5095e0..9bad0ee 100644
--- a/conn/bind_windows.go
+++ b/conn/bind_windows.go
@@ -17,7 +17,7 @@ import (
"golang.org/x/sys/windows"
- "golang.zx2c4.com/wireguard/conn/winrio"
+ "github.com/amnezia-vpn/amnezia-wg/conn/winrio"
)
const (
diff --git a/conn/bindtest/bindtest.go b/conn/bindtest/bindtest.go
index 74e7add..713c371 100644
--- a/conn/bindtest/bindtest.go
+++ b/conn/bindtest/bindtest.go
@@ -12,7 +12,7 @@ import (
"net/netip"
"os"
- "golang.zx2c4.com/wireguard/conn"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
)
type ChannelBind struct {
diff --git a/device/bind_test.go b/device/bind_test.go
index 302a521..eae36c2 100644
--- a/device/bind_test.go
+++ b/device/bind_test.go
@@ -8,7 +8,7 @@ package device
import (
"errors"
- "golang.zx2c4.com/wireguard/conn"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
)
type DummyDatagram struct {
diff --git a/device/device.go b/device/device.go
index 1af9fe0..10365d1 100644
--- a/device/device.go
+++ b/device/device.go
@@ -11,10 +11,12 @@ import (
"sync/atomic"
"time"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/ratelimiter"
- "golang.zx2c4.com/wireguard/rwcancel"
- "golang.zx2c4.com/wireguard/tun"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/ipc"
+ "github.com/amnezia-vpn/amnezia-wg/ratelimiter"
+ "github.com/amnezia-vpn/amnezia-wg/rwcancel"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/tevino/abool/v2"
)
type Device struct {
@@ -89,6 +91,22 @@ type Device struct {
ipcMutex sync.RWMutex
closed chan struct{}
log *Logger
+
+ isASecOn abool.AtomicBool
+ aSecMux sync.RWMutex
+ aSecCfg aSecCfgType
+}
+
+type aSecCfgType struct {
+ junkPacketCount int
+ junkPacketMinSize int
+ junkPacketMaxSize int
+ initPacketJunkSize int
+ responsePacketJunkSize int
+ initPacketMagicHeader uint32
+ responsePacketMagicHeader uint32
+ underloadPacketMagicHeader uint32
+ transportPacketMagicHeader uint32
}
// deviceState represents the state of a Device.
@@ -162,7 +180,8 @@ func (device *Device) changeState(want deviceState) (err error) {
err = errDown
}
}
- device.log.Verbosef("Interface state was %s, requested %s, now %s", old, want, device.deviceState())
+ device.log.Verbosef(
+ "Interface state was %s, requested %s, now %s", old, want, device.deviceState())
return
}
@@ -526,7 +545,7 @@ func (device *Device) BindUpdate() error {
// start receiving routines
device.net.stopping.Add(len(recvFns))
device.queue.decryption.wg.Add(len(recvFns)) // each RoutineReceiveIncoming goroutine writes to device.queue.decryption
- device.queue.handshake.wg.Add(len(recvFns)) // each RoutineReceiveIncoming goroutine writes to device.queue.handshake
+ device.queue.handshake.wg.Add(len(recvFns)) // each RoutineReceiveIncoming goroutine writes to device.queue.handshake
batchSize := netc.bind.BatchSize()
for _, fn := range recvFns {
go device.RoutineReceiveIncoming(batchSize, fn)
@@ -542,3 +561,250 @@ func (device *Device) BindClose() error {
device.net.Unlock()
return err
}
+func (device *Device) isAdvancedSecurityOn() bool {
+ return device.isASecOn.IsSet()
+}
+
+func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
+
+ if tempASecCfg.junkPacketCount == 0 &&
+ tempASecCfg.junkPacketMaxSize == 0 &&
+ tempASecCfg.junkPacketMinSize == 0 &&
+ tempASecCfg.initPacketJunkSize == 0 &&
+ tempASecCfg.responsePacketJunkSize == 0 &&
+ tempASecCfg.initPacketMagicHeader == 0 &&
+ tempASecCfg.responsePacketMagicHeader == 0 &&
+ tempASecCfg.underloadPacketMagicHeader == 0 &&
+ tempASecCfg.transportPacketMagicHeader == 0 {
+ return err
+ }
+
+ isASecOn := false
+ device.aSecMux.Lock()
+ if tempASecCfg.junkPacketCount < 0 {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "JunkPacketCount should be non negative",
+ )
+ }
+ device.aSecCfg.junkPacketCount = tempASecCfg.junkPacketCount
+ if tempASecCfg.junkPacketCount != 0 {
+ isASecOn = true
+ }
+
+ device.aSecCfg.junkPacketMinSize = tempASecCfg.junkPacketMinSize
+ if tempASecCfg.junkPacketMinSize != 0 {
+ isASecOn = true
+ }
+
+ if device.aSecCfg.junkPacketCount > 0 &&
+ tempASecCfg.junkPacketMaxSize == tempASecCfg.junkPacketMinSize {
+
+ tempASecCfg.junkPacketMaxSize++ // to make rand gen work
+ }
+
+ if tempASecCfg.junkPacketMaxSize >= MaxSegmentSize{
+ device.aSecCfg.junkPacketMinSize = 0
+ device.aSecCfg.junkPacketMaxSize = 1
+ if err != nil {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d; %w",
+ tempASecCfg.junkPacketMaxSize,
+ MaxSegmentSize,
+ err,
+ )
+ } else {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d",
+ tempASecCfg.junkPacketMaxSize,
+ MaxSegmentSize,
+ )
+ }
+ } else if tempASecCfg.junkPacketMaxSize < tempASecCfg.junkPacketMinSize {
+ if err != nil {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "maxSize: %d; should be greater than minSize: %d; %w",
+ tempASecCfg.junkPacketMaxSize,
+ tempASecCfg.junkPacketMinSize,
+ err,
+ )
+ } else {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "maxSize: %d; should be greater than minSize: %d",
+ tempASecCfg.junkPacketMaxSize,
+ tempASecCfg.junkPacketMinSize,
+ )
+ }
+ } else {
+ device.aSecCfg.junkPacketMaxSize = tempASecCfg.junkPacketMaxSize
+ }
+
+ if tempASecCfg.junkPacketMaxSize != 0 {
+ isASecOn = true
+ }
+
+ if MessageInitiationSize+tempASecCfg.initPacketJunkSize >= MaxSegmentSize {
+ if err != nil {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d; %w`,
+ tempASecCfg.initPacketJunkSize,
+ MaxSegmentSize,
+ err,
+ )
+ } else {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempASecCfg.initPacketJunkSize,
+ MaxSegmentSize,
+ )
+ }
+ } else {
+ device.aSecCfg.initPacketJunkSize = tempASecCfg.initPacketJunkSize
+ }
+
+ if tempASecCfg.initPacketJunkSize != 0 {
+ isASecOn = true
+ }
+
+ if MessageResponseSize+tempASecCfg.responsePacketJunkSize >= MaxSegmentSize {
+ if err != nil {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d; %w`,
+ tempASecCfg.responsePacketJunkSize,
+ MaxSegmentSize,
+ err,
+ )
+ } else {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempASecCfg.responsePacketJunkSize,
+ MaxSegmentSize,
+ )
+ }
+ } else {
+ device.aSecCfg.responsePacketJunkSize = tempASecCfg.responsePacketJunkSize
+ }
+
+ if tempASecCfg.responsePacketJunkSize != 0 {
+ isASecOn = true
+ }
+
+ if tempASecCfg.initPacketMagicHeader > 4 {
+ isASecOn = true
+ device.log.Verbosef("UAPI: Updating init_packet_magic_header")
+ device.aSecCfg.initPacketMagicHeader = tempASecCfg.initPacketMagicHeader
+ MessageInitiationType = device.aSecCfg.initPacketMagicHeader
+ } else {
+ device.log.Verbosef("UAPI: Using default init type")
+ MessageInitiationType = 1
+ }
+
+ if tempASecCfg.responsePacketMagicHeader > 4 {
+ isASecOn = true
+ device.log.Verbosef("UAPI: Updating response_packet_magic_header")
+ device.aSecCfg.responsePacketMagicHeader = tempASecCfg.responsePacketMagicHeader
+ MessageResponseType = device.aSecCfg.responsePacketMagicHeader
+ } else {
+ device.log.Verbosef("UAPI: Using default response type")
+ MessageResponseType = 2
+ }
+
+ if tempASecCfg.underloadPacketMagicHeader > 4 {
+ isASecOn = true
+ device.log.Verbosef("UAPI: Updating underload_packet_magic_header")
+ device.aSecCfg.underloadPacketMagicHeader = tempASecCfg.underloadPacketMagicHeader
+ MessageCookieReplyType = device.aSecCfg.underloadPacketMagicHeader
+ } else {
+ device.log.Verbosef("UAPI: Using default underload type")
+ MessageCookieReplyType = 3
+ }
+
+ if tempASecCfg.transportPacketMagicHeader > 4 {
+ isASecOn = true
+ device.log.Verbosef("UAPI: Updating transport_packet_magic_header")
+ device.aSecCfg.transportPacketMagicHeader = tempASecCfg.transportPacketMagicHeader
+ MessageTransportType = device.aSecCfg.transportPacketMagicHeader
+ } else {
+ device.log.Verbosef("UAPI: Using default transport type")
+ MessageTransportType = 4
+ }
+
+ isSameMap := map[uint32]bool{}
+ isSameMap[MessageInitiationType] = true
+ isSameMap[MessageResponseType] = true
+ isSameMap[MessageCookieReplyType] = true
+ isSameMap[MessageTransportType] = true
+
+ // size will be different if same values
+ if len(isSameMap) != 4 {
+ if err != nil {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `magic headers should differ; got: init:%d; recv:%d; unde:%d; tran:%d; %w`,
+ MessageInitiationType,
+ MessageResponseType,
+ MessageCookieReplyType,
+ MessageTransportType,
+ err,
+ )
+ } else {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `magic headers should differ; got: init:%d; recv:%d; unde:%d; tran:%d`,
+ MessageInitiationType,
+ MessageResponseType,
+ MessageCookieReplyType,
+ MessageTransportType,
+ )
+ }
+ }
+
+ newInitSize := MessageInitiationSize + device.aSecCfg.initPacketJunkSize
+ newResponseSize := MessageResponseSize + device.aSecCfg.responsePacketJunkSize
+
+ if newInitSize == newResponseSize {
+ if err != nil {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `new init size:%d; and new response size:%d; should differ; %w`,
+ newInitSize,
+ newResponseSize,
+ err,
+ )
+ } else {
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `new init size:%d; and new response size:%d; should differ`,
+ newInitSize,
+ newResponseSize,
+ )
+ }
+ } else {
+ packetSizeToMsgType = map[int]uint32{
+ newInitSize: MessageInitiationType,
+ newResponseSize: MessageResponseType,
+ MessageCookieReplySize: MessageCookieReplyType,
+ MessageTransportSize: MessageTransportType,
+ }
+
+ msgTypeToJunkSize = map[uint32]int{
+ MessageInitiationType: device.aSecCfg.initPacketJunkSize,
+ MessageResponseType: device.aSecCfg.responsePacketJunkSize,
+ MessageCookieReplyType: 0,
+ MessageTransportType: 0,
+ }
+ }
+
+ device.isASecOn.SetTo(isASecOn)
+ device.aSecMux.Unlock()
+
+ return err
+}
diff --git a/device/device_test.go b/device/device_test.go
index fff172b..afa1dc3 100644
--- a/device/device_test.go
+++ b/device/device_test.go
@@ -20,10 +20,10 @@ import (
"testing"
"time"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/conn/bindtest"
- "golang.zx2c4.com/wireguard/tun"
- "golang.zx2c4.com/wireguard/tun/tuntest"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/conn/bindtest"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amnezia-wg/tun/tuntest"
)
// uapiCfg returns a string that contains cfg formatted use with IpcSet.
@@ -91,6 +91,65 @@ func genConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
return
}
+func genASecurityConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
+ var key1, key2 NoisePrivateKey
+ _, err := rand.Read(key1[:])
+ if err != nil {
+ tb.Errorf("unable to generate private key random bytes: %v", err)
+ }
+ _, err = rand.Read(key2[:])
+ if err != nil {
+ tb.Errorf("unable to generate private key random bytes: %v", err)
+ }
+ pub1, pub2 := key1.publicKey(), key2.publicKey()
+
+ cfgs[0] = uapiCfg(
+ "private_key", hex.EncodeToString(key1[:]),
+ "listen_port", "0",
+ "replace_peers", "true",
+ "jc", "5",
+ "jmin", "500",
+ "jmax", "501",
+ "s1", "30",
+ "s2", "40",
+ "h1", "123456",
+ "h2", "67543",
+ "h4", "32345",
+ "h3", "123123",
+ "public_key", hex.EncodeToString(pub2[:]),
+ "protocol_version", "1",
+ "replace_allowed_ips", "true",
+ "allowed_ip", "1.0.0.2/32",
+ )
+ endpointCfgs[0] = uapiCfg(
+ "public_key", hex.EncodeToString(pub2[:]),
+ "endpoint", "127.0.0.1:%d",
+ )
+ cfgs[1] = uapiCfg(
+ "private_key", hex.EncodeToString(key2[:]),
+ "listen_port", "0",
+ "replace_peers", "true",
+ "jc", "5",
+ "jmin", "500",
+ "jmax", "501",
+ "s1", "30",
+ "s2", "40",
+ "h1", "123456",
+ "h2", "67543",
+ "h4", "32345",
+ "h3", "123123",
+ "public_key", hex.EncodeToString(pub1[:]),
+ "protocol_version", "1",
+ "replace_allowed_ips", "true",
+ "allowed_ip", "1.0.0.1/32",
+ )
+ endpointCfgs[1] = uapiCfg(
+ "public_key", hex.EncodeToString(pub1[:]),
+ "endpoint", "127.0.0.1:%d",
+ )
+ return
+}
+
// A testPair is a pair of testPeers.
type testPair [2]testPeer
@@ -115,7 +174,11 @@ func (d SendDirection) String() string {
return "pong"
}
-func (pair *testPair) Send(tb testing.TB, ping SendDirection, done chan struct{}) {
+func (pair *testPair) Send(
+ tb testing.TB,
+ ping SendDirection,
+ done chan struct{},
+) {
tb.Helper()
p0, p1 := pair[0], pair[1]
if !ping {
@@ -149,8 +212,16 @@ func (pair *testPair) Send(tb testing.TB, ping SendDirection, done chan struct{}
}
// genTestPair creates a testPair.
-func genTestPair(tb testing.TB, realSocket bool) (pair testPair) {
- cfg, endpointCfg := genConfigs(tb)
+func genTestPair(
+ tb testing.TB,
+ realSocket, withASecurity bool,
+) (pair testPair) {
+ var cfg, endpointCfg [2]string
+ if withASecurity {
+ cfg, endpointCfg = genASecurityConfigs(tb)
+ } else {
+ cfg, endpointCfg = genConfigs(tb)
+ }
var binds [2]conn.Bind
if realSocket {
binds[0], binds[1] = conn.NewDefaultBind(), conn.NewDefaultBind()
@@ -166,7 +237,7 @@ func genTestPair(tb testing.TB, realSocket bool) (pair testPair) {
if _, ok := tb.(*testing.B); ok && !testing.Verbose() {
level = LogLevelError
}
- p.dev = NewDevice(p.tun.TUN(), binds[i], NewLogger(level, fmt.Sprintf("dev%d: ", i)))
+ p.dev = NewDevice(p.tun.TUN(),binds[i],NewLogger(level, fmt.Sprintf("dev%d: ", i)))
if err := p.dev.IpcSet(cfg[i]); err != nil {
tb.Errorf("failed to configure device %d: %v", i, err)
p.dev.Close()
@@ -194,7 +265,18 @@ func genTestPair(tb testing.TB, realSocket bool) (pair testPair) {
func TestTwoDevicePing(t *testing.T) {
goroutineLeakCheck(t)
- pair := genTestPair(t, true)
+ pair := genTestPair(t, true, false)
+ t.Run("ping 1.0.0.1", func(t *testing.T) {
+ pair.Send(t, Ping, nil)
+ })
+ t.Run("ping 1.0.0.2", func(t *testing.T) {
+ pair.Send(t, Pong, nil)
+ })
+}
+
+func TestTwoDevicePingASecurity(t *testing.T) {
+ goroutineLeakCheck(t)
+ pair := genTestPair(t, true, true)
t.Run("ping 1.0.0.1", func(t *testing.T) {
pair.Send(t, Ping, nil)
})
@@ -209,10 +291,10 @@ func TestUpDown(t *testing.T) {
const otrials = 10
for n := 0; n < otrials; n++ {
- pair := genTestPair(t, false)
+ pair := genTestPair(t, false, false)
for i := range pair {
for k := range pair[i].dev.peers.keyMap {
- pair[i].dev.IpcSet(fmt.Sprintf("public_key=%s\npersistent_keepalive_interval=1\n", hex.EncodeToString(k[:])))
+ pair[i].dev.IpcSet(fmt.Sprintf("public_key=%s\npersistent_keepalive_interval=1\n",hex.EncodeToString(k[:])))
}
}
var wg sync.WaitGroup
@@ -243,7 +325,7 @@ func TestUpDown(t *testing.T) {
// TestConcurrencySafety does other things concurrently with tunnel use.
// It is intended to be used with the race detector to catch data races.
func TestConcurrencySafety(t *testing.T) {
- pair := genTestPair(t, true)
+ pair := genTestPair(t, true, false)
done := make(chan struct{})
const warmupIters = 10
@@ -324,7 +406,7 @@ func TestConcurrencySafety(t *testing.T) {
}
func BenchmarkLatency(b *testing.B) {
- pair := genTestPair(b, true)
+ pair := genTestPair(b, true, false)
// Establish a connection.
pair.Send(b, Ping, nil)
@@ -338,7 +420,7 @@ func BenchmarkLatency(b *testing.B) {
}
func BenchmarkThroughput(b *testing.B) {
- pair := genTestPair(b, true)
+ pair := genTestPair(b, true, false)
// Establish a connection.
pair.Send(b, Ping, nil)
@@ -382,7 +464,7 @@ func BenchmarkThroughput(b *testing.B) {
}
func BenchmarkUAPIGet(b *testing.B) {
- pair := genTestPair(b, true)
+ pair := genTestPair(b, true, false)
pair.Send(b, Ping, nil)
pair.Send(b, Pong, nil)
b.ReportAllocs()
@@ -423,29 +505,41 @@ type fakeBindSized struct {
size int
}
-func (b *fakeBindSized) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err error) {
+func (b *fakeBindSized) Open(
+ port uint16,
+) (fns []conn.ReceiveFunc, actualPort uint16, err error) {
return nil, 0, nil
}
-func (b *fakeBindSized) Close() error { return nil }
-func (b *fakeBindSized) SetMark(mark uint32) error { return nil }
-func (b *fakeBindSized) Send(bufs [][]byte, ep conn.Endpoint) error { return nil }
+
+func (b *fakeBindSized) Close() error { return nil }
+
+func (b *fakeBindSized) SetMark(mark uint32) error {return nil }
+
+func (b *fakeBindSized) Send(bufs [][]byte, ep conn.Endpoint) error { return nil }
+
func (b *fakeBindSized) ParseEndpoint(s string) (conn.Endpoint, error) { return nil, nil }
-func (b *fakeBindSized) BatchSize() int { return b.size }
+
+func (b *fakeBindSized) BatchSize() int { return b.size }
type fakeTUNDeviceSized struct {
size int
}
func (t *fakeTUNDeviceSized) File() *os.File { return nil }
-func (t *fakeTUNDeviceSized) Read(bufs [][]byte, sizes []int, offset int) (n int, err error) {
- return 0, nil
-}
+
+func (t *fakeTUNDeviceSized) Read(bufs [][]byte, sizes []int, offset int) (n int, err error) { return 0, nil }
+
func (t *fakeTUNDeviceSized) Write(bufs [][]byte, offset int) (int, error) { return 0, nil }
-func (t *fakeTUNDeviceSized) MTU() (int, error) { return 0, nil }
-func (t *fakeTUNDeviceSized) Name() (string, error) { return "", nil }
-func (t *fakeTUNDeviceSized) Events() <-chan tun.Event { return nil }
-func (t *fakeTUNDeviceSized) Close() error { return nil }
-func (t *fakeTUNDeviceSized) BatchSize() int { return t.size }
+
+func (t *fakeTUNDeviceSized) MTU() (int, error) { return 0, nil }
+
+func (t *fakeTUNDeviceSized) Name() (string, error) { return "", nil }
+
+func (t *fakeTUNDeviceSized) Events() <-chan tun.Event { return nil }
+
+func (t *fakeTUNDeviceSized) Close() error { return nil }
+
+func (t *fakeTUNDeviceSized) BatchSize() int { return t.size }
func TestBatchSize(t *testing.T) {
d := Device{}
diff --git a/device/keypair.go b/device/keypair.go
index e3540d7..73e69af 100644
--- a/device/keypair.go
+++ b/device/keypair.go
@@ -11,7 +11,7 @@ import (
"sync/atomic"
"time"
- "golang.zx2c4.com/wireguard/replay"
+ "github.com/amnezia-vpn/amnezia-wg/replay"
)
/* Due to limitations in Go and /x/crypto there is currently
diff --git a/device/noise-protocol.go b/device/noise-protocol.go
index e8f6145..75c1d87 100644
--- a/device/noise-protocol.go
+++ b/device/noise-protocol.go
@@ -15,7 +15,7 @@ import (
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/crypto/poly1305"
- "golang.zx2c4.com/wireguard/tai64n"
+ "github.com/amnezia-vpn/amnezia-wg/tai64n"
)
type handshakeState int
@@ -52,11 +52,11 @@ const (
WGLabelCookie = "cookie--"
)
-const (
- MessageInitiationType = 1
- MessageResponseType = 2
- MessageCookieReplyType = 3
- MessageTransportType = 4
+var (
+ MessageInitiationType uint32 = 1
+ MessageResponseType uint32 = 2
+ MessageCookieReplyType uint32 = 3
+ MessageTransportType uint32 = 4
)
const (
@@ -75,6 +75,10 @@ const (
MessageTransportOffsetContent = 16
)
+var packetSizeToMsgType map[int]uint32
+
+var msgTypeToJunkSize map[uint32]int
+
/* Type is an 8-bit field, followed by 3 nul bytes,
* by marshalling the messages in little-endian byteorder
* we can treat these as a 32-bit unsigned int (for now)
@@ -193,10 +197,12 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e
handshake.mixHash(handshake.remoteStatic[:])
+ device.aSecMux.RLock()
msg := MessageInitiation{
Type: MessageInitiationType,
Ephemeral: handshake.localEphemeral.publicKey(),
}
+ device.aSecMux.RUnlock()
handshake.mixKey(msg.Ephemeral[:])
handshake.mixHash(msg.Ephemeral[:])
@@ -250,9 +256,12 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer {
chainKey [blake2s.Size]byte
)
+ device.aSecMux.RLock()
if msg.Type != MessageInitiationType {
+ device.aSecMux.RUnlock()
return nil
}
+ device.aSecMux.RUnlock()
device.staticIdentity.RLock()
defer device.staticIdentity.RUnlock()
@@ -367,7 +376,9 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
var msg MessageResponse
+ device.aSecMux.RLock()
msg.Type = MessageResponseType
+ device.aSecMux.RUnlock()
msg.Sender = handshake.localIndex
msg.Receiver = handshake.remoteIndex
@@ -417,9 +428,12 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
+ device.aSecMux.RLock()
if msg.Type != MessageResponseType {
+ device.aSecMux.RUnlock()
return nil
}
+ device.aSecMux.RUnlock()
// lookup handshake by receiver
diff --git a/device/noise_test.go b/device/noise_test.go
index 2dd5324..2363365 100644
--- a/device/noise_test.go
+++ b/device/noise_test.go
@@ -10,8 +10,8 @@ import (
"encoding/binary"
"testing"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/tun/tuntest"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/tun/tuntest"
)
func TestCurveWrappers(t *testing.T) {
diff --git a/device/peer.go b/device/peer.go
index 0ac4896..72c7d1a 100644
--- a/device/peer.go
+++ b/device/peer.go
@@ -12,7 +12,7 @@ import (
"sync/atomic"
"time"
- "golang.zx2c4.com/wireguard/conn"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
)
type Peer struct {
diff --git a/device/queueconstants_android.go b/device/queueconstants_android.go
index 3d80ead..4adb687 100644
--- a/device/queueconstants_android.go
+++ b/device/queueconstants_android.go
@@ -5,7 +5,7 @@
package device
-import "golang.zx2c4.com/wireguard/conn"
+import "github.com/amnezia-vpn/amnezia-wg/conn"
/* Reduce memory consumption for Android */
diff --git a/device/queueconstants_default.go b/device/queueconstants_default.go
index ea763d0..4ee2966 100644
--- a/device/queueconstants_default.go
+++ b/device/queueconstants_default.go
@@ -7,7 +7,7 @@
package device
-import "golang.zx2c4.com/wireguard/conn"
+import "github.com/amnezia-vpn/amnezia-wg/conn"
const (
QueueStagedSize = conn.IdealBatchSize
diff --git a/device/receive.go b/device/receive.go
index e24d29f..ca71539 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -13,10 +13,10 @@ import (
"sync"
"time"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
- "golang.zx2c4.com/wireguard/conn"
)
type QueueHandshakeElement struct {
@@ -66,7 +66,10 @@ func (peer *Peer) keepKeyFreshReceiving() {
* Every time the bind is updated a new routine is started for
* IPv4 and IPv6 (separately)
*/
-func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.ReceiveFunc) {
+func (device *Device) RoutineReceiveIncoming(
+ maxBatchSize int,
+ recv conn.ReceiveFunc,
+) {
recvName := recv.PrettyName()
defer func() {
device.log.Verbosef("Routine: receive incoming %s - stopped", recvName)
@@ -122,6 +125,7 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive
}
deathSpiral = 0
+ device.aSecMux.RLock()
// handle each packet in the batch
for i, size := range sizes[:count] {
if size < MinMessageSize {
@@ -131,8 +135,29 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive
// check size of packet
packet := bufsArrs[i][:size]
- msgType := binary.LittleEndian.Uint32(packet[:4])
-
+ var msgType uint32
+ if device.isAdvancedSecurityOn() {
+ if assumedMsgType, ok := packetSizeToMsgType[size]; ok {
+ junkSize := msgTypeToJunkSize[assumedMsgType]
+ // transport size can align with other header types;
+ // making sure we have the right msgType
+ msgType = binary.LittleEndian.Uint32(packet[junkSize:junkSize+4])
+ if msgType == assumedMsgType {
+ packet = packet[junkSize:]
+ } else {
+ device.log.Verbosef("Transport packet lined up with another msg type")
+ msgType = binary.LittleEndian.Uint32(packet[:4])
+ }
+ } else {
+ msgType = binary.LittleEndian.Uint32(packet[:4])
+ if msgType != MessageTransportType {
+ device.log.Verbosef("ASec: Received message with unknown type")
+ continue
+ }
+ }
+ } else {
+ msgType = binary.LittleEndian.Uint32(packet[:4])
+ }
switch msgType {
// check if transport
@@ -217,6 +242,7 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive
default:
}
}
+ device.aSecMux.RUnlock()
for peer, elems := range elemsByPeer {
if peer.isRunning.Load() {
peer.queue.inbound.c <- elems
@@ -275,6 +301,8 @@ func (device *Device) RoutineHandshake(id int) {
for elem := range device.queue.handshake.c {
+ device.aSecMux.RLock()
+
// handle cookie fields and ratelimiting
switch elem.msgType {
@@ -302,9 +330,14 @@ func (device *Device) RoutineHandshake(id int) {
// consume reply
if peer := entry.peer; peer.isRunning.Load() {
- device.log.Verbosef("Receiving cookie response from %s", elem.endpoint.DstToString())
+ device.log.Verbosef(
+ "Receiving cookie response from %s",
+ elem.endpoint.DstToString(),
+ )
if !peer.cookieGenerator.ConsumeReply(&reply) {
- device.log.Verbosef("Could not decrypt invalid cookie response")
+ device.log.Verbosef(
+ "Could not decrypt invalid cookie response",
+ )
}
}
@@ -346,9 +379,7 @@ func (device *Device) RoutineHandshake(id int) {
switch elem.msgType {
case MessageInitiationType:
-
// unmarshal
-
var msg MessageInitiation
reader := bytes.NewReader(elem.packet)
err := binary.Read(reader, binary.LittleEndian, &msg)
@@ -358,7 +389,6 @@ func (device *Device) RoutineHandshake(id int) {
}
// consume initiation
-
peer := device.ConsumeMessageInitiation(&msg)
if peer == nil {
device.log.Verbosef("Received invalid initiation message from %s", elem.endpoint.DstToString())
@@ -423,6 +453,7 @@ func (device *Device) RoutineHandshake(id int) {
peer.SendKeepalive()
}
skip:
+ device.aSecMux.RUnlock()
device.PutMessageBuffer(elem.buffer)
}
}
@@ -503,11 +534,17 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
}
default:
- device.log.Verbosef("Packet with invalid IP version from %v", peer)
+ device.log.Verbosef(
+ "Packet with invalid IP version from %v",
+ peer,
+ )
continue
}
- bufs = append(bufs, elem.buffer[:MessageTransportOffsetContent+len(elem.packet)])
+ bufs = append(
+ bufs,
+ elem.buffer[:MessageTransportOffsetContent+len(elem.packet)],
+ )
}
if len(bufs) > 0 {
_, err := device.tun.device.Write(bufs, MessageTransportOffsetContent)
diff --git a/device/send.go b/device/send.go
index d22bf26..6f70d54 100644
--- a/device/send.go
+++ b/device/send.go
@@ -9,15 +9,16 @@ import (
"bytes"
"encoding/binary"
"errors"
+ "math/rand"
"net"
"os"
"sync"
"time"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
- "golang.zx2c4.com/wireguard/tun"
)
/* Outbound flow
@@ -119,17 +120,44 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
peer.device.log.Errorf("%v - Failed to create initiation message: %v", peer, err)
return err
}
-
+ var sendBuffer [][]byte
+ // so only packet processed for cookie generation
+ var junkedHeader []byte
+ if peer.device.isAdvancedSecurityOn() {
+ peer.device.aSecMux.RLock()
+ junks, err := peer.createJunkPackets()
+ if err != nil {
+ peer.device.aSecMux.RUnlock()
+ peer.device.log.Errorf("%v - %v", peer, err)
+ return err
+ }
+ sendBuffer = append(sendBuffer, junks...)
+ if peer.device.aSecCfg.initPacketJunkSize != 0 {
+ buf := make([]byte, 0, peer.device.aSecCfg.initPacketJunkSize)
+ writer := bytes.NewBuffer(buf[:0])
+ err = appendJunk(writer, peer.device.aSecCfg.initPacketJunkSize)
+ if err != nil {
+ peer.device.aSecMux.RUnlock()
+ peer.device.log.Errorf("%v - %v", peer, err)
+ return err
+ }
+ junkedHeader = writer.Bytes()
+ }
+ peer.device.aSecMux.RUnlock()
+ }
var buf [MessageInitiationSize]byte
writer := bytes.NewBuffer(buf[:0])
binary.Write(writer, binary.LittleEndian, msg)
packet := writer.Bytes()
peer.cookieGenerator.AddMacs(packet)
+ junkedHeader = append(junkedHeader, packet...)
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
-
- err = peer.SendBuffers([][]byte{packet})
+
+ sendBuffer = append(sendBuffer, junkedHeader)
+
+ err = peer.SendBuffers(sendBuffer)
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err)
}
@@ -150,12 +178,29 @@ func (peer *Peer) SendHandshakeResponse() error {
peer.device.log.Errorf("%v - Failed to create response message: %v", peer, err)
return err
}
-
+ var junkedHeader []byte
+ if peer.device.isAdvancedSecurityOn() {
+ peer.device.aSecMux.RLock()
+ if peer.device.aSecCfg.responsePacketJunkSize != 0 {
+ buf := make([]byte, 0, peer.device.aSecCfg.responsePacketJunkSize)
+ writer := bytes.NewBuffer(buf[:0])
+ err = appendJunk(writer, peer.device.aSecCfg.responsePacketJunkSize)
+ if err != nil {
+ peer.device.aSecMux.RUnlock()
+ peer.device.log.Errorf("%v - %v", peer, err)
+ return err
+ }
+ junkedHeader = writer.Bytes()
+ }
+ peer.device.aSecMux.RUnlock()
+ }
var buf [MessageResponseSize]byte
writer := bytes.NewBuffer(buf[:0])
+
binary.Write(writer, binary.LittleEndian, response)
packet := writer.Bytes()
peer.cookieGenerator.AddMacs(packet)
+ junkedHeader = append(junkedHeader, packet...)
err = peer.BeginSymmetricSession()
if err != nil {
@@ -168,18 +213,24 @@ func (peer *Peer) SendHandshakeResponse() error {
peer.timersAnyAuthenticatedPacketSent()
// TODO: allocation could be avoided
- err = peer.SendBuffers([][]byte{packet})
+ err = peer.SendBuffers([][]byte{junkedHeader})
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake response: %v", peer, err)
}
return err
}
-func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) error {
+func (device *Device) SendHandshakeCookie(
+ initiatingElem *QueueHandshakeElement,
+) error {
device.log.Verbosef("Sending cookie response for denied handshake message for %v", initiatingElem.endpoint.DstToString())
sender := binary.LittleEndian.Uint32(initiatingElem.packet[4:8])
- reply, err := device.cookieChecker.CreateReply(initiatingElem.packet, sender, initiatingElem.endpoint.DstToBytes())
+ reply, err := device.cookieChecker.CreateReply(
+ initiatingElem.packet,
+ sender,
+ initiatingElem.endpoint.DstToBytes(),
+ )
if err != nil {
device.log.Errorf("Failed to create cookie reply: %v", err)
return err
@@ -404,6 +455,31 @@ top:
}
}
+func (peer *Peer) createJunkPackets() ([][]byte, error) {
+ if peer.device.aSecCfg.junkPacketCount == 0 {
+ return nil, nil
+ }
+
+ junks := make([][]byte, 0, peer.device.aSecCfg.junkPacketCount)
+ for i := 0; i < peer.device.aSecCfg.junkPacketCount; i++ {
+ packetSize := rand.Intn(
+ peer.device.aSecCfg.junkPacketMaxSize-peer.device.aSecCfg.junkPacketMinSize,
+ ) + peer.device.aSecCfg.junkPacketMinSize
+
+ junk, err := randomJunkWithSize(packetSize)
+ if err != nil {
+ peer.device.log.Errorf(
+ "%v - Failed to create junk packet: %v",
+ peer,
+ err,
+ )
+ return nil, err
+ }
+ junks = append(junks, junk)
+ }
+ return junks, nil
+}
+
func (peer *Peer) FlushStagedPackets() {
for {
select {
@@ -459,18 +535,16 @@ func (device *Device) RoutineEncryption(id int) {
binary.LittleEndian.PutUint64(fieldNonce, elem.nonce)
// pad content to multiple of 16
- paddingSize := calculatePaddingSize(len(elem.packet), int(device.tun.mtu.Load()))
+ paddingSize := calculatePaddingSize(
+ len(elem.packet),
+ int(device.tun.mtu.Load()),
+ )
elem.packet = append(elem.packet, paddingZeros[:paddingSize]...)
// encrypt content and release to consumer
binary.LittleEndian.PutUint64(nonce[4:], elem.nonce)
- elem.packet = elem.keypair.send.Seal(
- header,
- nonce[:],
- elem.packet,
- nil,
- )
+ elem.packet = elem.keypair.send.Seal(header, nonce[:], elem.packet, nil)
elem.Unlock()
}
}
diff --git a/device/sticky_default.go b/device/sticky_default.go
index 1038256..940702c 100644
--- a/device/sticky_default.go
+++ b/device/sticky_default.go
@@ -3,8 +3,8 @@
package device
import (
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/rwcancel"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/rwcancel"
)
func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) {
diff --git a/device/sticky_linux.go b/device/sticky_linux.go
index f9230f8..5c17480 100644
--- a/device/sticky_linux.go
+++ b/device/sticky_linux.go
@@ -20,8 +20,8 @@ import (
"golang.org/x/sys/unix"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/rwcancel"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/rwcancel"
)
func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) {
diff --git a/device/tun.go b/device/tun.go
index 2a2ace9..efc543d 100644
--- a/device/tun.go
+++ b/device/tun.go
@@ -8,7 +8,7 @@ package device
import (
"fmt"
- "golang.zx2c4.com/wireguard/tun"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
)
const DefaultMTU = 1420
diff --git a/device/uapi.go b/device/uapi.go
index 617dcd3..bfd005a 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -18,7 +18,7 @@ import (
"sync"
"time"
- "golang.zx2c4.com/wireguard/ipc"
+ "github.com/amnezia-vpn/amnezia-wg/ipc"
)
type IPCError struct {
@@ -97,6 +97,36 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("fwmark=%d", device.net.fwmark)
}
+ if device.isAdvancedSecurityOn() {
+ if device.aSecCfg.junkPacketCount != 0 {
+ sendf("jc=%d", device.aSecCfg.junkPacketCount)
+ }
+ if device.aSecCfg.junkPacketMinSize != 0 {
+ sendf("jmin=%d", device.aSecCfg.junkPacketMinSize)
+ }
+ if device.aSecCfg.junkPacketMaxSize != 0 {
+ sendf("jmax=%d", device.aSecCfg.junkPacketMaxSize)
+ }
+ if device.aSecCfg.initPacketJunkSize != 0 {
+ sendf("s1=%d", device.aSecCfg.initPacketJunkSize)
+ }
+ if device.aSecCfg.responsePacketJunkSize != 0 {
+ sendf("s2=%d", device.aSecCfg.responsePacketJunkSize)
+ }
+ if device.aSecCfg.initPacketMagicHeader != 0 {
+ sendf("h1=%d", device.aSecCfg.initPacketMagicHeader)
+ }
+ if device.aSecCfg.responsePacketMagicHeader != 0 {
+ sendf("h2=%d", device.aSecCfg.responsePacketMagicHeader)
+ }
+ if device.aSecCfg.underloadPacketMagicHeader != 0 {
+ sendf("h3=%d", device.aSecCfg.underloadPacketMagicHeader)
+ }
+ if device.aSecCfg.transportPacketMagicHeader != 0 {
+ sendf("h4=%d", device.aSecCfg.transportPacketMagicHeader)
+ }
+ }
+
for _, peer := range device.peers.keyMap {
// Serialize peer state.
// Do the work in an anonymous function so that we can use defer.
@@ -121,10 +151,13 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("rx_bytes=%d", peer.rxBytes.Load())
sendf("persistent_keepalive_interval=%d", peer.persistentKeepaliveInterval.Load())
- device.allowedips.EntriesForPeer(peer, func(prefix netip.Prefix) bool {
- sendf("allowed_ip=%s", prefix.String())
- return true
- })
+ device.allowedips.EntriesForPeer(
+ peer,
+ func(prefix netip.Prefix) bool {
+ sendf("allowed_ip=%s", prefix.String())
+ return true
+ },
+ )
}()
}
}()
@@ -152,17 +185,26 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
peer := new(ipcSetPeer)
deviceConfig := true
+ tempASecCfg := aSecCfgType{}
scanner := bufio.NewScanner(r)
for scanner.Scan() {
line := scanner.Text()
if line == "" {
// Blank line means terminate operation.
+ err := device.handlePostConfig(&tempASecCfg)
+ if err != nil {
+ return err
+ }
peer.handlePostConfig()
return nil
}
key, value, ok := strings.Cut(line, "=")
if !ok {
- return ipcErrorf(ipc.IpcErrorProtocol, "failed to parse line %q", line)
+ return ipcErrorf(
+ ipc.IpcErrorProtocol,
+ "failed to parse line %q",
+ line,
+ )
}
if key == "public_key" {
@@ -180,7 +222,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
var err error
if deviceConfig {
- err = device.handleDeviceLine(key, value)
+ err = device.handleDeviceLine(key, value, &tempASecCfg)
} else {
err = device.handlePeerLine(peer, key, value)
}
@@ -188,6 +230,10 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return err
}
}
+ err = device.handlePostConfig(&tempASecCfg)
+ if err != nil {
+ return err
+ }
peer.handlePostConfig()
if err := scanner.Err(); err != nil {
@@ -196,7 +242,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return nil
}
-func (device *Device) handleDeviceLine(key, value string) error {
+func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgType) error {
switch key {
case "private_key":
var sk NoisePrivateKey
@@ -242,8 +288,75 @@ func (device *Device) handleDeviceLine(key, value string) error {
device.log.Verbosef("UAPI: Removing all peers")
device.RemoveAllPeers()
+ case "jc":
+ junkPacketCount, err := strconv.Atoi(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse junk_packet_count %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating junk_packet_count")
+ tempASecCfg.junkPacketCount = junkPacketCount
+
+ case "jmin":
+ junkPacketMinSize, err := strconv.Atoi(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse junk_packet_min_size %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating junk_packet_min_size")
+ tempASecCfg.junkPacketMinSize = junkPacketMinSize
+
+ case "jmax":
+ junkPacketMaxSize, err := strconv.Atoi(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse junk_packet_max_size %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating junk_packet_max_size")
+ tempASecCfg.junkPacketMaxSize = junkPacketMaxSize
+
+ case "s1":
+ initPacketJunkSize, err := strconv.Atoi(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse init_packet_junk_size %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating init_packet_junk_size")
+ tempASecCfg.initPacketJunkSize = initPacketJunkSize
+
+ case "s2":
+ responsePacketJunkSize, err := strconv.Atoi(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse response_packet_junk_size %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating response_packet_junk_size")
+ tempASecCfg.responsePacketJunkSize = responsePacketJunkSize
+
+ case "h1":
+ initPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse init_packet_magic_header %w", err)
+ }
+ tempASecCfg.initPacketMagicHeader = uint32(initPacketMagicHeader)
+
+ case "h2":
+ responsePacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse response_packet_magic_header %w", err)
+ }
+ tempASecCfg.responsePacketMagicHeader = uint32(responsePacketMagicHeader)
+
+ case "h3":
+ underloadPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse underload_packet_magic_header %w", err)
+ }
+ tempASecCfg.underloadPacketMagicHeader = uint32(underloadPacketMagicHeader)
+
+ case "h4":
+ transportPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse transport_packet_magic_header %w", err)
+ }
+ tempASecCfg.transportPacketMagicHeader = uint32(transportPacketMagicHeader)
default:
- return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
+ return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v",key)
}
return nil
@@ -262,7 +375,8 @@ func (peer *ipcSetPeer) handlePostConfig() {
return
}
if peer.created {
- peer.disableRoaming = peer.device.net.brokenRoaming && peer.endpoint != nil
+ peer.disableRoaming = peer.device.net.brokenRoaming &&
+ peer.endpoint != nil
}
if peer.device.isUp() {
peer.Start()
@@ -273,7 +387,10 @@ func (peer *ipcSetPeer) handlePostConfig() {
}
}
-func (device *Device) handlePublicKeyLine(peer *ipcSetPeer, value string) error {
+func (device *Device) handlePublicKeyLine(
+ peer *ipcSetPeer,
+ value string,
+) error {
// Load/create the peer we are configuring.
var publicKey NoisePublicKey
err := publicKey.FromHex(value)
@@ -303,7 +420,10 @@ func (device *Device) handlePublicKeyLine(peer *ipcSetPeer, value string) error
return nil
}
-func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error {
+func (device *Device) handlePeerLine(
+ peer *ipcSetPeer,
+ key, value string,
+) error {
switch key {
case "update_only":
// allow disabling of creation
@@ -343,7 +463,7 @@ func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error
device.log.Verbosef("%v - UAPI: Updating endpoint", peer.Peer)
endpoint, err := device.net.bind.ParseEndpoint(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "failed to set endpoint %v: %w", value, err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to set endpoint %v: %w", value, err)
}
peer.Lock()
defer peer.Unlock()
diff --git a/device/util.go b/device/util.go
new file mode 100644
index 0000000..aab8ab7
--- /dev/null
+++ b/device/util.go
@@ -0,0 +1,25 @@
+package device
+
+import (
+ "bytes"
+ crand "crypto/rand"
+ "fmt"
+)
+
+func appendJunk(writer *bytes.Buffer, size int) error {
+ headerJunk, err := randomJunkWithSize(size)
+ if err != nil {
+ return fmt.Errorf("failed to create header junk: %v", err)
+ }
+ _, err = writer.Write(headerJunk)
+ if err != nil {
+ return fmt.Errorf("failed to write header junk: %v", err)
+ }
+ return nil
+}
+
+func randomJunkWithSize(size int) ([]byte, error) {
+ junk := make([]byte, size)
+ _, err := crand.Read(junk)
+ return junk, err
+}
diff --git a/device/util_test.go b/device/util_test.go
new file mode 100644
index 0000000..c061eef
--- /dev/null
+++ b/device/util_test.go
@@ -0,0 +1,27 @@
+package device
+
+import (
+ "bytes"
+ "fmt"
+ "testing"
+)
+
+func Test_randomJunktWithSize(t *testing.T) {
+ junk, err := randomJunkWithSize(30)
+ fmt.Println(string(junk), len(junk), err)
+}
+
+func Test_appendJunk(t *testing.T) {
+ t.Run("", func(t *testing.T) {
+ s := "apple"
+ buffer := bytes.NewBuffer([]byte(s))
+ err := appendJunk(buffer, 30)
+ if err != nil &&
+ buffer.Len() != len(s)+30 {
+ t.Errorf("appendWithJunk() size don't match")
+ }
+ read := make([]byte, 50)
+ buffer.Read(read)
+ fmt.Println(string(read))
+ })
+}
diff --git a/go.mod b/go.mod
index c04e1bb..4a3c9c6 100644
--- a/go.mod
+++ b/go.mod
@@ -1,8 +1,9 @@
-module golang.zx2c4.com/wireguard
+module github.com/amnezia-vpn/amnezia-wg
go 1.20
require (
+ github.com/tevino/abool/v2 v2.1.0
golang.org/x/crypto v0.6.0
golang.org/x/net v0.7.0
golang.org/x/sys v0.5.1-0.20230222185716-a3b23cc77e89
diff --git a/go.sum b/go.sum
index cfeaee6..3707808 100644
--- a/go.sum
+++ b/go.sum
@@ -1,5 +1,7 @@
github.com/google/btree v1.0.1 h1:gK4Kx5IaGY9CD5sPJ36FHiBJ6ZXl0kilRiiCj+jdYp4=
github.com/google/btree v1.0.1/go.mod h1:xXMiIv4Fb/0kKde4SpL7qlzvu5cMJDRkFDxJfI9uaxA=
+github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
+github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
golang.org/x/crypto v0.6.0 h1:qfktjS5LUO+fFKeJXZ+ikTRijMmljikvG68fpMMruSc=
golang.org/x/crypto v0.6.0/go.mod h1:OFC/31mSvZgRz0V1QTNCzfAI1aIRzbiufJtkMIlEp58=
golang.org/x/net v0.7.0 h1:rJrUqqhjsgNp7KqAIc25s9pZnjU7TUcSY7HcVZjdn1g=
diff --git a/ipc/namedpipe/namedpipe_test.go b/ipc/namedpipe/namedpipe_test.go
index 998453b..d4799e1 100644
--- a/ipc/namedpipe/namedpipe_test.go
+++ b/ipc/namedpipe/namedpipe_test.go
@@ -20,8 +20,8 @@ import (
"testing"
"time"
+ "github.com/amnezia-vpn/amnezia-wg/ipc/namedpipe"
"golang.org/x/sys/windows"
- "golang.zx2c4.com/wireguard/ipc/namedpipe"
)
func randomPipePath() string {
diff --git a/ipc/uapi_linux.go b/ipc/uapi_linux.go
index 1562a18..721c404 100644
--- a/ipc/uapi_linux.go
+++ b/ipc/uapi_linux.go
@@ -9,8 +9,8 @@ import (
"net"
"os"
+ "github.com/amnezia-vpn/amnezia-wg/rwcancel"
"golang.org/x/sys/unix"
- "golang.zx2c4.com/wireguard/rwcancel"
)
type UAPIListener struct {
diff --git a/ipc/uapi_windows.go b/ipc/uapi_windows.go
index aa023c9..97a4123 100644
--- a/ipc/uapi_windows.go
+++ b/ipc/uapi_windows.go
@@ -8,8 +8,8 @@ package ipc
import (
"net"
+ "github.com/amnezia-vpn/amnezia-wg/ipc/namedpipe"
"golang.org/x/sys/windows"
- "golang.zx2c4.com/wireguard/ipc/namedpipe"
)
// TODO: replace these with actual standard windows error numbers from the win package
diff --git a/main.go b/main.go
index e016116..ea7ef4e 100644
--- a/main.go
+++ b/main.go
@@ -14,11 +14,11 @@ import (
"runtime"
"strconv"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/device"
+ "github.com/amnezia-vpn/amnezia-wg/ipc"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
"golang.org/x/sys/unix"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/device"
- "golang.zx2c4.com/wireguard/ipc"
- "golang.zx2c4.com/wireguard/tun"
)
const (
diff --git a/main_windows.go b/main_windows.go
index a4dc46f..d00b146 100644
--- a/main_windows.go
+++ b/main_windows.go
@@ -12,11 +12,11 @@ import (
"golang.org/x/sys/windows"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/device"
- "golang.zx2c4.com/wireguard/ipc"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/device"
+ "github.com/amnezia-vpn/amnezia-wg/ipc"
- "golang.zx2c4.com/wireguard/tun"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
)
const (
diff --git a/tun/netstack/examples/http_client.go b/tun/netstack/examples/http_client.go
index ccd32ed..ed40904 100644
--- a/tun/netstack/examples/http_client.go
+++ b/tun/netstack/examples/http_client.go
@@ -13,9 +13,9 @@ import (
"net/http"
"net/netip"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/device"
- "golang.zx2c4.com/wireguard/tun/netstack"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/device"
+ "github.com/amnezia-vpn/amnezia-wg/tun/netstack"
)
func main() {
diff --git a/tun/netstack/examples/http_server.go b/tun/netstack/examples/http_server.go
index f5b7a8f..d5e7094 100644
--- a/tun/netstack/examples/http_server.go
+++ b/tun/netstack/examples/http_server.go
@@ -14,9 +14,9 @@ import (
"net/http"
"net/netip"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/device"
- "golang.zx2c4.com/wireguard/tun/netstack"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/device"
+ "github.com/amnezia-vpn/amnezia-wg/tun/netstack"
)
func main() {
diff --git a/tun/netstack/examples/ping_client.go b/tun/netstack/examples/ping_client.go
index 2eef0fb..9f917db 100644
--- a/tun/netstack/examples/ping_client.go
+++ b/tun/netstack/examples/ping_client.go
@@ -17,9 +17,9 @@ import (
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/device"
- "golang.zx2c4.com/wireguard/tun/netstack"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/device"
+ "github.com/amnezia-vpn/amnezia-wg/tun/netstack"
)
func main() {
diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go
index 596cfcd..f5a40f5 100644
--- a/tun/netstack/tun.go
+++ b/tun/netstack/tun.go
@@ -22,7 +22,7 @@ import (
"syscall"
"time"
- "golang.zx2c4.com/wireguard/tun"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
"golang.org/x/net/dns/dnsmessage"
"gvisor.dev/gvisor/pkg/bufferv2"
diff --git a/tun/tcp_offload_linux.go b/tun/tcp_offload_linux.go
index 39a7180..a43f0df 100644
--- a/tun/tcp_offload_linux.go
+++ b/tun/tcp_offload_linux.go
@@ -12,8 +12,8 @@ import (
"io"
"unsafe"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
"golang.org/x/sys/unix"
- "golang.zx2c4.com/wireguard/conn"
)
const tcpFlagsOffset = 13
diff --git a/tun/tcp_offload_linux_test.go b/tun/tcp_offload_linux_test.go
index 9160e18..57c6a09 100644
--- a/tun/tcp_offload_linux_test.go
+++ b/tun/tcp_offload_linux_test.go
@@ -9,8 +9,8 @@ import (
"net/netip"
"testing"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
"golang.org/x/sys/unix"
- "golang.zx2c4.com/wireguard/conn"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header"
)
diff --git a/tun/tun_linux.go b/tun/tun_linux.go
index 12cd49f..31c1513 100644
--- a/tun/tun_linux.go
+++ b/tun/tun_linux.go
@@ -17,9 +17,9 @@ import (
"time"
"unsafe"
+ "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amnezia-wg/rwcancel"
"golang.org/x/sys/unix"
- "golang.zx2c4.com/wireguard/conn"
- "golang.zx2c4.com/wireguard/rwcancel"
)
const (
diff --git a/tun/tuntest/tuntest.go b/tun/tuntest/tuntest.go
index d07e860..7068d9b 100644
--- a/tun/tuntest/tuntest.go
+++ b/tun/tuntest/tuntest.go
@@ -11,7 +11,7 @@ import (
"net/netip"
"os"
- "golang.zx2c4.com/wireguard/tun"
+ "github.com/amnezia-vpn/amnezia-wg/tun"
)
func Ping(dst, src netip.Addr) []byte {
From f30419e0d14ba692e0974a65e0514ca4571feee4 Mon Sep 17 00:00:00 2001
From: Mazay B
Date: Mon, 9 Oct 2023 13:22:49 +0100
Subject: [PATCH 02/75] Manage advanced sec via uapi
---
device/device.go | 61 +++++++++++++++++++++---------------------------
device/send.go | 20 ++++++++++------
device/uapi.go | 14 +++++++++--
3 files changed, 52 insertions(+), 43 deletions(-)
diff --git a/device/device.go b/device/device.go
index 10365d1..a10187b 100644
--- a/device/device.go
+++ b/device/device.go
@@ -98,6 +98,7 @@ type Device struct {
}
type aSecCfgType struct {
+ isSet bool
junkPacketCount int
junkPacketMinSize int
junkPacketMaxSize int
@@ -545,7 +546,7 @@ func (device *Device) BindUpdate() error {
// start receiving routines
device.net.stopping.Add(len(recvFns))
device.queue.decryption.wg.Add(len(recvFns)) // each RoutineReceiveIncoming goroutine writes to device.queue.decryption
- device.queue.handshake.wg.Add(len(recvFns)) // each RoutineReceiveIncoming goroutine writes to device.queue.handshake
+ device.queue.handshake.wg.Add(len(recvFns)) // each RoutineReceiveIncoming goroutine writes to device.queue.handshake
batchSize := netc.bind.BatchSize()
for _, fn := range recvFns {
go device.RoutineReceiveIncoming(batchSize, fn)
@@ -565,25 +566,17 @@ func (device *Device) isAdvancedSecurityOn() bool {
return device.isASecOn.IsSet()
}
-func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
-
- if tempASecCfg.junkPacketCount == 0 &&
- tempASecCfg.junkPacketMaxSize == 0 &&
- tempASecCfg.junkPacketMinSize == 0 &&
- tempASecCfg.initPacketJunkSize == 0 &&
- tempASecCfg.responsePacketJunkSize == 0 &&
- tempASecCfg.initPacketMagicHeader == 0 &&
- tempASecCfg.responsePacketMagicHeader == 0 &&
- tempASecCfg.underloadPacketMagicHeader == 0 &&
- tempASecCfg.transportPacketMagicHeader == 0 {
+func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
+
+ if !tempASecCfg.isSet {
return err
}
-
+
isASecOn := false
device.aSecMux.Lock()
if tempASecCfg.junkPacketCount < 0 {
err = ipcErrorf(
- ipc.IpcErrorInvalid,
+ ipc.IpcErrorInvalid,
"JunkPacketCount should be non negative",
)
}
@@ -591,24 +584,24 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
if tempASecCfg.junkPacketCount != 0 {
isASecOn = true
}
-
+
device.aSecCfg.junkPacketMinSize = tempASecCfg.junkPacketMinSize
if tempASecCfg.junkPacketMinSize != 0 {
isASecOn = true
}
- if device.aSecCfg.junkPacketCount > 0 &&
+ if device.aSecCfg.junkPacketCount > 0 &&
tempASecCfg.junkPacketMaxSize == tempASecCfg.junkPacketMinSize {
-
+
tempASecCfg.junkPacketMaxSize++ // to make rand gen work
}
- if tempASecCfg.junkPacketMaxSize >= MaxSegmentSize{
+ if tempASecCfg.junkPacketMaxSize >= MaxSegmentSize {
device.aSecCfg.junkPacketMinSize = 0
device.aSecCfg.junkPacketMaxSize = 1
if err != nil {
err = ipcErrorf(
- ipc.IpcErrorInvalid,
+ ipc.IpcErrorInvalid,
"JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d; %w",
tempASecCfg.junkPacketMaxSize,
MaxSegmentSize,
@@ -616,7 +609,7 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
)
} else {
err = ipcErrorf(
- ipc.IpcErrorInvalid,
+ ipc.IpcErrorInvalid,
"JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d",
tempASecCfg.junkPacketMaxSize,
MaxSegmentSize,
@@ -625,18 +618,18 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
} else if tempASecCfg.junkPacketMaxSize < tempASecCfg.junkPacketMinSize {
if err != nil {
err = ipcErrorf(
- ipc.IpcErrorInvalid,
+ ipc.IpcErrorInvalid,
"maxSize: %d; should be greater than minSize: %d; %w",
tempASecCfg.junkPacketMaxSize,
- tempASecCfg.junkPacketMinSize,
+ tempASecCfg.junkPacketMinSize,
err,
)
} else {
err = ipcErrorf(
- ipc.IpcErrorInvalid,
+ ipc.IpcErrorInvalid,
"maxSize: %d; should be greater than minSize: %d",
tempASecCfg.junkPacketMaxSize,
- tempASecCfg.junkPacketMinSize,
+ tempASecCfg.junkPacketMinSize,
)
}
} else {
@@ -664,10 +657,10 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
MaxSegmentSize,
)
}
- } else {
+ } else {
device.aSecCfg.initPacketJunkSize = tempASecCfg.initPacketJunkSize
}
-
+
if tempASecCfg.initPacketJunkSize != 0 {
isASecOn = true
}
@@ -689,7 +682,7 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
MaxSegmentSize,
)
}
- } else {
+ } else {
device.aSecCfg.responsePacketJunkSize = tempASecCfg.responsePacketJunkSize
}
@@ -706,7 +699,7 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
device.log.Verbosef("UAPI: Using default init type")
MessageInitiationType = 1
}
-
+
if tempASecCfg.responsePacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating response_packet_magic_header")
@@ -716,7 +709,7 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
device.log.Verbosef("UAPI: Using default response type")
MessageResponseType = 2
}
-
+
if tempASecCfg.underloadPacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating underload_packet_magic_header")
@@ -787,14 +780,14 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
newResponseSize,
)
}
- } else {
+ } else {
packetSizeToMsgType = map[int]uint32{
- newInitSize: MessageInitiationType,
- newResponseSize: MessageResponseType,
+ newInitSize: MessageInitiationType,
+ newResponseSize: MessageResponseType,
MessageCookieReplySize: MessageCookieReplyType,
MessageTransportSize: MessageTransportType,
}
-
+
msgTypeToJunkSize = map[uint32]int{
MessageInitiationType: device.aSecCfg.initPacketJunkSize,
MessageResponseType: device.aSecCfg.responsePacketJunkSize,
@@ -805,6 +798,6 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
device.isASecOn.SetTo(isASecOn)
device.aSecMux.Unlock()
-
+
return err
}
diff --git a/device/send.go b/device/send.go
index 6f70d54..b5c8e10 100644
--- a/device/send.go
+++ b/device/send.go
@@ -126,25 +126,31 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
if peer.device.isAdvancedSecurityOn() {
peer.device.aSecMux.RLock()
junks, err := peer.createJunkPackets()
+ peer.device.aSecMux.RUnlock()
+
if err != nil {
- peer.device.aSecMux.RUnlock()
peer.device.log.Errorf("%v - %v", peer, err)
return err
}
- sendBuffer = append(sendBuffer, junks...)
+
+ err = peer.SendBuffers(junks)
+ if err != nil {
+ peer.device.log.Errorf("%v - Failed to send junk packets: %v", peer, err)
+ return err
+ }
+
if peer.device.aSecCfg.initPacketJunkSize != 0 {
buf := make([]byte, 0, peer.device.aSecCfg.initPacketJunkSize)
writer := bytes.NewBuffer(buf[:0])
err = appendJunk(writer, peer.device.aSecCfg.initPacketJunkSize)
if err != nil {
- peer.device.aSecMux.RUnlock()
peer.device.log.Errorf("%v - %v", peer, err)
return err
}
junkedHeader = writer.Bytes()
}
- peer.device.aSecMux.RUnlock()
}
+
var buf [MessageInitiationSize]byte
writer := bytes.NewBuffer(buf[:0])
binary.Write(writer, binary.LittleEndian, msg)
@@ -154,9 +160,9 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
-
+
sendBuffer = append(sendBuffer, junkedHeader)
-
+
err = peer.SendBuffers(sendBuffer)
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err)
@@ -191,7 +197,7 @@ func (peer *Peer) SendHandshakeResponse() error {
return err
}
junkedHeader = writer.Bytes()
- }
+ }
peer.device.aSecMux.RUnlock()
}
var buf [MessageResponseSize]byte
diff --git a/device/uapi.go b/device/uapi.go
index bfd005a..653803c 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -295,6 +295,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
}
device.log.Verbosef("UAPI: Updating junk_packet_count")
tempASecCfg.junkPacketCount = junkPacketCount
+ tempASecCfg.isSet = true
case "jmin":
junkPacketMinSize, err := strconv.Atoi(value)
@@ -303,6 +304,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
}
device.log.Verbosef("UAPI: Updating junk_packet_min_size")
tempASecCfg.junkPacketMinSize = junkPacketMinSize
+ tempASecCfg.isSet = true
case "jmax":
junkPacketMaxSize, err := strconv.Atoi(value)
@@ -311,6 +313,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
}
device.log.Verbosef("UAPI: Updating junk_packet_max_size")
tempASecCfg.junkPacketMaxSize = junkPacketMaxSize
+ tempASecCfg.isSet = true
case "s1":
initPacketJunkSize, err := strconv.Atoi(value)
@@ -319,6 +322,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
}
device.log.Verbosef("UAPI: Updating init_packet_junk_size")
tempASecCfg.initPacketJunkSize = initPacketJunkSize
+ tempASecCfg.isSet = true
case "s2":
responsePacketJunkSize, err := strconv.Atoi(value)
@@ -327,6 +331,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
}
device.log.Verbosef("UAPI: Updating response_packet_junk_size")
tempASecCfg.responsePacketJunkSize = responsePacketJunkSize
+ tempASecCfg.isSet = true
case "h1":
initPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
@@ -334,6 +339,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse init_packet_magic_header %w", err)
}
tempASecCfg.initPacketMagicHeader = uint32(initPacketMagicHeader)
+ tempASecCfg.isSet = true
case "h2":
responsePacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
@@ -341,6 +347,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse response_packet_magic_header %w", err)
}
tempASecCfg.responsePacketMagicHeader = uint32(responsePacketMagicHeader)
+ tempASecCfg.isSet = true
case "h3":
underloadPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
@@ -348,6 +355,7 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse underload_packet_magic_header %w", err)
}
tempASecCfg.underloadPacketMagicHeader = uint32(underloadPacketMagicHeader)
+ tempASecCfg.isSet = true
case "h4":
transportPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
@@ -355,8 +363,10 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse transport_packet_magic_header %w", err)
}
tempASecCfg.transportPacketMagicHeader = uint32(transportPacketMagicHeader)
+ tempASecCfg.isSet = true
+
default:
- return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v",key)
+ return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
}
return nil
@@ -463,7 +473,7 @@ func (device *Device) handlePeerLine(
device.log.Verbosef("%v - UAPI: Updating endpoint", peer.Peer)
endpoint, err := device.net.bind.ParseEndpoint(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "failed to set endpoint %v: %w", value, err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to set endpoint %v: %w", value, err)
}
peer.Lock()
defer peer.Unlock()
From 6a84778f2ca810f5fb5cb078e001494f08d9085f Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Mon, 2 Oct 2023 13:53:07 -0700
Subject: [PATCH 03/75] conn, device: use UDP GSO and GRO on Linux
StdNetBind probes for UDP GSO and GRO support at runtime. UDP GSO is
dependent on checksum offload support on the egress netdev. UDP GSO
will be disabled in the event sendmmsg() returns EIO, which is a strong
signal that the egress netdev does not support checksum offload.
The iperf3 results below demonstrate the effect of this commit between
two Linux computers with i5-12400 CPUs. There is roughly ~13us of round
trip latency between them.
The first result is from commit 052af4a without UDP GSO or GRO.
Starting Test: protocol: TCP, 1 streams, 131072 byte blocks
[ ID] Interval Transfer Bitrate Retr Cwnd
[ 5] 0.00-10.00 sec 9.85 GBytes 8.46 Gbits/sec 1139 3.01 MBytes
- - - - - - - - - - - - - - - - - - - - - - - - -
Test Complete. Summary Results:
[ ID] Interval Transfer Bitrate Retr
[ 5] 0.00-10.00 sec 9.85 GBytes 8.46 Gbits/sec 1139 sender
[ 5] 0.00-10.04 sec 9.85 GBytes 8.42 Gbits/sec receiver
The second result is with UDP GSO and GRO.
Starting Test: protocol: TCP, 1 streams, 131072 byte blocks
[ ID] Interval Transfer Bitrate Retr Cwnd
[ 5] 0.00-10.00 sec 12.3 GBytes 10.6 Gbits/sec 232 3.15 MBytes
- - - - - - - - - - - - - - - - - - - - - - - - -
Test Complete. Summary Results:
[ ID] Interval Transfer Bitrate Retr
[ 5] 0.00-10.00 sec 12.3 GBytes 10.6 Gbits/sec 232 sender
[ 5] 0.00-10.04 sec 12.3 GBytes 10.6 Gbits/sec receiver
Reviewed-by: Adrian Dewhurst
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
conn/bind_std.go | 397 ++++++++++++------
conn/bind_std_test.go | 230 +++++++++-
.../{sticky_default.go => control_default.go} | 20 +-
conn/{sticky_linux.go => control_linux.go} | 51 ++-
...ky_linux_test.go => control_linux_test.go} | 10 +-
conn/controlfns_linux.go | 8 +
conn/errors_default.go | 12 +
conn/errors_linux.go | 26 ++
conn/features_default.go | 15 +
conn/features_linux.go | 35 ++
device/send.go | 8 +
go.mod | 2 +-
go.sum | 4 +-
13 files changed, 672 insertions(+), 146 deletions(-)
rename conn/{sticky_default.go => control_default.go} (56%)
rename conn/{sticky_linux.go => control_linux.go} (65%)
rename conn/{sticky_linux_test.go => control_linux_test.go} (96%)
create mode 100644 conn/errors_default.go
create mode 100644 conn/errors_linux.go
create mode 100644 conn/features_default.go
create mode 100644 conn/features_linux.go
diff --git a/conn/bind_std.go b/conn/bind_std.go
index c701ef8..9886c91 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -8,6 +8,7 @@ package conn
import (
"context"
"errors"
+ "fmt"
"net"
"net/netip"
"runtime"
@@ -29,16 +30,19 @@ var (
// methods for sending and receiving multiple datagrams per-syscall. See the
// proposal in https://github.com/golang/go/issues/45886#issuecomment-1218301564.
type StdNetBind struct {
- mu sync.Mutex // protects all fields except as specified
- ipv4 *net.UDPConn
- ipv6 *net.UDPConn
- ipv4PC *ipv4.PacketConn // will be nil on non-Linux
- ipv6PC *ipv6.PacketConn // will be nil on non-Linux
+ mu sync.Mutex // protects all fields except as specified
+ ipv4 *net.UDPConn
+ ipv6 *net.UDPConn
+ ipv4PC *ipv4.PacketConn // will be nil on non-Linux
+ ipv6PC *ipv6.PacketConn // will be nil on non-Linux
+ ipv4TxOffload bool
+ ipv4RxOffload bool
+ ipv6TxOffload bool
+ ipv6RxOffload bool
- // these three fields are not guarded by mu
- udpAddrPool sync.Pool
- ipv4MsgsPool sync.Pool
- ipv6MsgsPool sync.Pool
+ // these two fields are not guarded by mu
+ udpAddrPool sync.Pool
+ msgsPool sync.Pool
blackhole4 bool
blackhole6 bool
@@ -54,23 +58,14 @@ func NewStdNetBind() Bind {
},
},
- ipv4MsgsPool: sync.Pool{
- New: func() any {
- msgs := make([]ipv4.Message, IdealBatchSize)
- for i := range msgs {
- msgs[i].Buffers = make(net.Buffers, 1)
- msgs[i].OOB = make([]byte, srcControlSize)
- }
- return &msgs
- },
- },
-
- ipv6MsgsPool: sync.Pool{
+ msgsPool: sync.Pool{
New: func() any {
+ // ipv6.Message and ipv4.Message are interchangeable as they are
+ // both aliases for x/net/internal/socket.Message.
msgs := make([]ipv6.Message, IdealBatchSize)
for i := range msgs {
msgs[i].Buffers = make(net.Buffers, 1)
- msgs[i].OOB = make([]byte, srcControlSize)
+ msgs[i].OOB = make([]byte, controlSize)
}
return &msgs
},
@@ -113,7 +108,7 @@ func (e *StdNetEndpoint) DstIP() netip.Addr {
return e.AddrPort.Addr()
}
-// See sticky_default,linux, etc for implementations of SrcIP and SrcIfidx.
+// See control_default,linux, etc for implementations of SrcIP and SrcIfidx.
func (e *StdNetEndpoint) DstToBytes() []byte {
b, _ := e.AddrPort.MarshalBinary()
@@ -179,19 +174,21 @@ again:
}
var fns []ReceiveFunc
if v4conn != nil {
+ s.ipv4TxOffload, s.ipv4RxOffload = supportsUDPOffload(v4conn)
if runtime.GOOS == "linux" {
v4pc = ipv4.NewPacketConn(v4conn)
s.ipv4PC = v4pc
}
- fns = append(fns, s.makeReceiveIPv4(v4pc, v4conn))
+ fns = append(fns, s.makeReceiveIPv4(v4pc, v4conn, s.ipv4RxOffload))
s.ipv4 = v4conn
}
if v6conn != nil {
+ s.ipv6TxOffload, s.ipv6RxOffload = supportsUDPOffload(v6conn)
if runtime.GOOS == "linux" {
v6pc = ipv6.NewPacketConn(v6conn)
s.ipv6PC = v6pc
}
- fns = append(fns, s.makeReceiveIPv6(v6pc, v6conn))
+ fns = append(fns, s.makeReceiveIPv6(v6pc, v6conn, s.ipv6RxOffload))
s.ipv6 = v6conn
}
if len(fns) == 0 {
@@ -201,69 +198,93 @@ again:
return fns, uint16(port), nil
}
-func (s *StdNetBind) makeReceiveIPv4(pc *ipv4.PacketConn, conn *net.UDPConn) ReceiveFunc {
- return func(bufs [][]byte, sizes []int, eps []Endpoint) (n int, err error) {
- msgs := s.ipv4MsgsPool.Get().(*[]ipv4.Message)
- defer s.ipv4MsgsPool.Put(msgs)
- for i := range bufs {
- (*msgs)[i].Buffers[0] = bufs[i]
- }
- var numMsgs int
- if runtime.GOOS == "linux" {
- numMsgs, err = pc.ReadBatch(*msgs, 0)
+func (s *StdNetBind) putMessages(msgs *[]ipv6.Message) {
+ for i := range *msgs {
+ (*msgs)[i] = ipv6.Message{Buffers: (*msgs)[i].Buffers, OOB: (*msgs)[i].OOB}
+ }
+ s.msgsPool.Put(msgs)
+}
+
+func (s *StdNetBind) getMessages() *[]ipv6.Message {
+ return s.msgsPool.Get().(*[]ipv6.Message)
+}
+
+var (
+ // If compilation fails here these are no longer the same underlying type.
+ _ ipv6.Message = ipv4.Message{}
+)
+
+type batchReader interface {
+ ReadBatch([]ipv6.Message, int) (int, error)
+}
+
+type batchWriter interface {
+ WriteBatch([]ipv6.Message, int) (int, error)
+}
+
+func (s *StdNetBind) receiveIP(
+ br batchReader,
+ conn *net.UDPConn,
+ rxOffload bool,
+ bufs [][]byte,
+ sizes []int,
+ eps []Endpoint,
+) (n int, err error) {
+ msgs := s.getMessages()
+ for i := range bufs {
+ (*msgs)[i].Buffers[0] = bufs[i]
+ (*msgs)[i].OOB = (*msgs)[i].OOB[:cap((*msgs)[i].OOB)]
+ }
+ defer s.putMessages(msgs)
+ var numMsgs int
+ if runtime.GOOS == "linux" {
+ if rxOffload {
+ readAt := len(*msgs) - (IdealBatchSize / udpSegmentMaxDatagrams)
+ numMsgs, err = br.ReadBatch((*msgs)[readAt:], 0)
+ if err != nil {
+ return 0, err
+ }
+ numMsgs, err = splitCoalescedMessages(*msgs, readAt, getGSOSize)
if err != nil {
return 0, err
}
} else {
- msg := &(*msgs)[0]
- msg.N, msg.NN, _, msg.Addr, err = conn.ReadMsgUDP(msg.Buffers[0], msg.OOB)
+ numMsgs, err = br.ReadBatch(*msgs, 0)
if err != nil {
return 0, err
}
- numMsgs = 1
}
- for i := 0; i < numMsgs; i++ {
- msg := &(*msgs)[i]
- sizes[i] = msg.N
- addrPort := msg.Addr.(*net.UDPAddr).AddrPort()
- ep := &StdNetEndpoint{AddrPort: addrPort} // TODO: remove allocation
- getSrcFromControl(msg.OOB[:msg.NN], ep)
- eps[i] = ep
+ } else {
+ msg := &(*msgs)[0]
+ msg.N, msg.NN, _, msg.Addr, err = conn.ReadMsgUDP(msg.Buffers[0], msg.OOB)
+ if err != nil {
+ return 0, err
}
- return numMsgs, nil
+ numMsgs = 1
+ }
+ for i := 0; i < numMsgs; i++ {
+ msg := &(*msgs)[i]
+ sizes[i] = msg.N
+ if sizes[i] == 0 {
+ continue
+ }
+ addrPort := msg.Addr.(*net.UDPAddr).AddrPort()
+ ep := &StdNetEndpoint{AddrPort: addrPort} // TODO: remove allocation
+ getSrcFromControl(msg.OOB[:msg.NN], ep)
+ eps[i] = ep
+ }
+ return numMsgs, nil
+}
+
+func (s *StdNetBind) makeReceiveIPv4(pc *ipv4.PacketConn, conn *net.UDPConn, rxOffload bool) ReceiveFunc {
+ return func(bufs [][]byte, sizes []int, eps []Endpoint) (n int, err error) {
+ return s.receiveIP(pc, conn, rxOffload, bufs, sizes, eps)
}
}
-func (s *StdNetBind) makeReceiveIPv6(pc *ipv6.PacketConn, conn *net.UDPConn) ReceiveFunc {
+func (s *StdNetBind) makeReceiveIPv6(pc *ipv6.PacketConn, conn *net.UDPConn, rxOffload bool) ReceiveFunc {
return func(bufs [][]byte, sizes []int, eps []Endpoint) (n int, err error) {
- msgs := s.ipv6MsgsPool.Get().(*[]ipv6.Message)
- defer s.ipv6MsgsPool.Put(msgs)
- for i := range bufs {
- (*msgs)[i].Buffers[0] = bufs[i]
- }
- var numMsgs int
- if runtime.GOOS == "linux" {
- numMsgs, err = pc.ReadBatch(*msgs, 0)
- if err != nil {
- return 0, err
- }
- } else {
- msg := &(*msgs)[0]
- msg.N, msg.NN, _, msg.Addr, err = conn.ReadMsgUDP(msg.Buffers[0], msg.OOB)
- if err != nil {
- return 0, err
- }
- numMsgs = 1
- }
- for i := 0; i < numMsgs; i++ {
- msg := &(*msgs)[i]
- sizes[i] = msg.N
- addrPort := msg.Addr.(*net.UDPAddr).AddrPort()
- ep := &StdNetEndpoint{AddrPort: addrPort} // TODO: remove allocation
- getSrcFromControl(msg.OOB[:msg.NN], ep)
- eps[i] = ep
- }
- return numMsgs, nil
+ return s.receiveIP(pc, conn, rxOffload, bufs, sizes, eps)
}
}
@@ -293,28 +314,42 @@ func (s *StdNetBind) Close() error {
}
s.blackhole4 = false
s.blackhole6 = false
+ s.ipv4TxOffload = false
+ s.ipv4RxOffload = false
+ s.ipv6TxOffload = false
+ s.ipv6RxOffload = false
if err1 != nil {
return err1
}
return err2
}
+type ErrUDPGSODisabled struct {
+ onLaddr string
+ RetryErr error
+}
+
+func (e ErrUDPGSODisabled) Error() string {
+ return fmt.Sprintf("disabled UDP GSO on %s, NIC(s) may not support checksum offload", e.onLaddr)
+}
+
+func (e ErrUDPGSODisabled) Unwrap() error {
+ return e.RetryErr
+}
+
func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint) error {
s.mu.Lock()
blackhole := s.blackhole4
conn := s.ipv4
- var (
- pc4 *ipv4.PacketConn
- pc6 *ipv6.PacketConn
- )
+ offload := s.ipv4TxOffload
+ br := batchWriter(s.ipv4PC)
is6 := false
if endpoint.DstIP().Is6() {
blackhole = s.blackhole6
conn = s.ipv6
- pc6 = s.ipv6PC
+ br = s.ipv6PC
is6 = true
- } else {
- pc4 = s.ipv4PC
+ offload = s.ipv6TxOffload
}
s.mu.Unlock()
@@ -324,25 +359,56 @@ func (s *StdNetBind) Send(bufs [][]byte, endpoint Endpoint) error {
if conn == nil {
return syscall.EAFNOSUPPORT
}
+
+ msgs := s.getMessages()
+ defer s.putMessages(msgs)
+ ua := s.udpAddrPool.Get().(*net.UDPAddr)
+ defer s.udpAddrPool.Put(ua)
if is6 {
- return s.send6(conn, pc6, endpoint, bufs)
+ as16 := endpoint.DstIP().As16()
+ copy(ua.IP, as16[:])
+ ua.IP = ua.IP[:16]
} else {
- return s.send4(conn, pc4, endpoint, bufs)
+ as4 := endpoint.DstIP().As4()
+ copy(ua.IP, as4[:])
+ ua.IP = ua.IP[:4]
}
+ ua.Port = int(endpoint.(*StdNetEndpoint).Port())
+ var (
+ retried bool
+ err error
+ )
+retry:
+ if offload {
+ n := coalesceMessages(ua, endpoint.(*StdNetEndpoint), bufs, *msgs, setGSOSize)
+ err = s.send(conn, br, (*msgs)[:n])
+ if err != nil && offload && errShouldDisableUDPGSO(err) {
+ offload = false
+ s.mu.Lock()
+ if is6 {
+ s.ipv6TxOffload = false
+ } else {
+ s.ipv4TxOffload = false
+ }
+ s.mu.Unlock()
+ retried = true
+ goto retry
+ }
+ } else {
+ for i := range bufs {
+ (*msgs)[i].Addr = ua
+ (*msgs)[i].Buffers[0] = bufs[i]
+ setSrcControl(&(*msgs)[i].OOB, endpoint.(*StdNetEndpoint))
+ }
+ err = s.send(conn, br, (*msgs)[:len(bufs)])
+ }
+ if retried {
+ return ErrUDPGSODisabled{onLaddr: conn.LocalAddr().String(), RetryErr: err}
+ }
+ return err
}
-func (s *StdNetBind) send4(conn *net.UDPConn, pc *ipv4.PacketConn, ep Endpoint, bufs [][]byte) error {
- ua := s.udpAddrPool.Get().(*net.UDPAddr)
- as4 := ep.DstIP().As4()
- copy(ua.IP, as4[:])
- ua.IP = ua.IP[:4]
- ua.Port = int(ep.(*StdNetEndpoint).Port())
- msgs := s.ipv4MsgsPool.Get().(*[]ipv4.Message)
- for i, buf := range bufs {
- (*msgs)[i].Buffers[0] = buf
- (*msgs)[i].Addr = ua
- setSrcControl(&(*msgs)[i].OOB, ep.(*StdNetEndpoint))
- }
+func (s *StdNetBind) send(conn *net.UDPConn, pc batchWriter, msgs []ipv6.Message) error {
var (
n int
err error
@@ -350,59 +416,128 @@ func (s *StdNetBind) send4(conn *net.UDPConn, pc *ipv4.PacketConn, ep Endpoint,
)
if runtime.GOOS == "linux" {
for {
- n, err = pc.WriteBatch((*msgs)[start:len(bufs)], 0)
- if err != nil || n == len((*msgs)[start:len(bufs)]) {
+ n, err = pc.WriteBatch(msgs[start:], 0)
+ if err != nil || n == len(msgs[start:]) {
break
}
start += n
}
} else {
- for i, buf := range bufs {
- _, _, err = conn.WriteMsgUDP(buf, (*msgs)[i].OOB, ua)
+ for _, msg := range msgs {
+ _, _, err = conn.WriteMsgUDP(msg.Buffers[0], msg.OOB, msg.Addr.(*net.UDPAddr))
if err != nil {
break
}
}
}
- s.udpAddrPool.Put(ua)
- s.ipv4MsgsPool.Put(msgs)
return err
}
-func (s *StdNetBind) send6(conn *net.UDPConn, pc *ipv6.PacketConn, ep Endpoint, bufs [][]byte) error {
- ua := s.udpAddrPool.Get().(*net.UDPAddr)
- as16 := ep.DstIP().As16()
- copy(ua.IP, as16[:])
- ua.IP = ua.IP[:16]
- ua.Port = int(ep.(*StdNetEndpoint).Port())
- msgs := s.ipv6MsgsPool.Get().(*[]ipv6.Message)
- for i, buf := range bufs {
- (*msgs)[i].Buffers[0] = buf
- (*msgs)[i].Addr = ua
- setSrcControl(&(*msgs)[i].OOB, ep.(*StdNetEndpoint))
- }
+const (
+ // Exceeding these values results in EMSGSIZE. They account for layer3 and
+ // layer4 headers. IPv6 does not need to account for itself as the payload
+ // length field is self excluding.
+ maxIPv4PayloadLen = 1<<16 - 1 - 20 - 8
+ maxIPv6PayloadLen = 1<<16 - 1 - 8
+
+ // This is a hard limit imposed by the kernel.
+ udpSegmentMaxDatagrams = 64
+)
+
+type setGSOFunc func(control *[]byte, gsoSize uint16)
+
+func coalesceMessages(addr *net.UDPAddr, ep *StdNetEndpoint, bufs [][]byte, msgs []ipv6.Message, setGSO setGSOFunc) int {
var (
- n int
- err error
- start int
+ base = -1 // index of msg we are currently coalescing into
+ gsoSize int // segmentation size of msgs[base]
+ dgramCnt int // number of dgrams coalesced into msgs[base]
+ endBatch bool // tracking flag to start a new batch on next iteration of bufs
)
- if runtime.GOOS == "linux" {
- for {
- n, err = pc.WriteBatch((*msgs)[start:len(bufs)], 0)
- if err != nil || n == len((*msgs)[start:len(bufs)]) {
- break
+ maxPayloadLen := maxIPv4PayloadLen
+ if ep.DstIP().Is6() {
+ maxPayloadLen = maxIPv6PayloadLen
+ }
+ for i, buf := range bufs {
+ if i > 0 {
+ msgLen := len(buf)
+ baseLenBefore := len(msgs[base].Buffers[0])
+ freeBaseCap := cap(msgs[base].Buffers[0]) - baseLenBefore
+ if msgLen+baseLenBefore <= maxPayloadLen &&
+ msgLen <= gsoSize &&
+ msgLen <= freeBaseCap &&
+ dgramCnt < udpSegmentMaxDatagrams &&
+ !endBatch {
+ msgs[base].Buffers[0] = append(msgs[base].Buffers[0], buf...)
+ if i == len(bufs)-1 {
+ setGSO(&msgs[base].OOB, uint16(gsoSize))
+ }
+ dgramCnt++
+ if msgLen < gsoSize {
+ // A smaller than gsoSize packet on the tail is legal, but
+ // it must end the batch.
+ endBatch = true
+ }
+ continue
}
- start += n
}
- } else {
- for i, buf := range bufs {
- _, _, err = conn.WriteMsgUDP(buf, (*msgs)[i].OOB, ua)
- if err != nil {
- break
+ if dgramCnt > 1 {
+ setGSO(&msgs[base].OOB, uint16(gsoSize))
+ }
+ // Reset prior to incrementing base since we are preparing to start a
+ // new potential batch.
+ endBatch = false
+ base++
+ gsoSize = len(buf)
+ setSrcControl(&msgs[base].OOB, ep)
+ msgs[base].Buffers[0] = buf
+ msgs[base].Addr = addr
+ dgramCnt = 1
+ }
+ return base + 1
+}
+
+type getGSOFunc func(control []byte) (int, error)
+
+func splitCoalescedMessages(msgs []ipv6.Message, firstMsgAt int, getGSO getGSOFunc) (n int, err error) {
+ for i := firstMsgAt; i < len(msgs); i++ {
+ msg := &msgs[i]
+ if msg.N == 0 {
+ return n, err
+ }
+ var (
+ gsoSize int
+ start int
+ end = msg.N
+ numToSplit = 1
+ )
+ gsoSize, err = getGSO(msg.OOB[:msg.NN])
+ if err != nil {
+ return n, err
+ }
+ if gsoSize > 0 {
+ numToSplit = (msg.N + gsoSize - 1) / gsoSize
+ end = gsoSize
+ }
+ for j := 0; j < numToSplit; j++ {
+ if n > i {
+ return n, errors.New("splitting coalesced packet resulted in overflow")
}
+ copied := copy(msgs[n].Buffers[0], msg.Buffers[0][start:end])
+ msgs[n].N = copied
+ msgs[n].Addr = msg.Addr
+ start = end
+ end += gsoSize
+ if end > msg.N {
+ end = msg.N
+ }
+ n++
+ }
+ if i != n-1 {
+ // It is legal for bytes to move within msg.Buffers[0] as a result
+ // of splitting, so we only zero the source msg len when it is not
+ // the destination of the last split operation above.
+ msg.N = 0
}
}
- s.udpAddrPool.Put(ua)
- s.ipv6MsgsPool.Put(msgs)
- return err
+ return n, nil
}
diff --git a/conn/bind_std_test.go b/conn/bind_std_test.go
index 1e46776..34a3c9a 100644
--- a/conn/bind_std_test.go
+++ b/conn/bind_std_test.go
@@ -1,6 +1,12 @@
package conn
-import "testing"
+import (
+ "encoding/binary"
+ "net"
+ "testing"
+
+ "golang.org/x/net/ipv6"
+)
func TestStdNetBindReceiveFuncAfterClose(t *testing.T) {
bind := NewStdNetBind().(*StdNetBind)
@@ -20,3 +26,225 @@ func TestStdNetBindReceiveFuncAfterClose(t *testing.T) {
fn(bufs, sizes, eps)
}
}
+
+func mockSetGSOSize(control *[]byte, gsoSize uint16) {
+ *control = (*control)[:cap(*control)]
+ binary.LittleEndian.PutUint16(*control, gsoSize)
+}
+
+func Test_coalesceMessages(t *testing.T) {
+ cases := []struct {
+ name string
+ buffs [][]byte
+ wantLens []int
+ wantGSO []int
+ }{
+ {
+ name: "one message no coalesce",
+ buffs: [][]byte{
+ make([]byte, 1, 1),
+ },
+ wantLens: []int{1},
+ wantGSO: []int{0},
+ },
+ {
+ name: "two messages equal len coalesce",
+ buffs: [][]byte{
+ make([]byte, 1, 2),
+ make([]byte, 1, 1),
+ },
+ wantLens: []int{2},
+ wantGSO: []int{1},
+ },
+ {
+ name: "two messages unequal len coalesce",
+ buffs: [][]byte{
+ make([]byte, 2, 3),
+ make([]byte, 1, 1),
+ },
+ wantLens: []int{3},
+ wantGSO: []int{2},
+ },
+ {
+ name: "three messages second unequal len coalesce",
+ buffs: [][]byte{
+ make([]byte, 2, 3),
+ make([]byte, 1, 1),
+ make([]byte, 2, 2),
+ },
+ wantLens: []int{3, 2},
+ wantGSO: []int{2, 0},
+ },
+ {
+ name: "three messages limited cap coalesce",
+ buffs: [][]byte{
+ make([]byte, 2, 4),
+ make([]byte, 2, 2),
+ make([]byte, 2, 2),
+ },
+ wantLens: []int{4, 2},
+ wantGSO: []int{2, 0},
+ },
+ }
+
+ for _, tt := range cases {
+ t.Run(tt.name, func(t *testing.T) {
+ addr := &net.UDPAddr{
+ IP: net.ParseIP("127.0.0.1").To4(),
+ Port: 1,
+ }
+ msgs := make([]ipv6.Message, len(tt.buffs))
+ for i := range msgs {
+ msgs[i].Buffers = make([][]byte, 1)
+ msgs[i].OOB = make([]byte, 0, 2)
+ }
+ got := coalesceMessages(addr, &StdNetEndpoint{AddrPort: addr.AddrPort()}, tt.buffs, msgs, mockSetGSOSize)
+ if got != len(tt.wantLens) {
+ t.Fatalf("got len %d want: %d", got, len(tt.wantLens))
+ }
+ for i := 0; i < got; i++ {
+ if msgs[i].Addr != addr {
+ t.Errorf("msgs[%d].Addr != passed addr", i)
+ }
+ gotLen := len(msgs[i].Buffers[0])
+ if gotLen != tt.wantLens[i] {
+ t.Errorf("len(msgs[%d].Buffers[0]) %d != %d", i, gotLen, tt.wantLens[i])
+ }
+ gotGSO, err := mockGetGSOSize(msgs[i].OOB)
+ if err != nil {
+ t.Fatalf("msgs[%d] getGSOSize err: %v", i, err)
+ }
+ if gotGSO != tt.wantGSO[i] {
+ t.Errorf("msgs[%d] gsoSize %d != %d", i, gotGSO, tt.wantGSO[i])
+ }
+ }
+ })
+ }
+}
+
+func mockGetGSOSize(control []byte) (int, error) {
+ if len(control) < 2 {
+ return 0, nil
+ }
+ return int(binary.LittleEndian.Uint16(control)), nil
+}
+
+func Test_splitCoalescedMessages(t *testing.T) {
+ newMsg := func(n, gso int) ipv6.Message {
+ msg := ipv6.Message{
+ Buffers: [][]byte{make([]byte, 1<<16-1)},
+ N: n,
+ OOB: make([]byte, 2),
+ }
+ binary.LittleEndian.PutUint16(msg.OOB, uint16(gso))
+ if gso > 0 {
+ msg.NN = 2
+ }
+ return msg
+ }
+
+ cases := []struct {
+ name string
+ msgs []ipv6.Message
+ firstMsgAt int
+ wantNumEval int
+ wantMsgLens []int
+ wantErr bool
+ }{
+ {
+ name: "second last split last empty",
+ msgs: []ipv6.Message{
+ newMsg(0, 0),
+ newMsg(0, 0),
+ newMsg(3, 1),
+ newMsg(0, 0),
+ },
+ firstMsgAt: 2,
+ wantNumEval: 3,
+ wantMsgLens: []int{1, 1, 1, 0},
+ wantErr: false,
+ },
+ {
+ name: "second last no split last empty",
+ msgs: []ipv6.Message{
+ newMsg(0, 0),
+ newMsg(0, 0),
+ newMsg(1, 0),
+ newMsg(0, 0),
+ },
+ firstMsgAt: 2,
+ wantNumEval: 1,
+ wantMsgLens: []int{1, 0, 0, 0},
+ wantErr: false,
+ },
+ {
+ name: "second last no split last no split",
+ msgs: []ipv6.Message{
+ newMsg(0, 0),
+ newMsg(0, 0),
+ newMsg(1, 0),
+ newMsg(1, 0),
+ },
+ firstMsgAt: 2,
+ wantNumEval: 2,
+ wantMsgLens: []int{1, 1, 0, 0},
+ wantErr: false,
+ },
+ {
+ name: "second last no split last split",
+ msgs: []ipv6.Message{
+ newMsg(0, 0),
+ newMsg(0, 0),
+ newMsg(1, 0),
+ newMsg(3, 1),
+ },
+ firstMsgAt: 2,
+ wantNumEval: 4,
+ wantMsgLens: []int{1, 1, 1, 1},
+ wantErr: false,
+ },
+ {
+ name: "second last split last split",
+ msgs: []ipv6.Message{
+ newMsg(0, 0),
+ newMsg(0, 0),
+ newMsg(2, 1),
+ newMsg(2, 1),
+ },
+ firstMsgAt: 2,
+ wantNumEval: 4,
+ wantMsgLens: []int{1, 1, 1, 1},
+ wantErr: false,
+ },
+ {
+ name: "second last no split last split overflow",
+ msgs: []ipv6.Message{
+ newMsg(0, 0),
+ newMsg(0, 0),
+ newMsg(1, 0),
+ newMsg(4, 1),
+ },
+ firstMsgAt: 2,
+ wantNumEval: 4,
+ wantMsgLens: []int{1, 1, 1, 1},
+ wantErr: true,
+ },
+ }
+
+ for _, tt := range cases {
+ t.Run(tt.name, func(t *testing.T) {
+ got, err := splitCoalescedMessages(tt.msgs, 2, mockGetGSOSize)
+ if err != nil && !tt.wantErr {
+ t.Fatalf("err: %v", err)
+ }
+ if got != tt.wantNumEval {
+ t.Fatalf("got to eval: %d want: %d", got, tt.wantNumEval)
+ }
+ for i, msg := range tt.msgs {
+ if msg.N != tt.wantMsgLens[i] {
+ t.Fatalf("msg[%d].N: %d want: %d", i, msg.N, tt.wantMsgLens[i])
+ }
+ }
+ })
+ }
+}
diff --git a/conn/sticky_default.go b/conn/control_default.go
similarity index 56%
rename from conn/sticky_default.go
rename to conn/control_default.go
index 1fa8a0c..9459da5 100644
--- a/conn/sticky_default.go
+++ b/conn/control_default.go
@@ -21,8 +21,9 @@ func (e *StdNetEndpoint) SrcToString() string {
return ""
}
-// TODO: macOS, FreeBSD and other BSDs likely do support this feature set, but
-// use alternatively named flags and need ports and require testing.
+// TODO: macOS, FreeBSD and other BSDs likely do support the sticky sockets
+// {get,set}srcControl feature set, but use alternatively named flags and need
+// ports and require testing.
// getSrcFromControl parses the control for PKTINFO and if found updates ep with
// the source information found.
@@ -34,8 +35,17 @@ func getSrcFromControl(control []byte, ep *StdNetEndpoint) {
func setSrcControl(control *[]byte, ep *StdNetEndpoint) {
}
-// srcControlSize returns the recommended buffer size for pooling sticky control
-// data.
-const srcControlSize = 0
+// getGSOSize parses control for UDP_GRO and if found returns its GSO size data.
+func getGSOSize(control []byte) (int, error) {
+ return 0, nil
+}
+
+// setGSOSize sets a UDP_SEGMENT in control based on gsoSize.
+func setGSOSize(control *[]byte, gsoSize uint16) {
+}
+
+// controlSize returns the recommended buffer size for pooling sticky and UDP
+// offloading control data.
+const controlSize = 0
const StdNetSupportsStickySockets = false
diff --git a/conn/sticky_linux.go b/conn/control_linux.go
similarity index 65%
rename from conn/sticky_linux.go
rename to conn/control_linux.go
index a30ccc7..44a94e6 100644
--- a/conn/sticky_linux.go
+++ b/conn/control_linux.go
@@ -8,6 +8,7 @@
package conn
import (
+ "fmt"
"net/netip"
"unsafe"
@@ -105,6 +106,54 @@ func setSrcControl(control *[]byte, ep *StdNetEndpoint) {
*control = append(*control, ep.src...)
}
-var srcControlSize = unix.CmsgSpace(unix.SizeofInet6Pktinfo)
+const (
+ sizeOfGSOData = 2
+)
+
+// getGSOSize parses control for UDP_GRO and if found returns its GSO size data.
+func getGSOSize(control []byte) (int, error) {
+ var (
+ hdr unix.Cmsghdr
+ data []byte
+ rem = control
+ err error
+ )
+
+ for len(rem) > unix.SizeofCmsghdr {
+ hdr, data, rem, err = unix.ParseOneSocketControlMessage(rem)
+ if err != nil {
+ return 0, fmt.Errorf("error parsing socket control message: %w", err)
+ }
+ if hdr.Level == unix.SOL_UDP && hdr.Type == unix.UDP_GRO && len(data) >= sizeOfGSOData {
+ var gso uint16
+ copy(unsafe.Slice((*byte)(unsafe.Pointer(&gso)), sizeOfGSOData), data[:sizeOfGSOData])
+ return int(gso), nil
+ }
+ }
+ return 0, nil
+}
+
+// setGSOSize sets a UDP_SEGMENT in control based on gsoSize. It leaves existing
+// data in control untouched.
+func setGSOSize(control *[]byte, gsoSize uint16) {
+ existingLen := len(*control)
+ avail := cap(*control) - existingLen
+ space := unix.CmsgSpace(sizeOfGSOData)
+ if avail < space {
+ return
+ }
+ *control = (*control)[:cap(*control)]
+ gsoControl := (*control)[existingLen:]
+ hdr := (*unix.Cmsghdr)(unsafe.Pointer(&(gsoControl)[0]))
+ hdr.Level = unix.SOL_UDP
+ hdr.Type = unix.UDP_SEGMENT
+ hdr.SetLen(unix.CmsgLen(sizeOfGSOData))
+ copy((gsoControl)[unix.SizeofCmsghdr:], unsafe.Slice((*byte)(unsafe.Pointer(&gsoSize)), sizeOfGSOData))
+ *control = (*control)[:existingLen+space]
+}
+
+// controlSize returns the recommended buffer size for pooling sticky and UDP
+// offloading control data.
+var controlSize = unix.CmsgSpace(unix.SizeofInet6Pktinfo) + unix.CmsgSpace(sizeOfGSOData)
const StdNetSupportsStickySockets = true
diff --git a/conn/sticky_linux_test.go b/conn/control_linux_test.go
similarity index 96%
rename from conn/sticky_linux_test.go
rename to conn/control_linux_test.go
index 679213a..96f9da2 100644
--- a/conn/sticky_linux_test.go
+++ b/conn/control_linux_test.go
@@ -60,7 +60,7 @@ func Test_setSrcControl(t *testing.T) {
}
setSrc(ep, netip.MustParseAddr("127.0.0.1"), 5)
- control := make([]byte, srcControlSize)
+ control := make([]byte, controlSize)
setSrcControl(&control, ep)
@@ -89,7 +89,7 @@ func Test_setSrcControl(t *testing.T) {
}
setSrc(ep, netip.MustParseAddr("::1"), 5)
- control := make([]byte, srcControlSize)
+ control := make([]byte, controlSize)
setSrcControl(&control, ep)
@@ -113,7 +113,7 @@ func Test_setSrcControl(t *testing.T) {
})
t.Run("ClearOnNoSrc", func(t *testing.T) {
- control := make([]byte, unix.CmsgLen(0))
+ control := make([]byte, controlSize)
hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0]))
hdr.Level = 1
hdr.Type = 2
@@ -129,7 +129,7 @@ func Test_setSrcControl(t *testing.T) {
func Test_getSrcFromControl(t *testing.T) {
t.Run("IPv4", func(t *testing.T) {
- control := make([]byte, unix.CmsgSpace(unix.SizeofInet4Pktinfo))
+ control := make([]byte, controlSize)
hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0]))
hdr.Level = unix.IPPROTO_IP
hdr.Type = unix.IP_PKTINFO
@@ -149,7 +149,7 @@ func Test_getSrcFromControl(t *testing.T) {
}
})
t.Run("IPv6", func(t *testing.T) {
- control := make([]byte, unix.CmsgSpace(unix.SizeofInet6Pktinfo))
+ control := make([]byte, controlSize)
hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0]))
hdr.Level = unix.IPPROTO_IPV6
hdr.Type = unix.IPV6_PKTINFO
diff --git a/conn/controlfns_linux.go b/conn/controlfns_linux.go
index a2396fe..f6ab1d2 100644
--- a/conn/controlfns_linux.go
+++ b/conn/controlfns_linux.go
@@ -57,5 +57,13 @@ func init() {
}
return err
},
+
+ // Attempt to enable UDP_GRO
+ func(network, address string, c syscall.RawConn) error {
+ c.Control(func(fd uintptr) {
+ _ = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1)
+ })
+ return nil
+ },
)
}
diff --git a/conn/errors_default.go b/conn/errors_default.go
new file mode 100644
index 0000000..f1e5b90
--- /dev/null
+++ b/conn/errors_default.go
@@ -0,0 +1,12 @@
+//go:build !linux
+
+/* SPDX-License-Identifier: MIT
+ *
+ * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ */
+
+package conn
+
+func errShouldDisableUDPGSO(err error) bool {
+ return false
+}
diff --git a/conn/errors_linux.go b/conn/errors_linux.go
new file mode 100644
index 0000000..8e61000
--- /dev/null
+++ b/conn/errors_linux.go
@@ -0,0 +1,26 @@
+/* SPDX-License-Identifier: MIT
+ *
+ * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ */
+
+package conn
+
+import (
+ "errors"
+ "os"
+
+ "golang.org/x/sys/unix"
+)
+
+func errShouldDisableUDPGSO(err error) bool {
+ var serr *os.SyscallError
+ if errors.As(err, &serr) {
+ // EIO is returned by udp_send_skb() if the device driver does not have
+ // tx checksumming enabled, which is a hard requirement of UDP_SEGMENT.
+ // See:
+ // https://git.kernel.org/pub/scm/docs/man-pages/man-pages.git/tree/man7/udp.7?id=806eabd74910447f21005160e90957bde4db0183#n228
+ // https://git.kernel.org/pub/scm/linux/kernel/git/torvalds/linux.git/tree/net/ipv4/udp.c?h=v6.2&id=c9c3395d5e3dcc6daee66c6908354d47bf98cb0c#n942
+ return serr.Err == unix.EIO
+ }
+ return false
+}
diff --git a/conn/features_default.go b/conn/features_default.go
new file mode 100644
index 0000000..d53ff5f
--- /dev/null
+++ b/conn/features_default.go
@@ -0,0 +1,15 @@
+//go:build !linux
+// +build !linux
+
+/* SPDX-License-Identifier: MIT
+ *
+ * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ */
+
+package conn
+
+import "net"
+
+func supportsUDPOffload(conn *net.UDPConn) (txOffload, rxOffload bool) {
+ return
+}
diff --git a/conn/features_linux.go b/conn/features_linux.go
new file mode 100644
index 0000000..e1fb57f
--- /dev/null
+++ b/conn/features_linux.go
@@ -0,0 +1,35 @@
+/* SPDX-License-Identifier: MIT
+ *
+ * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ */
+
+package conn
+
+import (
+ "net"
+
+ "golang.org/x/sys/unix"
+)
+
+func supportsUDPOffload(conn *net.UDPConn) (txOffload, rxOffload bool) {
+ rc, err := conn.SyscallConn()
+ if err != nil {
+ return
+ }
+ err = rc.Control(func(fd uintptr) {
+ _, errSyscall := unix.GetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_SEGMENT)
+ if errSyscall != nil {
+ return
+ }
+ txOffload = true
+ opt, errSyscall := unix.GetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO)
+ if errSyscall != nil {
+ return
+ }
+ rxOffload = opt == 1
+ })
+ if err != nil {
+ return false, false
+ }
+ return txOffload, rxOffload
+}
diff --git a/device/send.go b/device/send.go
index d22bf26..cd8a2a0 100644
--- a/device/send.go
+++ b/device/send.go
@@ -17,6 +17,7 @@ import (
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
+ "golang.zx2c4.com/wireguard/conn"
"golang.zx2c4.com/wireguard/tun"
)
@@ -525,6 +526,13 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
device.PutOutboundElement(elem)
}
device.PutOutboundElementsSlice(elems)
+ if err != nil {
+ var errGSO conn.ErrUDPGSODisabled
+ if errors.As(err, &errGSO) {
+ device.log.Verbosef(err.Error())
+ err = errGSO.RetryErr
+ }
+ }
if err != nil {
device.log.Errorf("%v - Failed to send data packets: %v", peer, err)
continue
diff --git a/go.mod b/go.mod
index c04e1bb..758dcde 100644
--- a/go.mod
+++ b/go.mod
@@ -5,7 +5,7 @@ go 1.20
require (
golang.org/x/crypto v0.6.0
golang.org/x/net v0.7.0
- golang.org/x/sys v0.5.1-0.20230222185716-a3b23cc77e89
+ golang.org/x/sys v0.12.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
gvisor.dev/gvisor v0.0.0-20221203005347-703fd9b7fbc0
)
diff --git a/go.sum b/go.sum
index cfeaee6..fe4ca7e 100644
--- a/go.sum
+++ b/go.sum
@@ -4,8 +4,8 @@ golang.org/x/crypto v0.6.0 h1:qfktjS5LUO+fFKeJXZ+ikTRijMmljikvG68fpMMruSc=
golang.org/x/crypto v0.6.0/go.mod h1:OFC/31mSvZgRz0V1QTNCzfAI1aIRzbiufJtkMIlEp58=
golang.org/x/net v0.7.0 h1:rJrUqqhjsgNp7KqAIc25s9pZnjU7TUcSY7HcVZjdn1g=
golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
-golang.org/x/sys v0.5.1-0.20230222185716-a3b23cc77e89 h1:260HNjMTPDya+jq5AM1zZLgG9pv9GASPAGiEEJUbRg4=
-golang.org/x/sys v0.5.1-0.20230222185716-a3b23cc77e89/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
+golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o=
+golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/time v0.0.0-20191024005414-555d28b269f0 h1:/5xXl8Y5W96D+TtHSlonuFqGHIWVuyCkGJLwGh9JJFs=
golang.org/x/time v0.0.0-20191024005414-555d28b269f0/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
From 4201e08f1dbb521e5555d96a3b6464a860466f5f Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Mon, 2 Oct 2023 14:41:04 -0700
Subject: [PATCH 04/75] device: distribute crypto work as slice of elements
After reducing UDP stack traversal overhead via GSO and GRO,
runtime.chanrecv() began to account for a high percentage (20% in one
environment) of perf samples during a throughput benchmark. The
individual packet channel ops with the crypto goroutines was the primary
contributor to this overhead.
Updating these channels to pass vectors, which the device package
already handles at its ends, reduced this overhead substantially, and
improved throughput.
The iperf3 results below demonstrate the effect of this commit between
two Linux computers with i5-12400 CPUs. There is roughly ~13us of round
trip latency between them.
The first result is with UDP GSO and GRO, and with single element
channels.
Starting Test: protocol: TCP, 1 streams, 131072 byte blocks
[ ID] Interval Transfer Bitrate Retr Cwnd
[ 5] 0.00-10.00 sec 12.3 GBytes 10.6 Gbits/sec 232 3.15 MBytes
- - - - - - - - - - - - - - - - - - - - - - - - -
Test Complete. Summary Results:
[ ID] Interval Transfer Bitrate Retr
[ 5] 0.00-10.00 sec 12.3 GBytes 10.6 Gbits/sec 232 sender
[ 5] 0.00-10.04 sec 12.3 GBytes 10.6 Gbits/sec receiver
The second result is with channels updated to pass a slice of
elements.
Starting Test: protocol: TCP, 1 streams, 131072 byte blocks
[ ID] Interval Transfer Bitrate Retr Cwnd
[ 5] 0.00-10.00 sec 13.2 GBytes 11.3 Gbits/sec 182 3.15 MBytes
- - - - - - - - - - - - - - - - - - - - - - - - -
Test Complete. Summary Results:
[ ID] Interval Transfer Bitrate Retr
[ 5] 0.00-10.00 sec 13.2 GBytes 11.3 Gbits/sec 182 sender
[ 5] 0.00-10.04 sec 13.2 GBytes 11.3 Gbits/sec receiver
Reviewed-by: Adrian Dewhurst
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
device/channels.go | 8 ++++----
device/receive.go | 42 ++++++++++++++++++++--------------------
device/send.go | 48 +++++++++++++++++++++++-----------------------
3 files changed, 49 insertions(+), 49 deletions(-)
diff --git a/device/channels.go b/device/channels.go
index 039d8df..40ee5c9 100644
--- a/device/channels.go
+++ b/device/channels.go
@@ -19,13 +19,13 @@ import (
// call wg.Done to remove the initial reference.
// When the refcount hits 0, the queue's channel is closed.
type outboundQueue struct {
- c chan *QueueOutboundElement
+ c chan *[]*QueueOutboundElement
wg sync.WaitGroup
}
func newOutboundQueue() *outboundQueue {
q := &outboundQueue{
- c: make(chan *QueueOutboundElement, QueueOutboundSize),
+ c: make(chan *[]*QueueOutboundElement, QueueOutboundSize),
}
q.wg.Add(1)
go func() {
@@ -37,13 +37,13 @@ func newOutboundQueue() *outboundQueue {
// A inboundQueue is similar to an outboundQueue; see those docs.
type inboundQueue struct {
- c chan *QueueInboundElement
+ c chan *[]*QueueInboundElement
wg sync.WaitGroup
}
func newInboundQueue() *inboundQueue {
q := &inboundQueue{
- c: make(chan *QueueInboundElement, QueueInboundSize),
+ c: make(chan *[]*QueueInboundElement, QueueInboundSize),
}
q.wg.Add(1)
go func() {
diff --git a/device/receive.go b/device/receive.go
index e24d29f..f0f37a1 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -220,9 +220,7 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive
for peer, elems := range elemsByPeer {
if peer.isRunning.Load() {
peer.queue.inbound.c <- elems
- for _, elem := range *elems {
- device.queue.decryption.c <- elem
- }
+ device.queue.decryption.c <- elems
} else {
for _, elem := range *elems {
device.PutMessageBuffer(elem.buffer)
@@ -241,26 +239,28 @@ func (device *Device) RoutineDecryption(id int) {
defer device.log.Verbosef("Routine: decryption worker %d - stopped", id)
device.log.Verbosef("Routine: decryption worker %d - started", id)
- for elem := range device.queue.decryption.c {
- // split message into fields
- counter := elem.packet[MessageTransportOffsetCounter:MessageTransportOffsetContent]
- content := elem.packet[MessageTransportOffsetContent:]
+ for elems := range device.queue.decryption.c {
+ for _, elem := range *elems {
+ // split message into fields
+ counter := elem.packet[MessageTransportOffsetCounter:MessageTransportOffsetContent]
+ content := elem.packet[MessageTransportOffsetContent:]
- // decrypt and release to consumer
- var err error
- elem.counter = binary.LittleEndian.Uint64(counter)
- // copy counter to nonce
- binary.LittleEndian.PutUint64(nonce[0x4:0xc], elem.counter)
- elem.packet, err = elem.keypair.receive.Open(
- content[:0],
- nonce[:],
- content,
- nil,
- )
- if err != nil {
- elem.packet = nil
+ // decrypt and release to consumer
+ var err error
+ elem.counter = binary.LittleEndian.Uint64(counter)
+ // copy counter to nonce
+ binary.LittleEndian.PutUint64(nonce[0x4:0xc], elem.counter)
+ elem.packet, err = elem.keypair.receive.Open(
+ content[:0],
+ nonce[:],
+ content,
+ nil,
+ )
+ if err != nil {
+ elem.packet = nil
+ }
+ elem.Unlock()
}
- elem.Unlock()
}
}
diff --git a/device/send.go b/device/send.go
index cd8a2a0..e838c4e 100644
--- a/device/send.go
+++ b/device/send.go
@@ -385,9 +385,7 @@ top:
// add to parallel and sequential queue
if peer.isRunning.Load() {
peer.queue.outbound.c <- elems
- for _, elem := range *elems {
- peer.device.queue.encryption.c <- elem
- }
+ peer.device.queue.encryption.c <- elems
} else {
for _, elem := range *elems {
peer.device.PutMessageBuffer(elem.buffer)
@@ -447,32 +445,34 @@ func (device *Device) RoutineEncryption(id int) {
defer device.log.Verbosef("Routine: encryption worker %d - stopped", id)
device.log.Verbosef("Routine: encryption worker %d - started", id)
- for elem := range device.queue.encryption.c {
- // populate header fields
- header := elem.buffer[:MessageTransportHeaderSize]
+ for elems := range device.queue.encryption.c {
+ for _, elem := range *elems {
+ // populate header fields
+ header := elem.buffer[:MessageTransportHeaderSize]
- fieldType := header[0:4]
- fieldReceiver := header[4:8]
- fieldNonce := header[8:16]
+ fieldType := header[0:4]
+ fieldReceiver := header[4:8]
+ fieldNonce := header[8:16]
- binary.LittleEndian.PutUint32(fieldType, MessageTransportType)
- binary.LittleEndian.PutUint32(fieldReceiver, elem.keypair.remoteIndex)
- binary.LittleEndian.PutUint64(fieldNonce, elem.nonce)
+ binary.LittleEndian.PutUint32(fieldType, MessageTransportType)
+ binary.LittleEndian.PutUint32(fieldReceiver, elem.keypair.remoteIndex)
+ binary.LittleEndian.PutUint64(fieldNonce, elem.nonce)
- // pad content to multiple of 16
- paddingSize := calculatePaddingSize(len(elem.packet), int(device.tun.mtu.Load()))
- elem.packet = append(elem.packet, paddingZeros[:paddingSize]...)
+ // pad content to multiple of 16
+ paddingSize := calculatePaddingSize(len(elem.packet), int(device.tun.mtu.Load()))
+ elem.packet = append(elem.packet, paddingZeros[:paddingSize]...)
- // encrypt content and release to consumer
+ // encrypt content and release to consumer
- binary.LittleEndian.PutUint64(nonce[4:], elem.nonce)
- elem.packet = elem.keypair.send.Seal(
- header,
- nonce[:],
- elem.packet,
- nil,
- )
- elem.Unlock()
+ binary.LittleEndian.PutUint64(nonce[4:], elem.nonce)
+ elem.packet = elem.keypair.send.Seal(
+ header,
+ nonce[:],
+ elem.packet,
+ nil,
+ )
+ elem.Unlock()
+ }
}
}
From 895d6c23cd60bd0522c5b6598a69ad6c5f1ab3a7 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Mon, 2 Oct 2023 14:43:56 -0700
Subject: [PATCH 05/75] tun: unwind summing loop in checksumNoFold()
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
$ benchstat old.txt new.txt
goos: linux
goarch: amd64
pkg: golang.zx2c4.com/wireguard/tun
cpu: 12th Gen Intel(R) Core(TM) i5-12400
│ old.txt │ new.txt │
│ sec/op │ sec/op vs base │
Checksum/64-12 10.670n ± 2% 4.769n ± 0% -55.30% (p=0.000 n=10)
Checksum/128-12 19.665n ± 2% 8.032n ± 0% -59.16% (p=0.000 n=10)
Checksum/256-12 37.68n ± 1% 16.06n ± 0% -57.37% (p=0.000 n=10)
Checksum/512-12 76.61n ± 3% 32.13n ± 0% -58.06% (p=0.000 n=10)
Checksum/1024-12 160.55n ± 4% 64.25n ± 0% -59.98% (p=0.000 n=10)
Checksum/1500-12 231.05n ± 7% 94.12n ± 0% -59.26% (p=0.000 n=10)
Checksum/2048-12 309.5n ± 3% 128.5n ± 0% -58.48% (p=0.000 n=10)
Checksum/4096-12 603.8n ± 4% 257.2n ± 0% -57.41% (p=0.000 n=10)
Checksum/8192-12 1185.0n ± 3% 515.5n ± 0% -56.50% (p=0.000 n=10)
Checksum/9000-12 1328.5n ± 5% 564.8n ± 0% -57.49% (p=0.000 n=10)
Checksum/9001-12 1340.5n ± 3% 564.8n ± 0% -57.87% (p=0.000 n=10)
geomean 185.3n 77.99n -57.92%
Reviewed-by: Adrian Dewhurst
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
tun/checksum.go | 100 +++++++++++++++++++++++++++++++++++++------
tun/checksum_test.go | 35 +++++++++++++++
2 files changed, 123 insertions(+), 12 deletions(-)
create mode 100644 tun/checksum_test.go
diff --git a/tun/checksum.go b/tun/checksum.go
index f4f8471..29a8fc8 100644
--- a/tun/checksum.go
+++ b/tun/checksum.go
@@ -3,23 +3,99 @@ package tun
import "encoding/binary"
// TODO: Explore SIMD and/or other assembly optimizations.
+// TODO: Test native endian loads. See RFC 1071 section 2 part B.
func checksumNoFold(b []byte, initial uint64) uint64 {
ac := initial
- i := 0
- n := len(b)
- for n >= 4 {
- ac += uint64(binary.BigEndian.Uint32(b[i : i+4]))
- n -= 4
- i += 4
+
+ for len(b) >= 128 {
+ ac += uint64(binary.BigEndian.Uint32(b[:4]))
+ ac += uint64(binary.BigEndian.Uint32(b[4:8]))
+ ac += uint64(binary.BigEndian.Uint32(b[8:12]))
+ ac += uint64(binary.BigEndian.Uint32(b[12:16]))
+ ac += uint64(binary.BigEndian.Uint32(b[16:20]))
+ ac += uint64(binary.BigEndian.Uint32(b[20:24]))
+ ac += uint64(binary.BigEndian.Uint32(b[24:28]))
+ ac += uint64(binary.BigEndian.Uint32(b[28:32]))
+ ac += uint64(binary.BigEndian.Uint32(b[32:36]))
+ ac += uint64(binary.BigEndian.Uint32(b[36:40]))
+ ac += uint64(binary.BigEndian.Uint32(b[40:44]))
+ ac += uint64(binary.BigEndian.Uint32(b[44:48]))
+ ac += uint64(binary.BigEndian.Uint32(b[48:52]))
+ ac += uint64(binary.BigEndian.Uint32(b[52:56]))
+ ac += uint64(binary.BigEndian.Uint32(b[56:60]))
+ ac += uint64(binary.BigEndian.Uint32(b[60:64]))
+ ac += uint64(binary.BigEndian.Uint32(b[64:68]))
+ ac += uint64(binary.BigEndian.Uint32(b[68:72]))
+ ac += uint64(binary.BigEndian.Uint32(b[72:76]))
+ ac += uint64(binary.BigEndian.Uint32(b[76:80]))
+ ac += uint64(binary.BigEndian.Uint32(b[80:84]))
+ ac += uint64(binary.BigEndian.Uint32(b[84:88]))
+ ac += uint64(binary.BigEndian.Uint32(b[88:92]))
+ ac += uint64(binary.BigEndian.Uint32(b[92:96]))
+ ac += uint64(binary.BigEndian.Uint32(b[96:100]))
+ ac += uint64(binary.BigEndian.Uint32(b[100:104]))
+ ac += uint64(binary.BigEndian.Uint32(b[104:108]))
+ ac += uint64(binary.BigEndian.Uint32(b[108:112]))
+ ac += uint64(binary.BigEndian.Uint32(b[112:116]))
+ ac += uint64(binary.BigEndian.Uint32(b[116:120]))
+ ac += uint64(binary.BigEndian.Uint32(b[120:124]))
+ ac += uint64(binary.BigEndian.Uint32(b[124:128]))
+ b = b[128:]
}
- for n >= 2 {
- ac += uint64(binary.BigEndian.Uint16(b[i : i+2]))
- n -= 2
- i += 2
+ if len(b) >= 64 {
+ ac += uint64(binary.BigEndian.Uint32(b[:4]))
+ ac += uint64(binary.BigEndian.Uint32(b[4:8]))
+ ac += uint64(binary.BigEndian.Uint32(b[8:12]))
+ ac += uint64(binary.BigEndian.Uint32(b[12:16]))
+ ac += uint64(binary.BigEndian.Uint32(b[16:20]))
+ ac += uint64(binary.BigEndian.Uint32(b[20:24]))
+ ac += uint64(binary.BigEndian.Uint32(b[24:28]))
+ ac += uint64(binary.BigEndian.Uint32(b[28:32]))
+ ac += uint64(binary.BigEndian.Uint32(b[32:36]))
+ ac += uint64(binary.BigEndian.Uint32(b[36:40]))
+ ac += uint64(binary.BigEndian.Uint32(b[40:44]))
+ ac += uint64(binary.BigEndian.Uint32(b[44:48]))
+ ac += uint64(binary.BigEndian.Uint32(b[48:52]))
+ ac += uint64(binary.BigEndian.Uint32(b[52:56]))
+ ac += uint64(binary.BigEndian.Uint32(b[56:60]))
+ ac += uint64(binary.BigEndian.Uint32(b[60:64]))
+ b = b[64:]
}
- if n == 1 {
- ac += uint64(b[i]) << 8
+ if len(b) >= 32 {
+ ac += uint64(binary.BigEndian.Uint32(b[:4]))
+ ac += uint64(binary.BigEndian.Uint32(b[4:8]))
+ ac += uint64(binary.BigEndian.Uint32(b[8:12]))
+ ac += uint64(binary.BigEndian.Uint32(b[12:16]))
+ ac += uint64(binary.BigEndian.Uint32(b[16:20]))
+ ac += uint64(binary.BigEndian.Uint32(b[20:24]))
+ ac += uint64(binary.BigEndian.Uint32(b[24:28]))
+ ac += uint64(binary.BigEndian.Uint32(b[28:32]))
+ b = b[32:]
}
+ if len(b) >= 16 {
+ ac += uint64(binary.BigEndian.Uint32(b[:4]))
+ ac += uint64(binary.BigEndian.Uint32(b[4:8]))
+ ac += uint64(binary.BigEndian.Uint32(b[8:12]))
+ ac += uint64(binary.BigEndian.Uint32(b[12:16]))
+ b = b[16:]
+ }
+ if len(b) >= 8 {
+ ac += uint64(binary.BigEndian.Uint32(b[:4]))
+ ac += uint64(binary.BigEndian.Uint32(b[4:8]))
+ b = b[8:]
+ }
+ if len(b) >= 4 {
+ ac += uint64(binary.BigEndian.Uint32(b))
+ b = b[4:]
+ }
+ if len(b) >= 2 {
+ ac += uint64(binary.BigEndian.Uint16(b))
+ b = b[2:]
+ }
+ if len(b) == 1 {
+ ac += uint64(b[0]) << 8
+ }
+
return ac
}
diff --git a/tun/checksum_test.go b/tun/checksum_test.go
new file mode 100644
index 0000000..c1ccff5
--- /dev/null
+++ b/tun/checksum_test.go
@@ -0,0 +1,35 @@
+package tun
+
+import (
+ "fmt"
+ "math/rand"
+ "testing"
+)
+
+func BenchmarkChecksum(b *testing.B) {
+ lengths := []int{
+ 64,
+ 128,
+ 256,
+ 512,
+ 1024,
+ 1500,
+ 2048,
+ 4096,
+ 8192,
+ 9000,
+ 9001,
+ }
+
+ for _, length := range lengths {
+ b.Run(fmt.Sprintf("%d", length), func(b *testing.B) {
+ buf := make([]byte, length)
+ rng := rand.New(rand.NewSource(1))
+ rng.Read(buf)
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ checksum(buf, 0)
+ }
+ })
+ }
+}
From 8a015f7c766564c21f6bef6fdddedce7e2ede830 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Mon, 2 Oct 2023 14:46:13 -0700
Subject: [PATCH 06/75] tun: reduce redundant checksumming in tcpGRO()
IPv4 header and pseudo header checksums were being computed on every
merge operation. Additionally, virtioNetHdr was being written at the
same time. This delays those operations until after all coalescing has
occurred.
Reviewed-by: Adrian Dewhurst
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
tun/tcp_offload_linux.go | 162 ++++++++++++++++++++++++---------------
1 file changed, 99 insertions(+), 63 deletions(-)
diff --git a/tun/tcp_offload_linux.go b/tun/tcp_offload_linux.go
index 39a7180..1afd27e 100644
--- a/tun/tcp_offload_linux.go
+++ b/tun/tcp_offload_linux.go
@@ -269,11 +269,11 @@ func tcpChecksumValid(pkt []byte, iphLen uint8, isV6 bool) bool {
type coalesceResult int
const (
- coalesceInsufficientCap coalesceResult = 0
- coalescePSHEnding coalesceResult = 1
- coalesceItemInvalidCSum coalesceResult = 2
- coalescePktInvalidCSum coalesceResult = 3
- coalesceSuccess coalesceResult = 4
+ coalesceInsufficientCap coalesceResult = iota
+ coalescePSHEnding
+ coalesceItemInvalidCSum
+ coalescePktInvalidCSum
+ coalesceSuccess
)
// coalesceTCPPackets attempts to coalesce pkt with the packet described by
@@ -339,42 +339,6 @@ func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize
if gsoSize > item.gsoSize {
item.gsoSize = gsoSize
}
- hdr := virtioNetHdr{
- flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, // this turns into CHECKSUM_PARTIAL in the skb
- hdrLen: uint16(headersLen),
- gsoSize: uint16(item.gsoSize),
- csumStart: uint16(item.iphLen),
- csumOffset: 16,
- }
-
- // Recalculate the total len (IPv4) or payload len (IPv6). Recalculate the
- // (IPv4) header checksum.
- if isV6 {
- hdr.gsoType = unix.VIRTIO_NET_HDR_GSO_TCPV6
- binary.BigEndian.PutUint16(pktHead[4:], uint16(coalescedLen)-uint16(item.iphLen)) // set new payload len
- } else {
- hdr.gsoType = unix.VIRTIO_NET_HDR_GSO_TCPV4
- pktHead[10], pktHead[11] = 0, 0 // clear checksum field
- binary.BigEndian.PutUint16(pktHead[2:], uint16(coalescedLen)) // set new total length
- iphCSum := ^checksum(pktHead[:item.iphLen], 0) // compute checksum
- binary.BigEndian.PutUint16(pktHead[10:], iphCSum) // set checksum field
- }
- hdr.encode(bufs[item.bufsIndex][bufsOffset-virtioNetHdrLen:])
-
- // Calculate the pseudo header checksum and place it at the TCP checksum
- // offset. Downstream checksum offloading will combine this with computation
- // of the tcp header and payload checksum.
- addrLen := 4
- addrOffset := ipv4SrcAddrOffset
- if isV6 {
- addrLen = 16
- addrOffset = ipv6SrcAddrOffset
- }
- srcAddrAt := bufsOffset + addrOffset
- srcAddr := bufs[item.bufsIndex][srcAddrAt : srcAddrAt+addrLen]
- dstAddr := bufs[item.bufsIndex][srcAddrAt+addrLen : srcAddrAt+addrLen*2]
- psum := pseudoHeaderChecksumNoFold(unix.IPPROTO_TCP, srcAddr, dstAddr, uint16(coalescedLen-int(item.iphLen)))
- binary.BigEndian.PutUint16(pktHead[hdr.csumStart+hdr.csumOffset:], checksum([]byte{}, psum))
item.numMerged++
return coalesceSuccess
@@ -390,43 +354,52 @@ const (
maxUint16 = 1<<16 - 1
)
+type tcpGROResult int
+
+const (
+ tcpGROResultNoop tcpGROResult = iota
+ tcpGROResultTableInsert
+ tcpGROResultCoalesced
+)
+
// tcpGRO evaluates the TCP packet at pktI in bufs for coalescing with
-// existing packets tracked in table. It will return false when pktI is not
-// coalesced, otherwise true. This indicates to the caller if bufs[pktI]
-// should be written to the Device.
-func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool) (pktCoalesced bool) {
+// existing packets tracked in table. It returns a tcpGROResultNoop when no
+// action was taken, tcpGROResultTableInsert when the evaluated packet was
+// inserted into table, and tcpGROResultCoalesced when the evaluated packet was
+// coalesced with another packet in table.
+func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool) tcpGROResult {
pkt := bufs[pktI][offset:]
if len(pkt) > maxUint16 {
// A valid IPv4 or IPv6 packet will never exceed this.
- return false
+ return tcpGROResultNoop
}
iphLen := int((pkt[0] & 0x0F) * 4)
if isV6 {
iphLen = 40
ipv6HPayloadLen := int(binary.BigEndian.Uint16(pkt[4:]))
if ipv6HPayloadLen != len(pkt)-iphLen {
- return false
+ return tcpGROResultNoop
}
} else {
totalLen := int(binary.BigEndian.Uint16(pkt[2:]))
if totalLen != len(pkt) {
- return false
+ return tcpGROResultNoop
}
}
if len(pkt) < iphLen {
- return false
+ return tcpGROResultNoop
}
tcphLen := int((pkt[iphLen+12] >> 4) * 4)
if tcphLen < 20 || tcphLen > 60 {
- return false
+ return tcpGROResultNoop
}
if len(pkt) < iphLen+tcphLen {
- return false
+ return tcpGROResultNoop
}
if !isV6 {
if pkt[6]&ipv4FlagMoreFragments != 0 || pkt[6]<<3 != 0 || pkt[7] != 0 {
// no GRO support for fragmented segments for now
- return false
+ return tcpGROResultNoop
}
}
tcpFlags := pkt[iphLen+tcpFlagsOffset]
@@ -434,14 +407,14 @@ func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool)
// not a candidate if any non-ACK flags (except PSH+ACK) are set
if tcpFlags != tcpFlagACK {
if pkt[iphLen+tcpFlagsOffset] != tcpFlagACK|tcpFlagPSH {
- return false
+ return tcpGROResultNoop
}
pshSet = true
}
gsoSize := uint16(len(pkt) - tcphLen - iphLen)
// not a candidate if payload len is 0
if gsoSize < 1 {
- return false
+ return tcpGROResultNoop
}
seq := binary.BigEndian.Uint32(pkt[iphLen+4:])
srcAddrOffset := ipv4SrcAddrOffset
@@ -452,7 +425,7 @@ func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool)
}
items, existing := table.lookupOrInsert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, tcphLen, pktI)
if !existing {
- return false
+ return tcpGROResultNoop
}
for i := len(items) - 1; i >= 0; i-- {
// In the best case of packets arriving in order iterating in reverse is
@@ -470,20 +443,20 @@ func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool)
switch result {
case coalesceSuccess:
table.updateAt(item, i)
- return true
+ return tcpGROResultCoalesced
case coalesceItemInvalidCSum:
// delete the item with an invalid csum
table.deleteAt(item.key, i)
case coalescePktInvalidCSum:
// no point in inserting an item that we can't coalesce
- return false
+ return tcpGROResultNoop
default:
}
}
}
// failed to coalesce with any other packets; store the item in the flow
table.insert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, tcphLen, pktI)
- return false
+ return tcpGROResultTableInsert
}
func isTCP4NoIPOptions(b []byte) bool {
@@ -515,6 +488,64 @@ func isTCP6NoEH(b []byte) bool {
return true
}
+// applyCoalesceAccounting updates bufs to account for coalescing based on the
+// metadata found in table.
+func applyCoalesceAccounting(bufs [][]byte, offset int, table *tcpGROTable, isV6 bool) error {
+ for _, items := range table.itemsByFlow {
+ for _, item := range items {
+ if item.numMerged > 0 {
+ hdr := virtioNetHdr{
+ flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, // this turns into CHECKSUM_PARTIAL in the skb
+ hdrLen: uint16(item.iphLen + item.tcphLen),
+ gsoSize: item.gsoSize,
+ csumStart: uint16(item.iphLen),
+ csumOffset: 16,
+ }
+ pkt := bufs[item.bufsIndex][offset:]
+
+ // Recalculate the total len (IPv4) or payload len (IPv6).
+ // Recalculate the (IPv4) header checksum.
+ if isV6 {
+ hdr.gsoType = unix.VIRTIO_NET_HDR_GSO_TCPV6
+ binary.BigEndian.PutUint16(pkt[4:], uint16(len(pkt))-uint16(item.iphLen)) // set new IPv6 header payload len
+ } else {
+ hdr.gsoType = unix.VIRTIO_NET_HDR_GSO_TCPV4
+ pkt[10], pkt[11] = 0, 0
+ binary.BigEndian.PutUint16(pkt[2:], uint16(len(pkt))) // set new total length
+ iphCSum := ^checksum(pkt[:item.iphLen], 0) // compute IPv4 header checksum
+ binary.BigEndian.PutUint16(pkt[10:], iphCSum) // set IPv4 header checksum field
+ }
+ err := hdr.encode(bufs[item.bufsIndex][offset-virtioNetHdrLen:])
+ if err != nil {
+ return err
+ }
+
+ // Calculate the pseudo header checksum and place it at the TCP
+ // checksum offset. Downstream checksum offloading will combine
+ // this with computation of the tcp header and payload checksum.
+ addrLen := 4
+ addrOffset := ipv4SrcAddrOffset
+ if isV6 {
+ addrLen = 16
+ addrOffset = ipv6SrcAddrOffset
+ }
+ srcAddrAt := offset + addrOffset
+ srcAddr := bufs[item.bufsIndex][srcAddrAt : srcAddrAt+addrLen]
+ dstAddr := bufs[item.bufsIndex][srcAddrAt+addrLen : srcAddrAt+addrLen*2]
+ psum := pseudoHeaderChecksumNoFold(unix.IPPROTO_TCP, srcAddr, dstAddr, uint16(len(pkt)-int(item.iphLen)))
+ binary.BigEndian.PutUint16(pkt[hdr.csumStart+hdr.csumOffset:], checksum([]byte{}, psum))
+ } else {
+ hdr := virtioNetHdr{}
+ err := hdr.encode(bufs[item.bufsIndex][offset-virtioNetHdrLen:])
+ if err != nil {
+ return err
+ }
+ }
+ }
+ }
+ return nil
+}
+
// handleGRO evaluates bufs for GRO, and writes the indices of the resulting
// packets into toWrite. toWrite, tcp4Table, and tcp6Table should initially be
// empty (but non-nil), and are passed in to save allocs as the caller may reset
@@ -524,23 +555,28 @@ func handleGRO(bufs [][]byte, offset int, tcp4Table, tcp6Table *tcpGROTable, toW
if offset < virtioNetHdrLen || offset > len(bufs[i])-1 {
return errors.New("invalid offset")
}
- var coalesced bool
+ var result tcpGROResult
switch {
case isTCP4NoIPOptions(bufs[i][offset:]): // ipv4 packets w/IP options do not coalesce
- coalesced = tcpGRO(bufs, offset, i, tcp4Table, false)
+ result = tcpGRO(bufs, offset, i, tcp4Table, false)
case isTCP6NoEH(bufs[i][offset:]): // ipv6 packets w/extension headers do not coalesce
- coalesced = tcpGRO(bufs, offset, i, tcp6Table, true)
+ result = tcpGRO(bufs, offset, i, tcp6Table, true)
}
- if !coalesced {
+ switch result {
+ case tcpGROResultNoop:
hdr := virtioNetHdr{}
err := hdr.encode(bufs[i][offset-virtioNetHdrLen:])
if err != nil {
return err
}
+ fallthrough
+ case tcpGROResultTableInsert:
*toWrite = append(*toWrite, i)
}
}
- return nil
+ err4 := applyCoalesceAccounting(bufs, offset, tcp4Table, false)
+ err6 := applyCoalesceAccounting(bufs, offset, tcp6Table, true)
+ return errors.Join(err4, err6)
}
// tcpTSO splits packets from in into outBuffs, writing the size of each
From 1ec454f253c068f74ba7a7aea34546c9819493c0 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Mon, 2 Oct 2023 14:48:28 -0700
Subject: [PATCH 07/75] device: move Queue{In,Out}boundElement Mutex to
container type
Queue{In,Out}boundElement locking can contribute to significant
overhead via sync.Mutex.lockSlow() in some environments. These types
are passed throughout the device package as elements in a slice, so
move the per-element Mutex to a container around the slice.
Reviewed-by: Maisem Ali
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
device/channels.go | 32 ++++++++--------
device/device.go | 10 ++---
device/peer.go | 8 ++--
device/pools.go | 44 +++++++++++----------
device/receive.go | 43 +++++++++++----------
device/send.go | 95 ++++++++++++++++++++++++----------------------
6 files changed, 121 insertions(+), 111 deletions(-)
diff --git a/device/channels.go b/device/channels.go
index 40ee5c9..e526f6b 100644
--- a/device/channels.go
+++ b/device/channels.go
@@ -19,13 +19,13 @@ import (
// call wg.Done to remove the initial reference.
// When the refcount hits 0, the queue's channel is closed.
type outboundQueue struct {
- c chan *[]*QueueOutboundElement
+ c chan *QueueOutboundElementsContainer
wg sync.WaitGroup
}
func newOutboundQueue() *outboundQueue {
q := &outboundQueue{
- c: make(chan *[]*QueueOutboundElement, QueueOutboundSize),
+ c: make(chan *QueueOutboundElementsContainer, QueueOutboundSize),
}
q.wg.Add(1)
go func() {
@@ -37,13 +37,13 @@ func newOutboundQueue() *outboundQueue {
// A inboundQueue is similar to an outboundQueue; see those docs.
type inboundQueue struct {
- c chan *[]*QueueInboundElement
+ c chan *QueueInboundElementsContainer
wg sync.WaitGroup
}
func newInboundQueue() *inboundQueue {
q := &inboundQueue{
- c: make(chan *[]*QueueInboundElement, QueueInboundSize),
+ c: make(chan *QueueInboundElementsContainer, QueueInboundSize),
}
q.wg.Add(1)
go func() {
@@ -72,7 +72,7 @@ func newHandshakeQueue() *handshakeQueue {
}
type autodrainingInboundQueue struct {
- c chan *[]*QueueInboundElement
+ c chan *QueueInboundElementsContainer
}
// newAutodrainingInboundQueue returns a channel that will be drained when it gets GC'd.
@@ -81,7 +81,7 @@ type autodrainingInboundQueue struct {
// some other means, such as sending a sentinel nil values.
func newAutodrainingInboundQueue(device *Device) *autodrainingInboundQueue {
q := &autodrainingInboundQueue{
- c: make(chan *[]*QueueInboundElement, QueueInboundSize),
+ c: make(chan *QueueInboundElementsContainer, QueueInboundSize),
}
runtime.SetFinalizer(q, device.flushInboundQueue)
return q
@@ -90,13 +90,13 @@ func newAutodrainingInboundQueue(device *Device) *autodrainingInboundQueue {
func (device *Device) flushInboundQueue(q *autodrainingInboundQueue) {
for {
select {
- case elems := <-q.c:
- for _, elem := range *elems {
- elem.Lock()
+ case elemsContainer := <-q.c:
+ elemsContainer.Lock()
+ for _, elem := range elemsContainer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutInboundElement(elem)
}
- device.PutInboundElementsSlice(elems)
+ device.PutInboundElementsContainer(elemsContainer)
default:
return
}
@@ -104,7 +104,7 @@ func (device *Device) flushInboundQueue(q *autodrainingInboundQueue) {
}
type autodrainingOutboundQueue struct {
- c chan *[]*QueueOutboundElement
+ c chan *QueueOutboundElementsContainer
}
// newAutodrainingOutboundQueue returns a channel that will be drained when it gets GC'd.
@@ -114,7 +114,7 @@ type autodrainingOutboundQueue struct {
// All sends to the channel must be best-effort, because there may be no receivers.
func newAutodrainingOutboundQueue(device *Device) *autodrainingOutboundQueue {
q := &autodrainingOutboundQueue{
- c: make(chan *[]*QueueOutboundElement, QueueOutboundSize),
+ c: make(chan *QueueOutboundElementsContainer, QueueOutboundSize),
}
runtime.SetFinalizer(q, device.flushOutboundQueue)
return q
@@ -123,13 +123,13 @@ func newAutodrainingOutboundQueue(device *Device) *autodrainingOutboundQueue {
func (device *Device) flushOutboundQueue(q *autodrainingOutboundQueue) {
for {
select {
- case elems := <-q.c:
- for _, elem := range *elems {
- elem.Lock()
+ case elemsContainer := <-q.c:
+ elemsContainer.Lock()
+ for _, elem := range elemsContainer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutOutboundElement(elem)
}
- device.PutOutboundElementsSlice(elems)
+ device.PutOutboundElementsContainer(elemsContainer)
default:
return
}
diff --git a/device/device.go b/device/device.go
index 1af9fe0..f9557a0 100644
--- a/device/device.go
+++ b/device/device.go
@@ -68,11 +68,11 @@ type Device struct {
cookieChecker CookieChecker
pool struct {
- outboundElementsSlice *WaitPool
- inboundElementsSlice *WaitPool
- messageBuffers *WaitPool
- inboundElements *WaitPool
- outboundElements *WaitPool
+ inboundElementsContainer *WaitPool
+ outboundElementsContainer *WaitPool
+ messageBuffers *WaitPool
+ inboundElements *WaitPool
+ outboundElements *WaitPool
}
queue struct {
diff --git a/device/peer.go b/device/peer.go
index 0ac4896..2fb5da6 100644
--- a/device/peer.go
+++ b/device/peer.go
@@ -45,9 +45,9 @@ type Peer struct {
}
queue struct {
- staged chan *[]*QueueOutboundElement // staged packets before a handshake is available
- outbound *autodrainingOutboundQueue // sequential ordering of udp transmission
- inbound *autodrainingInboundQueue // sequential ordering of tun writing
+ staged chan *QueueOutboundElementsContainer // staged packets before a handshake is available
+ outbound *autodrainingOutboundQueue // sequential ordering of udp transmission
+ inbound *autodrainingInboundQueue // sequential ordering of tun writing
}
cookieGenerator CookieGenerator
@@ -81,7 +81,7 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
peer.device = device
peer.queue.outbound = newAutodrainingOutboundQueue(device)
peer.queue.inbound = newAutodrainingInboundQueue(device)
- peer.queue.staged = make(chan *[]*QueueOutboundElement, QueueStagedSize)
+ peer.queue.staged = make(chan *QueueOutboundElementsContainer, QueueStagedSize)
// map public key
_, ok := device.peers.keyMap[pk]
diff --git a/device/pools.go b/device/pools.go
index 02a5d6a..94f3dc7 100644
--- a/device/pools.go
+++ b/device/pools.go
@@ -46,13 +46,13 @@ func (p *WaitPool) Put(x any) {
}
func (device *Device) PopulatePools() {
- device.pool.outboundElementsSlice = NewWaitPool(PreallocatedBuffersPerPool, func() any {
- s := make([]*QueueOutboundElement, 0, device.BatchSize())
- return &s
- })
- device.pool.inboundElementsSlice = NewWaitPool(PreallocatedBuffersPerPool, func() any {
+ device.pool.inboundElementsContainer = NewWaitPool(PreallocatedBuffersPerPool, func() any {
s := make([]*QueueInboundElement, 0, device.BatchSize())
- return &s
+ return &QueueInboundElementsContainer{elems: s}
+ })
+ device.pool.outboundElementsContainer = NewWaitPool(PreallocatedBuffersPerPool, func() any {
+ s := make([]*QueueOutboundElement, 0, device.BatchSize())
+ return &QueueOutboundElementsContainer{elems: s}
})
device.pool.messageBuffers = NewWaitPool(PreallocatedBuffersPerPool, func() any {
return new([MaxMessageSize]byte)
@@ -65,28 +65,32 @@ func (device *Device) PopulatePools() {
})
}
-func (device *Device) GetOutboundElementsSlice() *[]*QueueOutboundElement {
- return device.pool.outboundElementsSlice.Get().(*[]*QueueOutboundElement)
+func (device *Device) GetInboundElementsContainer() *QueueInboundElementsContainer {
+ c := device.pool.inboundElementsContainer.Get().(*QueueInboundElementsContainer)
+ c.Mutex = sync.Mutex{}
+ return c
}
-func (device *Device) PutOutboundElementsSlice(s *[]*QueueOutboundElement) {
- for i := range *s {
- (*s)[i] = nil
+func (device *Device) PutInboundElementsContainer(c *QueueInboundElementsContainer) {
+ for i := range c.elems {
+ c.elems[i] = nil
}
- *s = (*s)[:0]
- device.pool.outboundElementsSlice.Put(s)
+ c.elems = c.elems[:0]
+ device.pool.inboundElementsContainer.Put(c)
}
-func (device *Device) GetInboundElementsSlice() *[]*QueueInboundElement {
- return device.pool.inboundElementsSlice.Get().(*[]*QueueInboundElement)
+func (device *Device) GetOutboundElementsContainer() *QueueOutboundElementsContainer {
+ c := device.pool.outboundElementsContainer.Get().(*QueueOutboundElementsContainer)
+ c.Mutex = sync.Mutex{}
+ return c
}
-func (device *Device) PutInboundElementsSlice(s *[]*QueueInboundElement) {
- for i := range *s {
- (*s)[i] = nil
+func (device *Device) PutOutboundElementsContainer(c *QueueOutboundElementsContainer) {
+ for i := range c.elems {
+ c.elems[i] = nil
}
- *s = (*s)[:0]
- device.pool.inboundElementsSlice.Put(s)
+ c.elems = c.elems[:0]
+ device.pool.outboundElementsContainer.Put(c)
}
func (device *Device) GetMessageBuffer() *[MaxMessageSize]byte {
diff --git a/device/receive.go b/device/receive.go
index f0f37a1..4b32dc5 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -27,7 +27,6 @@ type QueueHandshakeElement struct {
}
type QueueInboundElement struct {
- sync.Mutex
buffer *[MaxMessageSize]byte
packet []byte
counter uint64
@@ -35,6 +34,11 @@ type QueueInboundElement struct {
endpoint conn.Endpoint
}
+type QueueInboundElementsContainer struct {
+ sync.Mutex
+ elems []*QueueInboundElement
+}
+
// clearPointers clears elem fields that contain pointers.
// This makes the garbage collector's life easier and
// avoids accidentally keeping other objects around unnecessarily.
@@ -87,7 +91,7 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive
count int
endpoints = make([]conn.Endpoint, maxBatchSize)
deathSpiral int
- elemsByPeer = make(map[*Peer]*[]*QueueInboundElement, maxBatchSize)
+ elemsByPeer = make(map[*Peer]*QueueInboundElementsContainer, maxBatchSize)
)
for i := range bufsArrs {
@@ -170,15 +174,14 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive
elem.keypair = keypair
elem.endpoint = endpoints[i]
elem.counter = 0
- elem.Mutex = sync.Mutex{}
- elem.Lock()
elemsForPeer, ok := elemsByPeer[peer]
if !ok {
- elemsForPeer = device.GetInboundElementsSlice()
+ elemsForPeer = device.GetInboundElementsContainer()
+ elemsForPeer.Lock()
elemsByPeer[peer] = elemsForPeer
}
- *elemsForPeer = append(*elemsForPeer, elem)
+ elemsForPeer.elems = append(elemsForPeer.elems, elem)
bufsArrs[i] = device.GetMessageBuffer()
bufs[i] = bufsArrs[i][:]
continue
@@ -217,16 +220,16 @@ func (device *Device) RoutineReceiveIncoming(maxBatchSize int, recv conn.Receive
default:
}
}
- for peer, elems := range elemsByPeer {
+ for peer, elemsContainer := range elemsByPeer {
if peer.isRunning.Load() {
- peer.queue.inbound.c <- elems
- device.queue.decryption.c <- elems
+ peer.queue.inbound.c <- elemsContainer
+ device.queue.decryption.c <- elemsContainer
} else {
- for _, elem := range *elems {
+ for _, elem := range elemsContainer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutInboundElement(elem)
}
- device.PutInboundElementsSlice(elems)
+ device.PutInboundElementsContainer(elemsContainer)
}
delete(elemsByPeer, peer)
}
@@ -239,8 +242,8 @@ func (device *Device) RoutineDecryption(id int) {
defer device.log.Verbosef("Routine: decryption worker %d - stopped", id)
device.log.Verbosef("Routine: decryption worker %d - started", id)
- for elems := range device.queue.decryption.c {
- for _, elem := range *elems {
+ for elemsContainer := range device.queue.decryption.c {
+ for _, elem := range elemsContainer.elems {
// split message into fields
counter := elem.packet[MessageTransportOffsetCounter:MessageTransportOffsetContent]
content := elem.packet[MessageTransportOffsetContent:]
@@ -259,8 +262,8 @@ func (device *Device) RoutineDecryption(id int) {
if err != nil {
elem.packet = nil
}
- elem.Unlock()
}
+ elemsContainer.Unlock()
}
}
@@ -437,12 +440,12 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
bufs := make([][]byte, 0, maxBatchSize)
- for elems := range peer.queue.inbound.c {
- if elems == nil {
+ for elemsContainer := range peer.queue.inbound.c {
+ if elemsContainer == nil {
return
}
- for _, elem := range *elems {
- elem.Lock()
+ elemsContainer.Lock()
+ for _, elem := range elemsContainer.elems {
if elem.packet == nil {
// decryption failed
continue
@@ -515,11 +518,11 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
device.log.Errorf("Failed to write packets to TUN device: %v", err)
}
}
- for _, elem := range *elems {
+ for _, elem := range elemsContainer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutInboundElement(elem)
}
bufs = bufs[:0]
- device.PutInboundElementsSlice(elems)
+ device.PutInboundElementsContainer(elemsContainer)
}
}
diff --git a/device/send.go b/device/send.go
index e838c4e..769720a 100644
--- a/device/send.go
+++ b/device/send.go
@@ -46,7 +46,6 @@ import (
*/
type QueueOutboundElement struct {
- sync.Mutex
buffer *[MaxMessageSize]byte // slice holding the packet data
packet []byte // slice of "buffer" (always!)
nonce uint64 // nonce for encryption
@@ -54,10 +53,14 @@ type QueueOutboundElement struct {
peer *Peer // related peer
}
+type QueueOutboundElementsContainer struct {
+ sync.Mutex
+ elems []*QueueOutboundElement
+}
+
func (device *Device) NewOutboundElement() *QueueOutboundElement {
elem := device.GetOutboundElement()
elem.buffer = device.GetMessageBuffer()
- elem.Mutex = sync.Mutex{}
elem.nonce = 0
// keypair and peer were cleared (if necessary) by clearPointers.
return elem
@@ -79,15 +82,15 @@ func (elem *QueueOutboundElement) clearPointers() {
func (peer *Peer) SendKeepalive() {
if len(peer.queue.staged) == 0 && peer.isRunning.Load() {
elem := peer.device.NewOutboundElement()
- elems := peer.device.GetOutboundElementsSlice()
- *elems = append(*elems, elem)
+ elemsContainer := peer.device.GetOutboundElementsContainer()
+ elemsContainer.elems = append(elemsContainer.elems, elem)
select {
- case peer.queue.staged <- elems:
+ case peer.queue.staged <- elemsContainer:
peer.device.log.Verbosef("%v - Sending keepalive packet", peer)
default:
peer.device.PutMessageBuffer(elem.buffer)
peer.device.PutOutboundElement(elem)
- peer.device.PutOutboundElementsSlice(elems)
+ peer.device.PutOutboundElementsContainer(elemsContainer)
}
}
peer.SendStagedPackets()
@@ -219,7 +222,7 @@ func (device *Device) RoutineReadFromTUN() {
readErr error
elems = make([]*QueueOutboundElement, batchSize)
bufs = make([][]byte, batchSize)
- elemsByPeer = make(map[*Peer]*[]*QueueOutboundElement, batchSize)
+ elemsByPeer = make(map[*Peer]*QueueOutboundElementsContainer, batchSize)
count = 0
sizes = make([]int, batchSize)
offset = MessageTransportHeaderSize
@@ -276,10 +279,10 @@ func (device *Device) RoutineReadFromTUN() {
}
elemsForPeer, ok := elemsByPeer[peer]
if !ok {
- elemsForPeer = device.GetOutboundElementsSlice()
+ elemsForPeer = device.GetOutboundElementsContainer()
elemsByPeer[peer] = elemsForPeer
}
- *elemsForPeer = append(*elemsForPeer, elem)
+ elemsForPeer.elems = append(elemsForPeer.elems, elem)
elems[i] = device.NewOutboundElement()
bufs[i] = elems[i].buffer[:]
}
@@ -289,11 +292,11 @@ func (device *Device) RoutineReadFromTUN() {
peer.StagePackets(elemsForPeer)
peer.SendStagedPackets()
} else {
- for _, elem := range *elemsForPeer {
+ for _, elem := range elemsForPeer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutOutboundElement(elem)
}
- device.PutOutboundElementsSlice(elemsForPeer)
+ device.PutOutboundElementsContainer(elemsForPeer)
}
delete(elemsByPeer, peer)
}
@@ -317,7 +320,7 @@ func (device *Device) RoutineReadFromTUN() {
}
}
-func (peer *Peer) StagePackets(elems *[]*QueueOutboundElement) {
+func (peer *Peer) StagePackets(elems *QueueOutboundElementsContainer) {
for {
select {
case peer.queue.staged <- elems:
@@ -326,11 +329,11 @@ func (peer *Peer) StagePackets(elems *[]*QueueOutboundElement) {
}
select {
case tooOld := <-peer.queue.staged:
- for _, elem := range *tooOld {
+ for _, elem := range tooOld.elems {
peer.device.PutMessageBuffer(elem.buffer)
peer.device.PutOutboundElement(elem)
}
- peer.device.PutOutboundElementsSlice(tooOld)
+ peer.device.PutOutboundElementsContainer(tooOld)
default:
}
}
@@ -349,52 +352,52 @@ top:
}
for {
- var elemsOOO *[]*QueueOutboundElement
+ var elemsContainerOOO *QueueOutboundElementsContainer
select {
- case elems := <-peer.queue.staged:
+ case elemsContainer := <-peer.queue.staged:
i := 0
- for _, elem := range *elems {
+ for _, elem := range elemsContainer.elems {
elem.peer = peer
elem.nonce = keypair.sendNonce.Add(1) - 1
if elem.nonce >= RejectAfterMessages {
keypair.sendNonce.Store(RejectAfterMessages)
- if elemsOOO == nil {
- elemsOOO = peer.device.GetOutboundElementsSlice()
+ if elemsContainerOOO == nil {
+ elemsContainerOOO = peer.device.GetOutboundElementsContainer()
}
- *elemsOOO = append(*elemsOOO, elem)
+ elemsContainerOOO.elems = append(elemsContainerOOO.elems, elem)
continue
} else {
- (*elems)[i] = elem
+ elemsContainer.elems[i] = elem
i++
}
elem.keypair = keypair
- elem.Lock()
}
- *elems = (*elems)[:i]
+ elemsContainer.Lock()
+ elemsContainer.elems = elemsContainer.elems[:i]
- if elemsOOO != nil {
- peer.StagePackets(elemsOOO) // XXX: Out of order, but we can't front-load go chans
+ if elemsContainerOOO != nil {
+ peer.StagePackets(elemsContainerOOO) // XXX: Out of order, but we can't front-load go chans
}
- if len(*elems) == 0 {
- peer.device.PutOutboundElementsSlice(elems)
+ if len(elemsContainer.elems) == 0 {
+ peer.device.PutOutboundElementsContainer(elemsContainer)
goto top
}
// add to parallel and sequential queue
if peer.isRunning.Load() {
- peer.queue.outbound.c <- elems
- peer.device.queue.encryption.c <- elems
+ peer.queue.outbound.c <- elemsContainer
+ peer.device.queue.encryption.c <- elemsContainer
} else {
- for _, elem := range *elems {
+ for _, elem := range elemsContainer.elems {
peer.device.PutMessageBuffer(elem.buffer)
peer.device.PutOutboundElement(elem)
}
- peer.device.PutOutboundElementsSlice(elems)
+ peer.device.PutOutboundElementsContainer(elemsContainer)
}
- if elemsOOO != nil {
+ if elemsContainerOOO != nil {
goto top
}
default:
@@ -406,12 +409,12 @@ top:
func (peer *Peer) FlushStagedPackets() {
for {
select {
- case elems := <-peer.queue.staged:
- for _, elem := range *elems {
+ case elemsContainer := <-peer.queue.staged:
+ for _, elem := range elemsContainer.elems {
peer.device.PutMessageBuffer(elem.buffer)
peer.device.PutOutboundElement(elem)
}
- peer.device.PutOutboundElementsSlice(elems)
+ peer.device.PutOutboundElementsContainer(elemsContainer)
default:
return
}
@@ -445,8 +448,8 @@ func (device *Device) RoutineEncryption(id int) {
defer device.log.Verbosef("Routine: encryption worker %d - stopped", id)
device.log.Verbosef("Routine: encryption worker %d - started", id)
- for elems := range device.queue.encryption.c {
- for _, elem := range *elems {
+ for elemsContainer := range device.queue.encryption.c {
+ for _, elem := range elemsContainer.elems {
// populate header fields
header := elem.buffer[:MessageTransportHeaderSize]
@@ -471,8 +474,8 @@ func (device *Device) RoutineEncryption(id int) {
elem.packet,
nil,
)
- elem.Unlock()
}
+ elemsContainer.Unlock()
}
}
@@ -486,9 +489,9 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
bufs := make([][]byte, 0, maxBatchSize)
- for elems := range peer.queue.outbound.c {
+ for elemsContainer := range peer.queue.outbound.c {
bufs = bufs[:0]
- if elems == nil {
+ if elemsContainer == nil {
return
}
if !peer.isRunning.Load() {
@@ -498,16 +501,16 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
// The timers and SendBuffers code are resilient to a few stragglers.
// TODO: rework peer shutdown order to ensure
// that we never accidentally keep timers alive longer than necessary.
- for _, elem := range *elems {
- elem.Lock()
+ elemsContainer.Lock()
+ for _, elem := range elemsContainer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutOutboundElement(elem)
}
continue
}
dataSent := false
- for _, elem := range *elems {
- elem.Lock()
+ elemsContainer.Lock()
+ for _, elem := range elemsContainer.elems {
if len(elem.packet) != MessageKeepaliveSize {
dataSent = true
}
@@ -521,11 +524,11 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
if dataSent {
peer.timersDataSent()
}
- for _, elem := range *elems {
+ for _, elem := range elemsContainer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutOutboundElement(elem)
}
- device.PutOutboundElementsSlice(elems)
+ device.PutOutboundElementsContainer(elemsContainer)
if err != nil {
var errGSO conn.ErrUDPGSODisabled
if errors.As(err, &errGSO) {
From ec8f6f82c20c617a3ea94478f2b5e4d49c6d3c2c Mon Sep 17 00:00:00 2001
From: James Tucker
Date: Wed, 27 Sep 2023 14:52:21 -0700
Subject: [PATCH 08/75] tun: fix crash when ForceMTU is called after close
Close closes the events channel, resulting in a panic from send on
closed channel.
Reported-By: Brad Fitzpatrick
Signed-off-by: James Tucker
Signed-off-by: Jason A. Donenfeld
---
tun/tun_windows.go | 3 +++
1 file changed, 3 insertions(+)
diff --git a/tun/tun_windows.go b/tun/tun_windows.go
index 0cb4ce1..34f2980 100644
--- a/tun/tun_windows.go
+++ b/tun/tun_windows.go
@@ -127,6 +127,9 @@ func (tun *NativeTun) MTU() (int, error) {
// TODO: This is a temporary hack. We really need to be monitoring the interface in real time and adapting to MTU changes.
func (tun *NativeTun) ForceMTU(mtu int) {
+ if tun.close.Load() {
+ return
+ }
update := tun.forcedMTU != mtu
tun.forcedMTU = mtu
if update {
From 42ec952eadc297efadc70b9911d5a59bcd2db4a6 Mon Sep 17 00:00:00 2001
From: James Tucker
Date: Wed, 27 Sep 2023 16:15:09 -0700
Subject: [PATCH 09/75] go.mod,tun/netstack: bump gvisor
Signed-off-by: James Tucker
Signed-off-by: Jason A. Donenfeld
---
go.mod | 8 ++++----
go.sum | 16 ++++++++--------
tun/netstack/tun.go | 14 +++++++-------
tun/tcp_offload_linux_test.go | 8 ++++----
4 files changed, 23 insertions(+), 23 deletions(-)
diff --git a/go.mod b/go.mod
index 758dcde..919dc49 100644
--- a/go.mod
+++ b/go.mod
@@ -3,14 +3,14 @@ module golang.zx2c4.com/wireguard
go 1.20
require (
- golang.org/x/crypto v0.6.0
- golang.org/x/net v0.7.0
+ golang.org/x/crypto v0.13.0
+ golang.org/x/net v0.15.0
golang.org/x/sys v0.12.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
- gvisor.dev/gvisor v0.0.0-20221203005347-703fd9b7fbc0
+ gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259
)
require (
github.com/google/btree v1.0.1 // indirect
- golang.org/x/time v0.0.0-20191024005414-555d28b269f0 // indirect
+ golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 // indirect
)
diff --git a/go.sum b/go.sum
index fe4ca7e..6bcecea 100644
--- a/go.sum
+++ b/go.sum
@@ -1,14 +1,14 @@
github.com/google/btree v1.0.1 h1:gK4Kx5IaGY9CD5sPJ36FHiBJ6ZXl0kilRiiCj+jdYp4=
github.com/google/btree v1.0.1/go.mod h1:xXMiIv4Fb/0kKde4SpL7qlzvu5cMJDRkFDxJfI9uaxA=
-golang.org/x/crypto v0.6.0 h1:qfktjS5LUO+fFKeJXZ+ikTRijMmljikvG68fpMMruSc=
-golang.org/x/crypto v0.6.0/go.mod h1:OFC/31mSvZgRz0V1QTNCzfAI1aIRzbiufJtkMIlEp58=
-golang.org/x/net v0.7.0 h1:rJrUqqhjsgNp7KqAIc25s9pZnjU7TUcSY7HcVZjdn1g=
-golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs=
+golang.org/x/crypto v0.13.0 h1:mvySKfSWJ+UKUii46M40LOvyWfN0s2U+46/jDd0e6Ck=
+golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc=
+golang.org/x/net v0.15.0 h1:ugBLEUaxABaB5AJqW9enI0ACdci2RUd4eP51NTBvuJ8=
+golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o=
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
-golang.org/x/time v0.0.0-20191024005414-555d28b269f0 h1:/5xXl8Y5W96D+TtHSlonuFqGHIWVuyCkGJLwGh9JJFs=
-golang.org/x/time v0.0.0-20191024005414-555d28b269f0/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
+golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 h1:vVKdlvoWBphwdxWKrFZEuM0kGgGLxUOYcY4U/2Vjg44=
+golang.org/x/time v0.0.0-20220210224613-90d013bbcef8/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
-gvisor.dev/gvisor v0.0.0-20221203005347-703fd9b7fbc0 h1:Wobr37noukisGxpKo5jAsLREcpj61RxrWYzD8uwveOY=
-gvisor.dev/gvisor v0.0.0-20221203005347-703fd9b7fbc0/go.mod h1:Dn5idtptoW1dIos9U6A2rpebLs/MtTwFacjKb8jLdQA=
+gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259 h1:TbRPT0HtzFP3Cno1zZo7yPzEEnfu8EjLfl6IU9VfqkQ=
+gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259/go.mod h1:AVgIgHMwK63XvmAzWG9vLQ41YnVHN0du0tEC46fI7yY=
diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go
index 596cfcd..2b73054 100644
--- a/tun/netstack/tun.go
+++ b/tun/netstack/tun.go
@@ -25,7 +25,7 @@ import (
"golang.zx2c4.com/wireguard/tun"
"golang.org/x/net/dns/dnsmessage"
- "gvisor.dev/gvisor/pkg/bufferv2"
+ "gvisor.dev/gvisor/pkg/buffer"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
"gvisor.dev/gvisor/pkg/tcpip/header"
@@ -43,7 +43,7 @@ type netTun struct {
ep *channel.Endpoint
stack *stack.Stack
events chan tun.Event
- incomingPacket chan *bufferv2.View
+ incomingPacket chan *buffer.View
mtu int
dnsServers []netip.Addr
hasV4, hasV6 bool
@@ -61,7 +61,7 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device,
ep: channel.New(1024, uint32(mtu), ""),
stack: stack.New(opts),
events: make(chan tun.Event, 10),
- incomingPacket: make(chan *bufferv2.View),
+ incomingPacket: make(chan *buffer.View),
dnsServers: dnsServers,
mtu: mtu,
}
@@ -84,7 +84,7 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device,
}
protoAddr := tcpip.ProtocolAddress{
Protocol: protoNumber,
- AddressWithPrefix: tcpip.Address(ip.AsSlice()).WithPrefix(),
+ AddressWithPrefix: tcpip.AddrFromSlice(ip.AsSlice()).WithPrefix(),
}
tcpipErr := dev.stack.AddProtocolAddress(1, protoAddr, stack.AddressProperties{})
if tcpipErr != nil {
@@ -140,7 +140,7 @@ func (tun *netTun) Write(buf [][]byte, offset int) (int, error) {
continue
}
- pkb := stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: bufferv2.MakeWithData(packet)})
+ pkb := stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buffer.MakeWithData(packet)})
switch packet[0] >> 4 {
case 4:
tun.ep.InjectInbound(header.IPv4ProtocolNumber, pkb)
@@ -198,7 +198,7 @@ func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.Networ
}
return tcpip.FullAddress{
NIC: 1,
- Addr: tcpip.Address(endpoint.Addr().AsSlice()),
+ Addr: tcpip.AddrFromSlice(endpoint.Addr().AsSlice()),
Port: endpoint.Port(),
}, protoNumber
}
@@ -453,7 +453,7 @@ func (pc *PingConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
return 0, nil, fmt.Errorf("ping read: %s", tcpipErr)
}
- remoteAddr, _ := netip.AddrFromSlice([]byte(res.RemoteAddr.Addr))
+ remoteAddr, _ := netip.AddrFromSlice(res.RemoteAddr.Addr.AsSlice())
return res.Count, &PingAddr{remoteAddr}, nil
}
diff --git a/tun/tcp_offload_linux_test.go b/tun/tcp_offload_linux_test.go
index 9160e18..ddddc48 100644
--- a/tun/tcp_offload_linux_test.go
+++ b/tun/tcp_offload_linux_test.go
@@ -35,8 +35,8 @@ func tcp4PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, flags header.
srcAs4 := srcIPPort.Addr().As4()
dstAs4 := dstIPPort.Addr().As4()
ipFields := &header.IPv4Fields{
- SrcAddr: tcpip.Address(srcAs4[:]),
- DstAddr: tcpip.Address(dstAs4[:]),
+ SrcAddr: tcpip.AddrFromSlice(srcAs4[:]),
+ DstAddr: tcpip.AddrFromSlice(dstAs4[:]),
Protocol: unix.IPPROTO_TCP,
TTL: 64,
TotalLength: uint16(totalLen),
@@ -72,8 +72,8 @@ func tcp6PacketMutateIPFields(srcIPPort, dstIPPort netip.AddrPort, flags header.
srcAs16 := srcIPPort.Addr().As16()
dstAs16 := dstIPPort.Addr().As16()
ipFields := &header.IPv6Fields{
- SrcAddr: tcpip.Address(srcAs16[:]),
- DstAddr: tcpip.Address(dstAs16[:]),
+ SrcAddr: tcpip.AddrFromSlice(srcAs16[:]),
+ DstAddr: tcpip.AddrFromSlice(dstAs16[:]),
TransportProtocol: unix.IPPROTO_TCP,
HopLimit: 64,
PayloadLength: uint16(segmentSize + 20),
From b81ca925dbeb9ba775dcf2ad38b54f24e256a6e2 Mon Sep 17 00:00:00 2001
From: Mazay B
Date: Sat, 14 Oct 2023 11:42:30 +0100
Subject: [PATCH 10/75] peer.device.aSecMux.RLock added
---
device/send.go | 3 +++
1 file changed, 3 insertions(+)
diff --git a/device/send.go b/device/send.go
index b5c8e10..c60342e 100644
--- a/device/send.go
+++ b/device/send.go
@@ -139,16 +139,19 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
return err
}
+ peer.device.aSecMux.RLock()
if peer.device.aSecCfg.initPacketJunkSize != 0 {
buf := make([]byte, 0, peer.device.aSecCfg.initPacketJunkSize)
writer := bytes.NewBuffer(buf[:0])
err = appendJunk(writer, peer.device.aSecCfg.initPacketJunkSize)
if err != nil {
peer.device.log.Errorf("%v - %v", peer, err)
+ peer.device.aSecMux.RUnlock()
return err
}
junkedHeader = writer.Bytes()
}
+ peer.device.aSecMux.RUnlock()
}
var buf [MessageInitiationSize]byte
From 177caa7e4419d1b95bbf0423f6be6230c7101504 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Wed, 18 Oct 2023 21:02:52 +0200
Subject: [PATCH 11/75] conn: simplify supportsUDPOffload
This allows a kernel to support UDP_GRO while not supporting
UDP_SEGMENT.
Signed-off-by: Jason A. Donenfeld
---
conn/features_linux.go | 10 ++--------
1 file changed, 2 insertions(+), 8 deletions(-)
diff --git a/conn/features_linux.go b/conn/features_linux.go
index e1fb57f..8959d93 100644
--- a/conn/features_linux.go
+++ b/conn/features_linux.go
@@ -18,15 +18,9 @@ func supportsUDPOffload(conn *net.UDPConn) (txOffload, rxOffload bool) {
}
err = rc.Control(func(fd uintptr) {
_, errSyscall := unix.GetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_SEGMENT)
- if errSyscall != nil {
- return
- }
- txOffload = true
+ txOffload = errSyscall == nil
opt, errSyscall := unix.GetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO)
- if errSyscall != nil {
- return
- }
- rxOffload = opt == 1
+ rxOffload = errSyscall == nil && opt == 1
})
if err != nil {
return false, false
From 24ea13351eb7a06c3760f2eae18a484ce009fcf9 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Wed, 18 Oct 2023 21:14:13 +0200
Subject: [PATCH 12/75] conn: harmonize GOOS checks between "linux" and
"android"
Otherwise GRO gets enabled on Android, but the conn doesn't use it,
resulting in bundled packets being discarded.
Signed-off-by: Jason A. Donenfeld
---
conn/bind_std.go | 10 +++++-----
1 file changed, 5 insertions(+), 5 deletions(-)
diff --git a/conn/bind_std.go b/conn/bind_std.go
index 9886c91..5a00f34 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -175,7 +175,7 @@ again:
var fns []ReceiveFunc
if v4conn != nil {
s.ipv4TxOffload, s.ipv4RxOffload = supportsUDPOffload(v4conn)
- if runtime.GOOS == "linux" {
+ if runtime.GOOS == "linux" || runtime.GOOS == "android" {
v4pc = ipv4.NewPacketConn(v4conn)
s.ipv4PC = v4pc
}
@@ -184,7 +184,7 @@ again:
}
if v6conn != nil {
s.ipv6TxOffload, s.ipv6RxOffload = supportsUDPOffload(v6conn)
- if runtime.GOOS == "linux" {
+ if runtime.GOOS == "linux" || runtime.GOOS == "android" {
v6pc = ipv6.NewPacketConn(v6conn)
s.ipv6PC = v6pc
}
@@ -237,7 +237,7 @@ func (s *StdNetBind) receiveIP(
}
defer s.putMessages(msgs)
var numMsgs int
- if runtime.GOOS == "linux" {
+ if runtime.GOOS == "linux" || runtime.GOOS == "android" {
if rxOffload {
readAt := len(*msgs) - (IdealBatchSize / udpSegmentMaxDatagrams)
numMsgs, err = br.ReadBatch((*msgs)[readAt:], 0)
@@ -291,7 +291,7 @@ func (s *StdNetBind) makeReceiveIPv6(pc *ipv6.PacketConn, conn *net.UDPConn, rxO
// TODO: When all Binds handle IdealBatchSize, remove this dynamic function and
// rename the IdealBatchSize constant to BatchSize.
func (s *StdNetBind) BatchSize() int {
- if runtime.GOOS == "linux" {
+ if runtime.GOOS == "linux" || runtime.GOOS == "android" {
return IdealBatchSize
}
return 1
@@ -414,7 +414,7 @@ func (s *StdNetBind) send(conn *net.UDPConn, pc batchWriter, msgs []ipv6.Message
err error
start int
)
- if runtime.GOOS == "linux" {
+ if runtime.GOOS == "linux" || runtime.GOOS == "android" {
for {
n, err = pc.WriteBatch(msgs[start:], 0)
if err != nil || n == len(msgs[start:]) {
From 5d37bd24e14e3fff6c1ce61e299480beb3d68c00 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Sat, 21 Oct 2023 18:41:27 +0200
Subject: [PATCH 13/75] conn: separate gso and sticky control
Android wants GSO but not sticky.
Signed-off-by: Jason A. Donenfeld
---
conn/bind_std.go | 2 +-
conn/gso_default.go | 21 ++++++
conn/gso_linux.go | 65 +++++++++++++++++++
.../{control_default.go => sticky_default.go} | 13 +---
conn/{control_linux.go => sticky_linux.go} | 51 +--------------
...rol_linux_test.go => sticky_linux_test.go} | 10 +--
6 files changed, 96 insertions(+), 66 deletions(-)
create mode 100644 conn/gso_default.go
create mode 100644 conn/gso_linux.go
rename conn/{control_default.go => sticky_default.go} (72%)
rename conn/{control_linux.go => sticky_linux.go} (66%)
rename conn/{control_linux_test.go => sticky_linux_test.go} (96%)
diff --git a/conn/bind_std.go b/conn/bind_std.go
index 5a00f34..e1bcbd1 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -65,7 +65,7 @@ func NewStdNetBind() Bind {
msgs := make([]ipv6.Message, IdealBatchSize)
for i := range msgs {
msgs[i].Buffers = make(net.Buffers, 1)
- msgs[i].OOB = make([]byte, controlSize)
+ msgs[i].OOB = make([]byte, stickyControlSize+gsoControlSize)
}
return &msgs
},
diff --git a/conn/gso_default.go b/conn/gso_default.go
new file mode 100644
index 0000000..57780db
--- /dev/null
+++ b/conn/gso_default.go
@@ -0,0 +1,21 @@
+//go:build !linux
+
+/* SPDX-License-Identifier: MIT
+ *
+ * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ */
+
+package conn
+
+// getGSOSize parses control for UDP_GRO and if found returns its GSO size data.
+func getGSOSize(control []byte) (int, error) {
+ return 0, nil
+}
+
+// setGSOSize sets a UDP_SEGMENT in control based on gsoSize.
+func setGSOSize(control *[]byte, gsoSize uint16) {
+}
+
+// gsoControlSize returns the recommended buffer size for pooling sticky and UDP
+// offloading control data.
+const gsoControlSize = 0
diff --git a/conn/gso_linux.go b/conn/gso_linux.go
new file mode 100644
index 0000000..b8599ce
--- /dev/null
+++ b/conn/gso_linux.go
@@ -0,0 +1,65 @@
+//go:build linux
+
+/* SPDX-License-Identifier: MIT
+ *
+ * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ */
+
+package conn
+
+import (
+ "fmt"
+ "unsafe"
+
+ "golang.org/x/sys/unix"
+)
+
+const (
+ sizeOfGSOData = 2
+)
+
+// getGSOSize parses control for UDP_GRO and if found returns its GSO size data.
+func getGSOSize(control []byte) (int, error) {
+ var (
+ hdr unix.Cmsghdr
+ data []byte
+ rem = control
+ err error
+ )
+
+ for len(rem) > unix.SizeofCmsghdr {
+ hdr, data, rem, err = unix.ParseOneSocketControlMessage(rem)
+ if err != nil {
+ return 0, fmt.Errorf("error parsing socket control message: %w", err)
+ }
+ if hdr.Level == unix.SOL_UDP && hdr.Type == unix.UDP_GRO && len(data) >= sizeOfGSOData {
+ var gso uint16
+ copy(unsafe.Slice((*byte)(unsafe.Pointer(&gso)), sizeOfGSOData), data[:sizeOfGSOData])
+ return int(gso), nil
+ }
+ }
+ return 0, nil
+}
+
+// setGSOSize sets a UDP_SEGMENT in control based on gsoSize. It leaves existing
+// data in control untouched.
+func setGSOSize(control *[]byte, gsoSize uint16) {
+ existingLen := len(*control)
+ avail := cap(*control) - existingLen
+ space := unix.CmsgSpace(sizeOfGSOData)
+ if avail < space {
+ return
+ }
+ *control = (*control)[:cap(*control)]
+ gsoControl := (*control)[existingLen:]
+ hdr := (*unix.Cmsghdr)(unsafe.Pointer(&(gsoControl)[0]))
+ hdr.Level = unix.SOL_UDP
+ hdr.Type = unix.UDP_SEGMENT
+ hdr.SetLen(unix.CmsgLen(sizeOfGSOData))
+ copy((gsoControl)[unix.SizeofCmsghdr:], unsafe.Slice((*byte)(unsafe.Pointer(&gsoSize)), sizeOfGSOData))
+ *control = (*control)[:existingLen+space]
+}
+
+// gsoControlSize returns the recommended buffer size for pooling UDP
+// offloading control data.
+var gsoControlSize = unix.CmsgSpace(sizeOfGSOData)
diff --git a/conn/control_default.go b/conn/sticky_default.go
similarity index 72%
rename from conn/control_default.go
rename to conn/sticky_default.go
index 9459da5..0b21386 100644
--- a/conn/control_default.go
+++ b/conn/sticky_default.go
@@ -35,17 +35,8 @@ func getSrcFromControl(control []byte, ep *StdNetEndpoint) {
func setSrcControl(control *[]byte, ep *StdNetEndpoint) {
}
-// getGSOSize parses control for UDP_GRO and if found returns its GSO size data.
-func getGSOSize(control []byte) (int, error) {
- return 0, nil
-}
-
-// setGSOSize sets a UDP_SEGMENT in control based on gsoSize.
-func setGSOSize(control *[]byte, gsoSize uint16) {
-}
-
-// controlSize returns the recommended buffer size for pooling sticky and UDP
+// stickyControlSize returns the recommended buffer size for pooling sticky
// offloading control data.
-const controlSize = 0
+const stickyControlSize = 0
const StdNetSupportsStickySockets = false
diff --git a/conn/control_linux.go b/conn/sticky_linux.go
similarity index 66%
rename from conn/control_linux.go
rename to conn/sticky_linux.go
index 44a94e6..8e206e9 100644
--- a/conn/control_linux.go
+++ b/conn/sticky_linux.go
@@ -8,7 +8,6 @@
package conn
import (
- "fmt"
"net/netip"
"unsafe"
@@ -106,54 +105,8 @@ func setSrcControl(control *[]byte, ep *StdNetEndpoint) {
*control = append(*control, ep.src...)
}
-const (
- sizeOfGSOData = 2
-)
-
-// getGSOSize parses control for UDP_GRO and if found returns its GSO size data.
-func getGSOSize(control []byte) (int, error) {
- var (
- hdr unix.Cmsghdr
- data []byte
- rem = control
- err error
- )
-
- for len(rem) > unix.SizeofCmsghdr {
- hdr, data, rem, err = unix.ParseOneSocketControlMessage(rem)
- if err != nil {
- return 0, fmt.Errorf("error parsing socket control message: %w", err)
- }
- if hdr.Level == unix.SOL_UDP && hdr.Type == unix.UDP_GRO && len(data) >= sizeOfGSOData {
- var gso uint16
- copy(unsafe.Slice((*byte)(unsafe.Pointer(&gso)), sizeOfGSOData), data[:sizeOfGSOData])
- return int(gso), nil
- }
- }
- return 0, nil
-}
-
-// setGSOSize sets a UDP_SEGMENT in control based on gsoSize. It leaves existing
-// data in control untouched.
-func setGSOSize(control *[]byte, gsoSize uint16) {
- existingLen := len(*control)
- avail := cap(*control) - existingLen
- space := unix.CmsgSpace(sizeOfGSOData)
- if avail < space {
- return
- }
- *control = (*control)[:cap(*control)]
- gsoControl := (*control)[existingLen:]
- hdr := (*unix.Cmsghdr)(unsafe.Pointer(&(gsoControl)[0]))
- hdr.Level = unix.SOL_UDP
- hdr.Type = unix.UDP_SEGMENT
- hdr.SetLen(unix.CmsgLen(sizeOfGSOData))
- copy((gsoControl)[unix.SizeofCmsghdr:], unsafe.Slice((*byte)(unsafe.Pointer(&gsoSize)), sizeOfGSOData))
- *control = (*control)[:existingLen+space]
-}
-
-// controlSize returns the recommended buffer size for pooling sticky and UDP
+// stickyControlSize returns the recommended buffer size for pooling sticky
// offloading control data.
-var controlSize = unix.CmsgSpace(unix.SizeofInet6Pktinfo) + unix.CmsgSpace(sizeOfGSOData)
+var stickyControlSize = unix.CmsgSpace(unix.SizeofInet6Pktinfo)
const StdNetSupportsStickySockets = true
diff --git a/conn/control_linux_test.go b/conn/sticky_linux_test.go
similarity index 96%
rename from conn/control_linux_test.go
rename to conn/sticky_linux_test.go
index 96f9da2..d2bd584 100644
--- a/conn/control_linux_test.go
+++ b/conn/sticky_linux_test.go
@@ -60,7 +60,7 @@ func Test_setSrcControl(t *testing.T) {
}
setSrc(ep, netip.MustParseAddr("127.0.0.1"), 5)
- control := make([]byte, controlSize)
+ control := make([]byte, stickyControlSize)
setSrcControl(&control, ep)
@@ -89,7 +89,7 @@ func Test_setSrcControl(t *testing.T) {
}
setSrc(ep, netip.MustParseAddr("::1"), 5)
- control := make([]byte, controlSize)
+ control := make([]byte, stickyControlSize)
setSrcControl(&control, ep)
@@ -113,7 +113,7 @@ func Test_setSrcControl(t *testing.T) {
})
t.Run("ClearOnNoSrc", func(t *testing.T) {
- control := make([]byte, controlSize)
+ control := make([]byte, stickyControlSize)
hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0]))
hdr.Level = 1
hdr.Type = 2
@@ -129,7 +129,7 @@ func Test_setSrcControl(t *testing.T) {
func Test_getSrcFromControl(t *testing.T) {
t.Run("IPv4", func(t *testing.T) {
- control := make([]byte, controlSize)
+ control := make([]byte, stickyControlSize)
hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0]))
hdr.Level = unix.IPPROTO_IP
hdr.Type = unix.IP_PKTINFO
@@ -149,7 +149,7 @@ func Test_getSrcFromControl(t *testing.T) {
}
})
t.Run("IPv6", func(t *testing.T) {
- control := make([]byte, controlSize)
+ control := make([]byte, stickyControlSize)
hdr := (*unix.Cmsghdr)(unsafe.Pointer(&control[0]))
hdr.Level = unix.IPPROTO_IPV6
hdr.Type = unix.IPV6_PKTINFO
From f502ec3fad116d11109529bcf283e464f4822c18 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Sat, 21 Oct 2023 19:06:38 +0200
Subject: [PATCH 14/75] conn: fix cmsg data padding calculation for gso
Signed-off-by: Jason A. Donenfeld
---
conn/gso_linux.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/conn/gso_linux.go b/conn/gso_linux.go
index b8599ce..8596b29 100644
--- a/conn/gso_linux.go
+++ b/conn/gso_linux.go
@@ -56,7 +56,7 @@ func setGSOSize(control *[]byte, gsoSize uint16) {
hdr.Level = unix.SOL_UDP
hdr.Type = unix.UDP_SEGMENT
hdr.SetLen(unix.CmsgLen(sizeOfGSOData))
- copy((gsoControl)[unix.SizeofCmsghdr:], unsafe.Slice((*byte)(unsafe.Pointer(&gsoSize)), sizeOfGSOData))
+ copy((gsoControl)[unix.CmsgLen(0):], unsafe.Slice((*byte)(unsafe.Pointer(&gsoSize)), sizeOfGSOData))
*control = (*control)[:existingLen+space]
}
From b3df23dcd40ba4568572f338f9fd16b87053fc29 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Sat, 21 Oct 2023 19:32:07 +0200
Subject: [PATCH 15/75] conn: set unused OOB to zero length
Otherwise in the event that we're using GSO without sticky sockets, we
pass garbage OOB buffers to sendmmsg, making a EINVAL, when GSO doesn't
set its header.
Signed-off-by: Jason A. Donenfeld
---
conn/bind_std.go | 3 ++-
1 file changed, 2 insertions(+), 1 deletion(-)
diff --git a/conn/bind_std.go b/conn/bind_std.go
index e1bcbd1..46df7fd 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -65,7 +65,7 @@ func NewStdNetBind() Bind {
msgs := make([]ipv6.Message, IdealBatchSize)
for i := range msgs {
msgs[i].Buffers = make(net.Buffers, 1)
- msgs[i].OOB = make([]byte, stickyControlSize+gsoControlSize)
+ msgs[i].OOB = make([]byte, 0, stickyControlSize+gsoControlSize)
}
return &msgs
},
@@ -200,6 +200,7 @@ again:
func (s *StdNetBind) putMessages(msgs *[]ipv6.Message) {
for i := range *msgs {
+ (*msgs)[i].OOB = (*msgs)[i].OOB[:0]
(*msgs)[i] = ipv6.Message{Buffers: (*msgs)[i].Buffers, OOB: (*msgs)[i].OOB}
}
s.msgsPool.Put(msgs)
From 2e0774f246fb4fc1bd5cb44584d033038c89174e Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Sun, 22 Oct 2023 02:12:13 +0200
Subject: [PATCH 16/75] device: ratchet up max segment size on android
GRO requires big allocations to be efficient. This isn't great, as there
might be Android memory usage issues. So we should revisit this commit.
But at least it gets things working again.
Signed-off-by: Jason A. Donenfeld
---
device/queueconstants_android.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/device/queueconstants_android.go b/device/queueconstants_android.go
index 3d80ead..25f700a 100644
--- a/device/queueconstants_android.go
+++ b/device/queueconstants_android.go
@@ -14,6 +14,6 @@ const (
QueueOutboundSize = 1024
QueueInboundSize = 1024
QueueHandshakeSize = 1024
- MaxSegmentSize = 2200
+ MaxSegmentSize = (1 << 16) - 1 // largest possible UDP datagram
PreallocatedBuffersPerPool = 4096
)
From c493b95f66b9ddad97fa782d49a74eaa185ed4a3 Mon Sep 17 00:00:00 2001
From: pokamest
Date: Wed, 25 Oct 2023 22:41:33 +0100
Subject: [PATCH 17/75] Update README.md
Signed-off-by: pokamest
---
README.md | 57 ++++++++++++++++---------------------------------------
1 file changed, 16 insertions(+), 41 deletions(-)
diff --git a/README.md b/README.md
index 074f7ec..717c4c5 100644
--- a/README.md
+++ b/README.md
@@ -1,24 +1,27 @@
-# Go Implementation of [WireGuard](https://www.wireguard.com/)
+# Go Implementation of AmneziaWG
-This is an implementation of WireGuard in Go.
+AmneziaWG is a contemporary version of the WireGuard protocol. It's a fork of WireGuard-Go and offers protection against detection by Deep Packet Inspection (DPI) systems. At the same time, it retains the simplified architecture and high performance of the original.
+
+The precursor, WireGuard, is known for its efficiency but had issues with detection due to its distinctive packet signatures.
+AmneziaWG addresses this problem by employing advanced obfuscation methods, allowing its traffic to blend seamlessly with regular internet traffic.
+As a result, AmneziaWG maintains high performance while adding an extra layer of stealth, making it a superb choice for those seeking a fast and discreet VPN connection.
## Usage
-Most Linux kernel WireGuard users are used to adding an interface with `ip link add wg0 type wireguard`. With wireguard-go, instead simply run:
+Simply run:
```
-$ wireguard-go wg0
+$ amnezia-wg wg0
```
This will create an interface and fork into the background. To remove the interface, use the usual `ip link del wg0`, or if your system does not support removing interfaces directly, you may instead remove the control socket via `rm -f /var/run/wireguard/wg0.sock`, which will result in wireguard-go shutting down.
-To run wireguard-go without forking to the background, pass `-f` or `--foreground`:
+To run amnezia-wg without forking to the background, pass `-f` or `--foreground`:
```
-$ wireguard-go -f wg0
+$ amnezia-wg -f wg0
```
-
-When an interface is running, you may use [`wg(8)`](https://git.zx2c4.com/wireguard-tools/about/src/man/wg.8) to configure it, as well as the usual `ip(8)` and `ifconfig(8)` commands.
+When an interface is running, you may use [`amnezia-wg-tools `](https://github.com/amnezia-vpn/amnezia-wg-tools) to configure it, as well as the usual `ip(8)` and `ifconfig(8)` commands.
To run with more logging you may set the environment variable `LOG_LEVEL=debug`.
@@ -26,52 +29,24 @@ To run with more logging you may set the environment variable `LOG_LEVEL=debug`.
### Linux
-This will run on Linux; however you should instead use the kernel module, which is faster and better integrated into the OS. See the [installation page](https://www.wireguard.com/install/) for instructions.
+This will run on Linux; you should run amnezia-wg instead of using default linux kernel module.
### macOS
This runs on macOS using the utun driver. It does not yet support sticky sockets, and won't support fwmarks because of Darwin limitations. Since the utun driver cannot have arbitrary interface names, you must either use `utun[0-9]+` for an explicit interface name or `utun` to have the kernel select one for you. If you choose `utun` as the interface name, and the environment variable `WG_TUN_NAME_FILE` is defined, then the actual name of the interface chosen by the kernel is written to the file specified by that variable.
+This runs on MacOS, you should use it from [awg-apple](https://github.com/amnezia-vpn/awg-apple)
### Windows
-This runs on Windows, but you should instead use it from the more [fully featured Windows app](https://git.zx2c4.com/wireguard-windows/about/), which uses this as a module.
+This runs on Windows, you should use it from [awg-windows](https://github.com/amnezia-vpn/awg-windows), which uses this as a module.
-### FreeBSD
-
-This will run on FreeBSD. It does not yet support sticky sockets. Fwmark is mapped to `SO_USER_COOKIE`.
-
-### OpenBSD
-
-This will run on OpenBSD. It does not yet support sticky sockets. Fwmark is mapped to `SO_RTABLE`. Since the tun driver cannot have arbitrary interface names, you must either use `tun[0-9]+` for an explicit interface name or `tun` to have the program select one for you. If you choose `tun` as the interface name, and the environment variable `WG_TUN_NAME_FILE` is defined, then the actual name of the interface chosen by the kernel is written to the file specified by that variable.
## Building
This requires an installation of the latest version of [Go](https://go.dev/).
```
-$ git clone https://git.zx2c4.com/wireguard-go
-$ cd wireguard-go
+$ git clone https://github.com/amnezia-vpn/amnezia-wg
+$ cd amnezia-wg
$ make
```
-
-## License
-
- Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
-
- Permission is hereby granted, free of charge, to any person obtaining a copy of
- this software and associated documentation files (the "Software"), to deal in
- the Software without restriction, including without limitation the rights to
- use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies
- of the Software, and to permit persons to whom the Software is furnished to do
- so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
From 1cf89f5339b549236f38ce5fbc40f7bf993d9626 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Wed, 8 Nov 2023 14:06:20 -0800
Subject: [PATCH 18/75] tun: fix Device.Read() buf length assumption on Windows
The length of a packet read from the underlying TUN device may exceed
the length of a supplied buffer when MTU exceeds device.MaxMessageSize.
Reviewed-by: Brad Fitzpatrick
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
tun/tun_windows.go | 7 +++----
1 file changed, 3 insertions(+), 4 deletions(-)
diff --git a/tun/tun_windows.go b/tun/tun_windows.go
index 34f2980..2af8e3e 100644
--- a/tun/tun_windows.go
+++ b/tun/tun_windows.go
@@ -160,11 +160,10 @@ retry:
packet, err := tun.session.ReceivePacket()
switch err {
case nil:
- packetSize := len(packet)
- copy(bufs[0][offset:], packet)
- sizes[0] = packetSize
+ n := copy(bufs[0][offset:], packet)
+ sizes[0] = n
tun.session.ReleaseReceivePacket(packet)
- tun.rate.update(uint64(packetSize))
+ tun.rate.update(uint64(n))
return 1, nil
case windows.ERROR_NO_MORE_ITEMS:
if !shouldSpin || uint64(nanotime()-start) >= spinloopDuration {
From d0bc03c707974a84a672716c718f99fab49e7eb8 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Tue, 31 Oct 2023 19:53:35 -0700
Subject: [PATCH 19/75] tun: implement UDP GSO/GRO for Linux
Implement UDP GSO and GRO for the Linux tun.Device, which is made
possible by virtio extensions in the kernel's TUN driver starting in
v6.2.
secnetperf, a QUIC benchmark utility from microsoft/msquic@8e1eb1a, is
used to demonstrate the effect of this commit between two Linux
computers with i5-12400 CPUs. There is roughly ~13us of round trip
latency between them. secnetperf was invoked with the following command
line options:
-stats:1 -exec:maxtput -test:tput -download:10000 -timed:1 -encrypt:0
The first result is from commit 2e0774f without UDP GSO/GRO on the TUN.
[conn][0x55739a144980] STATS: EcnCapable=0 RTT=3973 us
SendTotalPackets=55859 SendSuspectedLostPackets=61
SendSpuriousLostPackets=59 SendCongestionCount=27
SendEcnCongestionCount=0 RecvTotalPackets=2779122
RecvReorderedPackets=0 RecvDroppedPackets=0
RecvDuplicatePackets=0 RecvDecryptionFailures=0
Result: 3654977571 bytes @ 2922821 kbps (10003.972 ms).
The second result is with UDP GSO/GRO on the TUN.
[conn][0x56493dfd09a0] STATS: EcnCapable=0 RTT=1216 us
SendTotalPackets=165033 SendSuspectedLostPackets=64
SendSpuriousLostPackets=61 SendCongestionCount=53
SendEcnCongestionCount=0 RecvTotalPackets=11845268
RecvReorderedPackets=25267 RecvDroppedPackets=0
RecvDuplicatePackets=0 RecvDecryptionFailures=0
Result: 15574671184 bytes @ 12458214 kbps (10001.222 ms).
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
...{tcp_offload_linux.go => offload_linux.go} | 598 ++++++++++----
tun/offload_linux_test.go | 752 ++++++++++++++++++
tun/tcp_offload_linux_test.go | 411 ----------
...65e4830d6dc087cab24cd1e154c2e790589a309b77 | 8 -
...6784411a8ce2e8e03aa3384105e581f2c67494700d | 8 -
tun/tun_linux.go | 71 +-
6 files changed, 1258 insertions(+), 590 deletions(-)
rename tun/{tcp_offload_linux.go => offload_linux.go} (50%)
create mode 100644 tun/offload_linux_test.go
delete mode 100644 tun/tcp_offload_linux_test.go
delete mode 100644 tun/testdata/fuzz/Fuzz_handleGRO/032aec0105f26f709c118365e4830d6dc087cab24cd1e154c2e790589a309b77
delete mode 100644 tun/testdata/fuzz/Fuzz_handleGRO/0da283f9a2098dec30d1c86784411a8ce2e8e03aa3384105e581f2c67494700d
diff --git a/tun/tcp_offload_linux.go b/tun/offload_linux.go
similarity index 50%
rename from tun/tcp_offload_linux.go
rename to tun/offload_linux.go
index 1afd27e..9ff7fea 100644
--- a/tun/tcp_offload_linux.go
+++ b/tun/offload_linux.go
@@ -57,22 +57,23 @@ const (
virtioNetHdrLen = int(unsafe.Sizeof(virtioNetHdr{}))
)
-// flowKey represents the key for a flow.
-type flowKey struct {
+// tcpFlowKey represents the key for a TCP flow.
+type tcpFlowKey struct {
srcAddr, dstAddr [16]byte
srcPort, dstPort uint16
rxAck uint32 // varying ack values should not be coalesced. Treat them as separate flows.
+ isV6 bool
}
-// tcpGROTable holds flow and coalescing information for the purposes of GRO.
+// tcpGROTable holds flow and coalescing information for the purposes of TCP GRO.
type tcpGROTable struct {
- itemsByFlow map[flowKey][]tcpGROItem
+ itemsByFlow map[tcpFlowKey][]tcpGROItem
itemsPool [][]tcpGROItem
}
func newTCPGROTable() *tcpGROTable {
t := &tcpGROTable{
- itemsByFlow: make(map[flowKey][]tcpGROItem, conn.IdealBatchSize),
+ itemsByFlow: make(map[tcpFlowKey][]tcpGROItem, conn.IdealBatchSize),
itemsPool: make([][]tcpGROItem, conn.IdealBatchSize),
}
for i := range t.itemsPool {
@@ -81,14 +82,15 @@ func newTCPGROTable() *tcpGROTable {
return t
}
-func newFlowKey(pkt []byte, srcAddr, dstAddr, tcphOffset int) flowKey {
- key := flowKey{}
- addrSize := dstAddr - srcAddr
- copy(key.srcAddr[:], pkt[srcAddr:dstAddr])
- copy(key.dstAddr[:], pkt[dstAddr:dstAddr+addrSize])
+func newTCPFlowKey(pkt []byte, srcAddrOffset, dstAddrOffset, tcphOffset int) tcpFlowKey {
+ key := tcpFlowKey{}
+ addrSize := dstAddrOffset - srcAddrOffset
+ copy(key.srcAddr[:], pkt[srcAddrOffset:dstAddrOffset])
+ copy(key.dstAddr[:], pkt[dstAddrOffset:dstAddrOffset+addrSize])
key.srcPort = binary.BigEndian.Uint16(pkt[tcphOffset:])
key.dstPort = binary.BigEndian.Uint16(pkt[tcphOffset+2:])
key.rxAck = binary.BigEndian.Uint32(pkt[tcphOffset+8:])
+ key.isV6 = addrSize == 16
return key
}
@@ -96,7 +98,7 @@ func newFlowKey(pkt []byte, srcAddr, dstAddr, tcphOffset int) flowKey {
// returning the packets found for the flow, or inserting a new one if none
// is found.
func (t *tcpGROTable) lookupOrInsert(pkt []byte, srcAddrOffset, dstAddrOffset, tcphOffset, tcphLen, bufsIndex int) ([]tcpGROItem, bool) {
- key := newFlowKey(pkt, srcAddrOffset, dstAddrOffset, tcphOffset)
+ key := newTCPFlowKey(pkt, srcAddrOffset, dstAddrOffset, tcphOffset)
items, ok := t.itemsByFlow[key]
if ok {
return items, ok
@@ -108,7 +110,7 @@ func (t *tcpGROTable) lookupOrInsert(pkt []byte, srcAddrOffset, dstAddrOffset, t
// insert an item in the table for the provided packet and packet metadata.
func (t *tcpGROTable) insert(pkt []byte, srcAddrOffset, dstAddrOffset, tcphOffset, tcphLen, bufsIndex int) {
- key := newFlowKey(pkt, srcAddrOffset, dstAddrOffset, tcphOffset)
+ key := newTCPFlowKey(pkt, srcAddrOffset, dstAddrOffset, tcphOffset)
item := tcpGROItem{
key: key,
bufsIndex: uint16(bufsIndex),
@@ -131,7 +133,7 @@ func (t *tcpGROTable) updateAt(item tcpGROItem, i int) {
items[i] = item
}
-func (t *tcpGROTable) deleteAt(key flowKey, i int) {
+func (t *tcpGROTable) deleteAt(key tcpFlowKey, i int) {
items, _ := t.itemsByFlow[key]
items = append(items[:i], items[i+1:]...)
t.itemsByFlow[key] = items
@@ -140,7 +142,7 @@ func (t *tcpGROTable) deleteAt(key flowKey, i int) {
// tcpGROItem represents bookkeeping data for a TCP packet during the lifetime
// of a GRO evaluation across a vector of packets.
type tcpGROItem struct {
- key flowKey
+ key tcpFlowKey
sentSeq uint32 // the sequence number
bufsIndex uint16 // the index into the original bufs slice
numMerged uint16 // the number of packets merged into this item
@@ -164,6 +166,103 @@ func (t *tcpGROTable) reset() {
}
}
+// udpFlowKey represents the key for a UDP flow.
+type udpFlowKey struct {
+ srcAddr, dstAddr [16]byte
+ srcPort, dstPort uint16
+ isV6 bool
+}
+
+// udpGROTable holds flow and coalescing information for the purposes of UDP GRO.
+type udpGROTable struct {
+ itemsByFlow map[udpFlowKey][]udpGROItem
+ itemsPool [][]udpGROItem
+}
+
+func newUDPGROTable() *udpGROTable {
+ u := &udpGROTable{
+ itemsByFlow: make(map[udpFlowKey][]udpGROItem, conn.IdealBatchSize),
+ itemsPool: make([][]udpGROItem, conn.IdealBatchSize),
+ }
+ for i := range u.itemsPool {
+ u.itemsPool[i] = make([]udpGROItem, 0, conn.IdealBatchSize)
+ }
+ return u
+}
+
+func newUDPFlowKey(pkt []byte, srcAddrOffset, dstAddrOffset, udphOffset int) udpFlowKey {
+ key := udpFlowKey{}
+ addrSize := dstAddrOffset - srcAddrOffset
+ copy(key.srcAddr[:], pkt[srcAddrOffset:dstAddrOffset])
+ copy(key.dstAddr[:], pkt[dstAddrOffset:dstAddrOffset+addrSize])
+ key.srcPort = binary.BigEndian.Uint16(pkt[udphOffset:])
+ key.dstPort = binary.BigEndian.Uint16(pkt[udphOffset+2:])
+ key.isV6 = addrSize == 16
+ return key
+}
+
+// lookupOrInsert looks up a flow for the provided packet and metadata,
+// returning the packets found for the flow, or inserting a new one if none
+// is found.
+func (u *udpGROTable) lookupOrInsert(pkt []byte, srcAddrOffset, dstAddrOffset, udphOffset, bufsIndex int) ([]udpGROItem, bool) {
+ key := newUDPFlowKey(pkt, srcAddrOffset, dstAddrOffset, udphOffset)
+ items, ok := u.itemsByFlow[key]
+ if ok {
+ return items, ok
+ }
+ // TODO: insert() performs another map lookup. This could be rearranged to avoid.
+ u.insert(pkt, srcAddrOffset, dstAddrOffset, udphOffset, bufsIndex, false)
+ return nil, false
+}
+
+// insert an item in the table for the provided packet and packet metadata.
+func (u *udpGROTable) insert(pkt []byte, srcAddrOffset, dstAddrOffset, udphOffset, bufsIndex int, cSumKnownInvalid bool) {
+ key := newUDPFlowKey(pkt, srcAddrOffset, dstAddrOffset, udphOffset)
+ item := udpGROItem{
+ key: key,
+ bufsIndex: uint16(bufsIndex),
+ gsoSize: uint16(len(pkt[udphOffset+udphLen:])),
+ iphLen: uint8(udphOffset),
+ cSumKnownInvalid: cSumKnownInvalid,
+ }
+ items, ok := u.itemsByFlow[key]
+ if !ok {
+ items = u.newItems()
+ }
+ items = append(items, item)
+ u.itemsByFlow[key] = items
+}
+
+func (u *udpGROTable) updateAt(item udpGROItem, i int) {
+ items, _ := u.itemsByFlow[item.key]
+ items[i] = item
+}
+
+// udpGROItem represents bookkeeping data for a UDP packet during the lifetime
+// of a GRO evaluation across a vector of packets.
+type udpGROItem struct {
+ key udpFlowKey
+ bufsIndex uint16 // the index into the original bufs slice
+ numMerged uint16 // the number of packets merged into this item
+ gsoSize uint16 // payload size
+ iphLen uint8 // ip header len
+ cSumKnownInvalid bool // UDP header checksum validity; a false value DOES NOT imply valid, just unknown.
+}
+
+func (u *udpGROTable) newItems() []udpGROItem {
+ var items []udpGROItem
+ items, u.itemsPool = u.itemsPool[len(u.itemsPool)-1], u.itemsPool[:len(u.itemsPool)-1]
+ return items
+}
+
+func (u *udpGROTable) reset() {
+ for k, items := range u.itemsByFlow {
+ items = items[:0]
+ u.itemsPool = append(u.itemsPool, items)
+ delete(u.itemsByFlow, k)
+ }
+}
+
// canCoalesce represents the outcome of checking if two TCP packets are
// candidates for coalescing.
type canCoalesce int
@@ -174,6 +273,61 @@ const (
coalesceAppend canCoalesce = 1
)
+// ipHeadersCanCoalesce returns true if the IP headers found in pktA and pktB
+// meet all requirements to be merged as part of a GRO operation, otherwise it
+// returns false.
+func ipHeadersCanCoalesce(pktA, pktB []byte) bool {
+ if len(pktA) < 9 || len(pktB) < 9 {
+ return false
+ }
+ if pktA[0]>>4 == 6 {
+ if pktA[0] != pktB[0] || pktA[1]>>4 != pktB[1]>>4 {
+ // cannot coalesce with unequal Traffic class values
+ return false
+ }
+ if pktA[7] != pktB[7] {
+ // cannot coalesce with unequal Hop limit values
+ return false
+ }
+ } else {
+ if pktA[1] != pktB[1] {
+ // cannot coalesce with unequal ToS values
+ return false
+ }
+ if pktA[6]>>5 != pktB[6]>>5 {
+ // cannot coalesce with unequal DF or reserved bits. MF is checked
+ // further up the stack.
+ return false
+ }
+ if pktA[8] != pktB[8] {
+ // cannot coalesce with unequal TTL values
+ return false
+ }
+ }
+ return true
+}
+
+// udpPacketsCanCoalesce evaluates if pkt can be coalesced with the packet
+// described by item. iphLen and gsoSize describe pkt. bufs is the vector of
+// packets involved in the current GRO evaluation. bufsOffset is the offset at
+// which packet data begins within bufs.
+func udpPacketsCanCoalesce(pkt []byte, iphLen uint8, gsoSize uint16, item udpGROItem, bufs [][]byte, bufsOffset int) canCoalesce {
+ pktTarget := bufs[item.bufsIndex][bufsOffset:]
+ if !ipHeadersCanCoalesce(pkt, pktTarget) {
+ return coalesceUnavailable
+ }
+ if len(pktTarget[iphLen+udphLen:])%int(item.gsoSize) != 0 {
+ // A smaller than gsoSize packet has been appended previously.
+ // Nothing can come after a smaller packet on the end.
+ return coalesceUnavailable
+ }
+ if gsoSize > item.gsoSize {
+ // We cannot have a larger packet following a smaller one.
+ return coalesceUnavailable
+ }
+ return coalesceAppend
+}
+
// tcpPacketsCanCoalesce evaluates if pkt can be coalesced with the packet
// described by item. This function makes considerations that match the kernel's
// GRO self tests, which can be found in tools/testing/selftests/net/gro.c.
@@ -189,29 +343,8 @@ func tcpPacketsCanCoalesce(pkt []byte, iphLen, tcphLen uint8, seq uint32, pshSet
return coalesceUnavailable
}
}
- if pkt[0]>>4 == 6 {
- if pkt[0] != pktTarget[0] || pkt[1]>>4 != pktTarget[1]>>4 {
- // cannot coalesce with unequal Traffic class values
- return coalesceUnavailable
- }
- if pkt[7] != pktTarget[7] {
- // cannot coalesce with unequal Hop limit values
- return coalesceUnavailable
- }
- } else {
- if pkt[1] != pktTarget[1] {
- // cannot coalesce with unequal ToS values
- return coalesceUnavailable
- }
- if pkt[6]>>5 != pktTarget[6]>>5 {
- // cannot coalesce with unequal DF or reserved bits. MF is checked
- // further up the stack.
- return coalesceUnavailable
- }
- if pkt[8] != pktTarget[8] {
- // cannot coalesce with unequal TTL values
- return coalesceUnavailable
- }
+ if !ipHeadersCanCoalesce(pkt, pktTarget) {
+ return coalesceUnavailable
}
// seq adjacency
lhsLen := item.gsoSize
@@ -252,16 +385,16 @@ func tcpPacketsCanCoalesce(pkt []byte, iphLen, tcphLen uint8, seq uint32, pshSet
return coalesceUnavailable
}
-func tcpChecksumValid(pkt []byte, iphLen uint8, isV6 bool) bool {
+func checksumValid(pkt []byte, iphLen, proto uint8, isV6 bool) bool {
srcAddrAt := ipv4SrcAddrOffset
addrSize := 4
if isV6 {
srcAddrAt = ipv6SrcAddrOffset
addrSize = 16
}
- tcpTotalLen := uint16(len(pkt) - int(iphLen))
- tcpCSumNoFold := pseudoHeaderChecksumNoFold(unix.IPPROTO_TCP, pkt[srcAddrAt:srcAddrAt+addrSize], pkt[srcAddrAt+addrSize:srcAddrAt+addrSize*2], tcpTotalLen)
- return ^checksum(pkt[iphLen:], tcpCSumNoFold) == 0
+ lenForPseudo := uint16(len(pkt) - int(iphLen))
+ cSum := pseudoHeaderChecksumNoFold(proto, pkt[srcAddrAt:srcAddrAt+addrSize], pkt[srcAddrAt+addrSize:srcAddrAt+addrSize*2], lenForPseudo)
+ return ^checksum(pkt[iphLen:], cSum) == 0
}
// coalesceResult represents the result of attempting to coalesce two TCP
@@ -276,8 +409,36 @@ const (
coalesceSuccess
)
+// coalesceUDPPackets attempts to coalesce pkt with the packet described by
+// item, and returns the outcome.
+func coalesceUDPPackets(pkt []byte, item *udpGROItem, bufs [][]byte, bufsOffset int, isV6 bool) coalesceResult {
+ pktHead := bufs[item.bufsIndex][bufsOffset:] // the packet that will end up at the front
+ headersLen := item.iphLen + udphLen
+ coalescedLen := len(bufs[item.bufsIndex][bufsOffset:]) + len(pkt) - int(headersLen)
+
+ if cap(pktHead)-bufsOffset < coalescedLen {
+ // We don't want to allocate a new underlying array if capacity is
+ // too small.
+ return coalesceInsufficientCap
+ }
+ if item.numMerged == 0 {
+ if item.cSumKnownInvalid || !checksumValid(bufs[item.bufsIndex][bufsOffset:], item.iphLen, unix.IPPROTO_UDP, isV6) {
+ return coalesceItemInvalidCSum
+ }
+ }
+ if !checksumValid(pkt, item.iphLen, unix.IPPROTO_UDP, isV6) {
+ return coalescePktInvalidCSum
+ }
+ extendBy := len(pkt) - int(headersLen)
+ bufs[item.bufsIndex] = append(bufs[item.bufsIndex], make([]byte, extendBy)...)
+ copy(bufs[item.bufsIndex][bufsOffset+len(pktHead):], pkt[headersLen:])
+
+ item.numMerged++
+ return coalesceSuccess
+}
+
// coalesceTCPPackets attempts to coalesce pkt with the packet described by
-// item, returning the outcome. This function may swap bufs elements in the
+// item, and returns the outcome. This function may swap bufs elements in the
// event of a prepend as item's bufs index is already being tracked for writing
// to a Device.
func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize uint16, seq uint32, pshSet bool, item *tcpGROItem, bufs [][]byte, bufsOffset int, isV6 bool) coalesceResult {
@@ -297,11 +458,11 @@ func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize
return coalescePSHEnding
}
if item.numMerged == 0 {
- if !tcpChecksumValid(bufs[item.bufsIndex][bufsOffset:], item.iphLen, isV6) {
+ if !checksumValid(bufs[item.bufsIndex][bufsOffset:], item.iphLen, unix.IPPROTO_TCP, isV6) {
return coalesceItemInvalidCSum
}
}
- if !tcpChecksumValid(pkt, item.iphLen, isV6) {
+ if !checksumValid(pkt, item.iphLen, unix.IPPROTO_TCP, isV6) {
return coalescePktInvalidCSum
}
item.sentSeq = seq
@@ -319,11 +480,11 @@ func coalesceTCPPackets(mode canCoalesce, pkt []byte, pktBuffsIndex int, gsoSize
return coalesceInsufficientCap
}
if item.numMerged == 0 {
- if !tcpChecksumValid(bufs[item.bufsIndex][bufsOffset:], item.iphLen, isV6) {
+ if !checksumValid(bufs[item.bufsIndex][bufsOffset:], item.iphLen, unix.IPPROTO_TCP, isV6) {
return coalesceItemInvalidCSum
}
}
- if !tcpChecksumValid(pkt, item.iphLen, isV6) {
+ if !checksumValid(pkt, item.iphLen, unix.IPPROTO_TCP, isV6) {
return coalescePktInvalidCSum
}
if pshSet {
@@ -354,52 +515,52 @@ const (
maxUint16 = 1<<16 - 1
)
-type tcpGROResult int
+type groResult int
const (
- tcpGROResultNoop tcpGROResult = iota
- tcpGROResultTableInsert
- tcpGROResultCoalesced
+ groResultNoop groResult = iota
+ groResultTableInsert
+ groResultCoalesced
)
// tcpGRO evaluates the TCP packet at pktI in bufs for coalescing with
-// existing packets tracked in table. It returns a tcpGROResultNoop when no
-// action was taken, tcpGROResultTableInsert when the evaluated packet was
-// inserted into table, and tcpGROResultCoalesced when the evaluated packet was
+// existing packets tracked in table. It returns a groResultNoop when no
+// action was taken, groResultTableInsert when the evaluated packet was
+// inserted into table, and groResultCoalesced when the evaluated packet was
// coalesced with another packet in table.
-func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool) tcpGROResult {
+func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool) groResult {
pkt := bufs[pktI][offset:]
if len(pkt) > maxUint16 {
// A valid IPv4 or IPv6 packet will never exceed this.
- return tcpGROResultNoop
+ return groResultNoop
}
iphLen := int((pkt[0] & 0x0F) * 4)
if isV6 {
iphLen = 40
ipv6HPayloadLen := int(binary.BigEndian.Uint16(pkt[4:]))
if ipv6HPayloadLen != len(pkt)-iphLen {
- return tcpGROResultNoop
+ return groResultNoop
}
} else {
totalLen := int(binary.BigEndian.Uint16(pkt[2:]))
if totalLen != len(pkt) {
- return tcpGROResultNoop
+ return groResultNoop
}
}
if len(pkt) < iphLen {
- return tcpGROResultNoop
+ return groResultNoop
}
tcphLen := int((pkt[iphLen+12] >> 4) * 4)
if tcphLen < 20 || tcphLen > 60 {
- return tcpGROResultNoop
+ return groResultNoop
}
if len(pkt) < iphLen+tcphLen {
- return tcpGROResultNoop
+ return groResultNoop
}
if !isV6 {
if pkt[6]&ipv4FlagMoreFragments != 0 || pkt[6]<<3 != 0 || pkt[7] != 0 {
// no GRO support for fragmented segments for now
- return tcpGROResultNoop
+ return groResultNoop
}
}
tcpFlags := pkt[iphLen+tcpFlagsOffset]
@@ -407,14 +568,14 @@ func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool)
// not a candidate if any non-ACK flags (except PSH+ACK) are set
if tcpFlags != tcpFlagACK {
if pkt[iphLen+tcpFlagsOffset] != tcpFlagACK|tcpFlagPSH {
- return tcpGROResultNoop
+ return groResultNoop
}
pshSet = true
}
gsoSize := uint16(len(pkt) - tcphLen - iphLen)
// not a candidate if payload len is 0
if gsoSize < 1 {
- return tcpGROResultNoop
+ return groResultNoop
}
seq := binary.BigEndian.Uint32(pkt[iphLen+4:])
srcAddrOffset := ipv4SrcAddrOffset
@@ -425,7 +586,7 @@ func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool)
}
items, existing := table.lookupOrInsert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, tcphLen, pktI)
if !existing {
- return tcpGROResultNoop
+ return groResultTableInsert
}
for i := len(items) - 1; i >= 0; i-- {
// In the best case of packets arriving in order iterating in reverse is
@@ -443,54 +604,25 @@ func tcpGRO(bufs [][]byte, offset int, pktI int, table *tcpGROTable, isV6 bool)
switch result {
case coalesceSuccess:
table.updateAt(item, i)
- return tcpGROResultCoalesced
+ return groResultCoalesced
case coalesceItemInvalidCSum:
// delete the item with an invalid csum
table.deleteAt(item.key, i)
case coalescePktInvalidCSum:
// no point in inserting an item that we can't coalesce
- return tcpGROResultNoop
+ return groResultNoop
default:
}
}
}
// failed to coalesce with any other packets; store the item in the flow
table.insert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, tcphLen, pktI)
- return tcpGROResultTableInsert
+ return groResultTableInsert
}
-func isTCP4NoIPOptions(b []byte) bool {
- if len(b) < 40 {
- return false
- }
- if b[0]>>4 != 4 {
- return false
- }
- if b[0]&0x0F != 5 {
- return false
- }
- if b[9] != unix.IPPROTO_TCP {
- return false
- }
- return true
-}
-
-func isTCP6NoEH(b []byte) bool {
- if len(b) < 60 {
- return false
- }
- if b[0]>>4 != 6 {
- return false
- }
- if b[6] != unix.IPPROTO_TCP {
- return false
- }
- return true
-}
-
-// applyCoalesceAccounting updates bufs to account for coalescing based on the
+// applyTCPCoalesceAccounting updates bufs to account for coalescing based on the
// metadata found in table.
-func applyCoalesceAccounting(bufs [][]byte, offset int, table *tcpGROTable, isV6 bool) error {
+func applyTCPCoalesceAccounting(bufs [][]byte, offset int, table *tcpGROTable) error {
for _, items := range table.itemsByFlow {
for _, item := range items {
if item.numMerged > 0 {
@@ -505,7 +637,7 @@ func applyCoalesceAccounting(bufs [][]byte, offset int, table *tcpGROTable, isV6
// Recalculate the total len (IPv4) or payload len (IPv6).
// Recalculate the (IPv4) header checksum.
- if isV6 {
+ if item.key.isV6 {
hdr.gsoType = unix.VIRTIO_NET_HDR_GSO_TCPV6
binary.BigEndian.PutUint16(pkt[4:], uint16(len(pkt))-uint16(item.iphLen)) // set new IPv6 header payload len
} else {
@@ -525,7 +657,7 @@ func applyCoalesceAccounting(bufs [][]byte, offset int, table *tcpGROTable, isV6
// this with computation of the tcp header and payload checksum.
addrLen := 4
addrOffset := ipv4SrcAddrOffset
- if isV6 {
+ if item.key.isV6 {
addrLen = 16
addrOffset = ipv6SrcAddrOffset
}
@@ -546,54 +678,245 @@ func applyCoalesceAccounting(bufs [][]byte, offset int, table *tcpGROTable, isV6
return nil
}
+// applyUDPCoalesceAccounting updates bufs to account for coalescing based on the
+// metadata found in table.
+func applyUDPCoalesceAccounting(bufs [][]byte, offset int, table *udpGROTable) error {
+ for _, items := range table.itemsByFlow {
+ for _, item := range items {
+ if item.numMerged > 0 {
+ hdr := virtioNetHdr{
+ flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM, // this turns into CHECKSUM_PARTIAL in the skb
+ hdrLen: uint16(item.iphLen + udphLen),
+ gsoSize: item.gsoSize,
+ csumStart: uint16(item.iphLen),
+ csumOffset: 6,
+ }
+ pkt := bufs[item.bufsIndex][offset:]
+
+ // Recalculate the total len (IPv4) or payload len (IPv6).
+ // Recalculate the (IPv4) header checksum.
+ hdr.gsoType = unix.VIRTIO_NET_HDR_GSO_UDP_L4
+ if item.key.isV6 {
+ binary.BigEndian.PutUint16(pkt[4:], uint16(len(pkt))-uint16(item.iphLen)) // set new IPv6 header payload len
+ } else {
+ pkt[10], pkt[11] = 0, 0
+ binary.BigEndian.PutUint16(pkt[2:], uint16(len(pkt))) // set new total length
+ iphCSum := ^checksum(pkt[:item.iphLen], 0) // compute IPv4 header checksum
+ binary.BigEndian.PutUint16(pkt[10:], iphCSum) // set IPv4 header checksum field
+ }
+ err := hdr.encode(bufs[item.bufsIndex][offset-virtioNetHdrLen:])
+ if err != nil {
+ return err
+ }
+
+ // Recalculate the UDP len field value
+ binary.BigEndian.PutUint16(pkt[item.iphLen+4:], uint16(len(pkt[item.iphLen:])))
+
+ // Calculate the pseudo header checksum and place it at the UDP
+ // checksum offset. Downstream checksum offloading will combine
+ // this with computation of the udp header and payload checksum.
+ addrLen := 4
+ addrOffset := ipv4SrcAddrOffset
+ if item.key.isV6 {
+ addrLen = 16
+ addrOffset = ipv6SrcAddrOffset
+ }
+ srcAddrAt := offset + addrOffset
+ srcAddr := bufs[item.bufsIndex][srcAddrAt : srcAddrAt+addrLen]
+ dstAddr := bufs[item.bufsIndex][srcAddrAt+addrLen : srcAddrAt+addrLen*2]
+ psum := pseudoHeaderChecksumNoFold(unix.IPPROTO_UDP, srcAddr, dstAddr, uint16(len(pkt)-int(item.iphLen)))
+ binary.BigEndian.PutUint16(pkt[hdr.csumStart+hdr.csumOffset:], checksum([]byte{}, psum))
+ } else {
+ hdr := virtioNetHdr{}
+ err := hdr.encode(bufs[item.bufsIndex][offset-virtioNetHdrLen:])
+ if err != nil {
+ return err
+ }
+ }
+ }
+ }
+ return nil
+}
+
+type groCandidateType uint8
+
+const (
+ notGROCandidate groCandidateType = iota
+ tcp4GROCandidate
+ tcp6GROCandidate
+ udp4GROCandidate
+ udp6GROCandidate
+)
+
+func packetIsGROCandidate(b []byte, canUDPGRO bool) groCandidateType {
+ if len(b) < 28 {
+ return notGROCandidate
+ }
+ if b[0]>>4 == 4 {
+ if b[0]&0x0F != 5 {
+ // IPv4 packets w/IP options do not coalesce
+ return notGROCandidate
+ }
+ if b[9] == unix.IPPROTO_TCP && len(b) >= 40 {
+ return tcp4GROCandidate
+ }
+ if b[9] == unix.IPPROTO_UDP && canUDPGRO {
+ return udp4GROCandidate
+ }
+ } else if b[0]>>4 == 6 {
+ if b[6] == unix.IPPROTO_TCP && len(b) >= 60 {
+ return tcp6GROCandidate
+ }
+ if b[6] == unix.IPPROTO_UDP && len(b) >= 48 && canUDPGRO {
+ return udp6GROCandidate
+ }
+ }
+ return notGROCandidate
+}
+
+const (
+ udphLen = 8
+)
+
+// udpGRO evaluates the UDP packet at pktI in bufs for coalescing with
+// existing packets tracked in table. It returns a groResultNoop when no
+// action was taken, groResultTableInsert when the evaluated packet was
+// inserted into table, and groResultCoalesced when the evaluated packet was
+// coalesced with another packet in table.
+func udpGRO(bufs [][]byte, offset int, pktI int, table *udpGROTable, isV6 bool) groResult {
+ pkt := bufs[pktI][offset:]
+ if len(pkt) > maxUint16 {
+ // A valid IPv4 or IPv6 packet will never exceed this.
+ return groResultNoop
+ }
+ iphLen := int((pkt[0] & 0x0F) * 4)
+ if isV6 {
+ iphLen = 40
+ ipv6HPayloadLen := int(binary.BigEndian.Uint16(pkt[4:]))
+ if ipv6HPayloadLen != len(pkt)-iphLen {
+ return groResultNoop
+ }
+ } else {
+ totalLen := int(binary.BigEndian.Uint16(pkt[2:]))
+ if totalLen != len(pkt) {
+ return groResultNoop
+ }
+ }
+ if len(pkt) < iphLen {
+ return groResultNoop
+ }
+ if len(pkt) < iphLen+udphLen {
+ return groResultNoop
+ }
+ if !isV6 {
+ if pkt[6]&ipv4FlagMoreFragments != 0 || pkt[6]<<3 != 0 || pkt[7] != 0 {
+ // no GRO support for fragmented segments for now
+ return groResultNoop
+ }
+ }
+ gsoSize := uint16(len(pkt) - udphLen - iphLen)
+ // not a candidate if payload len is 0
+ if gsoSize < 1 {
+ return groResultNoop
+ }
+ srcAddrOffset := ipv4SrcAddrOffset
+ addrLen := 4
+ if isV6 {
+ srcAddrOffset = ipv6SrcAddrOffset
+ addrLen = 16
+ }
+ items, existing := table.lookupOrInsert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, pktI)
+ if !existing {
+ return groResultTableInsert
+ }
+ // With UDP we only check the last item, otherwise we could reorder packets
+ // for a given flow. We must also always insert a new item, or successfully
+ // coalesce with an existing item, for the same reason.
+ item := items[len(items)-1]
+ can := udpPacketsCanCoalesce(pkt, uint8(iphLen), gsoSize, item, bufs, offset)
+ var pktCSumKnownInvalid bool
+ if can == coalesceAppend {
+ result := coalesceUDPPackets(pkt, &item, bufs, offset, isV6)
+ switch result {
+ case coalesceSuccess:
+ table.updateAt(item, len(items)-1)
+ return groResultCoalesced
+ case coalesceItemInvalidCSum:
+ // If the existing item has an invalid csum we take no action. A new
+ // item will be stored after it, and the existing item will never be
+ // revisited as part of future coalescing candidacy checks.
+ case coalescePktInvalidCSum:
+ // We must insert a new item, but we also mark it as invalid csum
+ // to prevent a repeat checksum validation.
+ pktCSumKnownInvalid = true
+ default:
+ }
+ }
+ // failed to coalesce with any other packets; store the item in the flow
+ table.insert(pkt, srcAddrOffset, srcAddrOffset+addrLen, iphLen, pktI, pktCSumKnownInvalid)
+ return groResultTableInsert
+}
+
// handleGRO evaluates bufs for GRO, and writes the indices of the resulting
-// packets into toWrite. toWrite, tcp4Table, and tcp6Table should initially be
+// packets into toWrite. toWrite, tcpTable, and udpTable should initially be
// empty (but non-nil), and are passed in to save allocs as the caller may reset
-// and recycle them across vectors of packets.
-func handleGRO(bufs [][]byte, offset int, tcp4Table, tcp6Table *tcpGROTable, toWrite *[]int) error {
+// and recycle them across vectors of packets. canUDPGRO indicates if UDP GRO is
+// supported.
+func handleGRO(bufs [][]byte, offset int, tcpTable *tcpGROTable, udpTable *udpGROTable, canUDPGRO bool, toWrite *[]int) error {
for i := range bufs {
if offset < virtioNetHdrLen || offset > len(bufs[i])-1 {
return errors.New("invalid offset")
}
- var result tcpGROResult
- switch {
- case isTCP4NoIPOptions(bufs[i][offset:]): // ipv4 packets w/IP options do not coalesce
- result = tcpGRO(bufs, offset, i, tcp4Table, false)
- case isTCP6NoEH(bufs[i][offset:]): // ipv6 packets w/extension headers do not coalesce
- result = tcpGRO(bufs, offset, i, tcp6Table, true)
+ var result groResult
+ switch packetIsGROCandidate(bufs[i][offset:], canUDPGRO) {
+ case tcp4GROCandidate:
+ result = tcpGRO(bufs, offset, i, tcpTable, false)
+ case tcp6GROCandidate:
+ result = tcpGRO(bufs, offset, i, tcpTable, true)
+ case udp4GROCandidate:
+ result = udpGRO(bufs, offset, i, udpTable, false)
+ case udp6GROCandidate:
+ result = udpGRO(bufs, offset, i, udpTable, true)
}
switch result {
- case tcpGROResultNoop:
+ case groResultNoop:
hdr := virtioNetHdr{}
err := hdr.encode(bufs[i][offset-virtioNetHdrLen:])
if err != nil {
return err
}
fallthrough
- case tcpGROResultTableInsert:
+ case groResultTableInsert:
*toWrite = append(*toWrite, i)
}
}
- err4 := applyCoalesceAccounting(bufs, offset, tcp4Table, false)
- err6 := applyCoalesceAccounting(bufs, offset, tcp6Table, true)
- return errors.Join(err4, err6)
+ errTCP := applyTCPCoalesceAccounting(bufs, offset, tcpTable)
+ errUDP := applyUDPCoalesceAccounting(bufs, offset, udpTable)
+ return errors.Join(errTCP, errUDP)
}
-// tcpTSO splits packets from in into outBuffs, writing the size of each
+// gsoSplit splits packets from in into outBuffs, writing the size of each
// element into sizes. It returns the number of buffers populated, and/or an
// error.
-func tcpTSO(in []byte, hdr virtioNetHdr, outBuffs [][]byte, sizes []int, outOffset int) (int, error) {
+func gsoSplit(in []byte, hdr virtioNetHdr, outBuffs [][]byte, sizes []int, outOffset int, isV6 bool) (int, error) {
iphLen := int(hdr.csumStart)
srcAddrOffset := ipv6SrcAddrOffset
addrLen := 16
- if hdr.gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 {
+ if !isV6 {
in[10], in[11] = 0, 0 // clear ipv4 header checksum
srcAddrOffset = ipv4SrcAddrOffset
addrLen = 4
}
- tcpCSumAt := int(hdr.csumStart + hdr.csumOffset)
- in[tcpCSumAt], in[tcpCSumAt+1] = 0, 0 // clear tcp checksum
- firstTCPSeqNum := binary.BigEndian.Uint32(in[hdr.csumStart+4:])
+ transportCsumAt := int(hdr.csumStart + hdr.csumOffset)
+ in[transportCsumAt], in[transportCsumAt+1] = 0, 0 // clear tcp/udp checksum
+ var firstTCPSeqNum uint32
+ var protocol uint8
+ if hdr.gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 || hdr.gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV6 {
+ protocol = unix.IPPROTO_TCP
+ firstTCPSeqNum = binary.BigEndian.Uint32(in[hdr.csumStart+4:])
+ } else {
+ protocol = unix.IPPROTO_UDP
+ }
nextSegmentDataAt := int(hdr.hdrLen)
i := 0
for ; nextSegmentDataAt < len(in); i++ {
@@ -610,7 +933,7 @@ func tcpTSO(in []byte, hdr virtioNetHdr, outBuffs [][]byte, sizes []int, outOffs
out := outBuffs[i][outOffset:]
copy(out, in[:iphLen])
- if hdr.gsoType == unix.VIRTIO_NET_HDR_GSO_TCPV4 {
+ if !isV6 {
// For IPv4 we are responsible for incrementing the ID field,
// updating the total len field, and recalculating the header
// checksum.
@@ -627,25 +950,32 @@ func tcpTSO(in []byte, hdr virtioNetHdr, outBuffs [][]byte, sizes []int, outOffs
binary.BigEndian.PutUint16(out[4:], uint16(totalLen-iphLen))
}
- // TCP header
+ // copy transport header
copy(out[hdr.csumStart:hdr.hdrLen], in[hdr.csumStart:hdr.hdrLen])
- tcpSeq := firstTCPSeqNum + uint32(hdr.gsoSize*uint16(i))
- binary.BigEndian.PutUint32(out[hdr.csumStart+4:], tcpSeq)
- if nextSegmentEnd != len(in) {
- // FIN and PSH should only be set on last segment
- clearFlags := tcpFlagFIN | tcpFlagPSH
- out[hdr.csumStart+tcpFlagsOffset] &^= clearFlags
+
+ if protocol == unix.IPPROTO_TCP {
+ // set TCP seq and adjust TCP flags
+ tcpSeq := firstTCPSeqNum + uint32(hdr.gsoSize*uint16(i))
+ binary.BigEndian.PutUint32(out[hdr.csumStart+4:], tcpSeq)
+ if nextSegmentEnd != len(in) {
+ // FIN and PSH should only be set on last segment
+ clearFlags := tcpFlagFIN | tcpFlagPSH
+ out[hdr.csumStart+tcpFlagsOffset] &^= clearFlags
+ }
+ } else {
+ // set UDP header len
+ binary.BigEndian.PutUint16(out[hdr.csumStart+4:], uint16(segmentDataLen)+(hdr.hdrLen-hdr.csumStart))
}
// payload
copy(out[hdr.hdrLen:], in[nextSegmentDataAt:nextSegmentEnd])
- // TCP checksum
- tcpHLen := int(hdr.hdrLen - hdr.csumStart)
- tcpLenForPseudo := uint16(tcpHLen + segmentDataLen)
- tcpCSumNoFold := pseudoHeaderChecksumNoFold(unix.IPPROTO_TCP, in[srcAddrOffset:srcAddrOffset+addrLen], in[srcAddrOffset+addrLen:srcAddrOffset+addrLen*2], tcpLenForPseudo)
- tcpCSum := ^checksum(out[hdr.csumStart:totalLen], tcpCSumNoFold)
- binary.BigEndian.PutUint16(out[hdr.csumStart+hdr.csumOffset:], tcpCSum)
+ // transport checksum
+ transportHeaderLen := int(hdr.hdrLen - hdr.csumStart)
+ lenForPseudo := uint16(transportHeaderLen + segmentDataLen)
+ transportCSumNoFold := pseudoHeaderChecksumNoFold(protocol, in[srcAddrOffset:srcAddrOffset+addrLen], in[srcAddrOffset+addrLen:srcAddrOffset+addrLen*2], lenForPseudo)
+ transportCSum := ^checksum(out[hdr.csumStart:totalLen], transportCSumNoFold)
+ binary.BigEndian.PutUint16(out[hdr.csumStart+hdr.csumOffset:], transportCSum)
nextSegmentDataAt += int(hdr.gsoSize)
}
diff --git a/tun/offload_linux_test.go b/tun/offload_linux_test.go
new file mode 100644
index 0000000..ae55c8c
--- /dev/null
+++ b/tun/offload_linux_test.go
@@ -0,0 +1,752 @@
+/* SPDX-License-Identifier: MIT
+ *
+ * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ */
+
+package tun
+
+import (
+ "net/netip"
+ "testing"
+
+ "golang.org/x/sys/unix"
+ "golang.zx2c4.com/wireguard/conn"
+ "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, true, offset)
+ f.Fuzz(func(t *testing.T, pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, pkt6, pkt7, pkt8, pkt9, pkt10, pkt11 []byte, canUDPGRO bool, 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(), canUDPGRO, &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
+ canUDPGRO bool
+ 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
+ },
+ true,
+ []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
+ },
+ false,
+ []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
+ },
+ true,
+ []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),
+ },
+ true,
+ []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
+ },
+ true,
+ []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++
+ }),
+ },
+ true,
+ []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++
+ }),
+ },
+ true,
+ []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
+ }),
+ },
+ true,
+ []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
+ }),
+ },
+ true,
+ []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++
+ }),
+ },
+ true,
+ []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++
+ }),
+ },
+ true,
+ []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.canUDPGRO, &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
+ canUDPGRO bool
+ want groCandidateType
+ }{
+ {
+ "tcp4",
+ tcp4,
+ true,
+ tcp4GROCandidate,
+ },
+ {
+ "tcp6",
+ tcp6,
+ true,
+ tcp6GROCandidate,
+ },
+ {
+ "udp4",
+ udp4,
+ true,
+ udp4GROCandidate,
+ },
+ {
+ "udp4 no support",
+ udp4,
+ false,
+ notGROCandidate,
+ },
+ {
+ "udp6",
+ udp6,
+ true,
+ udp6GROCandidate,
+ },
+ {
+ "udp6 no support",
+ udp6,
+ false,
+ notGROCandidate,
+ },
+ {
+ "udp4 too short",
+ udp4TooShort,
+ true,
+ notGROCandidate,
+ },
+ {
+ "udp6 too short",
+ udp6TooShort,
+ true,
+ notGROCandidate,
+ },
+ {
+ "tcp4 too short",
+ tcp4TooShort,
+ true,
+ notGROCandidate,
+ },
+ {
+ "tcp6 too short",
+ tcp6TooShort,
+ true,
+ notGROCandidate,
+ },
+ {
+ "invalid IP version",
+ []byte{0x00},
+ true,
+ notGROCandidate,
+ },
+ {
+ "invalid IP header len",
+ ip4InvalidHeaderLen,
+ true,
+ notGROCandidate,
+ },
+ {
+ "ip4 invalid protocol",
+ ip4InvalidProtocol,
+ true,
+ notGROCandidate,
+ },
+ {
+ "ip6 invalid protocol",
+ ip6InvalidProtocol,
+ true,
+ notGROCandidate,
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ if got := packetIsGROCandidate(tt.b, tt.canUDPGRO); 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)
+ }
+ })
+ }
+}
diff --git a/tun/tcp_offload_linux_test.go b/tun/tcp_offload_linux_test.go
deleted file mode 100644
index ddddc48..0000000
--- a/tun/tcp_offload_linux_test.go
+++ /dev/null
@@ -1,411 +0,0 @@
-/* SPDX-License-Identifier: MIT
- *
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
- */
-
-package tun
-
-import (
- "net/netip"
- "testing"
-
- "golang.org/x/sys/unix"
- "golang.zx2c4.com/wireguard/conn"
- "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 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,
- },
- }
-
- 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 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)
- f.Add(pkt0, pkt1, pkt2, pkt3, pkt4, pkt5, offset)
- f.Fuzz(func(t *testing.T, pkt0, pkt1, pkt2, pkt3, pkt4, pkt5 []byte, offset int) {
- pkts := [][]byte{pkt0, pkt1, pkt2, pkt3, pkt4, pkt5}
- toWrite := make([]int, 0, len(pkts))
- handleGRO(pkts, offset, newTCPGROTable(), newTCPGROTable(), &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
- wantToWrite []int
- wantLens []int
- wantErr bool
- }{
- {
- "multiple flows",
- [][]byte{
- tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1), // v4 flow 1
- tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101), // v4 flow 1
- tcp4Packet(ip4PortA, ip4PortC, header.TCPFlagAck, 100, 201), // v4 flow 2
- tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1), // v6 flow 1
- tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101), // v6 flow 1
- tcp6Packet(ip6PortA, ip6PortC, header.TCPFlagAck, 100, 201), // v6 flow 2
- },
- []int{0, 2, 3, 5},
- []int{240, 140, 260, 160},
- 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
- },
- []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
- },
- []int{0, 1},
- []int{140, 240},
- 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
- },
- []int{0},
- []int{340},
- false,
- },
- {
- "tcp4 unequal TTL",
- [][]byte{
- tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1),
- tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) {
- fields.TTL++
- }),
- },
- []int{0, 1},
- []int{140, 140},
- false,
- },
- {
- "tcp4 unequal ToS",
- [][]byte{
- tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1),
- tcp4PacketMutateIPFields(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv4Fields) {
- fields.TOS++
- }),
- },
- []int{0, 1},
- []int{140, 140},
- false,
- },
- {
- "tcp4 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
- }),
- },
- []int{0, 1},
- []int{140, 140},
- false,
- },
- {
- "tcp4 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
- }),
- },
- []int{0, 1},
- []int{140, 140},
- false,
- },
- {
- "tcp6 unequal hop limit",
- [][]byte{
- tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1),
- tcp6PacketMutateIPFields(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv6Fields) {
- fields.HopLimit++
- }),
- },
- []int{0, 1},
- []int{160, 160},
- false,
- },
- {
- "tcp6 unequal traffic class",
- [][]byte{
- tcp6Packet(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 1),
- tcp6PacketMutateIPFields(ip6PortA, ip6PortB, header.TCPFlagAck, 100, 101, func(fields *header.IPv6Fields) {
- fields.TrafficClass++
- }),
- },
- []int{0, 1},
- []int{160, 160},
- 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(), newTCPGROTable(), &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_isTCP4NoIPOptions(t *testing.T) {
- valid := tcp4Packet(ip4PortA, ip4PortB, header.TCPFlagAck, 100, 1)[virtioNetHdrLen:]
- invalidLen := valid[:39]
- invalidHeaderLen := make([]byte, len(valid))
- copy(invalidHeaderLen, valid)
- invalidHeaderLen[0] = 0x46
- invalidProtocol := make([]byte, len(valid))
- copy(invalidProtocol, valid)
- invalidProtocol[9] = unix.IPPROTO_TCP + 1
-
- tests := []struct {
- name string
- b []byte
- want bool
- }{
- {
- "valid",
- valid,
- true,
- },
- {
- "invalid length",
- invalidLen,
- false,
- },
- {
- "invalid version",
- []byte{0x00},
- false,
- },
- {
- "invalid header len",
- invalidHeaderLen,
- false,
- },
- {
- "invalid protocol",
- invalidProtocol,
- false,
- },
- }
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- if got := isTCP4NoIPOptions(tt.b); got != tt.want {
- t.Errorf("isTCP4NoIPOptions() = %v, want %v", got, tt.want)
- }
- })
- }
-}
diff --git a/tun/testdata/fuzz/Fuzz_handleGRO/032aec0105f26f709c118365e4830d6dc087cab24cd1e154c2e790589a309b77 b/tun/testdata/fuzz/Fuzz_handleGRO/032aec0105f26f709c118365e4830d6dc087cab24cd1e154c2e790589a309b77
deleted file mode 100644
index 5461e79..0000000
--- a/tun/testdata/fuzz/Fuzz_handleGRO/032aec0105f26f709c118365e4830d6dc087cab24cd1e154c2e790589a309b77
+++ /dev/null
@@ -1,8 +0,0 @@
-go test fuzz v1
-[]byte("0")
-[]byte("0")
-[]byte("0")
-[]byte("0")
-[]byte("0")
-[]byte("0")
-int(34)
diff --git a/tun/testdata/fuzz/Fuzz_handleGRO/0da283f9a2098dec30d1c86784411a8ce2e8e03aa3384105e581f2c67494700d b/tun/testdata/fuzz/Fuzz_handleGRO/0da283f9a2098dec30d1c86784411a8ce2e8e03aa3384105e581f2c67494700d
deleted file mode 100644
index b441819..0000000
--- a/tun/testdata/fuzz/Fuzz_handleGRO/0da283f9a2098dec30d1c86784411a8ce2e8e03aa3384105e581f2c67494700d
+++ /dev/null
@@ -1,8 +0,0 @@
-go test fuzz v1
-[]byte("0")
-[]byte("0")
-[]byte("0")
-[]byte("0")
-[]byte("0")
-[]byte("0")
-int(-48)
diff --git a/tun/tun_linux.go b/tun/tun_linux.go
index 12cd49f..bd69cb5 100644
--- a/tun/tun_linux.go
+++ b/tun/tun_linux.go
@@ -38,6 +38,7 @@ type NativeTun struct {
statusListenersShutdown chan struct{}
batchSize int
vnetHdr bool
+ udpGSO bool
closeOnce sync.Once
@@ -48,9 +49,10 @@ type NativeTun struct {
readOpMu sync.Mutex // readOpMu guards readBuff
readBuff [virtioNetHdrLen + 65535]byte // if vnetHdr every read() is prefixed by virtioNetHdr
- writeOpMu sync.Mutex // writeOpMu guards toWrite, tcp4GROTable, tcp6GROTable
- toWrite []int
- tcp4GROTable, tcp6GROTable *tcpGROTable
+ writeOpMu sync.Mutex // writeOpMu guards toWrite, tcpGROTable
+ toWrite []int
+ tcpGROTable *tcpGROTable
+ udpGROTable *udpGROTable
}
func (tun *NativeTun) File() *os.File {
@@ -333,8 +335,8 @@ func (tun *NativeTun) nameSlow() (string, error) {
func (tun *NativeTun) Write(bufs [][]byte, offset int) (int, error) {
tun.writeOpMu.Lock()
defer func() {
- tun.tcp4GROTable.reset()
- tun.tcp6GROTable.reset()
+ tun.tcpGROTable.reset()
+ tun.udpGROTable.reset()
tun.writeOpMu.Unlock()
}()
var (
@@ -343,7 +345,7 @@ func (tun *NativeTun) Write(bufs [][]byte, offset int) (int, error) {
)
tun.toWrite = tun.toWrite[:0]
if tun.vnetHdr {
- err := handleGRO(bufs, offset, tun.tcp4GROTable, tun.tcp6GROTable, &tun.toWrite)
+ err := handleGRO(bufs, offset, tun.tcpGROTable, tun.udpGROTable, tun.udpGSO, &tun.toWrite)
if err != nil {
return 0, err
}
@@ -394,37 +396,42 @@ func handleVirtioRead(in []byte, bufs [][]byte, sizes []int, offset int) (int, e
sizes[0] = n
return 1, nil
}
- if hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV4 && hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV6 {
+ if hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV4 && hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV6 && hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
return 0, fmt.Errorf("unsupported virtio GSO type: %d", hdr.gsoType)
}
ipVersion := in[0] >> 4
switch ipVersion {
case 4:
- if hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV4 {
+ if hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV4 && hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
return 0, fmt.Errorf("ip header version: %d, GSO type: %d", ipVersion, hdr.gsoType)
}
case 6:
- if hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV6 {
+ if hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_TCPV6 && hdr.gsoType != unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
return 0, fmt.Errorf("ip header version: %d, GSO type: %d", ipVersion, hdr.gsoType)
}
default:
return 0, fmt.Errorf("invalid ip header version: %d", ipVersion)
}
- if len(in) <= int(hdr.csumStart+12) {
- return 0, errors.New("packet is too short")
- }
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
// of the entire first packet when the kernel is handling it as part of a
- // FORWARD path. Instead, parse the TCP header length and add it onto
+ // FORWARD path. Instead, parse the transport header length and add it onto
// csumStart, which is synonymous for IP header length.
- tcpHLen := uint16(in[hdr.csumStart+12] >> 4 * 4)
- if tcpHLen < 20 || tcpHLen > 60 {
- // A TCP header must be between 20 and 60 bytes in length.
- return 0, fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
+ if hdr.gsoType == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
+ hdr.hdrLen = hdr.csumStart + 8
+ } else {
+ if len(in) <= int(hdr.csumStart+12) {
+ return 0, errors.New("packet is too short")
+ }
+
+ tcpHLen := uint16(in[hdr.csumStart+12] >> 4 * 4)
+ if tcpHLen < 20 || tcpHLen > 60 {
+ // A TCP header must be between 20 and 60 bytes in length.
+ return 0, fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
+ }
+ hdr.hdrLen = hdr.csumStart + tcpHLen
}
- hdr.hdrLen = hdr.csumStart + tcpHLen
if len(in) < int(hdr.hdrLen) {
return 0, fmt.Errorf("length of packet (%d) < virtioNetHdr.hdrLen (%d)", len(in), hdr.hdrLen)
@@ -438,7 +445,7 @@ func handleVirtioRead(in []byte, bufs [][]byte, sizes []int, offset int) (int, e
return 0, fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(in))
}
- return tcpTSO(in, hdr, bufs, sizes, offset)
+ return gsoSplit(in, hdr, bufs, sizes, offset, ipVersion == 6)
}
func (tun *NativeTun) Read(bufs [][]byte, sizes []int, offset int) (int, error) {
@@ -497,7 +504,8 @@ func (tun *NativeTun) BatchSize() int {
const (
// TODO: support TSO with ECN bits
- tunOffloads = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6
+ tunTCPOffloads = unix.TUN_F_CSUM | unix.TUN_F_TSO4 | unix.TUN_F_TSO6
+ tunUDPOffloads = unix.TUN_F_USO4 | unix.TUN_F_USO6
)
func (tun *NativeTun) initFromFlags(name string) error {
@@ -519,12 +527,17 @@ func (tun *NativeTun) initFromFlags(name string) error {
}
got := ifr.Uint16()
if got&unix.IFF_VNET_HDR != 0 {
- err = unix.IoctlSetInt(int(fd), unix.TUNSETOFFLOAD, tunOffloads)
+ // tunTCPOffloads were added in Linux v2.6. We require their support
+ // if IFF_VNET_HDR is set.
+ err = unix.IoctlSetInt(int(fd), unix.TUNSETOFFLOAD, tunTCPOffloads)
if err != nil {
return
}
tun.vnetHdr = true
tun.batchSize = conn.IdealBatchSize
+ // tunUDPOffloads were added in Linux v6.2. We do not return an
+ // error if they are unsupported at runtime.
+ tun.udpGSO = unix.IoctlSetInt(int(fd), unix.TUNSETOFFLOAD, tunTCPOffloads|tunUDPOffloads) == nil
} else {
tun.batchSize = 1
}
@@ -575,8 +588,8 @@ func CreateTUNFromFile(file *os.File, mtu int) (Device, error) {
events: make(chan Event, 5),
errors: make(chan error, 5),
statusListenersShutdown: make(chan struct{}),
- tcp4GROTable: newTCPGROTable(),
- tcp6GROTable: newTCPGROTable(),
+ tcpGROTable: newTCPGROTable(),
+ udpGROTable: newUDPGROTable(),
toWrite: make([]int, 0, conn.IdealBatchSize),
}
@@ -628,12 +641,12 @@ func CreateUnmonitoredTUNFromFD(fd int) (Device, string, error) {
}
file := os.NewFile(uintptr(fd), "/dev/tun")
tun := &NativeTun{
- tunFile: file,
- events: make(chan Event, 5),
- errors: make(chan error, 5),
- tcp4GROTable: newTCPGROTable(),
- tcp6GROTable: newTCPGROTable(),
- toWrite: make([]int, 0, conn.IdealBatchSize),
+ tunFile: file,
+ events: make(chan Event, 5),
+ errors: make(chan error, 5),
+ tcpGROTable: newTCPGROTable(),
+ udpGROTable: newUDPGROTable(),
+ toWrite: make([]int, 0, conn.IdealBatchSize),
}
name, err := tun.Name()
if err != nil {
From 4ffa9c20327b9471c3eeb142347f679b69f84648 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Mon, 20 Nov 2023 16:49:06 -0800
Subject: [PATCH 20/75] device: change Peer.endpoint locking to reduce
contention
Access to Peer.endpoint was previously synchronized by Peer.RWMutex.
This has now moved to Peer.endpoint.Mutex. Peer.SendBuffers() is now the
sole caller of Endpoint.ClearSrc(), which is signaled via a new bool,
Peer.endpoint.clearSrcOnTx. Previous Callers of Endpoint.ClearSrc() now
set this bool, primarily via peer.markEndpointSrcForClearing().
Peer.SetEndpointFromPacket() clears Peer.endpoint.clearSrcOnTx when an
updated conn.Endpoint is stored. This maintains the same event order as
before, i.e. a conn.Endpoint received after peer.endpoint.clearSrcOnTx
is set, but before the next Peer.SendBuffers() call results in the
latest conn.Endpoint source being used for the next packet transmission.
These changes result in throughput improvements for single flow,
parallel (-P n) flow, and bidirectional (--bidir) flow iperf3 TCP/UDP
tests as measured on both Linux and Windows. Latency under load improves
especially for high throughput Linux scenarios. These improvements are
likely realized on all platforms to some degree, as the changes are not
platform-specific.
Co-authored-by: James Tucker
Signed-off-by: James Tucker
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
device/device.go | 12 ++--------
device/mobilequirks.go | 6 ++---
device/peer.go | 50 ++++++++++++++++++++++++++------------
device/sticky_linux.go | 30 +++++++++++------------
device/timers.go | 12 ++--------
device/uapi.go | 54 ++++++++++++++++++++----------------------
6 files changed, 83 insertions(+), 81 deletions(-)
diff --git a/device/device.go b/device/device.go
index f9557a0..ca26d00 100644
--- a/device/device.go
+++ b/device/device.go
@@ -461,11 +461,7 @@ func (device *Device) BindSetMark(mark uint32) error {
// clear cached source addresses
device.peers.RLock()
for _, peer := range device.peers.keyMap {
- peer.Lock()
- defer peer.Unlock()
- if peer.endpoint != nil {
- peer.endpoint.ClearSrc()
- }
+ peer.markEndpointSrcForClearing()
}
device.peers.RUnlock()
@@ -515,11 +511,7 @@ func (device *Device) BindUpdate() error {
// clear cached source addresses
device.peers.RLock()
for _, peer := range device.peers.keyMap {
- peer.Lock()
- defer peer.Unlock()
- if peer.endpoint != nil {
- peer.endpoint.ClearSrc()
- }
+ peer.markEndpointSrcForClearing()
}
device.peers.RUnlock()
diff --git a/device/mobilequirks.go b/device/mobilequirks.go
index 4e5051d..0a0080e 100644
--- a/device/mobilequirks.go
+++ b/device/mobilequirks.go
@@ -11,9 +11,9 @@ func (device *Device) DisableSomeRoamingForBrokenMobileSemantics() {
device.net.brokenRoaming = true
device.peers.RLock()
for _, peer := range device.peers.keyMap {
- peer.Lock()
- peer.disableRoaming = peer.endpoint != nil
- peer.Unlock()
+ peer.endpoint.Lock()
+ peer.endpoint.disableRoaming = peer.endpoint.val != nil
+ peer.endpoint.Unlock()
}
device.peers.RUnlock()
}
diff --git a/device/peer.go b/device/peer.go
index 2fb5da6..47a2f14 100644
--- a/device/peer.go
+++ b/device/peer.go
@@ -17,17 +17,20 @@ import (
type Peer struct {
isRunning atomic.Bool
- sync.RWMutex // Mostly protects endpoint, but is generally taken whenever we modify peer
keypairs Keypairs
handshake Handshake
device *Device
- endpoint conn.Endpoint
stopping sync.WaitGroup // routines pending stop
txBytes atomic.Uint64 // bytes send to peer (endpoint)
rxBytes atomic.Uint64 // bytes received from peer
lastHandshakeNano atomic.Int64 // nano seconds since epoch
- disableRoaming bool
+ endpoint struct {
+ sync.Mutex
+ val conn.Endpoint
+ clearSrcOnTx bool // signal to val.ClearSrc() prior to next packet transmission
+ disableRoaming bool
+ }
timers struct {
retransmitHandshake *Timer
@@ -74,8 +77,6 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
// create peer
peer := new(Peer)
- peer.Lock()
- defer peer.Unlock()
peer.cookieGenerator.Init(pk)
peer.device = device
@@ -97,7 +98,11 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
handshake.mutex.Unlock()
// reset endpoint
- peer.endpoint = nil
+ peer.endpoint.Lock()
+ peer.endpoint.val = nil
+ peer.endpoint.disableRoaming = false
+ peer.endpoint.clearSrcOnTx = false
+ peer.endpoint.Unlock()
// init timers
peer.timersInit()
@@ -116,14 +121,19 @@ func (peer *Peer) SendBuffers(buffers [][]byte) error {
return nil
}
- peer.RLock()
- defer peer.RUnlock()
-
- if peer.endpoint == nil {
+ peer.endpoint.Lock()
+ endpoint := peer.endpoint.val
+ if endpoint == nil {
+ peer.endpoint.Unlock()
return errors.New("no known endpoint for peer")
}
+ if peer.endpoint.clearSrcOnTx {
+ endpoint.ClearSrc()
+ peer.endpoint.clearSrcOnTx = false
+ }
+ peer.endpoint.Unlock()
- err := peer.device.net.bind.Send(buffers, peer.endpoint)
+ err := peer.device.net.bind.Send(buffers, endpoint)
if err == nil {
var totalLen uint64
for _, b := range buffers {
@@ -267,10 +277,20 @@ func (peer *Peer) Stop() {
}
func (peer *Peer) SetEndpointFromPacket(endpoint conn.Endpoint) {
- if peer.disableRoaming {
+ peer.endpoint.Lock()
+ defer peer.endpoint.Unlock()
+ if peer.endpoint.disableRoaming {
return
}
- peer.Lock()
- peer.endpoint = endpoint
- peer.Unlock()
+ peer.endpoint.clearSrcOnTx = false
+ peer.endpoint.val = endpoint
+}
+
+func (peer *Peer) markEndpointSrcForClearing() {
+ peer.endpoint.Lock()
+ defer peer.endpoint.Unlock()
+ if peer.endpoint.val == nil {
+ return
+ }
+ peer.endpoint.clearSrcOnTx = true
}
diff --git a/device/sticky_linux.go b/device/sticky_linux.go
index f9230f8..6057ff1 100644
--- a/device/sticky_linux.go
+++ b/device/sticky_linux.go
@@ -110,17 +110,17 @@ func (device *Device) routineRouteListener(bind conn.Bind, netlinkSock int, netl
if !ok {
break
}
- pePtr.peer.Lock()
- if &pePtr.peer.endpoint != pePtr.endpoint {
- pePtr.peer.Unlock()
+ pePtr.peer.endpoint.Lock()
+ if &pePtr.peer.endpoint.val != pePtr.endpoint {
+ pePtr.peer.endpoint.Unlock()
break
}
- if uint32(pePtr.peer.endpoint.(*conn.StdNetEndpoint).SrcIfidx()) == ifidx {
- pePtr.peer.Unlock()
+ if uint32(pePtr.peer.endpoint.val.(*conn.StdNetEndpoint).SrcIfidx()) == ifidx {
+ pePtr.peer.endpoint.Unlock()
break
}
- pePtr.peer.endpoint.(*conn.StdNetEndpoint).ClearSrc()
- pePtr.peer.Unlock()
+ pePtr.peer.endpoint.clearSrcOnTx = true
+ pePtr.peer.endpoint.Unlock()
}
attr = attr[attrhdr.Len:]
}
@@ -134,18 +134,18 @@ func (device *Device) routineRouteListener(bind conn.Bind, netlinkSock int, netl
device.peers.RLock()
i := uint32(1)
for _, peer := range device.peers.keyMap {
- peer.RLock()
- if peer.endpoint == nil {
- peer.RUnlock()
+ peer.endpoint.Lock()
+ if peer.endpoint.val == nil {
+ peer.endpoint.Unlock()
continue
}
- nativeEP, _ := peer.endpoint.(*conn.StdNetEndpoint)
+ nativeEP, _ := peer.endpoint.val.(*conn.StdNetEndpoint)
if nativeEP == nil {
- peer.RUnlock()
+ peer.endpoint.Unlock()
continue
}
if nativeEP.DstIP().Is6() || nativeEP.SrcIfidx() == 0 {
- peer.RUnlock()
+ peer.endpoint.Unlock()
break
}
nlmsg := struct {
@@ -188,10 +188,10 @@ func (device *Device) routineRouteListener(bind conn.Bind, netlinkSock int, netl
reqPeerLock.Lock()
reqPeer[i] = peerEndpointPtr{
peer: peer,
- endpoint: &peer.endpoint,
+ endpoint: &peer.endpoint.val,
}
reqPeerLock.Unlock()
- peer.RUnlock()
+ peer.endpoint.Unlock()
i++
_, err := netlinkCancel.Write((*[unsafe.Sizeof(nlmsg)]byte)(unsafe.Pointer(&nlmsg))[:])
if err != nil {
diff --git a/device/timers.go b/device/timers.go
index e28732c..d4a4ed4 100644
--- a/device/timers.go
+++ b/device/timers.go
@@ -100,11 +100,7 @@ func expiredRetransmitHandshake(peer *Peer) {
peer.device.log.Verbosef("%s - Handshake did not complete after %d seconds, retrying (try %d)", peer, int(RekeyTimeout.Seconds()), peer.timers.handshakeAttempts.Load()+1)
/* We clear the endpoint address src address, in case this is the cause of trouble. */
- peer.Lock()
- if peer.endpoint != nil {
- peer.endpoint.ClearSrc()
- }
- peer.Unlock()
+ peer.markEndpointSrcForClearing()
peer.SendHandshakeInitiation(true)
}
@@ -123,11 +119,7 @@ func expiredSendKeepalive(peer *Peer) {
func expiredNewHandshake(peer *Peer) {
peer.device.log.Verbosef("%s - Retrying handshake because we stopped hearing back after %d seconds", peer, int((KeepaliveTimeout + RekeyTimeout).Seconds()))
/* We clear the endpoint address src address, in case this is the cause of trouble. */
- peer.Lock()
- if peer.endpoint != nil {
- peer.endpoint.ClearSrc()
- }
- peer.Unlock()
+ peer.markEndpointSrcForClearing()
peer.SendHandshakeInitiation(false)
}
diff --git a/device/uapi.go b/device/uapi.go
index 617dcd3..d81dae3 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -99,33 +99,31 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
for _, peer := range device.peers.keyMap {
// Serialize peer state.
- // Do the work in an anonymous function so that we can use defer.
- func() {
- peer.RLock()
- defer peer.RUnlock()
+ peer.handshake.mutex.RLock()
+ keyf("public_key", (*[32]byte)(&peer.handshake.remoteStatic))
+ keyf("preshared_key", (*[32]byte)(&peer.handshake.presharedKey))
+ peer.handshake.mutex.RUnlock()
+ sendf("protocol_version=1")
+ peer.endpoint.Lock()
+ if peer.endpoint.val != nil {
+ sendf("endpoint=%s", peer.endpoint.val.DstToString())
+ }
+ peer.endpoint.Unlock()
- keyf("public_key", (*[32]byte)(&peer.handshake.remoteStatic))
- keyf("preshared_key", (*[32]byte)(&peer.handshake.presharedKey))
- sendf("protocol_version=1")
- if peer.endpoint != nil {
- sendf("endpoint=%s", peer.endpoint.DstToString())
- }
+ nano := peer.lastHandshakeNano.Load()
+ secs := nano / time.Second.Nanoseconds()
+ nano %= time.Second.Nanoseconds()
- nano := peer.lastHandshakeNano.Load()
- secs := nano / time.Second.Nanoseconds()
- nano %= time.Second.Nanoseconds()
+ sendf("last_handshake_time_sec=%d", secs)
+ sendf("last_handshake_time_nsec=%d", nano)
+ sendf("tx_bytes=%d", peer.txBytes.Load())
+ sendf("rx_bytes=%d", peer.rxBytes.Load())
+ sendf("persistent_keepalive_interval=%d", peer.persistentKeepaliveInterval.Load())
- sendf("last_handshake_time_sec=%d", secs)
- sendf("last_handshake_time_nsec=%d", nano)
- sendf("tx_bytes=%d", peer.txBytes.Load())
- sendf("rx_bytes=%d", peer.rxBytes.Load())
- sendf("persistent_keepalive_interval=%d", peer.persistentKeepaliveInterval.Load())
-
- device.allowedips.EntriesForPeer(peer, func(prefix netip.Prefix) bool {
- sendf("allowed_ip=%s", prefix.String())
- return true
- })
- }()
+ device.allowedips.EntriesForPeer(peer, func(prefix netip.Prefix) bool {
+ sendf("allowed_ip=%s", prefix.String())
+ return true
+ })
}
}()
@@ -262,7 +260,7 @@ func (peer *ipcSetPeer) handlePostConfig() {
return
}
if peer.created {
- peer.disableRoaming = peer.device.net.brokenRoaming && peer.endpoint != nil
+ peer.endpoint.disableRoaming = peer.device.net.brokenRoaming && peer.endpoint.val != nil
}
if peer.device.isUp() {
peer.Start()
@@ -345,9 +343,9 @@ func (device *Device) handlePeerLine(peer *ipcSetPeer, key, value string) error
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to set endpoint %v: %w", value, err)
}
- peer.Lock()
- defer peer.Unlock()
- peer.endpoint = endpoint
+ peer.endpoint.Lock()
+ defer peer.endpoint.Unlock()
+ peer.endpoint.val = endpoint
case "persistent_keepalive_interval":
device.log.Verbosef("%v - UAPI: Updating persistent keepalive interval", peer.Peer)
From 7c20311b3d30b96576a95fec31f58e4d5e0d3234 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Tue, 7 Nov 2023 15:24:21 -0800
Subject: [PATCH 21/75] device: reduce redundant per-packet overhead in RX path
Peer.RoutineSequentialReceiver() deals with packet vectors and does not
need to perform timer and endpoint operations for every packet in a
given vector. Changing these per-packet operations to per-vector
improves throughput by as much as 10% in some environments.
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
device/receive.go | 21 +++++++++++++++------
1 file changed, 15 insertions(+), 6 deletions(-)
diff --git a/device/receive.go b/device/receive.go
index 4b32dc5..98e2024 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -445,7 +445,9 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
return
}
elemsContainer.Lock()
- for _, elem := range elemsContainer.elems {
+ validTailPacket := -1
+ dataPacketReceived := false
+ for i, elem := range elemsContainer.elems {
if elem.packet == nil {
// decryption failed
continue
@@ -455,21 +457,19 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
continue
}
- peer.SetEndpointFromPacket(elem.endpoint)
+ validTailPacket = i
if peer.ReceivedWithKeypair(elem.keypair) {
+ peer.SetEndpointFromPacket(elem.endpoint)
peer.timersHandshakeComplete()
peer.SendStagedPackets()
}
- peer.keepKeyFreshReceiving()
- peer.timersAnyAuthenticatedPacketTraversal()
- peer.timersAnyAuthenticatedPacketReceived()
peer.rxBytes.Add(uint64(len(elem.packet) + MinMessageSize))
if len(elem.packet) == 0 {
device.log.Verbosef("%v - Receiving keepalive packet", peer)
continue
}
- peer.timersDataReceived()
+ dataPacketReceived = true
switch elem.packet[0] >> 4 {
case 4:
@@ -512,6 +512,15 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
bufs = append(bufs, elem.buffer[:MessageTransportOffsetContent+len(elem.packet)])
}
+ if validTailPacket >= 0 {
+ peer.SetEndpointFromPacket(elemsContainer.elems[validTailPacket].endpoint)
+ peer.keepKeyFreshReceiving()
+ peer.timersAnyAuthenticatedPacketTraversal()
+ peer.timersAnyAuthenticatedPacketReceived()
+ }
+ if dataPacketReceived {
+ peer.timersDataReceived()
+ }
if len(bufs) > 0 {
_, err := device.tun.device.Write(bufs, MessageTransportOffsetContent)
if err != nil && !device.isClosed() {
From 542e565baa776ed4c5c55b73ef9aa38d33d55197 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Mon, 11 Dec 2023 16:35:57 +0100
Subject: [PATCH 22/75] device: do atomic 64-bit add outside of vector loop
Only bother updating the rxBytes counter once we've processed a whole
vector, since additions are atomic.
Signed-off-by: Jason A. Donenfeld
---
device/receive.go | 5 ++++-
1 file changed, 4 insertions(+), 1 deletion(-)
diff --git a/device/receive.go b/device/receive.go
index 98e2024..1ab3e29 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -447,6 +447,7 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
elemsContainer.Lock()
validTailPacket := -1
dataPacketReceived := false
+ rxBytesLen := uint64(0)
for i, elem := range elemsContainer.elems {
if elem.packet == nil {
// decryption failed
@@ -463,7 +464,7 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
peer.timersHandshakeComplete()
peer.SendStagedPackets()
}
- peer.rxBytes.Add(uint64(len(elem.packet) + MinMessageSize))
+ rxBytesLen += uint64(len(elem.packet) + MinMessageSize)
if len(elem.packet) == 0 {
device.log.Verbosef("%v - Receiving keepalive packet", peer)
@@ -512,6 +513,8 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
bufs = append(bufs, elem.buffer[:MessageTransportOffsetContent+len(elem.packet)])
}
+
+ peer.rxBytes.Add(rxBytesLen)
if validTailPacket >= 0 {
peer.SetEndpointFromPacket(elemsContainer.elems[validTailPacket].endpoint)
peer.keepKeyFreshReceiving()
From 12269c2761734b15625017d8565745096325392f Mon Sep 17 00:00:00 2001
From: Martin Basovnik
Date: Fri, 10 Nov 2023 11:10:12 +0100
Subject: [PATCH 23/75] device: fix possible deadlock in close method
There is a possible deadlock in `device.Close()` when you try to close
the device very soon after its start. The problem is that two different
methods acquire the same locks in different order:
1. device.Close()
- device.ipcMutex.Lock()
- device.state.Lock()
2. device.changeState(deviceState)
- device.state.Lock()
- device.ipcMutex.Lock()
Reproducer:
func TestDevice_deadlock(t *testing.T) {
d := randDevice(t)
d.Close()
}
Problem:
$ go clean -testcache && go test -race -timeout 3s -run TestDevice_deadlock ./device | grep -A 10 sync.runtime_SemacquireMutex
sync.runtime_SemacquireMutex(0xc000117d20?, 0x94?, 0x0?)
/usr/local/opt/go/libexec/src/runtime/sema.go:77 +0x25
sync.(*Mutex).lockSlow(0xc000130518)
/usr/local/opt/go/libexec/src/sync/mutex.go:171 +0x213
sync.(*Mutex).Lock(0xc000130518)
/usr/local/opt/go/libexec/src/sync/mutex.go:90 +0x55
golang.zx2c4.com/wireguard/device.(*Device).Close(0xc000130500)
/Users/martin.basovnik/git/basovnik/wireguard-go/device/device.go:373 +0xb6
golang.zx2c4.com/wireguard/device.TestDevice_deadlock(0x0?)
/Users/martin.basovnik/git/basovnik/wireguard-go/device/device_test.go:480 +0x2c
testing.tRunner(0xc00014c000, 0x131d7b0)
--
sync.runtime_SemacquireMutex(0xc000130564?, 0x60?, 0xc000130548?)
/usr/local/opt/go/libexec/src/runtime/sema.go:77 +0x25
sync.(*Mutex).lockSlow(0xc000130750)
/usr/local/opt/go/libexec/src/sync/mutex.go:171 +0x213
sync.(*Mutex).Lock(0xc000130750)
/usr/local/opt/go/libexec/src/sync/mutex.go:90 +0x55
sync.(*RWMutex).Lock(0xc000130750)
/usr/local/opt/go/libexec/src/sync/rwmutex.go:147 +0x45
golang.zx2c4.com/wireguard/device.(*Device).upLocked(0xc000130500)
/Users/martin.basovnik/git/basovnik/wireguard-go/device/device.go:179 +0x72
golang.zx2c4.com/wireguard/device.(*Device).changeState(0xc000130500, 0x1)
Signed-off-by: Martin Basovnik
Signed-off-by: Jason A. Donenfeld
---
device/device.go | 4 ++--
1 file changed, 2 insertions(+), 2 deletions(-)
diff --git a/device/device.go b/device/device.go
index ca26d00..83c33ee 100644
--- a/device/device.go
+++ b/device/device.go
@@ -368,10 +368,10 @@ func (device *Device) RemoveAllPeers() {
}
func (device *Device) Close() {
- device.ipcMutex.Lock()
- defer device.ipcMutex.Unlock()
device.state.Lock()
defer device.state.Unlock()
+ device.ipcMutex.Lock()
+ defer device.ipcMutex.Unlock()
if device.isClosed() {
return
}
From 015e11875d52955de1d627142159535a6001fe1c Mon Sep 17 00:00:00 2001
From: tiaga
Date: Tue, 19 Dec 2023 18:46:04 +0700
Subject: [PATCH 24/75] Add Dockerfile
Build Docker image with the corresponding wg-tools version.
---
Dockerfile | 20 ++++++++++++++++++++
1 file changed, 20 insertions(+)
create mode 100644 Dockerfile
diff --git a/Dockerfile b/Dockerfile
new file mode 100644
index 0000000..a0d63da
--- /dev/null
+++ b/Dockerfile
@@ -0,0 +1,20 @@
+FROM golang:1.20 as awg
+COPY . /awg
+WORKDIR /awg
+RUN go mod download && \
+ go mod verify && \
+ go build -ldflags '-linkmode external -extldflags "-fno-PIC -static"' -v -o /usr/bin
+
+FROM alpine:3.15 as awg-tools
+ARG AWGTOOLS_RELEASE="1.0.20231215"
+RUN apk --no-cache add linux-headers build-base bash && \
+ wget https://github.com/amnezia-vpn/amnezia-wg-tools/archive/refs/tags/v${AWGTOOLS_RELEASE}.zip && \
+ unzip v${AWGTOOLS_RELEASE}.zip && \
+ cd amnezia-wg-tools-${AWGTOOLS_RELEASE}/src && \
+ make -e LDFLAGS=-static && \
+ make install
+
+FROM alpine:3.15
+RUN apk --no-cache add iproute2 bash
+COPY --from=awg /usr/bin/amnezia-wg /usr/bin/wireguard-go
+COPY --from=awg-tools /usr/bin/wg /usr/bin/wg-quick /usr/bin/
From e5f355e843a71a0492b9201884f028a01197473b Mon Sep 17 00:00:00 2001
From: Iurii Egorov
Date: Sun, 14 Jan 2024 18:22:02 +0300
Subject: [PATCH 25/75] Fix incorrect configuration handling for zero-valued Jc
---
device/send.go | 11 +++++++----
1 file changed, 7 insertions(+), 4 deletions(-)
diff --git a/device/send.go b/device/send.go
index 8191f07..db4e1a1 100644
--- a/device/send.go
+++ b/device/send.go
@@ -137,10 +137,13 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
return err
}
- err = peer.SendBuffers(junks)
- if err != nil {
- peer.device.log.Errorf("%v - Failed to send junk packets: %v", peer, err)
- return err
+ if len(junks) > 0 {
+ err = peer.SendBuffers(junks)
+
+ if err != nil {
+ peer.device.log.Errorf("%v - Failed to send junk packets: %v", peer, err)
+ return err
+ }
}
peer.device.aSecMux.RLock()
From e3c9ec801293387e5fb761c078a11a20cd9a6a5c Mon Sep 17 00:00:00 2001
From: Iurii Egorov
Date: Fri, 19 Jan 2024 15:08:27 +0300
Subject: [PATCH 26/75] Naming unify
---
.gitignore | 2 +-
Dockerfile | 8 ++++----
Makefile | 10 +++++-----
README.md | 14 +++++++-------
conn/bind_windows.go | 2 +-
conn/bindtest/bindtest.go | 2 +-
device/bind_test.go | 2 +-
device/device.go | 10 +++++-----
device/device_test.go | 18 ++++++++++--------
device/keypair.go | 2 +-
device/noise-protocol.go | 2 +-
device/noise_test.go | 4 ++--
device/peer.go | 2 +-
device/queueconstants_android.go | 2 +-
device/queueconstants_default.go | 2 +-
device/receive.go | 4 ++--
device/send.go | 4 ++--
device/sticky_default.go | 4 ++--
device/sticky_linux.go | 4 ++--
device/tun.go | 2 +-
device/uapi.go | 2 +-
go.mod | 2 +-
ipc/namedpipe/namedpipe_test.go | 2 +-
ipc/uapi_linux.go | 2 +-
ipc/uapi_unix.go | 2 +-
ipc/uapi_windows.go | 2 +-
main.go | 8 ++++----
main_windows.go | 8 ++++----
tun/netstack/examples/http_client.go | 6 +++---
tun/netstack/examples/http_server.go | 6 +++---
tun/netstack/examples/ping_client.go | 6 +++---
tun/netstack/tun.go | 2 +-
tun/offload_linux.go | 2 +-
tun/offload_linux_test.go | 2 +-
tun/tun_linux.go | 4 ++--
tun/tuntest/tuntest.go | 2 +-
36 files changed, 80 insertions(+), 78 deletions(-)
diff --git a/.gitignore b/.gitignore
index 71549f4..c6bbd9c 100644
--- a/.gitignore
+++ b/.gitignore
@@ -1 +1 @@
-wireguard-go
\ No newline at end of file
+amneziawg-go
\ No newline at end of file
diff --git a/Dockerfile b/Dockerfile
index a0d63da..cbf05b4 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -8,13 +8,13 @@ RUN go mod download && \
FROM alpine:3.15 as awg-tools
ARG AWGTOOLS_RELEASE="1.0.20231215"
RUN apk --no-cache add linux-headers build-base bash && \
- wget https://github.com/amnezia-vpn/amnezia-wg-tools/archive/refs/tags/v${AWGTOOLS_RELEASE}.zip && \
+ wget https://github.com/amnezia-vpn/amneziawg-tools/archive/refs/tags/v${AWGTOOLS_RELEASE}.zip && \
unzip v${AWGTOOLS_RELEASE}.zip && \
- cd amnezia-wg-tools-${AWGTOOLS_RELEASE}/src && \
+ cd amneziawg-tools-${AWGTOOLS_RELEASE}/src && \
make -e LDFLAGS=-static && \
make install
FROM alpine:3.15
RUN apk --no-cache add iproute2 bash
-COPY --from=awg /usr/bin/amnezia-wg /usr/bin/wireguard-go
-COPY --from=awg-tools /usr/bin/wg /usr/bin/wg-quick /usr/bin/
+COPY --from=awg /usr/bin/amneziawg-go /usr/bin/amneziawg-go
+COPY --from=awg-tools /usr/bin/awg /usr/bin/awg-quick /usr/bin/
diff --git a/Makefile b/Makefile
index 3f6e407..4087cba 100644
--- a/Makefile
+++ b/Makefile
@@ -14,18 +14,18 @@ generate-version-and-build:
[ "$$(cat version.go 2>/dev/null)" != "$$ver" ] && \
echo "$$ver" > version.go && \
git update-index --assume-unchanged version.go || true
- @$(MAKE) wireguard-go
+ @$(MAKE) amneziawg-go
-wireguard-go: $(wildcard *.go) $(wildcard */*.go)
+amneziawg-go: $(wildcard *.go) $(wildcard */*.go)
go build -v -o "$@"
-install: wireguard-go
- @install -v -d "$(DESTDIR)$(BINDIR)" && install -v -m 0755 "$<" "$(DESTDIR)$(BINDIR)/wireguard-go"
+install: amneziawg-go
+ @install -v -d "$(DESTDIR)$(BINDIR)" && install -v -m 0755 "$<" "$(DESTDIR)$(BINDIR)/amneziawg-go"
test:
go test ./...
clean:
- rm -f wireguard-go
+ rm -f amneziawg-go
.PHONY: all clean test install generate-version-and-build
diff --git a/README.md b/README.md
index 717c4c5..ab6f62b 100644
--- a/README.md
+++ b/README.md
@@ -11,17 +11,17 @@ As a result, AmneziaWG maintains high performance while adding an extra layer of
Simply run:
```
-$ amnezia-wg wg0
+$ amneziawg-go wg0
```
-This will create an interface and fork into the background. To remove the interface, use the usual `ip link del wg0`, or if your system does not support removing interfaces directly, you may instead remove the control socket via `rm -f /var/run/wireguard/wg0.sock`, which will result in wireguard-go shutting down.
+This will create an interface and fork into the background. To remove the interface, use the usual `ip link del wg0`, or if your system does not support removing interfaces directly, you may instead remove the control socket via `rm -f /var/run/amneziawg/wg0.sock`, which will result in amneziawg-go shutting down.
-To run amnezia-wg without forking to the background, pass `-f` or `--foreground`:
+To run amneziawg-go without forking to the background, pass `-f` or `--foreground`:
```
-$ amnezia-wg -f wg0
+$ amneziawg-go -f wg0
```
-When an interface is running, you may use [`amnezia-wg-tools `](https://github.com/amnezia-vpn/amnezia-wg-tools) to configure it, as well as the usual `ip(8)` and `ifconfig(8)` commands.
+When an interface is running, you may use [`amnezia-wg-tools `](https://github.com/amnezia-vpn/amneziawg-go-tools) to configure it, as well as the usual `ip(8)` and `ifconfig(8)` commands.
To run with more logging you may set the environment variable `LOG_LEVEL=debug`.
@@ -46,7 +46,7 @@ This runs on Windows, you should use it from [awg-windows](https://github.com/am
This requires an installation of the latest version of [Go](https://go.dev/).
```
-$ git clone https://github.com/amnezia-vpn/amnezia-wg
-$ cd amnezia-wg
+$ git clone https://github.com/amnezia-vpn/amneziawg-go
+$ cd amneziawg-go
$ make
```
diff --git a/conn/bind_windows.go b/conn/bind_windows.go
index 9bad0ee..6cfa099 100644
--- a/conn/bind_windows.go
+++ b/conn/bind_windows.go
@@ -17,7 +17,7 @@ import (
"golang.org/x/sys/windows"
- "github.com/amnezia-vpn/amnezia-wg/conn/winrio"
+ "github.com/amnezia-vpn/amneziawg-go/conn/winrio"
)
const (
diff --git a/conn/bindtest/bindtest.go b/conn/bindtest/bindtest.go
index 713c371..42b0bb7 100644
--- a/conn/bindtest/bindtest.go
+++ b/conn/bindtest/bindtest.go
@@ -12,7 +12,7 @@ import (
"net/netip"
"os"
- "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
)
type ChannelBind struct {
diff --git a/device/bind_test.go b/device/bind_test.go
index eae36c2..34d1c4a 100644
--- a/device/bind_test.go
+++ b/device/bind_test.go
@@ -8,7 +8,7 @@ package device
import (
"errors"
- "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
)
type DummyDatagram struct {
diff --git a/device/device.go b/device/device.go
index eded424..a9d6281 100644
--- a/device/device.go
+++ b/device/device.go
@@ -11,11 +11,11 @@ import (
"sync/atomic"
"time"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/ipc"
- "github.com/amnezia-vpn/amnezia-wg/ratelimiter"
- "github.com/amnezia-vpn/amnezia-wg/rwcancel"
- "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/ipc"
+ "github.com/amnezia-vpn/amneziawg-go/ratelimiter"
+ "github.com/amnezia-vpn/amneziawg-go/rwcancel"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
"github.com/tevino/abool/v2"
)
diff --git a/device/device_test.go b/device/device_test.go
index afa1dc3..e6664a6 100644
--- a/device/device_test.go
+++ b/device/device_test.go
@@ -20,10 +20,10 @@ import (
"testing"
"time"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/conn/bindtest"
- "github.com/amnezia-vpn/amnezia-wg/tun"
- "github.com/amnezia-vpn/amnezia-wg/tun/tuntest"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/conn/bindtest"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
+ "github.com/amnezia-vpn/amneziawg-go/tun/tuntest"
)
// uapiCfg returns a string that contains cfg formatted use with IpcSet.
@@ -237,7 +237,7 @@ func genTestPair(
if _, ok := tb.(*testing.B); ok && !testing.Verbose() {
level = LogLevelError
}
- p.dev = NewDevice(p.tun.TUN(),binds[i],NewLogger(level, fmt.Sprintf("dev%d: ", i)))
+ p.dev = NewDevice(p.tun.TUN(), binds[i], NewLogger(level, fmt.Sprintf("dev%d: ", i)))
if err := p.dev.IpcSet(cfg[i]); err != nil {
tb.Errorf("failed to configure device %d: %v", i, err)
p.dev.Close()
@@ -294,7 +294,7 @@ func TestUpDown(t *testing.T) {
pair := genTestPair(t, false, false)
for i := range pair {
for k := range pair[i].dev.peers.keyMap {
- pair[i].dev.IpcSet(fmt.Sprintf("public_key=%s\npersistent_keepalive_interval=1\n",hex.EncodeToString(k[:])))
+ pair[i].dev.IpcSet(fmt.Sprintf("public_key=%s\npersistent_keepalive_interval=1\n", hex.EncodeToString(k[:])))
}
}
var wg sync.WaitGroup
@@ -513,7 +513,7 @@ func (b *fakeBindSized) Open(
func (b *fakeBindSized) Close() error { return nil }
-func (b *fakeBindSized) SetMark(mark uint32) error {return nil }
+func (b *fakeBindSized) SetMark(mark uint32) error { return nil }
func (b *fakeBindSized) Send(bufs [][]byte, ep conn.Endpoint) error { return nil }
@@ -527,7 +527,9 @@ type fakeTUNDeviceSized struct {
func (t *fakeTUNDeviceSized) File() *os.File { return nil }
-func (t *fakeTUNDeviceSized) Read(bufs [][]byte, sizes []int, offset int) (n int, err error) { return 0, nil }
+func (t *fakeTUNDeviceSized) Read(bufs [][]byte, sizes []int, offset int) (n int, err error) {
+ return 0, nil
+}
func (t *fakeTUNDeviceSized) Write(bufs [][]byte, offset int) (int, error) { return 0, nil }
diff --git a/device/keypair.go b/device/keypair.go
index 73e69af..cc2941a 100644
--- a/device/keypair.go
+++ b/device/keypair.go
@@ -11,7 +11,7 @@ import (
"sync/atomic"
"time"
- "github.com/amnezia-vpn/amnezia-wg/replay"
+ "github.com/amnezia-vpn/amneziawg-go/replay"
)
/* Due to limitations in Go and /x/crypto there is currently
diff --git a/device/noise-protocol.go b/device/noise-protocol.go
index 75c1d87..1289249 100644
--- a/device/noise-protocol.go
+++ b/device/noise-protocol.go
@@ -15,7 +15,7 @@ import (
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/crypto/poly1305"
- "github.com/amnezia-vpn/amnezia-wg/tai64n"
+ "github.com/amnezia-vpn/amneziawg-go/tai64n"
)
type handshakeState int
diff --git a/device/noise_test.go b/device/noise_test.go
index 2363365..075b6d3 100644
--- a/device/noise_test.go
+++ b/device/noise_test.go
@@ -10,8 +10,8 @@ import (
"encoding/binary"
"testing"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/tun/tuntest"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/tun/tuntest"
)
func TestCurveWrappers(t *testing.T) {
diff --git a/device/peer.go b/device/peer.go
index 98bc0ec..5bc8ca4 100644
--- a/device/peer.go
+++ b/device/peer.go
@@ -12,7 +12,7 @@ import (
"sync/atomic"
"time"
- "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
)
type Peer struct {
diff --git a/device/queueconstants_android.go b/device/queueconstants_android.go
index d29dbc8..1bff95a 100644
--- a/device/queueconstants_android.go
+++ b/device/queueconstants_android.go
@@ -5,7 +5,7 @@
package device
-import "github.com/amnezia-vpn/amnezia-wg/conn"
+import "github.com/amnezia-vpn/amneziawg-go/conn"
/* Reduce memory consumption for Android */
diff --git a/device/queueconstants_default.go b/device/queueconstants_default.go
index 4ee2966..0061b63 100644
--- a/device/queueconstants_default.go
+++ b/device/queueconstants_default.go
@@ -7,7 +7,7 @@
package device
-import "github.com/amnezia-vpn/amnezia-wg/conn"
+import "github.com/amnezia-vpn/amneziawg-go/conn"
const (
QueueStagedSize = conn.IdealBatchSize
diff --git a/device/receive.go b/device/receive.go
index 06d092e..66c1a32 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -13,7 +13,7 @@ import (
"sync"
"time"
- "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
@@ -145,7 +145,7 @@ func (device *Device) RoutineReceiveIncoming(
junkSize := msgTypeToJunkSize[assumedMsgType]
// transport size can align with other header types;
// making sure we have the right msgType
- msgType = binary.LittleEndian.Uint32(packet[junkSize:junkSize+4])
+ msgType = binary.LittleEndian.Uint32(packet[junkSize : junkSize+4])
if msgType == assumedMsgType {
packet = packet[junkSize:]
} else {
diff --git a/device/send.go b/device/send.go
index db4e1a1..1b4406d 100644
--- a/device/send.go
+++ b/device/send.go
@@ -15,8 +15,8 @@ import (
"sync"
"time"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
"golang.org/x/crypto/chacha20poly1305"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
diff --git a/device/sticky_default.go b/device/sticky_default.go
index 940702c..da776e8 100644
--- a/device/sticky_default.go
+++ b/device/sticky_default.go
@@ -3,8 +3,8 @@
package device
import (
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/rwcancel"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/rwcancel"
)
func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) {
diff --git a/device/sticky_linux.go b/device/sticky_linux.go
index 070986c..63164a7 100644
--- a/device/sticky_linux.go
+++ b/device/sticky_linux.go
@@ -20,8 +20,8 @@ import (
"golang.org/x/sys/unix"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/rwcancel"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/rwcancel"
)
func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) {
diff --git a/device/tun.go b/device/tun.go
index efc543d..600a5e5 100644
--- a/device/tun.go
+++ b/device/tun.go
@@ -8,7 +8,7 @@ package device
import (
"fmt"
- "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
)
const DefaultMTU = 1420
diff --git a/device/uapi.go b/device/uapi.go
index 02a9fb7..777bdda 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -18,7 +18,7 @@ import (
"sync"
"time"
- "github.com/amnezia-vpn/amnezia-wg/ipc"
+ "github.com/amnezia-vpn/amneziawg-go/ipc"
)
type IPCError struct {
diff --git a/go.mod b/go.mod
index 97cba1c..2df4282 100644
--- a/go.mod
+++ b/go.mod
@@ -1,4 +1,4 @@
-module github.com/amnezia-vpn/amnezia-wg
+module github.com/amnezia-vpn/amneziawg-go
go 1.20
diff --git a/ipc/namedpipe/namedpipe_test.go b/ipc/namedpipe/namedpipe_test.go
index d4799e1..9f9cd6a 100644
--- a/ipc/namedpipe/namedpipe_test.go
+++ b/ipc/namedpipe/namedpipe_test.go
@@ -20,7 +20,7 @@ import (
"testing"
"time"
- "github.com/amnezia-vpn/amnezia-wg/ipc/namedpipe"
+ "github.com/amnezia-vpn/amneziawg-go/ipc/namedpipe"
"golang.org/x/sys/windows"
)
diff --git a/ipc/uapi_linux.go b/ipc/uapi_linux.go
index 721c404..9738aea 100644
--- a/ipc/uapi_linux.go
+++ b/ipc/uapi_linux.go
@@ -9,7 +9,7 @@ import (
"net"
"os"
- "github.com/amnezia-vpn/amnezia-wg/rwcancel"
+ "github.com/amnezia-vpn/amneziawg-go/rwcancel"
"golang.org/x/sys/unix"
)
diff --git a/ipc/uapi_unix.go b/ipc/uapi_unix.go
index e67be26..0da452a 100644
--- a/ipc/uapi_unix.go
+++ b/ipc/uapi_unix.go
@@ -26,7 +26,7 @@ const (
// socketDirectory is variable because it is modified by a linker
// flag in wireguard-android.
-var socketDirectory = "/var/run/wireguard"
+var socketDirectory = "/var/run/amneziawg"
func sockPath(iface string) string {
return fmt.Sprintf("%s/%s.sock", socketDirectory, iface)
diff --git a/ipc/uapi_windows.go b/ipc/uapi_windows.go
index 97a4123..bfe7965 100644
--- a/ipc/uapi_windows.go
+++ b/ipc/uapi_windows.go
@@ -8,7 +8,7 @@ package ipc
import (
"net"
- "github.com/amnezia-vpn/amnezia-wg/ipc/namedpipe"
+ "github.com/amnezia-vpn/amneziawg-go/ipc/namedpipe"
"golang.org/x/sys/windows"
)
diff --git a/main.go b/main.go
index ea7ef4e..775372c 100644
--- a/main.go
+++ b/main.go
@@ -14,10 +14,10 @@ import (
"runtime"
"strconv"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/device"
- "github.com/amnezia-vpn/amnezia-wg/ipc"
- "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device"
+ "github.com/amnezia-vpn/amneziawg-go/ipc"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
"golang.org/x/sys/unix"
)
diff --git a/main_windows.go b/main_windows.go
index d00b146..807f6e2 100644
--- a/main_windows.go
+++ b/main_windows.go
@@ -12,11 +12,11 @@ import (
"golang.org/x/sys/windows"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/device"
- "github.com/amnezia-vpn/amnezia-wg/ipc"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device"
+ "github.com/amnezia-vpn/amneziawg-go/ipc"
- "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
)
const (
diff --git a/tun/netstack/examples/http_client.go b/tun/netstack/examples/http_client.go
index ed40904..4c4ea12 100644
--- a/tun/netstack/examples/http_client.go
+++ b/tun/netstack/examples/http_client.go
@@ -13,9 +13,9 @@ import (
"net/http"
"net/netip"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/device"
- "github.com/amnezia-vpn/amnezia-wg/tun/netstack"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device"
+ "github.com/amnezia-vpn/amneziawg-go/tun/netstack"
)
func main() {
diff --git a/tun/netstack/examples/http_server.go b/tun/netstack/examples/http_server.go
index d5e7094..09929e0 100644
--- a/tun/netstack/examples/http_server.go
+++ b/tun/netstack/examples/http_server.go
@@ -14,9 +14,9 @@ import (
"net/http"
"net/netip"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/device"
- "github.com/amnezia-vpn/amnezia-wg/tun/netstack"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device"
+ "github.com/amnezia-vpn/amneziawg-go/tun/netstack"
)
func main() {
diff --git a/tun/netstack/examples/ping_client.go b/tun/netstack/examples/ping_client.go
index 9f917db..d7897b2 100644
--- a/tun/netstack/examples/ping_client.go
+++ b/tun/netstack/examples/ping_client.go
@@ -17,9 +17,9 @@ import (
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/device"
- "github.com/amnezia-vpn/amnezia-wg/tun/netstack"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device"
+ "github.com/amnezia-vpn/amneziawg-go/tun/netstack"
)
func main() {
diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go
index b5e6145..2275173 100644
--- a/tun/netstack/tun.go
+++ b/tun/netstack/tun.go
@@ -22,7 +22,7 @@ import (
"syscall"
"time"
- "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
"golang.org/x/net/dns/dnsmessage"
"gvisor.dev/gvisor/pkg/buffer"
diff --git a/tun/offload_linux.go b/tun/offload_linux.go
index 551f14d..89cf024 100644
--- a/tun/offload_linux.go
+++ b/tun/offload_linux.go
@@ -12,7 +12,7 @@ import (
"io"
"unsafe"
- "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
"golang.org/x/sys/unix"
)
diff --git a/tun/offload_linux_test.go b/tun/offload_linux_test.go
index 71dfba3..a68cd98 100644
--- a/tun/offload_linux_test.go
+++ b/tun/offload_linux_test.go
@@ -9,7 +9,7 @@ import (
"net/netip"
"testing"
- "github.com/amnezia-vpn/amnezia-wg/conn"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header"
diff --git a/tun/tun_linux.go b/tun/tun_linux.go
index d57b167..011e56a 100644
--- a/tun/tun_linux.go
+++ b/tun/tun_linux.go
@@ -17,8 +17,8 @@ import (
"time"
"unsafe"
- "github.com/amnezia-vpn/amnezia-wg/conn"
- "github.com/amnezia-vpn/amnezia-wg/rwcancel"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/rwcancel"
"golang.org/x/sys/unix"
)
diff --git a/tun/tuntest/tuntest.go b/tun/tuntest/tuntest.go
index 7068d9b..f620e0a 100644
--- a/tun/tuntest/tuntest.go
+++ b/tun/tuntest/tuntest.go
@@ -11,7 +11,7 @@ import (
"net/netip"
"os"
- "github.com/amnezia-vpn/amnezia-wg/tun"
+ "github.com/amnezia-vpn/amneziawg-go/tun"
)
func Ping(dst, src netip.Addr) []byte {
From bfeb3954f693dc045db4ac8232ee5176152f0d98 Mon Sep 17 00:00:00 2001
From: tiaga
Date: Fri, 2 Feb 2024 22:56:00 +0700
Subject: [PATCH 27/75] Update Dockerfile
- update Alpine version
- improve `Dockerfile` to use pre-built AmneziaWG tools
---
Dockerfile | 18 ++++++------------
1 file changed, 6 insertions(+), 12 deletions(-)
diff --git a/Dockerfile b/Dockerfile
index cbf05b4..136ebc8 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -5,16 +5,10 @@ RUN go mod download && \
go mod verify && \
go build -ldflags '-linkmode external -extldflags "-fno-PIC -static"' -v -o /usr/bin
-FROM alpine:3.15 as awg-tools
-ARG AWGTOOLS_RELEASE="1.0.20231215"
-RUN apk --no-cache add linux-headers build-base bash && \
- wget https://github.com/amnezia-vpn/amneziawg-tools/archive/refs/tags/v${AWGTOOLS_RELEASE}.zip && \
- unzip v${AWGTOOLS_RELEASE}.zip && \
- cd amneziawg-tools-${AWGTOOLS_RELEASE}/src && \
- make -e LDFLAGS=-static && \
- make install
-
-FROM alpine:3.15
-RUN apk --no-cache add iproute2 bash
+FROM alpine:3.19
+ARG AWGTOOLS_RELEASE="1.0.20240202"
+RUN apk --no-cache add iproute2 bash && \
+ wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
+ unzip alpine-3.19-amneziawg-tools.zip -d /usr/bin/ && \
+ chmod +x /usr/bin/wg /usr/bin/wg-quick
COPY --from=awg /usr/bin/amneziawg-go /usr/bin/amneziawg-go
-COPY --from=awg-tools /usr/bin/awg /usr/bin/awg-quick /usr/bin/
From cbd414dfecfcd711bde99b8696aba58582580732 Mon Sep 17 00:00:00 2001
From: tiaga
Date: Wed, 7 Feb 2024 18:44:59 +0700
Subject: [PATCH 28/75] Add pipeline
Build and push Docker image on a tag push.
---
.github/workflows/build-if-tag.yml | 41 +++++++++++++++++++++++++++++
1 file changed, 41 insertions(+)
create mode 100644 .github/workflows/build-if-tag.yml
diff --git a/ .github/workflows/build-if-tag.yml b/ .github/workflows/build-if-tag.yml
new file mode 100644
index 0000000..4fa0198
--- /dev/null
+++ b/ .github/workflows/build-if-tag.yml
@@ -0,0 +1,41 @@
+name: build-if-tag
+
+on:
+ push:
+ tags:
+ - 'v[0-9]+.[0-9]+.[0-9]+'
+
+env:
+ APP: amneziawg-go
+
+jobs:
+ build:
+ runs-on: ubuntu-latest
+ name: build
+ steps:
+ - name: Checkout
+ uses: actions/checkout@v4
+ with:
+ ref: ${{ github.ref_name }}
+
+ - name: Login to Docker Hub
+ uses: docker/login-action@v3
+ with:
+ username: ${{ secrets.DOCKERHUB_USERNAME }}
+ password: ${{ secrets.DOCKERHUB_TOKEN }}
+
+ - name: Setup metadata
+ uses: docker/metadata-action@v5
+ id: metadata
+ with:
+ images: amneziavpn/${{ env.APP }}
+ tags: type=semver,pattern={{version}}
+
+ - name: Set up Docker Buildx
+ uses: docker/setup-buildx-action@v3
+
+ - name: Build
+ uses: docker/build-push-action@v5
+ with:
+ push: true
+ tags: ${{ steps.metadata.outputs.tags }}
From f0dfb5eaccf52b021cf000a8004266b7faa3881b Mon Sep 17 00:00:00 2001
From: tiaga
Date: Wed, 7 Feb 2024 18:53:55 +0700
Subject: [PATCH 29/75] Fix pipeline
Fix path to GitHub Actions workflow.
---
{ .github => .github}/workflows/build-if-tag.yml | 0
1 file changed, 0 insertions(+), 0 deletions(-)
rename { .github => .github}/workflows/build-if-tag.yml (100%)
diff --git a/ .github/workflows/build-if-tag.yml b/.github/workflows/build-if-tag.yml
similarity index 100%
rename from .github/workflows/build-if-tag.yml
rename to .github/workflows/build-if-tag.yml
From 59101fd202067ae42362dab5494b4aaded3a3456 Mon Sep 17 00:00:00 2001
From: albexk
Date: Sat, 10 Feb 2024 16:02:05 +0300
Subject: [PATCH 30/75] Bump crypto, net, sys modules to the latest versions
---
go.mod | 6 +++---
go.sum | 5 +++++
2 files changed, 8 insertions(+), 3 deletions(-)
diff --git a/go.mod b/go.mod
index 2df4282..33182ee 100644
--- a/go.mod
+++ b/go.mod
@@ -4,9 +4,9 @@ go 1.20
require (
github.com/tevino/abool/v2 v2.1.0
- golang.org/x/crypto v0.13.0
- golang.org/x/net v0.15.0
- golang.org/x/sys v0.12.0
+ golang.org/x/crypto v0.19.0
+ golang.org/x/net v0.21.0
+ golang.org/x/sys v0.17.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259
)
diff --git a/go.sum b/go.sum
index 71c64b6..eb7f470 100644
--- a/go.sum
+++ b/go.sum
@@ -4,10 +4,15 @@ github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
golang.org/x/crypto v0.13.0 h1:mvySKfSWJ+UKUii46M40LOvyWfN0s2U+46/jDd0e6Ck=
golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc=
+golang.org/x/crypto v0.19.0 h1:ENy+Az/9Y1vSrlrvBSyna3PITt4tiZLf7sgCjZBX7Wo=
+golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
golang.org/x/net v0.15.0 h1:ugBLEUaxABaB5AJqW9enI0ACdci2RUd4eP51NTBvuJ8=
golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
+golang.org/x/net v0.21.0 h1:AQyQV4dYCvJ7vGmJyKki9+PBdyvhkSd8EIx/qb0AYv4=
+golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o=
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
+golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 h1:vVKdlvoWBphwdxWKrFZEuM0kGgGLxUOYcY4U/2Vjg44=
golang.org/x/time v0.0.0-20220210224613-90d013bbcef8/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
From 032e33f5776a569168a0bd02b92f2d34fb0345f4 Mon Sep 17 00:00:00 2001
From: albexk
Date: Sat, 10 Feb 2024 17:14:51 +0300
Subject: [PATCH 31/75] Fix Android UDP GRO check
---
conn/controlfns_linux.go | 8 --------
conn/features_linux.go | 6 ++++--
2 files changed, 4 insertions(+), 10 deletions(-)
diff --git a/conn/controlfns_linux.go b/conn/controlfns_linux.go
index f6ab1d2..a2396fe 100644
--- a/conn/controlfns_linux.go
+++ b/conn/controlfns_linux.go
@@ -57,13 +57,5 @@ func init() {
}
return err
},
-
- // Attempt to enable UDP_GRO
- func(network, address string, c syscall.RawConn) error {
- c.Control(func(fd uintptr) {
- _ = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1)
- })
- return nil
- },
)
}
diff --git a/conn/features_linux.go b/conn/features_linux.go
index 8959d93..a6de8c1 100644
--- a/conn/features_linux.go
+++ b/conn/features_linux.go
@@ -19,8 +19,10 @@ func supportsUDPOffload(conn *net.UDPConn) (txOffload, rxOffload bool) {
err = rc.Control(func(fd uintptr) {
_, errSyscall := unix.GetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_SEGMENT)
txOffload = errSyscall == nil
- opt, errSyscall := unix.GetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO)
- rxOffload = errSyscall == nil && opt == 1
+ // getsockopt(IPPROTO_UDP, UDP_GRO) is not supported in android
+ // use setsockopt workaround
+ errSyscall = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1)
+ rxOffload = errSyscall == nil
})
if err != nil {
return false, false
From 6705978fc86f48f6d882535c3180eaa7653f298f Mon Sep 17 00:00:00 2001
From: albexk
Date: Sat, 10 Feb 2024 17:47:33 +0300
Subject: [PATCH 32/75] Add debug udp offload info
---
conn/bind_std.go | 5 +++++
conn/bind_windows.go | 4 ++++
conn/bindtest/bindtest.go | 2 ++
conn/conn.go | 2 ++
device/device.go | 1 +
5 files changed, 14 insertions(+)
diff --git a/conn/bind_std.go b/conn/bind_std.go
index 46df7fd..b416e22 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -298,6 +298,11 @@ func (s *StdNetBind) BatchSize() int {
return 1
}
+func (s *StdNetBind) GetOffloadInfo() string {
+ return fmt.Sprintf("ipv4TxOffload: %v, ipv4RxOffload: %v\nipv6TxOffload: %v, ipv6RxOffload: %v",
+ s.ipv4TxOffload, s.ipv4RxOffload, s.ipv6TxOffload, s.ipv6RxOffload)
+}
+
func (s *StdNetBind) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
diff --git a/conn/bind_windows.go b/conn/bind_windows.go
index 6cfa099..3481f00 100644
--- a/conn/bind_windows.go
+++ b/conn/bind_windows.go
@@ -328,6 +328,10 @@ func (bind *WinRingBind) BatchSize() int {
return 1
}
+func (bind *WinRingBind) GetOffloadInfo() string {
+ return ""
+}
+
func (bind *WinRingBind) SetMark(mark uint32) error {
return nil
}
diff --git a/conn/bindtest/bindtest.go b/conn/bindtest/bindtest.go
index 42b0bb7..0df1420 100644
--- a/conn/bindtest/bindtest.go
+++ b/conn/bindtest/bindtest.go
@@ -91,6 +91,8 @@ func (c *ChannelBind) Close() error {
func (c *ChannelBind) BatchSize() int { return 1 }
+func (c *ChannelBind) GetOffloadInfo() string { return "" }
+
func (c *ChannelBind) SetMark(mark uint32) error { return nil }
func (c *ChannelBind) makeReceiveFunc(ch chan []byte) conn.ReceiveFunc {
diff --git a/conn/conn.go b/conn/conn.go
index a1f57d2..489cb35 100644
--- a/conn/conn.go
+++ b/conn/conn.go
@@ -55,6 +55,8 @@ type Bind interface {
// BatchSize is the number of buffers expected to be passed to
// the ReceiveFuncs, and the maximum expected to be passed to SendBatch.
BatchSize() int
+
+ GetOffloadInfo() string
}
// BindSocketToInterface is implemented by Bind objects that support being
diff --git a/device/device.go b/device/device.go
index a9d6281..24ae1ea 100644
--- a/device/device.go
+++ b/device/device.go
@@ -545,6 +545,7 @@ func (device *Device) BindUpdate() error {
}
device.log.Verbosef("UDP bind has been updated")
+ device.log.Verbosef(netc.bind.GetOffloadInfo())
return nil
}
From 0c347529b8f752fde34b261a510a8d3d597ae75c Mon Sep 17 00:00:00 2001
From: albexk
Date: Mon, 12 Feb 2024 16:27:56 +0300
Subject: [PATCH 33/75] Fix go.sum
---
go.sum | 7 +------
1 file changed, 1 insertion(+), 6 deletions(-)
diff --git a/go.sum b/go.sum
index eb7f470..0e6f733 100644
--- a/go.sum
+++ b/go.sum
@@ -2,16 +2,11 @@ github.com/google/btree v1.0.1 h1:gK4Kx5IaGY9CD5sPJ36FHiBJ6ZXl0kilRiiCj+jdYp4=
github.com/google/btree v1.0.1/go.mod h1:xXMiIv4Fb/0kKde4SpL7qlzvu5cMJDRkFDxJfI9uaxA=
github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
-golang.org/x/crypto v0.13.0 h1:mvySKfSWJ+UKUii46M40LOvyWfN0s2U+46/jDd0e6Ck=
-golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc=
golang.org/x/crypto v0.19.0 h1:ENy+Az/9Y1vSrlrvBSyna3PITt4tiZLf7sgCjZBX7Wo=
golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
-golang.org/x/net v0.15.0 h1:ugBLEUaxABaB5AJqW9enI0ACdci2RUd4eP51NTBvuJ8=
-golang.org/x/net v0.15.0/go.mod h1:idbUs1IY1+zTqbi8yxTbhexhEEk5ur9LInksu6HrEpk=
golang.org/x/net v0.21.0 h1:AQyQV4dYCvJ7vGmJyKki9+PBdyvhkSd8EIx/qb0AYv4=
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
-golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o=
-golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
+golang.org/x/sys v0.17.0 h1:25cE3gD+tdBA7lp7QfhuV+rJiE9YXTcS3VG1SqssI/Y=
golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 h1:vVKdlvoWBphwdxWKrFZEuM0kGgGLxUOYcY4U/2Vjg44=
golang.org/x/time v0.0.0-20220210224613-90d013bbcef8/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
From 9c6b3ff332ab47cfa5613bc3d0f582cc887c6348 Mon Sep 17 00:00:00 2001
From: tiaga
Date: Tue, 13 Feb 2024 21:27:34 +0700
Subject: [PATCH 34/75] Update Dockerfile
- rename `wg` and `wg-quick` to `awg` and `awg-quick` accordingly
- add iptables
- update AmneziaWG tools version
---
Dockerfile | 6 +++---
1 file changed, 3 insertions(+), 3 deletions(-)
diff --git a/Dockerfile b/Dockerfile
index 136ebc8..9d41002 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -6,9 +6,9 @@ RUN go mod download && \
go build -ldflags '-linkmode external -extldflags "-fno-PIC -static"' -v -o /usr/bin
FROM alpine:3.19
-ARG AWGTOOLS_RELEASE="1.0.20240202"
-RUN apk --no-cache add iproute2 bash && \
+ARG AWGTOOLS_RELEASE="1.0.20240213"
+RUN apk --no-cache add iproute2 iptables bash && \
wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
unzip alpine-3.19-amneziawg-tools.zip -d /usr/bin/ && \
- chmod +x /usr/bin/wg /usr/bin/wg-quick
+ chmod +x /usr/bin/awg /usr/bin/awg-quick
COPY --from=awg /usr/bin/amneziawg-go /usr/bin/amneziawg-go
From 92e28a0d14c643f620f945147d826a55b3a082ce Mon Sep 17 00:00:00 2001
From: tiaga
Date: Tue, 13 Feb 2024 21:44:41 +0700
Subject: [PATCH 35/75] Fix Dockerfile
Fix AmneziaWG tools installation.
---
Dockerfile | 3 ++-
1 file changed, 2 insertions(+), 1 deletion(-)
diff --git a/Dockerfile b/Dockerfile
index 9d41002..6586268 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -8,7 +8,8 @@ RUN go mod download && \
FROM alpine:3.19
ARG AWGTOOLS_RELEASE="1.0.20240213"
RUN apk --no-cache add iproute2 iptables bash && \
+ cd /usr/bin/ && \
wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
- unzip alpine-3.19-amneziawg-tools.zip -d /usr/bin/ && \
+ unzip -j alpine-3.19-amneziawg-tools.zip && \
chmod +x /usr/bin/awg /usr/bin/awg-quick
COPY --from=awg /usr/bin/amneziawg-go /usr/bin/amneziawg-go
From 4dddf62e576b5acc703b3fd354cf5e3d2cedfbca Mon Sep 17 00:00:00 2001
From: AlexanderGalkov <143902290+AlexanderGalkov@users.noreply.github.com>
Date: Tue, 20 Feb 2024 20:29:36 +0700
Subject: [PATCH 36/75] Update Dockerfile
add wg and wg-quick symlinks
Signed-off-by: AlexanderGalkov <143902290+AlexanderGalkov@users.noreply.github.com>
---
Dockerfile | 4 +++-
1 file changed, 3 insertions(+), 1 deletion(-)
diff --git a/Dockerfile b/Dockerfile
index 6586268..caeb333 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -11,5 +11,7 @@ RUN apk --no-cache add iproute2 iptables bash && \
cd /usr/bin/ && \
wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
unzip -j alpine-3.19-amneziawg-tools.zip && \
- chmod +x /usr/bin/awg /usr/bin/awg-quick
+ chmod +x /usr/bin/awg /usr/bin/awg-quick && \
+ ln -s /usr/bin/awg /usr/bin/wg && \
+ ln -s /usr/bin/awg-quick /usr/bin/wg-quick
COPY --from=awg /usr/bin/amneziawg-go /usr/bin/amneziawg-go
From 3f0a3bcfa0da940e6ded34e654fdc9186038dd83 Mon Sep 17 00:00:00 2001
From: albexk
Date: Sat, 16 Mar 2024 14:43:32 +0300
Subject: [PATCH 37/75] Fix wg reconnection problem after awg connection
---
device/device.go | 5 +++++
1 file changed, 5 insertions(+)
diff --git a/device/device.go b/device/device.go
index 24ae1ea..21a6546 100644
--- a/device/device.go
+++ b/device/device.go
@@ -562,6 +562,11 @@ func (device *Device) isAdvancedSecurityOn() bool {
func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
if !tempASecCfg.isSet {
+ // restore default values
+ MessageInitiationType = 1
+ MessageResponseType = 2
+ MessageCookieReplyType = 3
+ MessageTransportType = 4
return err
}
From 3ddf952973fac2e4d30f4aa61a7c6dd49bfc8fcb Mon Sep 17 00:00:00 2001
From: RomikB
Date: Sat, 11 May 2024 22:16:22 +0200
Subject: [PATCH 38/75] unsafe rebranding: change pipe name
---
ipc/uapi_windows.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/ipc/uapi_windows.go b/ipc/uapi_windows.go
index bfe7965..31d2a63 100644
--- a/ipc/uapi_windows.go
+++ b/ipc/uapi_windows.go
@@ -62,7 +62,7 @@ func init() {
func UAPIListen(name string) (net.Listener, error) {
listener, err := (&namedpipe.ListenConfig{
SecurityDescriptor: UAPISecurityDescriptor,
- }).Listen(`\\.\pipe\ProtectedPrefix\Administrators\WireGuard\` + name)
+ }).Listen(`\\.\pipe\ProtectedPrefix\Administrators\AmneziaWG\` + name)
if err != nil {
return nil, err
}
From e433d13df6ac99fa317b5ce27d04c7dad307cab3 Mon Sep 17 00:00:00 2001
From: albexk
Date: Wed, 3 Apr 2024 18:42:37 +0300
Subject: [PATCH 39/75] Add disabling UDP GSO when an error occurs due to
inconsistent peer mtu
---
conn/bind_std.go | 2 +-
conn/errors_linux.go | 4 +++-
2 files changed, 4 insertions(+), 2 deletions(-)
diff --git a/conn/bind_std.go b/conn/bind_std.go
index b416e22..ea06cd5 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -336,7 +336,7 @@ type ErrUDPGSODisabled struct {
}
func (e ErrUDPGSODisabled) Error() string {
- return fmt.Sprintf("disabled UDP GSO on %s, NIC(s) may not support checksum offload", e.onLaddr)
+ return fmt.Sprintf("disabled UDP GSO on %s, NIC(s) may not support checksum offload or peer MTU with protocol headers is greater than path MTU", e.onLaddr)
}
func (e ErrUDPGSODisabled) Unwrap() error {
diff --git a/conn/errors_linux.go b/conn/errors_linux.go
index 8e61000..7548a8a 100644
--- a/conn/errors_linux.go
+++ b/conn/errors_linux.go
@@ -20,7 +20,9 @@ func errShouldDisableUDPGSO(err error) bool {
// See:
// https://git.kernel.org/pub/scm/docs/man-pages/man-pages.git/tree/man7/udp.7?id=806eabd74910447f21005160e90957bde4db0183#n228
// https://git.kernel.org/pub/scm/linux/kernel/git/torvalds/linux.git/tree/net/ipv4/udp.c?h=v6.2&id=c9c3395d5e3dcc6daee66c6908354d47bf98cb0c#n942
- return serr.Err == unix.EIO
+ // If gso_size + udp + ip headers > fragment size EINVAL is returned.
+ // It occurs when the peer mtu + wg headers is greater than path mtu.
+ return serr.Err == unix.EIO || serr.Err == unix.EINVAL
}
return false
}
From 77d39ff3b9b1144d0106bd968807c62a033be8dc Mon Sep 17 00:00:00 2001
From: albexk
Date: Wed, 3 Apr 2024 18:45:26 +0300
Subject: [PATCH 40/75] Minor naming changes
---
README.md | 6 +++---
main.go | 22 +++++++++++-----------
main_windows.go | 4 ++--
3 files changed, 16 insertions(+), 16 deletions(-)
diff --git a/README.md b/README.md
index ab6f62b..853d318 100644
--- a/README.md
+++ b/README.md
@@ -21,7 +21,7 @@ To run amneziawg-go without forking to the background, pass `-f` or `--foregroun
```
$ amneziawg-go -f wg0
```
-When an interface is running, you may use [`amnezia-wg-tools `](https://github.com/amnezia-vpn/amneziawg-go-tools) to configure it, as well as the usual `ip(8)` and `ifconfig(8)` commands.
+When an interface is running, you may use [`amneziawg-tools `](https://github.com/amnezia-vpn/amneziawg-tools) to configure it, as well as the usual `ip(8)` and `ifconfig(8)` commands.
To run with more logging you may set the environment variable `LOG_LEVEL=debug`.
@@ -34,11 +34,11 @@ This will run on Linux; you should run amnezia-wg instead of using default linux
### macOS
This runs on macOS using the utun driver. It does not yet support sticky sockets, and won't support fwmarks because of Darwin limitations. Since the utun driver cannot have arbitrary interface names, you must either use `utun[0-9]+` for an explicit interface name or `utun` to have the kernel select one for you. If you choose `utun` as the interface name, and the environment variable `WG_TUN_NAME_FILE` is defined, then the actual name of the interface chosen by the kernel is written to the file specified by that variable.
-This runs on MacOS, you should use it from [awg-apple](https://github.com/amnezia-vpn/awg-apple)
+This runs on MacOS, you should use it from [amneziawg-apple](https://github.com/amnezia-vpn/amneziawg-apple)
### Windows
-This runs on Windows, you should use it from [awg-windows](https://github.com/amnezia-vpn/awg-windows), which uses this as a module.
+This runs on Windows, you should use it from [amneziawg-windows](https://github.com/amnezia-vpn/amneziawg-windows), which uses this as a module.
## Building
diff --git a/main.go b/main.go
index 775372c..c17c405 100644
--- a/main.go
+++ b/main.go
@@ -46,20 +46,20 @@ func warning() {
return
}
- fmt.Fprintln(os.Stderr, "┌──────────────────────────────────────────────────────┐")
- fmt.Fprintln(os.Stderr, "│ │")
- fmt.Fprintln(os.Stderr, "│ Running wireguard-go is not required because this │")
- fmt.Fprintln(os.Stderr, "│ kernel has first class support for WireGuard. For │")
- fmt.Fprintln(os.Stderr, "│ information on installing the kernel module, │")
- fmt.Fprintln(os.Stderr, "│ please visit: │")
- fmt.Fprintln(os.Stderr, "│ https://www.wireguard.com/install/ │")
- fmt.Fprintln(os.Stderr, "│ │")
- fmt.Fprintln(os.Stderr, "└──────────────────────────────────────────────────────┘")
+ fmt.Fprintln(os.Stderr, "┌──────────────────────────────────────────────────────────────┐")
+ fmt.Fprintln(os.Stderr, "│ │")
+ fmt.Fprintln(os.Stderr, "│ Running amneziawg-go is not required because this │")
+ fmt.Fprintln(os.Stderr, "│ kernel has first class support for AmneziaWG. For │")
+ fmt.Fprintln(os.Stderr, "│ information on installing the kernel module, │")
+ fmt.Fprintln(os.Stderr, "│ please visit: │")
+ fmt.Fprintln(os.Stderr, "| https://github.com/amnezia-vpn/amneziawg-linux-kernel-module │")
+ fmt.Fprintln(os.Stderr, "│ │")
+ fmt.Fprintln(os.Stderr, "└──────────────────────────────────────────────────────────────┘")
}
func main() {
if len(os.Args) == 2 && os.Args[1] == "--version" {
- fmt.Printf("wireguard-go v%s\n\nUserspace WireGuard daemon for %s-%s.\nInformation available at https://www.wireguard.com.\nCopyright (C) Jason A. Donenfeld .\n", Version, runtime.GOOS, runtime.GOARCH)
+ fmt.Printf("amneziawg-go v%s\n\nUserspace AmneziaWG daemon for %s-%s.\nInformation available at https://amnezia.org\n", Version, runtime.GOOS, runtime.GOARCH)
return
}
@@ -145,7 +145,7 @@ func main() {
fmt.Sprintf("(%s) ", interfaceName),
)
- logger.Verbosef("Starting wireguard-go version %s", Version)
+ logger.Verbosef("Starting amneziawg-go version %s", Version)
if err != nil {
logger.Errorf("Failed to create TUN device: %v", err)
diff --git a/main_windows.go b/main_windows.go
index 807f6e2..bbfa690 100644
--- a/main_windows.go
+++ b/main_windows.go
@@ -30,13 +30,13 @@ func main() {
}
interfaceName := os.Args[1]
- fmt.Fprintln(os.Stderr, "Warning: this is a test program for Windows, mainly used for debugging this Go package. For a real WireGuard for Windows client, the repo you want is , which includes this code as a module.")
+ fmt.Fprintln(os.Stderr, "Warning: this is a test program for Windows, mainly used for debugging this Go package. For a real AmneziaWG for Windows client, please visit: https://amnezia.org")
logger := device.NewLogger(
device.LogLevelVerbose,
fmt.Sprintf("(%s) ", interfaceName),
)
- logger.Verbosef("Starting wireguard-go version %s", Version)
+ logger.Verbosef("Starting amneziawg-go version %s", Version)
tun, err := tun.CreateTUN(interfaceName, 0)
if err == nil {
From d2b0fc97892bbdd9802a5152a15a52c59fe80800 Mon Sep 17 00:00:00 2001
From: albexk
Date: Tue, 9 Apr 2024 21:45:50 +0300
Subject: [PATCH 41/75] Add resetting of message types when closing the device
---
device/device.go | 15 ++++++++++-----
1 file changed, 10 insertions(+), 5 deletions(-)
diff --git a/device/device.go b/device/device.go
index 21a6546..a8a9e2f 100644
--- a/device/device.go
+++ b/device/device.go
@@ -415,6 +415,8 @@ func (device *Device) Close() {
device.rate.limiter.Close()
+ device.resetProtocol()
+
device.log.Verbosef("Device closed")
close(device.closed)
}
@@ -559,14 +561,17 @@ func (device *Device) isAdvancedSecurityOn() bool {
return device.isASecOn.IsSet()
}
+func (device *Device) resetProtocol() {
+ // restore default message type values
+ MessageInitiationType = 1
+ MessageResponseType = 2
+ MessageCookieReplyType = 3
+ MessageTransportType = 4
+}
+
func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
if !tempASecCfg.isSet {
- // restore default values
- MessageInitiationType = 1
- MessageResponseType = 2
- MessageCookieReplyType = 3
- MessageTransportType = 4
return err
}
From c00bda9200364d05b071dde04f0157c7a72c39b8 Mon Sep 17 00:00:00 2001
From: albexk
Date: Wed, 10 Apr 2024 16:05:31 +0300
Subject: [PATCH 42/75] Fix output of the version command
---
main.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/main.go b/main.go
index c17c405..5a3dfef 100644
--- a/main.go
+++ b/main.go
@@ -59,7 +59,7 @@ func warning() {
func main() {
if len(os.Args) == 2 && os.Args[1] == "--version" {
- fmt.Printf("amneziawg-go v%s\n\nUserspace AmneziaWG daemon for %s-%s.\nInformation available at https://amnezia.org\n", Version, runtime.GOOS, runtime.GOARCH)
+ fmt.Printf("amneziawg-go %s\n\nUserspace AmneziaWG daemon for %s-%s.\nInformation available at https://amnezia.org\n", Version, runtime.GOOS, runtime.GOARCH)
return
}
From 87d8c00f869645293c9ecd449488e151d33bde2f Mon Sep 17 00:00:00 2001
From: albexk
Date: Tue, 21 May 2024 18:03:30 +0300
Subject: [PATCH 43/75] Up go to 1.22.3, up crypto to 0.21.0
---
go.mod | 6 +++---
go.sum | 8 ++++----
2 files changed, 7 insertions(+), 7 deletions(-)
diff --git a/go.mod b/go.mod
index 33182ee..115ae88 100644
--- a/go.mod
+++ b/go.mod
@@ -1,12 +1,12 @@
module github.com/amnezia-vpn/amneziawg-go
-go 1.20
+go 1.22.3
require (
github.com/tevino/abool/v2 v2.1.0
- golang.org/x/crypto v0.19.0
+ golang.org/x/crypto v0.21.0
golang.org/x/net v0.21.0
- golang.org/x/sys v0.17.0
+ golang.org/x/sys v0.18.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259
)
diff --git a/go.sum b/go.sum
index 0e6f733..7b53725 100644
--- a/go.sum
+++ b/go.sum
@@ -2,12 +2,12 @@ github.com/google/btree v1.0.1 h1:gK4Kx5IaGY9CD5sPJ36FHiBJ6ZXl0kilRiiCj+jdYp4=
github.com/google/btree v1.0.1/go.mod h1:xXMiIv4Fb/0kKde4SpL7qlzvu5cMJDRkFDxJfI9uaxA=
github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
-golang.org/x/crypto v0.19.0 h1:ENy+Az/9Y1vSrlrvBSyna3PITt4tiZLf7sgCjZBX7Wo=
-golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU=
+golang.org/x/crypto v0.21.0 h1:X31++rzVUdKhX5sWmSOFZxx8UW/ldWx55cbf08iNAMA=
+golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOMs=
golang.org/x/net v0.21.0 h1:AQyQV4dYCvJ7vGmJyKki9+PBdyvhkSd8EIx/qb0AYv4=
golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
-golang.org/x/sys v0.17.0 h1:25cE3gD+tdBA7lp7QfhuV+rJiE9YXTcS3VG1SqssI/Y=
-golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
+golang.org/x/sys v0.18.0 h1:DBdB3niSjOA/O0blCZBqDefyWNYveAYMNF1Wum0DYQ4=
+golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 h1:vVKdlvoWBphwdxWKrFZEuM0kGgGLxUOYcY4U/2Vjg44=
golang.org/x/time v0.0.0-20220210224613-90d013bbcef8/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
From 2e7780471af8efd13345692c11dac0cdb9cd8d35 Mon Sep 17 00:00:00 2001
From: Iurii Egorov
Date: Fri, 24 May 2024 18:18:23 +0300
Subject: [PATCH 44/75] Remove GetOffloadInfo() (#32)
* Remove GetOffloadInfo()
* Remove GetOffloadInfo() from bind_windows as well
* Allow lightweight tags to be used in the version
---
Makefile | 2 +-
conn/bind_std.go | 5 -----
conn/bind_windows.go | 4 ----
conn/bindtest/bindtest.go | 2 --
conn/conn.go | 2 --
device/device.go | 1 -
6 files changed, 1 insertion(+), 15 deletions(-)
diff --git a/Makefile b/Makefile
index 4087cba..7a88647 100644
--- a/Makefile
+++ b/Makefile
@@ -9,7 +9,7 @@ MAKEFLAGS += --no-print-directory
generate-version-and-build:
@export GIT_CEILING_DIRECTORIES="$(realpath $(CURDIR)/..)" && \
- tag="$$(git describe --dirty 2>/dev/null)" && \
+ tag="$$(git describe --tags --dirty 2>/dev/null)" && \
ver="$$(printf 'package main\n\nconst Version = "%s"\n' "$$tag")" && \
[ "$$(cat version.go 2>/dev/null)" != "$$ver" ] && \
echo "$$ver" > version.go && \
diff --git a/conn/bind_std.go b/conn/bind_std.go
index ea06cd5..312a538 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -298,11 +298,6 @@ func (s *StdNetBind) BatchSize() int {
return 1
}
-func (s *StdNetBind) GetOffloadInfo() string {
- return fmt.Sprintf("ipv4TxOffload: %v, ipv4RxOffload: %v\nipv6TxOffload: %v, ipv6RxOffload: %v",
- s.ipv4TxOffload, s.ipv4RxOffload, s.ipv6TxOffload, s.ipv6RxOffload)
-}
-
func (s *StdNetBind) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
diff --git a/conn/bind_windows.go b/conn/bind_windows.go
index 3481f00..6cfa099 100644
--- a/conn/bind_windows.go
+++ b/conn/bind_windows.go
@@ -328,10 +328,6 @@ func (bind *WinRingBind) BatchSize() int {
return 1
}
-func (bind *WinRingBind) GetOffloadInfo() string {
- return ""
-}
-
func (bind *WinRingBind) SetMark(mark uint32) error {
return nil
}
diff --git a/conn/bindtest/bindtest.go b/conn/bindtest/bindtest.go
index 0df1420..42b0bb7 100644
--- a/conn/bindtest/bindtest.go
+++ b/conn/bindtest/bindtest.go
@@ -91,8 +91,6 @@ func (c *ChannelBind) Close() error {
func (c *ChannelBind) BatchSize() int { return 1 }
-func (c *ChannelBind) GetOffloadInfo() string { return "" }
-
func (c *ChannelBind) SetMark(mark uint32) error { return nil }
func (c *ChannelBind) makeReceiveFunc(ch chan []byte) conn.ReceiveFunc {
diff --git a/conn/conn.go b/conn/conn.go
index 489cb35..a1f57d2 100644
--- a/conn/conn.go
+++ b/conn/conn.go
@@ -55,8 +55,6 @@ type Bind interface {
// BatchSize is the number of buffers expected to be passed to
// the ReceiveFuncs, and the maximum expected to be passed to SendBatch.
BatchSize() int
-
- GetOffloadInfo() string
}
// BindSocketToInterface is implemented by Bind objects that support being
diff --git a/device/device.go b/device/device.go
index a8a9e2f..80e3793 100644
--- a/device/device.go
+++ b/device/device.go
@@ -547,7 +547,6 @@ func (device *Device) BindUpdate() error {
}
device.log.Verbosef("UDP bind has been updated")
- device.log.Verbosef(netc.bind.GetOffloadInfo())
return nil
}
From 2e3f7d122ca8ef61e403fddc48a9db8fccd95dbf Mon Sep 17 00:00:00 2001
From: Iurii Egorov
Date: Mon, 1 Jul 2024 13:39:57 +0300
Subject: [PATCH 45/75] Update Go version in Dockerfile
---
Dockerfile | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/Dockerfile b/Dockerfile
index caeb333..590ec5a 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -1,4 +1,4 @@
-FROM golang:1.20 as awg
+FROM golang:1.22.3 as awg
COPY . /awg
WORKDIR /awg
RUN go mod download && \
From b8da08c1067a827c4537a6dc2d0655764792208b Mon Sep 17 00:00:00 2001
From: drkivi <115035277+drkivi@users.noreply.github.com>
Date: Mon, 10 Feb 2025 21:43:02 +0330
Subject: [PATCH 46/75] Update Dockerfile
golang -> 1.23.6
AWGTOOLS_RELEASE -> 1.0.20241018
Signed-off-by: drkivi <115035277+drkivi@users.noreply.github.com>
---
Dockerfile | 4 ++--
1 file changed, 2 insertions(+), 2 deletions(-)
diff --git a/Dockerfile b/Dockerfile
index 590ec5a..73016f7 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -1,4 +1,4 @@
-FROM golang:1.22.3 as awg
+FROM golang:1.23.6 as awg
COPY . /awg
WORKDIR /awg
RUN go mod download && \
@@ -6,7 +6,7 @@ RUN go mod download && \
go build -ldflags '-linkmode external -extldflags "-fno-PIC -static"' -v -o /usr/bin
FROM alpine:3.19
-ARG AWGTOOLS_RELEASE="1.0.20240213"
+ARG AWGTOOLS_RELEASE="1.0.20241018"
RUN apk --no-cache add iproute2 iptables bash && \
cd /usr/bin/ && \
wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
From 668ddfd455a55091b262325c99d85be473a751f7 Mon Sep 17 00:00:00 2001
From: drkivi <115035277+drkivi@users.noreply.github.com>
Date: Mon, 10 Feb 2025 21:44:17 +0330
Subject: [PATCH 47/75] Update go.mod
Submodules Version Up
Signed-off-by: drkivi <115035277+drkivi@users.noreply.github.com>
---
go.mod | 14 +++++++-------
1 file changed, 7 insertions(+), 7 deletions(-)
diff --git a/go.mod b/go.mod
index 115ae88..4575bc8 100644
--- a/go.mod
+++ b/go.mod
@@ -1,17 +1,17 @@
module github.com/amnezia-vpn/amneziawg-go
-go 1.22.3
+go 1.23.6
require (
github.com/tevino/abool/v2 v2.1.0
- golang.org/x/crypto v0.21.0
- golang.org/x/net v0.21.0
- golang.org/x/sys v0.18.0
+ golang.org/x/crypto v0.32.0
+ golang.org/x/net v0.34.0
+ golang.org/x/sys v0.29.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
- gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259
+ gvisor.dev/gvisor v0.0.0-20250130013005-04f9204697c6
)
require (
- github.com/google/btree v1.0.1 // indirect
- golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 // indirect
+ github.com/google/btree v1.1.3 // indirect
+ golang.org/x/time v0.9.0 // indirect
)
From c97b5b76158fd85b1d461c9937ba5ff9186912d9 Mon Sep 17 00:00:00 2001
From: drkivi <115035277+drkivi@users.noreply.github.com>
Date: Mon, 10 Feb 2025 21:44:58 +0330
Subject: [PATCH 48/75] Update go.sum
Signed-off-by: drkivi <115035277+drkivi@users.noreply.github.com>
---
go.sum | 28 ++++++++++++++++------------
1 file changed, 16 insertions(+), 12 deletions(-)
diff --git a/go.sum b/go.sum
index 7b53725..10f1f2a 100644
--- a/go.sum
+++ b/go.sum
@@ -1,16 +1,20 @@
-github.com/google/btree v1.0.1 h1:gK4Kx5IaGY9CD5sPJ36FHiBJ6ZXl0kilRiiCj+jdYp4=
-github.com/google/btree v1.0.1/go.mod h1:xXMiIv4Fb/0kKde4SpL7qlzvu5cMJDRkFDxJfI9uaxA=
+github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
+github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
+github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
+github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
-golang.org/x/crypto v0.21.0 h1:X31++rzVUdKhX5sWmSOFZxx8UW/ldWx55cbf08iNAMA=
-golang.org/x/crypto v0.21.0/go.mod h1:0BP7YvVV9gBbVKyeTG0Gyn+gZm94bibOW5BjDEYAOMs=
-golang.org/x/net v0.21.0 h1:AQyQV4dYCvJ7vGmJyKki9+PBdyvhkSd8EIx/qb0AYv4=
-golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44=
-golang.org/x/sys v0.18.0 h1:DBdB3niSjOA/O0blCZBqDefyWNYveAYMNF1Wum0DYQ4=
-golang.org/x/sys v0.18.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
-golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 h1:vVKdlvoWBphwdxWKrFZEuM0kGgGLxUOYcY4U/2Vjg44=
-golang.org/x/time v0.0.0-20220210224613-90d013bbcef8/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
+golang.org/x/crypto v0.32.0 h1:euUpcYgM8WcP71gNpTqQCn6rC2t6ULUPiOzfWaXVVfc=
+golang.org/x/crypto v0.32.0/go.mod h1:ZnnJkOaASj8g0AjIduWNlq2NRxL0PlBrbKVyZ6V/Ugc=
+golang.org/x/mod v0.21.0 h1:vvrHzRwRfVKSiLrG+d4FMl/Qi4ukBCE6kZlTUkDYRT0=
+golang.org/x/mod v0.21.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
+golang.org/x/net v0.34.0 h1:Mb7Mrk043xzHgnRM88suvJFwzVrRfHEHJEl5/71CKw0=
+golang.org/x/net v0.34.0/go.mod h1:di0qlW3YNM5oh6GqDGQr92MyTozJPmybPK4Ev/Gm31k=
+golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU=
+golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
+golang.org/x/time v0.9.0 h1:EsRrnYcQiGH+5FfbgvV4AP7qEZstoyrHB0DzarOQ4ZY=
+golang.org/x/time v0.9.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
-gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259 h1:TbRPT0HtzFP3Cno1zZo7yPzEEnfu8EjLfl6IU9VfqkQ=
-gvisor.dev/gvisor v0.0.0-20230927004350-cbd86285d259/go.mod h1:AVgIgHMwK63XvmAzWG9vLQ41YnVHN0du0tEC46fI7yY=
+gvisor.dev/gvisor v0.0.0-20250130013005-04f9204697c6 h1:6B7MdW3OEbJqOMr7cEYU9bkzvCjUBX/JlXk12xcANuQ=
+gvisor.dev/gvisor v0.0.0-20250130013005-04f9204697c6/go.mod h1:5DMfjtclAbTIjbXqO1qCe2K5GKKxWz2JHvCChuTcJEM=
From 71be0eb3a6547f172d17ce8b831b89e48052dc27 Mon Sep 17 00:00:00 2001
From: Mark Puha
Date: Tue, 18 Mar 2025 08:34:23 +0100
Subject: [PATCH 49/75] faster and more secure junk creation
---
Dockerfile | 2 +-
device/device.go | 2 +
device/device_test.go | 6 +-
device/junk_creator.go | 69 ++++++++++++++++++++
device/junk_creator_test.go | 124 ++++++++++++++++++++++++++++++++++++
device/send.go | 32 +---------
device/util.go | 25 --------
device/util_test.go | 27 --------
go.mod | 8 +--
go.sum | 12 ++--
10 files changed, 212 insertions(+), 95 deletions(-)
create mode 100644 device/junk_creator.go
create mode 100644 device/junk_creator_test.go
delete mode 100644 device/util.go
delete mode 100644 device/util_test.go
diff --git a/Dockerfile b/Dockerfile
index 73016f7..12159be 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -1,4 +1,4 @@
-FROM golang:1.23.6 as awg
+FROM golang:1.24 as awg
COPY . /awg
WORKDIR /awg
RUN go mod download && \
diff --git a/device/device.go b/device/device.go
index 80e3793..1be15d0 100644
--- a/device/device.go
+++ b/device/device.go
@@ -95,6 +95,7 @@ type Device struct {
isASecOn abool.AtomicBool
aSecMux sync.RWMutex
aSecCfg aSecCfgType
+ junkCreator junkCreator
}
type aSecCfgType struct {
@@ -799,6 +800,7 @@ func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
}
device.isASecOn.SetTo(isASecOn)
+ device.junkCreator, err = NewJunkCreator(device)
device.aSecMux.Unlock()
return err
diff --git a/device/device_test.go b/device/device_test.go
index e6664a6..d03610f 100644
--- a/device/device_test.go
+++ b/device/device_test.go
@@ -109,7 +109,7 @@ func genASecurityConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
"replace_peers", "true",
"jc", "5",
"jmin", "500",
- "jmax", "501",
+ "jmax", "1000",
"s1", "30",
"s2", "40",
"h1", "123456",
@@ -131,7 +131,7 @@ func genASecurityConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
"replace_peers", "true",
"jc", "5",
"jmin", "500",
- "jmax", "501",
+ "jmax", "1000",
"s1", "30",
"s2", "40",
"h1", "123456",
@@ -274,7 +274,7 @@ func TestTwoDevicePing(t *testing.T) {
})
}
-func TestTwoDevicePingASecurity(t *testing.T) {
+func TestASecurityTwoDevicePing(t *testing.T) {
goroutineLeakCheck(t)
pair := genTestPair(t, true, true)
t.Run("ping 1.0.0.1", func(t *testing.T) {
diff --git a/device/junk_creator.go b/device/junk_creator.go
new file mode 100644
index 0000000..3a2d3b4
--- /dev/null
+++ b/device/junk_creator.go
@@ -0,0 +1,69 @@
+package device
+
+import (
+ "bytes"
+ crand "crypto/rand"
+ "fmt"
+ v2 "math/rand/v2"
+)
+
+type junkCreator struct {
+ device *Device
+ cha8Rand *v2.ChaCha8
+}
+
+func NewJunkCreator(d *Device) (junkCreator, error) {
+ buf := make([]byte, 32)
+ _, err := crand.Read(buf)
+ if err != nil {
+ return junkCreator{}, err
+ }
+ return junkCreator{device: d, cha8Rand: v2.NewChaCha8([32]byte(buf))}, nil
+}
+
+// Should be called with aSecMux RLocked
+func (jc *junkCreator) createJunkPackets() ([][]byte, error) {
+ if jc.device.aSecCfg.junkPacketCount == 0 {
+ return nil, nil
+ }
+
+ junks := make([][]byte, 0, jc.device.aSecCfg.junkPacketCount)
+ for i := 0; i < jc.device.aSecCfg.junkPacketCount; i++ {
+ packetSize := jc.randomPacketSize()
+ junk, err := jc.randomJunkWithSize(packetSize)
+ if err != nil {
+ return nil, fmt.Errorf("Failed to create junk packet: %v", err)
+ }
+ junks = append(junks, junk)
+ }
+ return junks, nil
+}
+
+// Should be called with aSecMux RLocked
+func (jc *junkCreator) randomPacketSize() int {
+ return int(
+ jc.cha8Rand.Uint64()%uint64(
+ jc.device.aSecCfg.junkPacketMaxSize-jc.device.aSecCfg.junkPacketMinSize,
+ ),
+ ) + jc.device.aSecCfg.junkPacketMinSize
+}
+
+// Should be called with aSecMux RLocked
+func (jc *junkCreator) appendJunk(writer *bytes.Buffer, size int) error {
+ headerJunk, err := jc.randomJunkWithSize(size)
+ if err != nil {
+ return fmt.Errorf("failed to create header junk: %v", err)
+ }
+ _, err = writer.Write(headerJunk)
+ if err != nil {
+ return fmt.Errorf("failed to write header junk: %v", err)
+ }
+ return nil
+}
+
+// Should be called with aSecMux RLocked
+func (jc *junkCreator) randomJunkWithSize(size int) ([]byte, error) {
+ junk := make([]byte, size)
+ _, err := jc.cha8Rand.Read(junk)
+ return junk, err
+}
diff --git a/device/junk_creator_test.go b/device/junk_creator_test.go
new file mode 100644
index 0000000..d3cf2b3
--- /dev/null
+++ b/device/junk_creator_test.go
@@ -0,0 +1,124 @@
+package device
+
+import (
+ "bytes"
+ "fmt"
+ "testing"
+
+ "github.com/amnezia-vpn/amneziawg-go/conn/bindtest"
+ "github.com/amnezia-vpn/amneziawg-go/tun/tuntest"
+)
+
+func setUpJunkCreator(t *testing.T) (junkCreator, error) {
+ cfg, _ := genASecurityConfigs(t)
+ tun := tuntest.NewChannelTUN()
+ binds := bindtest.NewChannelBinds()
+ level := LogLevelVerbose
+ dev := NewDevice(
+ tun.TUN(),
+ binds[0],
+ NewLogger(level, ""),
+ )
+
+ if err := dev.IpcSet(cfg[0]); err != nil {
+ t.Errorf("failed to configure device %v", err)
+ dev.Close()
+ return junkCreator{}, err
+ }
+
+ jc, err := NewJunkCreator(dev)
+
+ if err != nil {
+ t.Errorf("failed to create junk creator %v", err)
+ dev.Close()
+ return junkCreator{}, err
+ }
+
+ return jc, nil
+}
+
+func Test_junkCreator_createJunkPackets(t *testing.T) {
+ jc, err := setUpJunkCreator(t)
+ if err != nil {
+ return
+ }
+ t.Run("", func(t *testing.T) {
+ got, err := jc.createJunkPackets()
+ if err != nil {
+ t.Errorf(
+ "junkCreator.createJunkPackets() = %v; failed",
+ err,
+ )
+ return
+ }
+ seen := make(map[string]bool)
+ for _, junk := range got {
+ key := string(junk)
+ if seen[key] {
+ t.Errorf(
+ "junkCreator.createJunkPackets() = %v, duplicate key: %v",
+ got,
+ junk,
+ )
+ return
+ }
+ seen[key] = true
+ }
+ })
+}
+
+func Test_junkCreator_randomJunkWithSize(t *testing.T) {
+ t.Run("", func(t *testing.T) {
+ jc, err := setUpJunkCreator(t)
+ if err != nil {
+ return
+ }
+ r1, _ := jc.randomJunkWithSize(10)
+ r2, _ := jc.randomJunkWithSize(10)
+ fmt.Printf("%v\n%v\n", r1, r2)
+ if bytes.Equal(r1, r2) {
+ t.Errorf("same junks %v", err)
+ jc.device.Close()
+ return
+ }
+ })
+}
+
+func Test_junkCreator_randomPacketSize(t *testing.T) {
+ jc, err := setUpJunkCreator(t)
+ if err != nil {
+ return
+ }
+ for range [30]struct{}{} {
+ t.Run("", func(t *testing.T) {
+ if got := jc.randomPacketSize(); jc.device.aSecCfg.junkPacketMinSize > got ||
+ got > jc.device.aSecCfg.junkPacketMaxSize {
+ t.Errorf(
+ "junkCreator.randomPacketSize() = %v, not between range [%v,%v]",
+ got,
+ jc.device.aSecCfg.junkPacketMinSize,
+ jc.device.aSecCfg.junkPacketMaxSize,
+ )
+ }
+ })
+ }
+}
+
+func Test_junkCreator_appendJunk(t *testing.T) {
+ jc, err := setUpJunkCreator(t)
+ if err != nil {
+ return
+ }
+ t.Run("", func(t *testing.T) {
+ s := "apple"
+ buffer := bytes.NewBuffer([]byte(s))
+ err := jc.appendJunk(buffer, 30)
+ if err != nil &&
+ buffer.Len() != len(s)+30 {
+ t.Errorf("appendWithJunk() size don't match")
+ }
+ read := make([]byte, 50)
+ buffer.Read(read)
+ fmt.Println(string(read))
+ })
+}
diff --git a/device/send.go b/device/send.go
index 1b4406d..7eca099 100644
--- a/device/send.go
+++ b/device/send.go
@@ -9,7 +9,6 @@ import (
"bytes"
"encoding/binary"
"errors"
- "math/rand"
"net"
"os"
"sync"
@@ -129,7 +128,7 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
var junkedHeader []byte
if peer.device.isAdvancedSecurityOn() {
peer.device.aSecMux.RLock()
- junks, err := peer.createJunkPackets()
+ junks, err := peer.device.junkCreator.createJunkPackets()
peer.device.aSecMux.RUnlock()
if err != nil {
@@ -150,7 +149,7 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
if peer.device.aSecCfg.initPacketJunkSize != 0 {
buf := make([]byte, 0, peer.device.aSecCfg.initPacketJunkSize)
writer := bytes.NewBuffer(buf[:0])
- err = appendJunk(writer, peer.device.aSecCfg.initPacketJunkSize)
+ err = peer.device.junkCreator.appendJunk(writer, peer.device.aSecCfg.initPacketJunkSize)
if err != nil {
peer.device.log.Errorf("%v - %v", peer, err)
peer.device.aSecMux.RUnlock()
@@ -200,7 +199,7 @@ func (peer *Peer) SendHandshakeResponse() error {
if peer.device.aSecCfg.responsePacketJunkSize != 0 {
buf := make([]byte, 0, peer.device.aSecCfg.responsePacketJunkSize)
writer := bytes.NewBuffer(buf[:0])
- err = appendJunk(writer, peer.device.aSecCfg.responsePacketJunkSize)
+ err = peer.device.junkCreator.appendJunk(writer, peer.device.aSecCfg.responsePacketJunkSize)
if err != nil {
peer.device.aSecMux.RUnlock()
peer.device.log.Errorf("%v - %v", peer, err)
@@ -469,31 +468,6 @@ top:
}
}
-func (peer *Peer) createJunkPackets() ([][]byte, error) {
- if peer.device.aSecCfg.junkPacketCount == 0 {
- return nil, nil
- }
-
- junks := make([][]byte, 0, peer.device.aSecCfg.junkPacketCount)
- for i := 0; i < peer.device.aSecCfg.junkPacketCount; i++ {
- packetSize := rand.Intn(
- peer.device.aSecCfg.junkPacketMaxSize-peer.device.aSecCfg.junkPacketMinSize,
- ) + peer.device.aSecCfg.junkPacketMinSize
-
- junk, err := randomJunkWithSize(packetSize)
- if err != nil {
- peer.device.log.Errorf(
- "%v - Failed to create junk packet: %v",
- peer,
- err,
- )
- return nil, err
- }
- junks = append(junks, junk)
- }
- return junks, nil
-}
-
func (peer *Peer) FlushStagedPackets() {
for {
select {
diff --git a/device/util.go b/device/util.go
deleted file mode 100644
index aab8ab7..0000000
--- a/device/util.go
+++ /dev/null
@@ -1,25 +0,0 @@
-package device
-
-import (
- "bytes"
- crand "crypto/rand"
- "fmt"
-)
-
-func appendJunk(writer *bytes.Buffer, size int) error {
- headerJunk, err := randomJunkWithSize(size)
- if err != nil {
- return fmt.Errorf("failed to create header junk: %v", err)
- }
- _, err = writer.Write(headerJunk)
- if err != nil {
- return fmt.Errorf("failed to write header junk: %v", err)
- }
- return nil
-}
-
-func randomJunkWithSize(size int) ([]byte, error) {
- junk := make([]byte, size)
- _, err := crand.Read(junk)
- return junk, err
-}
diff --git a/device/util_test.go b/device/util_test.go
deleted file mode 100644
index c061eef..0000000
--- a/device/util_test.go
+++ /dev/null
@@ -1,27 +0,0 @@
-package device
-
-import (
- "bytes"
- "fmt"
- "testing"
-)
-
-func Test_randomJunktWithSize(t *testing.T) {
- junk, err := randomJunkWithSize(30)
- fmt.Println(string(junk), len(junk), err)
-}
-
-func Test_appendJunk(t *testing.T) {
- t.Run("", func(t *testing.T) {
- s := "apple"
- buffer := bytes.NewBuffer([]byte(s))
- err := appendJunk(buffer, 30)
- if err != nil &&
- buffer.Len() != len(s)+30 {
- t.Errorf("appendWithJunk() size don't match")
- }
- read := make([]byte, 50)
- buffer.Read(read)
- fmt.Println(string(read))
- })
-}
diff --git a/go.mod b/go.mod
index 4575bc8..608969f 100644
--- a/go.mod
+++ b/go.mod
@@ -1,12 +1,12 @@
module github.com/amnezia-vpn/amneziawg-go
-go 1.23.6
+go 1.24
require (
github.com/tevino/abool/v2 v2.1.0
- golang.org/x/crypto v0.32.0
- golang.org/x/net v0.34.0
- golang.org/x/sys v0.29.0
+ golang.org/x/crypto v0.36.0
+ golang.org/x/net v0.37.0
+ golang.org/x/sys v0.31.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
gvisor.dev/gvisor v0.0.0-20250130013005-04f9204697c6
)
diff --git a/go.sum b/go.sum
index 10f1f2a..497f949 100644
--- a/go.sum
+++ b/go.sum
@@ -4,14 +4,14 @@ github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
-golang.org/x/crypto v0.32.0 h1:euUpcYgM8WcP71gNpTqQCn6rC2t6ULUPiOzfWaXVVfc=
-golang.org/x/crypto v0.32.0/go.mod h1:ZnnJkOaASj8g0AjIduWNlq2NRxL0PlBrbKVyZ6V/Ugc=
+golang.org/x/crypto v0.36.0 h1:AnAEvhDddvBdpY+uR+MyHmuZzzNqXSe/GvuDeob5L34=
+golang.org/x/crypto v0.36.0/go.mod h1:Y4J0ReaxCR1IMaabaSMugxJES1EpwhBHhv2bDHklZvc=
golang.org/x/mod v0.21.0 h1:vvrHzRwRfVKSiLrG+d4FMl/Qi4ukBCE6kZlTUkDYRT0=
golang.org/x/mod v0.21.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
-golang.org/x/net v0.34.0 h1:Mb7Mrk043xzHgnRM88suvJFwzVrRfHEHJEl5/71CKw0=
-golang.org/x/net v0.34.0/go.mod h1:di0qlW3YNM5oh6GqDGQr92MyTozJPmybPK4Ev/Gm31k=
-golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU=
-golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
+golang.org/x/net v0.37.0 h1:1zLorHbz+LYj7MQlSf1+2tPIIgibq2eL5xkrGk6f+2c=
+golang.org/x/net v0.37.0/go.mod h1:ivrbrMbzFq5J41QOQh0siUuly180yBYtLp+CKbEaFx8=
+golang.org/x/sys v0.31.0 h1:ioabZlmFYtWhL+TRYpcnNlLwhyxaM9kWTDEmfnprqik=
+golang.org/x/sys v0.31.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/time v0.9.0 h1:EsRrnYcQiGH+5FfbgvV4AP7qEZstoyrHB0DzarOQ4ZY=
golang.org/x/time v0.9.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
From deedce495a1616d4be20d8a12d78275abeb52d16 Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Thu, 27 Jun 2024 08:43:41 -0700
Subject: [PATCH 50/75] device: fix WaitPool sync.Cond usage
The sync.Locker used with a sync.Cond must be acquired when changing
the associated condition, otherwise there is a window within
sync.Cond.Wait() where a wake-up may be missed.
Fixes: 4846070 ("device: use a waiting sync.Pool instead of a channel")
Reviewed-by: Brad Fitzpatrick
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
device/pools.go | 11 ++++++-----
device/pools_test.go | 4 +++-
2 files changed, 9 insertions(+), 6 deletions(-)
diff --git a/device/pools.go b/device/pools.go
index 94f3dc7..55d2be7 100644
--- a/device/pools.go
+++ b/device/pools.go
@@ -7,14 +7,13 @@ package device
import (
"sync"
- "sync/atomic"
)
type WaitPool struct {
pool sync.Pool
cond sync.Cond
lock sync.Mutex
- count atomic.Uint32
+ count uint32 // Get calls not yet Put back
max uint32
}
@@ -27,10 +26,10 @@ func NewWaitPool(max uint32, new func() any) *WaitPool {
func (p *WaitPool) Get() any {
if p.max != 0 {
p.lock.Lock()
- for p.count.Load() >= p.max {
+ for p.count >= p.max {
p.cond.Wait()
}
- p.count.Add(1)
+ p.count++
p.lock.Unlock()
}
return p.pool.Get()
@@ -41,7 +40,9 @@ func (p *WaitPool) Put(x any) {
if p.max == 0 {
return
}
- p.count.Add(^uint32(0))
+ p.lock.Lock()
+ defer p.lock.Unlock()
+ p.count--
p.cond.Signal()
}
diff --git a/device/pools_test.go b/device/pools_test.go
index 82d7493..538230b 100644
--- a/device/pools_test.go
+++ b/device/pools_test.go
@@ -32,7 +32,9 @@ func TestWaitPool(t *testing.T) {
wg.Add(workers)
var max atomic.Uint32
updateMax := func() {
- count := p.count.Load()
+ p.lock.Lock()
+ count := p.count
+ p.lock.Unlock()
if count > p.max {
t.Errorf("count (%d) > max (%d)", count, p.max)
}
From c803ce1e5bd7723274500226ee56395a50c3ab8f Mon Sep 17 00:00:00 2001
From: Jordan Whited
Date: Thu, 27 Jun 2024 09:06:40 -0700
Subject: [PATCH 51/75] device: fix missed return of
QueueOutboundElementsContainer to its WaitPool
Fixes: 3bb8fec ("conn, device, tun: implement vectorized I/O plumbing")
Reviewed-by: Brad Fitzpatrick
Signed-off-by: Jordan Whited
Signed-off-by: Jason A. Donenfeld
---
device/send.go | 1 +
1 file changed, 1 insertion(+)
diff --git a/device/send.go b/device/send.go
index 7eca099..a00e2bb 100644
--- a/device/send.go
+++ b/device/send.go
@@ -568,6 +568,7 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
device.PutMessageBuffer(elem.buffer)
device.PutOutboundElement(elem)
}
+ device.PutOutboundElementsContainer(elemsContainer)
continue
}
dataSent := false
From c0b6e6a2001c1ad6529e30cb4a289467a5684b3c Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Sun, 4 May 2025 17:48:53 +0200
Subject: [PATCH 52/75] global: bump copyright notice
Signed-off-by: Jason A. Donenfeld
---
README.md | 1 +
conn/bind_std.go | 2 +-
conn/bind_windows.go | 2 +-
conn/bindtest/bindtest.go | 2 +-
conn/boundif_android.go | 2 +-
conn/conn.go | 2 +-
conn/conn_test.go | 2 +-
conn/controlfns.go | 2 +-
conn/controlfns_linux.go | 2 +-
conn/controlfns_unix.go | 2 +-
conn/controlfns_windows.go | 2 +-
conn/default.go | 2 +-
conn/errors_default.go | 2 +-
conn/errors_linux.go | 2 +-
conn/features_default.go | 2 +-
conn/features_linux.go | 2 +-
conn/gso_default.go | 2 +-
conn/gso_linux.go | 2 +-
conn/mark_default.go | 2 +-
conn/mark_unix.go | 2 +-
conn/sticky_default.go | 2 +-
conn/sticky_linux.go | 2 +-
conn/sticky_linux_test.go | 2 +-
conn/winrio/rio_windows.go | 2 +-
device/allowedips.go | 2 +-
device/allowedips_rand_test.go | 2 +-
device/allowedips_test.go | 2 +-
device/bind_test.go | 2 +-
device/channels.go | 2 +-
device/constants.go | 2 +-
device/cookie.go | 2 +-
device/cookie_test.go | 2 +-
device/device.go | 2 +-
device/device_test.go | 2 +-
device/endpoint_test.go | 2 +-
device/indextable.go | 2 +-
device/ip.go | 2 +-
device/kdf_test.go | 2 +-
device/keypair.go | 2 +-
device/logger.go | 2 +-
device/mobilequirks.go | 2 +-
device/noise-helpers.go | 2 +-
device/noise-protocol.go | 2 +-
device/noise-types.go | 2 +-
device/noise_test.go | 2 +-
device/peer.go | 2 +-
device/pools.go | 2 +-
device/pools_test.go | 2 +-
device/queueconstants_android.go | 2 +-
device/queueconstants_default.go | 2 +-
device/queueconstants_ios.go | 2 +-
device/queueconstants_windows.go | 2 +-
device/race_disabled_test.go | 2 +-
device/race_enabled_test.go | 2 +-
device/receive.go | 2 +-
device/send.go | 2 +-
device/sticky_linux.go | 2 +-
device/timers.go | 2 +-
device/tun.go | 2 +-
device/uapi.go | 2 +-
format_test.go | 2 +-
ipc/uapi_bsd.go | 2 +-
ipc/uapi_linux.go | 2 +-
ipc/uapi_unix.go | 2 +-
ipc/uapi_wasm.go | 2 +-
ipc/uapi_windows.go | 2 +-
main.go | 2 +-
main_windows.go | 2 +-
ratelimiter/ratelimiter.go | 2 +-
ratelimiter/ratelimiter_test.go | 2 +-
replay/replay.go | 2 +-
replay/replay_test.go | 2 +-
rwcancel/rwcancel.go | 2 +-
tai64n/tai64n.go | 2 +-
tai64n/tai64n_test.go | 2 +-
tun/alignment_windows_test.go | 2 +-
tun/netstack/examples/http_client.go | 2 +-
tun/netstack/examples/http_server.go | 2 +-
tun/netstack/examples/ping_client.go | 2 +-
tun/netstack/tun.go | 2 +-
tun/offload_linux.go | 2 +-
tun/offload_linux_test.go | 2 +-
tun/operateonfd.go | 2 +-
tun/tun.go | 2 +-
tun/tun_darwin.go | 2 +-
tun/tun_freebsd.go | 2 +-
tun/tun_linux.go | 2 +-
tun/tun_openbsd.go | 2 +-
tun/tun_windows.go | 2 +-
tun/tuntest/tuntest.go | 2 +-
90 files changed, 90 insertions(+), 89 deletions(-)
diff --git a/README.md b/README.md
index 853d318..428b752 100644
--- a/README.md
+++ b/README.md
@@ -50,3 +50,4 @@ $ git clone https://github.com/amnezia-vpn/amneziawg-go
$ cd amneziawg-go
$ make
```
+
diff --git a/conn/bind_std.go b/conn/bind_std.go
index 312a538..6908ba8 100644
--- a/conn/bind_std.go
+++ b/conn/bind_std.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/bind_windows.go b/conn/bind_windows.go
index 6cfa099..1a0e021 100644
--- a/conn/bind_windows.go
+++ b/conn/bind_windows.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/bindtest/bindtest.go b/conn/bindtest/bindtest.go
index 42b0bb7..25b5eab 100644
--- a/conn/bindtest/bindtest.go
+++ b/conn/bindtest/bindtest.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package bindtest
diff --git a/conn/boundif_android.go b/conn/boundif_android.go
index dd3ca5b..be69b2a 100644
--- a/conn/boundif_android.go
+++ b/conn/boundif_android.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/conn.go b/conn/conn.go
index a1f57d2..1304657 100644
--- a/conn/conn.go
+++ b/conn/conn.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
// Package conn implements WireGuard's network connections.
diff --git a/conn/conn_test.go b/conn/conn_test.go
index c6194ee..618d02b 100644
--- a/conn/conn_test.go
+++ b/conn/conn_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/controlfns.go b/conn/controlfns.go
index 4f7d90f..27421bd 100644
--- a/conn/controlfns.go
+++ b/conn/controlfns.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/controlfns_linux.go b/conn/controlfns_linux.go
index a2396fe..7bd3917 100644
--- a/conn/controlfns_linux.go
+++ b/conn/controlfns_linux.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/controlfns_unix.go b/conn/controlfns_unix.go
index 91692c0..b2e7570 100644
--- a/conn/controlfns_unix.go
+++ b/conn/controlfns_unix.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/controlfns_windows.go b/conn/controlfns_windows.go
index c3bdf7d..5e38305 100644
--- a/conn/controlfns_windows.go
+++ b/conn/controlfns_windows.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/default.go b/conn/default.go
index b6f761b..2ce1579 100644
--- a/conn/default.go
+++ b/conn/default.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/errors_default.go b/conn/errors_default.go
index f1e5b90..d967518 100644
--- a/conn/errors_default.go
+++ b/conn/errors_default.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/errors_linux.go b/conn/errors_linux.go
index 7548a8a..9ed7d76 100644
--- a/conn/errors_linux.go
+++ b/conn/errors_linux.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/features_default.go b/conn/features_default.go
index d53ff5f..cae2bea 100644
--- a/conn/features_default.go
+++ b/conn/features_default.go
@@ -3,7 +3,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/features_linux.go b/conn/features_linux.go
index a6de8c1..936029e 100644
--- a/conn/features_linux.go
+++ b/conn/features_linux.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/gso_default.go b/conn/gso_default.go
index 57780db..a9a3e80 100644
--- a/conn/gso_default.go
+++ b/conn/gso_default.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/gso_linux.go b/conn/gso_linux.go
index 8596b29..4ee31fa 100644
--- a/conn/gso_linux.go
+++ b/conn/gso_linux.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/mark_default.go b/conn/mark_default.go
index 3102384..72b266e 100644
--- a/conn/mark_default.go
+++ b/conn/mark_default.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/mark_unix.go b/conn/mark_unix.go
index d9e46ee..d0580d5 100644
--- a/conn/mark_unix.go
+++ b/conn/mark_unix.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/sticky_default.go b/conn/sticky_default.go
index 0b21386..15b65af 100644
--- a/conn/sticky_default.go
+++ b/conn/sticky_default.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/sticky_linux.go b/conn/sticky_linux.go
index 8e206e9..adfedc1 100644
--- a/conn/sticky_linux.go
+++ b/conn/sticky_linux.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/sticky_linux_test.go b/conn/sticky_linux_test.go
index d2bd584..1b1ee68 100644
--- a/conn/sticky_linux_test.go
+++ b/conn/sticky_linux_test.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package conn
diff --git a/conn/winrio/rio_windows.go b/conn/winrio/rio_windows.go
index d1037bb..c396658 100644
--- a/conn/winrio/rio_windows.go
+++ b/conn/winrio/rio_windows.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package winrio
diff --git a/device/allowedips.go b/device/allowedips.go
index fa46f97..b40c817 100644
--- a/device/allowedips.go
+++ b/device/allowedips.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/allowedips_rand_test.go b/device/allowedips_rand_test.go
index 07065c3..8dd9b67 100644
--- a/device/allowedips_rand_test.go
+++ b/device/allowedips_rand_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/allowedips_test.go b/device/allowedips_test.go
index cde068e..9ef8a76 100644
--- a/device/allowedips_test.go
+++ b/device/allowedips_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/bind_test.go b/device/bind_test.go
index 34d1c4a..24dec1f 100644
--- a/device/bind_test.go
+++ b/device/bind_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/channels.go b/device/channels.go
index e526f6b..be15d1c 100644
--- a/device/channels.go
+++ b/device/channels.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/constants.go b/device/constants.go
index 59854a1..41da618 100644
--- a/device/constants.go
+++ b/device/constants.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/cookie.go b/device/cookie.go
index 876f05d..a093c8b 100644
--- a/device/cookie.go
+++ b/device/cookie.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/cookie_test.go b/device/cookie_test.go
index 4f1e50a..c937290 100644
--- a/device/cookie_test.go
+++ b/device/cookie_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/device.go b/device/device.go
index 1be15d0..2a37321 100644
--- a/device/device.go
+++ b/device/device.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/device_test.go b/device/device_test.go
index d03610f..f66d326 100644
--- a/device/device_test.go
+++ b/device/device_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/endpoint_test.go b/device/endpoint_test.go
index 93a4998..85482d8 100644
--- a/device/endpoint_test.go
+++ b/device/endpoint_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/indextable.go b/device/indextable.go
index 00ade7d..2460fa6 100644
--- a/device/indextable.go
+++ b/device/indextable.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/ip.go b/device/ip.go
index eaf2363..f558744 100644
--- a/device/ip.go
+++ b/device/ip.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/kdf_test.go b/device/kdf_test.go
index f9c76d6..325db59 100644
--- a/device/kdf_test.go
+++ b/device/kdf_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/keypair.go b/device/keypair.go
index cc2941a..05bce68 100644
--- a/device/keypair.go
+++ b/device/keypair.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/logger.go b/device/logger.go
index 22b0df0..a2adea3 100644
--- a/device/logger.go
+++ b/device/logger.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/mobilequirks.go b/device/mobilequirks.go
index 0a0080e..af4be31 100644
--- a/device/mobilequirks.go
+++ b/device/mobilequirks.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/noise-helpers.go b/device/noise-helpers.go
index c2f356b..35dd907 100644
--- a/device/noise-helpers.go
+++ b/device/noise-helpers.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/noise-protocol.go b/device/noise-protocol.go
index 1289249..789eb16 100644
--- a/device/noise-protocol.go
+++ b/device/noise-protocol.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/noise-types.go b/device/noise-types.go
index e850359..41c944e 100644
--- a/device/noise-types.go
+++ b/device/noise-types.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/noise_test.go b/device/noise_test.go
index 075b6d3..8f72f29 100644
--- a/device/noise_test.go
+++ b/device/noise_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/peer.go b/device/peer.go
index 5bc8ca4..8f88b2a 100644
--- a/device/peer.go
+++ b/device/peer.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/pools.go b/device/pools.go
index 55d2be7..2c18f41 100644
--- a/device/pools.go
+++ b/device/pools.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/pools_test.go b/device/pools_test.go
index 538230b..8381d5a 100644
--- a/device/pools_test.go
+++ b/device/pools_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/queueconstants_android.go b/device/queueconstants_android.go
index 1bff95a..741fcf3 100644
--- a/device/queueconstants_android.go
+++ b/device/queueconstants_android.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/queueconstants_default.go b/device/queueconstants_default.go
index 0061b63..f19e9b1 100644
--- a/device/queueconstants_default.go
+++ b/device/queueconstants_default.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/queueconstants_ios.go b/device/queueconstants_ios.go
index acd3cec..632e29d 100644
--- a/device/queueconstants_ios.go
+++ b/device/queueconstants_ios.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/queueconstants_windows.go b/device/queueconstants_windows.go
index 1eee32b..9a296d6 100644
--- a/device/queueconstants_windows.go
+++ b/device/queueconstants_windows.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/race_disabled_test.go b/device/race_disabled_test.go
index bb5c450..14b3284 100644
--- a/device/race_disabled_test.go
+++ b/device/race_disabled_test.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/race_enabled_test.go b/device/race_enabled_test.go
index 4e9daea..f1ea5cf 100644
--- a/device/race_enabled_test.go
+++ b/device/race_enabled_test.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/receive.go b/device/receive.go
index 66c1a32..0a4910a 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/send.go b/device/send.go
index a00e2bb..7f0faa3 100644
--- a/device/send.go
+++ b/device/send.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/sticky_linux.go b/device/sticky_linux.go
index 63164a7..5ff9dd6 100644
--- a/device/sticky_linux.go
+++ b/device/sticky_linux.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*
* This implements userspace semantics of "sticky sockets", modeled after
* WireGuard's kernelspace implementation. This is more or less a straight port
diff --git a/device/timers.go b/device/timers.go
index d4a4ed4..32519aa 100644
--- a/device/timers.go
+++ b/device/timers.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*
* This is based heavily on timers.c from the kernel implementation.
*/
diff --git a/device/tun.go b/device/tun.go
index 600a5e5..42178b2 100644
--- a/device/tun.go
+++ b/device/tun.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/device/uapi.go b/device/uapi.go
index 777bdda..1b5e357 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package device
diff --git a/format_test.go b/format_test.go
index 6f6cab7..4d02c48 100644
--- a/format_test.go
+++ b/format_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package main
diff --git a/ipc/uapi_bsd.go b/ipc/uapi_bsd.go
index ddcaf27..fd433a5 100644
--- a/ipc/uapi_bsd.go
+++ b/ipc/uapi_bsd.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package ipc
diff --git a/ipc/uapi_linux.go b/ipc/uapi_linux.go
index 9738aea..058e8e7 100644
--- a/ipc/uapi_linux.go
+++ b/ipc/uapi_linux.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package ipc
diff --git a/ipc/uapi_unix.go b/ipc/uapi_unix.go
index 0da452a..79604ee 100644
--- a/ipc/uapi_unix.go
+++ b/ipc/uapi_unix.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package ipc
diff --git a/ipc/uapi_wasm.go b/ipc/uapi_wasm.go
index fa84684..50ac091 100644
--- a/ipc/uapi_wasm.go
+++ b/ipc/uapi_wasm.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package ipc
diff --git a/ipc/uapi_windows.go b/ipc/uapi_windows.go
index 31d2a63..321fe60 100644
--- a/ipc/uapi_windows.go
+++ b/ipc/uapi_windows.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package ipc
diff --git a/main.go b/main.go
index 5a3dfef..f8fded9 100644
--- a/main.go
+++ b/main.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package main
diff --git a/main_windows.go b/main_windows.go
index bbfa690..d3e2fe6 100644
--- a/main_windows.go
+++ b/main_windows.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package main
diff --git a/ratelimiter/ratelimiter.go b/ratelimiter/ratelimiter.go
index f7d05ef..ac69e3a 100644
--- a/ratelimiter/ratelimiter.go
+++ b/ratelimiter/ratelimiter.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package ratelimiter
diff --git a/ratelimiter/ratelimiter_test.go b/ratelimiter/ratelimiter_test.go
index 0bfa3af..71140da 100644
--- a/ratelimiter/ratelimiter_test.go
+++ b/ratelimiter/ratelimiter_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package ratelimiter
diff --git a/replay/replay.go b/replay/replay.go
index 8b99e23..46e224d 100644
--- a/replay/replay.go
+++ b/replay/replay.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
// Package replay implements an efficient anti-replay algorithm as specified in RFC 6479.
diff --git a/replay/replay_test.go b/replay/replay_test.go
index 9a9e4a8..8378ec3 100644
--- a/replay/replay_test.go
+++ b/replay/replay_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package replay
diff --git a/rwcancel/rwcancel.go b/rwcancel/rwcancel.go
index e397c0e..793e764 100644
--- a/rwcancel/rwcancel.go
+++ b/rwcancel/rwcancel.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
// Package rwcancel implements cancelable read/write operations on
diff --git a/tai64n/tai64n.go b/tai64n/tai64n.go
index 8f10b39..e1a97a5 100644
--- a/tai64n/tai64n.go
+++ b/tai64n/tai64n.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tai64n
diff --git a/tai64n/tai64n_test.go b/tai64n/tai64n_test.go
index c70fc1a..d0b4425 100644
--- a/tai64n/tai64n_test.go
+++ b/tai64n/tai64n_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tai64n
diff --git a/tun/alignment_windows_test.go b/tun/alignment_windows_test.go
index 67a785e..e3252b2 100644
--- a/tun/alignment_windows_test.go
+++ b/tun/alignment_windows_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/netstack/examples/http_client.go b/tun/netstack/examples/http_client.go
index 4c4ea12..8b12ecc 100644
--- a/tun/netstack/examples/http_client.go
+++ b/tun/netstack/examples/http_client.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package main
diff --git a/tun/netstack/examples/http_server.go b/tun/netstack/examples/http_server.go
index 09929e0..80cd036 100644
--- a/tun/netstack/examples/http_server.go
+++ b/tun/netstack/examples/http_server.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package main
diff --git a/tun/netstack/examples/ping_client.go b/tun/netstack/examples/ping_client.go
index d7897b2..b243b5c 100644
--- a/tun/netstack/examples/ping_client.go
+++ b/tun/netstack/examples/ping_client.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package main
diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go
index 2275173..13d1f11 100644
--- a/tun/netstack/tun.go
+++ b/tun/netstack/tun.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package netstack
diff --git a/tun/offload_linux.go b/tun/offload_linux.go
index 89cf024..b61654b 100644
--- a/tun/offload_linux.go
+++ b/tun/offload_linux.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/offload_linux_test.go b/tun/offload_linux_test.go
index a68cd98..c04e003 100644
--- a/tun/offload_linux_test.go
+++ b/tun/offload_linux_test.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/operateonfd.go b/tun/operateonfd.go
index f1beb6d..343f754 100644
--- a/tun/operateonfd.go
+++ b/tun/operateonfd.go
@@ -2,7 +2,7 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/tun.go b/tun/tun.go
index 0ae53d0..336d642 100644
--- a/tun/tun.go
+++ b/tun/tun.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/tun_darwin.go b/tun/tun_darwin.go
index c9a6c0b..407b6f2 100644
--- a/tun/tun_darwin.go
+++ b/tun/tun_darwin.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/tun_freebsd.go b/tun/tun_freebsd.go
index 7c65fd9..4adf3a1 100644
--- a/tun/tun_freebsd.go
+++ b/tun/tun_freebsd.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/tun_linux.go b/tun/tun_linux.go
index 011e56a..bc6e7c1 100644
--- a/tun/tun_linux.go
+++ b/tun/tun_linux.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/tun_openbsd.go b/tun/tun_openbsd.go
index ae571b9..5aa9070 100644
--- a/tun/tun_openbsd.go
+++ b/tun/tun_openbsd.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/tun_windows.go b/tun/tun_windows.go
index 2af8e3e..de65fb4 100644
--- a/tun/tun_windows.go
+++ b/tun/tun_windows.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tun
diff --git a/tun/tuntest/tuntest.go b/tun/tuntest/tuntest.go
index f620e0a..0fa70b0 100644
--- a/tun/tuntest/tuntest.go
+++ b/tun/tuntest/tuntest.go
@@ -1,6 +1,6 @@
/* SPDX-License-Identifier: MIT
*
- * Copyright (C) 2017-2023 WireGuard LLC. All Rights Reserved.
+ * Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
*/
package tuntest
From 704d57c27a6421df4fe5382eb9a524fc446511bd Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Sun, 4 May 2025 17:50:41 +0200
Subject: [PATCH 53/75] mod: bump deps
Signed-off-by: Jason A. Donenfeld
---
go.mod | 8 ++++----
go.sum | 20 ++++++++------------
2 files changed, 12 insertions(+), 16 deletions(-)
diff --git a/go.mod b/go.mod
index 608969f..99569f3 100644
--- a/go.mod
+++ b/go.mod
@@ -4,11 +4,11 @@ go 1.24
require (
github.com/tevino/abool/v2 v2.1.0
- golang.org/x/crypto v0.36.0
- golang.org/x/net v0.37.0
- golang.org/x/sys v0.31.0
+ golang.org/x/crypto v0.37.0
+ golang.org/x/net v0.39.0
+ golang.org/x/sys v0.32.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
- gvisor.dev/gvisor v0.0.0-20250130013005-04f9204697c6
+ gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c
)
require (
diff --git a/go.sum b/go.sum
index 497f949..b8ac0bd 100644
--- a/go.sum
+++ b/go.sum
@@ -1,20 +1,16 @@
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
-github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
-github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
-golang.org/x/crypto v0.36.0 h1:AnAEvhDddvBdpY+uR+MyHmuZzzNqXSe/GvuDeob5L34=
-golang.org/x/crypto v0.36.0/go.mod h1:Y4J0ReaxCR1IMaabaSMugxJES1EpwhBHhv2bDHklZvc=
-golang.org/x/mod v0.21.0 h1:vvrHzRwRfVKSiLrG+d4FMl/Qi4ukBCE6kZlTUkDYRT0=
-golang.org/x/mod v0.21.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
-golang.org/x/net v0.37.0 h1:1zLorHbz+LYj7MQlSf1+2tPIIgibq2eL5xkrGk6f+2c=
-golang.org/x/net v0.37.0/go.mod h1:ivrbrMbzFq5J41QOQh0siUuly180yBYtLp+CKbEaFx8=
-golang.org/x/sys v0.31.0 h1:ioabZlmFYtWhL+TRYpcnNlLwhyxaM9kWTDEmfnprqik=
-golang.org/x/sys v0.31.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
+golang.org/x/crypto v0.37.0 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE=
+golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc=
+golang.org/x/net v0.39.0 h1:ZCu7HMWDxpXpaiKdhzIfaltL9Lp31x/3fCP11bc6/fY=
+golang.org/x/net v0.39.0/go.mod h1:X7NRbYVEA+ewNkCNyJ513WmMdQ3BineSwVtN2zD/d+E=
+golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20=
+golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/time v0.9.0 h1:EsRrnYcQiGH+5FfbgvV4AP7qEZstoyrHB0DzarOQ4ZY=
golang.org/x/time v0.9.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
-gvisor.dev/gvisor v0.0.0-20250130013005-04f9204697c6 h1:6B7MdW3OEbJqOMr7cEYU9bkzvCjUBX/JlXk12xcANuQ=
-gvisor.dev/gvisor v0.0.0-20250130013005-04f9204697c6/go.mod h1:5DMfjtclAbTIjbXqO1qCe2K5GKKxWz2JHvCChuTcJEM=
+gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI=
+gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
From 6a7c878409f32dc39a82bc597766c81304ab9840 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Sun, 4 May 2025 17:54:57 +0200
Subject: [PATCH 54/75] tun/netstack: remove usage of pkt.IsNil()
Since 3c75945fd ("netstack: remove PacketBuffer.IsNil()") this has been
invalid. Follow the replacement pattern of that commit.
The old definition inlined to the same code anyway:
func (pk *PacketBuffer) IsNil() bool {
return pk == nil
}
Signed-off-by: Jason A. Donenfeld
---
tun/netstack/tun.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go
index 13d1f11..2c25649 100644
--- a/tun/netstack/tun.go
+++ b/tun/netstack/tun.go
@@ -155,7 +155,7 @@ func (tun *netTun) Write(buf [][]byte, offset int) (int, error) {
func (tun *netTun) WriteNotify() {
pkt := tun.ep.Read()
- if pkt.IsNil() {
+ if pkt == nil {
return
}
From ac8a885a0361332602c51164aa87da2607146e27 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Mon, 5 May 2025 15:09:09 +0200
Subject: [PATCH 55/75] tun/netstack: cleanup network stack at closing time
Colin's commit went a step further and protected tun.incomingPacket with
a lock on shutdown, but let's see if the tun.stack.Close() call actually
solves that on its own.
Suggested-by: kshangx
Suggested-by: Colin Adler
Signed-off-by: Jason A. Donenfeld
---
tun/netstack/tun.go | 8 +++++---
1 file changed, 5 insertions(+), 3 deletions(-)
diff --git a/tun/netstack/tun.go b/tun/netstack/tun.go
index 2c25649..48a428b 100644
--- a/tun/netstack/tun.go
+++ b/tun/netstack/tun.go
@@ -43,6 +43,7 @@ type netTun struct {
ep *channel.Endpoint
stack *stack.Stack
events chan tun.Event
+ notifyHandle *channel.NotificationHandle
incomingPacket chan *buffer.View
mtu int
dnsServers []netip.Addr
@@ -70,7 +71,7 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int) (tun.Device,
if tcpipErr != nil {
return nil, nil, fmt.Errorf("could not enable TCP SACK: %v", tcpipErr)
}
- dev.ep.AddNotify(dev)
+ dev.notifyHandle = dev.ep.AddNotify(dev)
tcpipErr = dev.stack.CreateNIC(1, dev.ep)
if tcpipErr != nil {
return nil, nil, fmt.Errorf("CreateNIC: %v", tcpipErr)
@@ -167,13 +168,14 @@ func (tun *netTun) WriteNotify() {
func (tun *netTun) Close() error {
tun.stack.RemoveNIC(1)
+ tun.stack.Close()
+ tun.ep.RemoveNotify(tun.notifyHandle)
+ tun.ep.Close()
if tun.events != nil {
close(tun.events)
}
- tun.ep.Close()
-
if tun.incomingPacket != nil {
close(tun.incomingPacket)
}
From 75d6c67a6711190bdee63770f7796446c5bbad01 Mon Sep 17 00:00:00 2001
From: Tu Dinh Ngoc
Date: Thu, 20 Jun 2024 13:28:38 +0000
Subject: [PATCH 56/75] tun: use add-with-carry in checksumNoFold()
Use parallel summation with native byte order per RFC 1071.
add-with-carry operation is used to add 4 words per operation. Byteswap
is performed before and after checksumming for compatibility with old
`checksumNoFold()`. With this we get a 30-80% speedup in `checksum()`
depending on packet sizes.
Add unit tests with comparison to a per-word implementation.
**Intel(R) Xeon(R) Silver 4210R CPU @ 2.40GHz**
| Size | OldTime | NewTime | Speedup |
|------|---------|---------|----------|
| 64 | 12.64 | 9.183 | 1.376456 |
| 128 | 18.52 | 12.72 | 1.455975 |
| 256 | 31.01 | 18.13 | 1.710425 |
| 512 | 54.46 | 29.03 | 1.87599 |
| 1024 | 102 | 52.2 | 1.954023 |
| 1500 | 146.8 | 81.36 | 1.804326 |
| 2048 | 196.9 | 102.5 | 1.920976 |
| 4096 | 389.8 | 200.8 | 1.941235 |
| 8192 | 767.3 | 413.3 | 1.856521 |
| 9000 | 851.7 | 448.8 | 1.897727 |
| 9001 | 854.8 | 451.9 | 1.891569 |
**AMD EPYC 7352 24-Core Processor**
| Size | OldTime | NewTime | Speedup |
|------|---------|---------|----------|
| 64 | 9.159 | 6.949 | 1.318031 |
| 128 | 13.59 | 10.59 | 1.283286 |
| 256 | 22.37 | 14.91 | 1.500335 |
| 512 | 41.42 | 24.22 | 1.710157 |
| 1024 | 81.59 | 45.05 | 1.811099 |
| 1500 | 120.4 | 68.35 | 1.761522 |
| 2048 | 162.8 | 90.14 | 1.806079 |
| 4096 | 321.4 | 180.3 | 1.782585 |
| 8192 | 650.4 | 360.8 | 1.802661 |
| 9000 | 706.3 | 398.1 | 1.774177 |
| 9001 | 712.4 | 398.2 | 1.789051 |
Signed-off-by: Tu Dinh Ngoc
[Jason: simplified and cleaned up unit tests]
Signed-off-by: Jason A. Donenfeld
---
tun/checksum.go | 122 +++++++++++++++++++------------------------
tun/checksum_test.go | 63 ++++++++++++++++++++++
2 files changed, 116 insertions(+), 69 deletions(-)
diff --git a/tun/checksum.go b/tun/checksum.go
index 29a8fc8..b489c56 100644
--- a/tun/checksum.go
+++ b/tun/checksum.go
@@ -1,102 +1,86 @@
package tun
-import "encoding/binary"
+import (
+ "encoding/binary"
+ "math/bits"
+)
// TODO: Explore SIMD and/or other assembly optimizations.
-// TODO: Test native endian loads. See RFC 1071 section 2 part B.
func checksumNoFold(b []byte, initial uint64) uint64 {
- ac := initial
+ tmp := make([]byte, 8)
+ binary.NativeEndian.PutUint64(tmp, initial)
+ ac := binary.BigEndian.Uint64(tmp)
+ var carry uint64
for len(b) >= 128 {
- ac += uint64(binary.BigEndian.Uint32(b[:4]))
- ac += uint64(binary.BigEndian.Uint32(b[4:8]))
- ac += uint64(binary.BigEndian.Uint32(b[8:12]))
- ac += uint64(binary.BigEndian.Uint32(b[12:16]))
- ac += uint64(binary.BigEndian.Uint32(b[16:20]))
- ac += uint64(binary.BigEndian.Uint32(b[20:24]))
- ac += uint64(binary.BigEndian.Uint32(b[24:28]))
- ac += uint64(binary.BigEndian.Uint32(b[28:32]))
- ac += uint64(binary.BigEndian.Uint32(b[32:36]))
- ac += uint64(binary.BigEndian.Uint32(b[36:40]))
- ac += uint64(binary.BigEndian.Uint32(b[40:44]))
- ac += uint64(binary.BigEndian.Uint32(b[44:48]))
- ac += uint64(binary.BigEndian.Uint32(b[48:52]))
- ac += uint64(binary.BigEndian.Uint32(b[52:56]))
- ac += uint64(binary.BigEndian.Uint32(b[56:60]))
- ac += uint64(binary.BigEndian.Uint32(b[60:64]))
- ac += uint64(binary.BigEndian.Uint32(b[64:68]))
- ac += uint64(binary.BigEndian.Uint32(b[68:72]))
- ac += uint64(binary.BigEndian.Uint32(b[72:76]))
- ac += uint64(binary.BigEndian.Uint32(b[76:80]))
- ac += uint64(binary.BigEndian.Uint32(b[80:84]))
- ac += uint64(binary.BigEndian.Uint32(b[84:88]))
- ac += uint64(binary.BigEndian.Uint32(b[88:92]))
- ac += uint64(binary.BigEndian.Uint32(b[92:96]))
- ac += uint64(binary.BigEndian.Uint32(b[96:100]))
- ac += uint64(binary.BigEndian.Uint32(b[100:104]))
- ac += uint64(binary.BigEndian.Uint32(b[104:108]))
- ac += uint64(binary.BigEndian.Uint32(b[108:112]))
- ac += uint64(binary.BigEndian.Uint32(b[112:116]))
- ac += uint64(binary.BigEndian.Uint32(b[116:120]))
- ac += uint64(binary.BigEndian.Uint32(b[120:124]))
- ac += uint64(binary.BigEndian.Uint32(b[124:128]))
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[16:24]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[24:32]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[32:40]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[40:48]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[48:56]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[56:64]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[64:72]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[72:80]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[80:88]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[88:96]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[96:104]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[104:112]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[112:120]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[120:128]), carry)
+ ac += carry
b = b[128:]
}
if len(b) >= 64 {
- ac += uint64(binary.BigEndian.Uint32(b[:4]))
- ac += uint64(binary.BigEndian.Uint32(b[4:8]))
- ac += uint64(binary.BigEndian.Uint32(b[8:12]))
- ac += uint64(binary.BigEndian.Uint32(b[12:16]))
- ac += uint64(binary.BigEndian.Uint32(b[16:20]))
- ac += uint64(binary.BigEndian.Uint32(b[20:24]))
- ac += uint64(binary.BigEndian.Uint32(b[24:28]))
- ac += uint64(binary.BigEndian.Uint32(b[28:32]))
- ac += uint64(binary.BigEndian.Uint32(b[32:36]))
- ac += uint64(binary.BigEndian.Uint32(b[36:40]))
- ac += uint64(binary.BigEndian.Uint32(b[40:44]))
- ac += uint64(binary.BigEndian.Uint32(b[44:48]))
- ac += uint64(binary.BigEndian.Uint32(b[48:52]))
- ac += uint64(binary.BigEndian.Uint32(b[52:56]))
- ac += uint64(binary.BigEndian.Uint32(b[56:60]))
- ac += uint64(binary.BigEndian.Uint32(b[60:64]))
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[16:24]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[24:32]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[32:40]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[40:48]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[48:56]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[56:64]), carry)
+ ac += carry
b = b[64:]
}
if len(b) >= 32 {
- ac += uint64(binary.BigEndian.Uint32(b[:4]))
- ac += uint64(binary.BigEndian.Uint32(b[4:8]))
- ac += uint64(binary.BigEndian.Uint32(b[8:12]))
- ac += uint64(binary.BigEndian.Uint32(b[12:16]))
- ac += uint64(binary.BigEndian.Uint32(b[16:20]))
- ac += uint64(binary.BigEndian.Uint32(b[20:24]))
- ac += uint64(binary.BigEndian.Uint32(b[24:28]))
- ac += uint64(binary.BigEndian.Uint32(b[28:32]))
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[16:24]), carry)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[24:32]), carry)
+ ac += carry
b = b[32:]
}
if len(b) >= 16 {
- ac += uint64(binary.BigEndian.Uint32(b[:4]))
- ac += uint64(binary.BigEndian.Uint32(b[4:8]))
- ac += uint64(binary.BigEndian.Uint32(b[8:12]))
- ac += uint64(binary.BigEndian.Uint32(b[12:16]))
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0)
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[8:16]), carry)
+ ac += carry
b = b[16:]
}
if len(b) >= 8 {
- ac += uint64(binary.BigEndian.Uint32(b[:4]))
- ac += uint64(binary.BigEndian.Uint32(b[4:8]))
+ ac, carry = bits.Add64(ac, binary.NativeEndian.Uint64(b[:8]), 0)
+ ac += carry
b = b[8:]
}
if len(b) >= 4 {
- ac += uint64(binary.BigEndian.Uint32(b))
+ ac, carry = bits.Add64(ac, uint64(binary.NativeEndian.Uint32(b[:4])), 0)
+ ac += carry
b = b[4:]
}
if len(b) >= 2 {
- ac += uint64(binary.BigEndian.Uint16(b))
+ ac, carry = bits.Add64(ac, uint64(binary.NativeEndian.Uint16(b[:2])), 0)
+ ac += carry
b = b[2:]
}
if len(b) == 1 {
- ac += uint64(b[0]) << 8
+ tmp := binary.NativeEndian.Uint16([]byte{b[0], 0})
+ ac, carry = bits.Add64(ac, uint64(tmp), 0)
+ ac += carry
}
- return ac
+ binary.NativeEndian.PutUint64(tmp, ac)
+ return binary.BigEndian.Uint64(tmp)
}
func checksum(b []byte, initial uint64) uint16 {
diff --git a/tun/checksum_test.go b/tun/checksum_test.go
index c1ccff5..4ea9b8b 100644
--- a/tun/checksum_test.go
+++ b/tun/checksum_test.go
@@ -1,11 +1,74 @@
package tun
import (
+ "encoding/binary"
"fmt"
"math/rand"
"testing"
+
+ "golang.org/x/sys/unix"
)
+func checksumRef(b []byte, initial uint16) uint16 {
+ ac := uint64(initial)
+
+ for len(b) >= 2 {
+ ac += uint64(binary.BigEndian.Uint16(b))
+ b = b[2:]
+ }
+ if len(b) == 1 {
+ ac += uint64(b[0]) << 8
+ }
+
+ for (ac >> 16) > 0 {
+ ac = (ac >> 16) + (ac & 0xffff)
+ }
+ return uint16(ac)
+}
+
+func pseudoHeaderChecksumRefNoFold(protocol uint8, srcAddr, dstAddr []byte, totalLen uint16) uint16 {
+ sum := checksumRef(srcAddr, 0)
+ sum = checksumRef(dstAddr, sum)
+ sum = checksumRef([]byte{0, protocol}, sum)
+ tmp := make([]byte, 2)
+ binary.BigEndian.PutUint16(tmp, totalLen)
+ return checksumRef(tmp, sum)
+}
+
+func TestChecksum(t *testing.T) {
+ for length := 0; length <= 9001; length++ {
+ buf := make([]byte, length)
+ rng := rand.New(rand.NewSource(1))
+ rng.Read(buf)
+ csum := checksum(buf, 0x1234)
+ csumRef := checksumRef(buf, 0x1234)
+ if csum != csumRef {
+ t.Error("Expected checksum", csumRef, "got", csum)
+ }
+ }
+}
+
+func TestPseudoHeaderChecksum(t *testing.T) {
+ for _, addrLen := range []int{4, 16} {
+ for length := 0; length <= 9001; length++ {
+ srcAddr := make([]byte, addrLen)
+ dstAddr := make([]byte, addrLen)
+ buf := make([]byte, length)
+ rng := rand.New(rand.NewSource(1))
+ rng.Read(srcAddr)
+ rng.Read(dstAddr)
+ rng.Read(buf)
+ phSum := pseudoHeaderChecksumNoFold(unix.IPPROTO_TCP, srcAddr, dstAddr, uint16(length))
+ csum := checksum(buf, phSum)
+ phSumRef := pseudoHeaderChecksumRefNoFold(unix.IPPROTO_TCP, srcAddr, dstAddr, uint16(length))
+ csumRef := checksumRef(buf, phSumRef)
+ if csum != csumRef {
+ t.Error("Expected checksumRef", csumRef, "got", csum)
+ }
+ }
+ }
+}
+
func BenchmarkChecksum(b *testing.B) {
lengths := []int{
64,
From 8a2b2bf4f49f56ef379dfb2104d1d312e4479182 Mon Sep 17 00:00:00 2001
From: ruokeqx
Date: Thu, 2 Jan 2025 20:28:33 +0800
Subject: [PATCH 57/75] tun: darwin: fetch flags and mtu from if_msghdr
directly
Signed-off-by: ruokeqx
Signed-off-by: Jason A. Donenfeld
---
tun/tun_darwin.go | 34 +++++++++-------------------------
1 file changed, 9 insertions(+), 25 deletions(-)
diff --git a/tun/tun_darwin.go b/tun/tun_darwin.go
index 407b6f2..341afe3 100644
--- a/tun/tun_darwin.go
+++ b/tun/tun_darwin.go
@@ -6,14 +6,12 @@
package tun
import (
- "errors"
"fmt"
"io"
"net"
"os"
"sync"
"syscall"
- "time"
"unsafe"
"golang.org/x/sys/unix"
@@ -30,18 +28,6 @@ type NativeTun struct {
closeOnce sync.Once
}
-func retryInterfaceByIndex(index int) (iface *net.Interface, err error) {
- for i := 0; i < 20; i++ {
- iface, err = net.InterfaceByIndex(index)
- if err != nil && errors.Is(err, unix.ENOMEM) {
- time.Sleep(time.Duration(i) * time.Second / 3)
- continue
- }
- return iface, err
- }
- return nil, err
-}
-
func (tun *NativeTun) routineRouteListener(tunIfindex int) {
var (
statusUp bool
@@ -62,26 +48,22 @@ func (tun *NativeTun) routineRouteListener(tunIfindex int) {
return
}
- if n < 14 {
+ if n < 28 {
continue
}
- if data[3 /* type */] != unix.RTM_IFINFO {
+ if data[3 /* ifm_type */] != unix.RTM_IFINFO {
continue
}
- ifindex := int(*(*uint16)(unsafe.Pointer(&data[12 /* ifindex */])))
+ ifindex := int(*(*uint16)(unsafe.Pointer(&data[12 /* ifm_index */])))
if ifindex != tunIfindex {
continue
}
- iface, err := retryInterfaceByIndex(ifindex)
- if err != nil {
- tun.errors <- err
- return
- }
+ flags := int(*(*uint32)(unsafe.Pointer(&data[8 /* ifm_flags */])))
// Up / Down event
- up := (iface.Flags & net.FlagUp) != 0
+ up := (flags & syscall.IFF_UP) != 0
if up != statusUp && up {
tun.events <- EventUp
}
@@ -90,11 +72,13 @@ func (tun *NativeTun) routineRouteListener(tunIfindex int) {
}
statusUp = up
+ mtu := int(*(*uint32)(unsafe.Pointer(&data[24 /* ifm_data.ifi_mtu */])))
+
// MTU changes
- if iface.MTU != statusMTU {
+ if mtu != statusMTU {
tun.events <- EventMTUUpdate
}
- statusMTU = iface.MTU
+ statusMTU = mtu
}
}
From ace3e11ef24195c2670e619d0a743c13af89ecf4 Mon Sep 17 00:00:00 2001
From: Tom Holford
Date: Sun, 4 May 2025 18:49:03 +0200
Subject: [PATCH 58/75] global: replaced unused function params with _
Signed-off-by: Jason A. Donenfeld
---
conn/errors_default.go | 2 +-
conn/features_default.go | 2 +-
device/allowedips_test.go | 2 +-
device/sticky_default.go | 2 +-
device/sticky_linux.go | 4 ++--
5 files changed, 6 insertions(+), 6 deletions(-)
diff --git a/conn/errors_default.go b/conn/errors_default.go
index d967518..3c9b223 100644
--- a/conn/errors_default.go
+++ b/conn/errors_default.go
@@ -7,6 +7,6 @@
package conn
-func errShouldDisableUDPGSO(err error) bool {
+func errShouldDisableUDPGSO(_ error) bool {
return false
}
diff --git a/conn/features_default.go b/conn/features_default.go
index cae2bea..9fc5088 100644
--- a/conn/features_default.go
+++ b/conn/features_default.go
@@ -10,6 +10,6 @@ package conn
import "net"
-func supportsUDPOffload(conn *net.UDPConn) (txOffload, rxOffload bool) {
+func supportsUDPOffload(_ *net.UDPConn) (txOffload, rxOffload bool) {
return
}
diff --git a/device/allowedips_test.go b/device/allowedips_test.go
index 9ef8a76..0ce45af 100644
--- a/device/allowedips_test.go
+++ b/device/allowedips_test.go
@@ -39,7 +39,7 @@ func TestCommonBits(t *testing.T) {
}
}
-func benchmarkTrie(peerNumber, addressNumber, addressLength int, b *testing.B) {
+func benchmarkTrie(peerNumber, addressNumber, _ int, b *testing.B) {
var trie *trieEntry
var peers []*Peer
root := parentIndirection{&trie, 2}
diff --git a/device/sticky_default.go b/device/sticky_default.go
index da776e8..1751927 100644
--- a/device/sticky_default.go
+++ b/device/sticky_default.go
@@ -7,6 +7,6 @@ import (
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
)
-func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, error) {
+func (device *Device) startRouteListener(_ conn.Bind) (*rwcancel.RWCancel, error) {
return nil, nil
}
diff --git a/device/sticky_linux.go b/device/sticky_linux.go
index 5ff9dd6..2edb628 100644
--- a/device/sticky_linux.go
+++ b/device/sticky_linux.go
@@ -9,7 +9,7 @@
*
* Currently there is no way to achieve this within the net package:
* See e.g. https://github.com/golang/go/issues/17930
- * So this code is remains platform dependent.
+ * So this code remains platform dependent.
*/
package device
@@ -47,7 +47,7 @@ func (device *Device) startRouteListener(bind conn.Bind) (*rwcancel.RWCancel, er
return netlinkCancel, nil
}
-func (device *Device) routineRouteListener(bind conn.Bind, netlinkSock int, netlinkCancel *rwcancel.RWCancel) {
+func (device *Device) routineRouteListener(_ conn.Bind, netlinkSock int, netlinkCancel *rwcancel.RWCancel) {
type peerEndpointPtr struct {
peer *Peer
endpoint *conn.Endpoint
From 8051f1714771201e1c1daccfb93f0a0847a29b21 Mon Sep 17 00:00:00 2001
From: Tom Holford
Date: Sun, 4 May 2025 18:49:49 +0200
Subject: [PATCH 59/75] device: use rand.NewSource instead of rand.Seed
Signed-off-by: Jason A. Donenfeld
---
device/allowedips_rand_test.go | 10 +++++-----
device/allowedips_test.go | 10 +++++-----
2 files changed, 10 insertions(+), 10 deletions(-)
diff --git a/device/allowedips_rand_test.go b/device/allowedips_rand_test.go
index 8dd9b67..b863696 100644
--- a/device/allowedips_rand_test.go
+++ b/device/allowedips_rand_test.go
@@ -83,7 +83,7 @@ func TestTrieRandom(t *testing.T) {
var peers []*Peer
var allowedIPs AllowedIPs
- rand.Seed(1)
+ rng := rand.New(rand.NewSource(1))
for n := 0; n < NumberOfPeers; n++ {
peers = append(peers, &Peer{})
@@ -91,14 +91,14 @@ func TestTrieRandom(t *testing.T) {
for n := 0; n < NumberOfAddresses; n++ {
var addr4 [4]byte
- rand.Read(addr4[:])
+ rng.Read(addr4[:])
cidr := uint8(rand.Intn(32) + 1)
index := rand.Intn(NumberOfPeers)
allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom4(addr4), int(cidr)), peers[index])
slow4 = slow4.Insert(addr4[:], cidr, peers[index])
var addr6 [16]byte
- rand.Read(addr6[:])
+ rng.Read(addr6[:])
cidr = uint8(rand.Intn(128) + 1)
index = rand.Intn(NumberOfPeers)
allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom16(addr6), int(cidr)), peers[index])
@@ -109,7 +109,7 @@ func TestTrieRandom(t *testing.T) {
for p = 0; ; p++ {
for n := 0; n < NumberOfTests; n++ {
var addr4 [4]byte
- rand.Read(addr4[:])
+ rng.Read(addr4[:])
peer1 := slow4.Lookup(addr4[:])
peer2 := allowedIPs.Lookup(addr4[:])
if peer1 != peer2 {
@@ -117,7 +117,7 @@ func TestTrieRandom(t *testing.T) {
}
var addr6 [16]byte
- rand.Read(addr6[:])
+ rng.Read(addr6[:])
peer1 = slow6.Lookup(addr6[:])
peer2 = allowedIPs.Lookup(addr6[:])
if peer1 != peer2 {
diff --git a/device/allowedips_test.go b/device/allowedips_test.go
index 0ce45af..7df7da5 100644
--- a/device/allowedips_test.go
+++ b/device/allowedips_test.go
@@ -44,7 +44,7 @@ func benchmarkTrie(peerNumber, addressNumber, _ int, b *testing.B) {
var peers []*Peer
root := parentIndirection{&trie, 2}
- rand.Seed(1)
+ rng := rand.New(rand.NewSource(1))
const AddressLength = 4
@@ -54,15 +54,15 @@ func benchmarkTrie(peerNumber, addressNumber, _ int, b *testing.B) {
for n := 0; n < addressNumber; n++ {
var addr [AddressLength]byte
- rand.Read(addr[:])
- cidr := uint8(rand.Uint32() % (AddressLength * 8))
- index := rand.Int() % peerNumber
+ rng.Read(addr[:])
+ cidr := uint8(rng.Uint32() % (AddressLength * 8))
+ index := rng.Int() % peerNumber
root.insert(addr[:], cidr, peers[index])
}
for n := 0; n < b.N; n++ {
var addr [AddressLength]byte
- rand.Read(addr[:])
+ rng.Read(addr[:])
trie.lookup(addr[:])
}
}
From 2cad62c40bca27495120f9a5c3c5bff795124621 Mon Sep 17 00:00:00 2001
From: Kurnia D Win
Date: Wed, 7 Jun 2023 12:41:02 +0700
Subject: [PATCH 60/75] rwcancel: fix wrong poll event flag on ReadyWrite
It should be POLLIN because closeFd is read-only file.
Signed-off-by: Kurnia D Win
Signed-off-by: Jason A. Donenfeld
---
rwcancel/rwcancel.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/rwcancel/rwcancel.go b/rwcancel/rwcancel.go
index 793e764..4372453 100644
--- a/rwcancel/rwcancel.go
+++ b/rwcancel/rwcancel.go
@@ -64,7 +64,7 @@ func (rw *RWCancel) ReadyRead() bool {
func (rw *RWCancel) ReadyWrite() bool {
closeFd := int32(rw.closingReader.Fd())
- pollFds := []unix.PollFd{{Fd: int32(rw.fd), Events: unix.POLLOUT}, {Fd: closeFd, Events: unix.POLLOUT}}
+ pollFds := []unix.PollFd{{Fd: int32(rw.fd), Events: unix.POLLOUT}, {Fd: closeFd, Events: unix.POLLIN}}
var err error
for {
_, err = unix.Poll(pollFds, -1)
From 676809066782e22df9e17013169bb1e90fb653bb Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Thu, 15 May 2025 16:54:03 +0200
Subject: [PATCH 61/75] version: bump snapshot
Signed-off-by: Jason A. Donenfeld
---
version.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/version.go b/version.go
index db75bb9..80f2d4b 100644
--- a/version.go
+++ b/version.go
@@ -1,3 +1,3 @@
package main
-const Version = "0.0.20230223"
+const Version = "0.0.20250515"
From d5359f52f098b5a500ff348a577cff2d1321422c Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Tue, 20 May 2025 23:03:06 +0200
Subject: [PATCH 62/75] device: add support for removing allowedips
individually
This pairs with the recent change in wireguard-tools.
Signed-off-by: Jason A. Donenfeld
---
device/allowedips.go | 87 +++++++++++++++++++++++++--------------
device/allowedips_test.go | 57 +++++++++++++++++++++++++
device/uapi.go | 15 ++++++-
3 files changed, 125 insertions(+), 34 deletions(-)
diff --git a/device/allowedips.go b/device/allowedips.go
index b40c817..d15373c 100644
--- a/device/allowedips.go
+++ b/device/allowedips.go
@@ -223,6 +223,60 @@ func (table *AllowedIPs) EntriesForPeer(peer *Peer, cb func(prefix netip.Prefix)
}
}
+func (node *trieEntry) remove() {
+ node.removeFromPeerEntries()
+ node.peer = nil
+ if node.child[0] != nil && node.child[1] != nil {
+ return
+ }
+ bit := 0
+ if node.child[0] == nil {
+ bit = 1
+ }
+ child := node.child[bit]
+ if child != nil {
+ child.parent = node.parent
+ }
+ *node.parent.parentBit = child
+ if node.child[0] != nil || node.child[1] != nil || node.parent.parentBitType > 1 {
+ node.zeroizePointers()
+ return
+ }
+ parent := (*trieEntry)(unsafe.Pointer(uintptr(unsafe.Pointer(node.parent.parentBit)) - unsafe.Offsetof(node.child) - unsafe.Sizeof(node.child[0])*uintptr(node.parent.parentBitType)))
+ if parent.peer != nil {
+ node.zeroizePointers()
+ return
+ }
+ child = parent.child[node.parent.parentBitType^1]
+ if child != nil {
+ child.parent = parent.parent
+ }
+ *parent.parent.parentBit = child
+ node.zeroizePointers()
+ parent.zeroizePointers()
+}
+
+func (table *AllowedIPs) Remove(prefix netip.Prefix, peer *Peer) {
+ table.mutex.Lock()
+ defer table.mutex.Unlock()
+ var node *trieEntry
+ var exact bool
+
+ if prefix.Addr().Is6() {
+ ip := prefix.Addr().As16()
+ node, exact = table.IPv6.nodePlacement(ip[:], uint8(prefix.Bits()))
+ } else if prefix.Addr().Is4() {
+ ip := prefix.Addr().As4()
+ node, exact = table.IPv4.nodePlacement(ip[:], uint8(prefix.Bits()))
+ } else {
+ panic(errors.New("removing unknown address type"))
+ }
+ if !exact || node == nil || peer != node.peer {
+ return
+ }
+ node.remove()
+}
+
func (table *AllowedIPs) RemoveByPeer(peer *Peer) {
table.mutex.Lock()
defer table.mutex.Unlock()
@@ -230,38 +284,7 @@ func (table *AllowedIPs) RemoveByPeer(peer *Peer) {
var next *list.Element
for elem := peer.trieEntries.Front(); elem != nil; elem = next {
next = elem.Next()
- node := elem.Value.(*trieEntry)
-
- node.removeFromPeerEntries()
- node.peer = nil
- if node.child[0] != nil && node.child[1] != nil {
- continue
- }
- bit := 0
- if node.child[0] == nil {
- bit = 1
- }
- child := node.child[bit]
- if child != nil {
- child.parent = node.parent
- }
- *node.parent.parentBit = child
- if node.child[0] != nil || node.child[1] != nil || node.parent.parentBitType > 1 {
- node.zeroizePointers()
- continue
- }
- parent := (*trieEntry)(unsafe.Pointer(uintptr(unsafe.Pointer(node.parent.parentBit)) - unsafe.Offsetof(node.child) - unsafe.Sizeof(node.child[0])*uintptr(node.parent.parentBitType)))
- if parent.peer != nil {
- node.zeroizePointers()
- continue
- }
- child = parent.child[node.parent.parentBitType^1]
- if child != nil {
- child.parent = parent.parent
- }
- *parent.parent.parentBit = child
- node.zeroizePointers()
- parent.zeroizePointers()
+ elem.Value.(*trieEntry).remove()
}
}
diff --git a/device/allowedips_test.go b/device/allowedips_test.go
index 7df7da5..a4b08a3 100644
--- a/device/allowedips_test.go
+++ b/device/allowedips_test.go
@@ -101,6 +101,10 @@ func TestTrieIPv4(t *testing.T) {
allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom4([4]byte{a, b, c, d}), int(cidr)), peer)
}
+ remove := func(peer *Peer, a, b, c, d byte, cidr uint8) {
+ allowedIPs.Remove(netip.PrefixFrom(netip.AddrFrom4([4]byte{a, b, c, d}), int(cidr)), peer)
+ }
+
assertEQ := func(peer *Peer, a, b, c, d byte) {
p := allowedIPs.Lookup([]byte{a, b, c, d})
if p != peer {
@@ -176,6 +180,21 @@ func TestTrieIPv4(t *testing.T) {
allowedIPs.RemoveByPeer(a)
assertNEQ(a, 192, 168, 0, 1)
+
+ insert(a, 1, 0, 0, 0, 32)
+ insert(a, 192, 0, 0, 0, 24)
+ assertEQ(a, 1, 0, 0, 0)
+ assertEQ(a, 192, 0, 0, 1)
+ remove(a, 192, 0, 0, 0, 32)
+ assertEQ(a, 192, 0, 0, 1)
+ remove(nil, 192, 0, 0, 0, 24)
+ assertEQ(a, 192, 0, 0, 1)
+ remove(b, 192, 0, 0, 0, 24)
+ assertEQ(a, 192, 0, 0, 1)
+ remove(a, 192, 0, 0, 0, 24)
+ assertNEQ(a, 192, 0, 0, 1)
+ remove(a, 1, 0, 0, 0, 32)
+ assertNEQ(a, 1, 0, 0, 0)
}
/* Test ported from kernel implementation:
@@ -211,6 +230,15 @@ func TestTrieIPv6(t *testing.T) {
allowedIPs.Insert(netip.PrefixFrom(netip.AddrFrom16(*(*[16]byte)(addr)), int(cidr)), peer)
}
+ remove := func(peer *Peer, a, b, c, d uint32, cidr uint8) {
+ var addr []byte
+ addr = append(addr, expand(a)...)
+ addr = append(addr, expand(b)...)
+ addr = append(addr, expand(c)...)
+ addr = append(addr, expand(d)...)
+ allowedIPs.Remove(netip.PrefixFrom(netip.AddrFrom16(*(*[16]byte)(addr)), int(cidr)), peer)
+ }
+
assertEQ := func(peer *Peer, a, b, c, d uint32) {
var addr []byte
addr = append(addr, expand(a)...)
@@ -223,6 +251,18 @@ func TestTrieIPv6(t *testing.T) {
}
}
+ assertNEQ := func(peer *Peer, a, b, c, d uint32) {
+ var addr []byte
+ addr = append(addr, expand(a)...)
+ addr = append(addr, expand(b)...)
+ addr = append(addr, expand(c)...)
+ addr = append(addr, expand(d)...)
+ p := allowedIPs.Lookup(addr)
+ if p == peer {
+ t.Error("Assert NEQ failed")
+ }
+ }
+
insert(d, 0x26075300, 0x60006b00, 0, 0xc05f0543, 128)
insert(c, 0x26075300, 0x60006b00, 0, 0, 64)
insert(e, 0, 0, 0, 0, 0)
@@ -244,4 +284,21 @@ func TestTrieIPv6(t *testing.T) {
assertEQ(h, 0x24046800, 0x40040800, 0, 0)
assertEQ(h, 0x24046800, 0x40040800, 0x10101010, 0x10101010)
assertEQ(a, 0x24046800, 0x40040800, 0xdeadbeef, 0xdeadbeef)
+
+ insert(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128)
+ insert(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0, 98)
+ assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef)
+ assertEQ(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0x10101010)
+ remove(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 96)
+ assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef)
+ remove(nil, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128)
+ assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef)
+ remove(b, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128)
+ assertEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef)
+ remove(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef, 128)
+ assertNEQ(a, 0x24446801, 0x40e40800, 0xdeaebeef, 0xdefbeef)
+ remove(b, 0x24446800, 0xf0e40800, 0xeeaebeef, 0, 98)
+ assertEQ(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0x10101010)
+ remove(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0, 98)
+ assertNEQ(a, 0x24446800, 0xf0e40800, 0xeeaebeef, 0x10101010)
}
diff --git a/device/uapi.go b/device/uapi.go
index 1b5e357..870bddc 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -497,7 +497,14 @@ func (device *Device) handlePeerLine(
device.allowedips.RemoveByPeer(peer.Peer)
case "allowed_ip":
- device.log.Verbosef("%v - UAPI: Adding allowedip", peer.Peer)
+ add := true
+ verb := "Adding"
+ if len(value) > 0 && value[0] == '-' {
+ add = false
+ verb = "Removing"
+ value = value[1:]
+ }
+ device.log.Verbosef("%v - UAPI: %s allowedip", peer.Peer, verb)
prefix, err := netip.ParsePrefix(value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "failed to set allowed ip: %w", err)
@@ -505,7 +512,11 @@ func (device *Device) handlePeerLine(
if peer.dummy {
return nil
}
- device.allowedips.Insert(prefix, peer.Peer)
+ if add {
+ device.allowedips.Insert(prefix, peer.Peer)
+ } else {
+ device.allowedips.Remove(prefix, peer.Peer)
+ }
case "protocol_version":
if value != "1" {
From 99f2e6d66f79dfc087bb6957738149900b44616d Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Thu, 22 May 2025 01:33:55 +0200
Subject: [PATCH 63/75] conn: don't enable GRO on Linux < 5.12
Kernels below 5.12 are missing this:
commit 98184612aca0a9ee42b8eb0262a49900ee9eef0d
Author: Norman Maurer
Date: Thu Apr 1 08:59:17 2021
net: udp: Add support for getsockopt(..., ..., UDP_GRO, ..., ...);
Support for UDP_GRO was added in the past but the implementation for
getsockopt was missed which did lead to an error when we tried to
retrieve the setting for UDP_GRO. This patch adds the missing switch
case for UDP_GRO
Fixes: e20cf8d3f1f7 ("udp: implement GRO for plain UDP sockets.")
Signed-off-by: Norman Maurer
Reviewed-by: David Ahern
Signed-off-by: David S. Miller
That means we can't set the option and then read it back later. Given
how buggy UDP_GRO is in general on odd kernels, just disable it on older
kernels all together.
Signed-off-by: Jason A. Donenfeld
---
conn/controlfns_linux.go | 48 ++++++++++++++++++++++++++++++++++++++++
1 file changed, 48 insertions(+)
diff --git a/conn/controlfns_linux.go b/conn/controlfns_linux.go
index 7bd3917..f0deefa 100644
--- a/conn/controlfns_linux.go
+++ b/conn/controlfns_linux.go
@@ -13,6 +13,35 @@ import (
"golang.org/x/sys/unix"
)
+// Taken from go/src/internal/syscall/unix/kernel_version_linux.go
+func kernelVersion() (major, minor int) {
+ var uname unix.Utsname
+ if err := unix.Uname(&uname); err != nil {
+ return
+ }
+
+ var (
+ values [2]int
+ value, vi int
+ )
+ for _, c := range uname.Release {
+ if '0' <= c && c <= '9' {
+ value = (value * 10) + int(c-'0')
+ } else {
+ // Note that we're assuming N.N.N here.
+ // If we see anything else, we are likely to mis-parse it.
+ values[vi] = value
+ vi++
+ if vi >= len(values) {
+ break
+ }
+ value = 0
+ }
+ }
+
+ return values[0], values[1]
+}
+
func init() {
controlFns = append(controlFns,
@@ -57,5 +86,24 @@ func init() {
}
return err
},
+
+ // Attempt to enable UDP_GRO
+ func(network, address string, c syscall.RawConn) error {
+ // Kernels below 5.12 are missing 98184612aca0 ("net:
+ // udp: Add support for getsockopt(..., ..., UDP_GRO,
+ // ..., ...);"), which means we can't read this back
+ // later. We could pipe the return value through to
+ // the rest of the code, but UDP_GRO is kind of buggy
+ // anyway, so just gate this here.
+ major, minor := kernelVersion()
+ if major < 5 || (major == 5 && minor < 12) {
+ return nil
+ }
+
+ c.Control(func(fd uintptr) {
+ _ = unix.SetsockoptInt(int(fd), unix.IPPROTO_UDP, unix.UDP_GRO, 1)
+ })
+ return nil
+ },
)
}
From eeb8aae13eedffd851789ee56f1c80d7cb4382e3 Mon Sep 17 00:00:00 2001
From: "Jason A. Donenfeld"
Date: Thu, 22 May 2025 01:45:02 +0200
Subject: [PATCH 64/75] version: bump snapshot
Signed-off-by: Jason A. Donenfeld
---
version.go | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/version.go b/version.go
index 80f2d4b..d5524e8 100644
--- a/version.go
+++ b/version.go
@@ -1,3 +1,3 @@
package main
-const Version = "0.0.20250515"
+const Version = "0.0.20250522"
From 169ed49a469bf5b05775bc1102174fd00e20def7 Mon Sep 17 00:00:00 2001
From: jmwample
Date: Mon, 23 Jun 2025 14:37:49 -0600
Subject: [PATCH 65/75] fix formatting discrepancy
---
device/device.go | 6 +++---
1 file changed, 3 insertions(+), 3 deletions(-)
diff --git a/device/device.go b/device/device.go
index 2a37321..124b74e 100644
--- a/device/device.go
+++ b/device/device.go
@@ -92,9 +92,9 @@ type Device struct {
closed chan struct{}
log *Logger
- isASecOn abool.AtomicBool
- aSecMux sync.RWMutex
- aSecCfg aSecCfgType
+ isASecOn abool.AtomicBool
+ aSecMux sync.RWMutex
+ aSecCfg aSecCfgType
junkCreator junkCreator
}
From c20789848019fb494dbe9d280eb246f29b95ab85 Mon Sep 17 00:00:00 2001
From: Mykola Baibuz
Date: Mon, 7 Jul 2025 05:34:51 -0700
Subject: [PATCH 66/75] AmneziaWG v1.5 (#84)
---
Dockerfile | 20 +-
device/awg/awg.go | 144 ++++++
device/awg/internal/mock.go | 37 ++
device/{ => awg}/junk_creator.go | 35 +-
device/{ => awg}/junk_creator_test.go | 59 ++-
device/awg/special_handshake_handler.go | 73 ++++
device/awg/tag_generator.go | 190 ++++++++
device/awg/tag_generator_test.go | 189 ++++++++
device/awg/tag_junk_packet_generator.go | 59 +++
device/awg/tag_junk_packet_generator_test.go | 210 +++++++++
device/awg/tag_junk_packet_generators.go | 66 +++
device/awg/tag_junk_packet_generators_test.go | 149 +++++++
device/awg/tag_parser.go | 112 +++++
device/awg/tag_parser_test.go | 77 ++++
device/device.go | 409 ++++++++++--------
device/device_test.go | 167 +++----
device/noise-protocol.go | 42 +-
device/peer.go | 11 +
device/receive.go | 32 +-
device/send.go | 88 ++--
device/uapi.go | 209 ++++++---
go.mod | 16 +-
go.sum | 40 +-
23 files changed, 1982 insertions(+), 452 deletions(-)
create mode 100644 device/awg/awg.go
create mode 100644 device/awg/internal/mock.go
rename device/{ => awg}/junk_creator.go (52%)
rename device/{ => awg}/junk_creator_test.go (61%)
create mode 100644 device/awg/special_handshake_handler.go
create mode 100644 device/awg/tag_generator.go
create mode 100644 device/awg/tag_generator_test.go
create mode 100644 device/awg/tag_junk_packet_generator.go
create mode 100644 device/awg/tag_junk_packet_generator_test.go
create mode 100644 device/awg/tag_junk_packet_generators.go
create mode 100644 device/awg/tag_junk_packet_generators_test.go
create mode 100644 device/awg/tag_parser.go
create mode 100644 device/awg/tag_parser_test.go
diff --git a/Dockerfile b/Dockerfile
index 12159be..6d60440 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -1,4 +1,4 @@
-FROM golang:1.24 as awg
+FROM golang:1.24.4 as awg
COPY . /awg
WORKDIR /awg
RUN go mod download && \
@@ -7,10 +7,24 @@ RUN go mod download && \
FROM alpine:3.19
ARG AWGTOOLS_RELEASE="1.0.20241018"
+
+RUN apk add linux-headers build-base
+COPY awg-tools /awg-tools
+RUN pwd && ls -la / && ls -la /awg-tools
+WORKDIR /awg-tools/src
+# RUN ls -la && pwd && ls awg-tools
+RUN make
+RUN mkdir -p build && \
+ cp wg ./build/awg && \
+ cp wg-quick/linux.bash ./build/awg-quick
+
+RUN cp build/awg /usr/bin/awg
+RUN cp build/awg-quick /usr/bin/awg-quick
+
RUN apk --no-cache add iproute2 iptables bash && \
cd /usr/bin/ && \
- wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
- unzip -j alpine-3.19-amneziawg-tools.zip && \
+ # wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
+ # unzip -j alpine-3.19-amneziawg-tools.zip && \
chmod +x /usr/bin/awg /usr/bin/awg-quick && \
ln -s /usr/bin/awg /usr/bin/wg && \
ln -s /usr/bin/awg-quick /usr/bin/wg-quick
diff --git a/device/awg/awg.go b/device/awg/awg.go
new file mode 100644
index 0000000..fd5a96d
--- /dev/null
+++ b/device/awg/awg.go
@@ -0,0 +1,144 @@
+package awg
+
+import (
+ "bytes"
+ "fmt"
+ "slices"
+ "strconv"
+ "strings"
+ "sync"
+
+ "github.com/tevino/abool"
+)
+
+type aSecCfgType struct {
+ IsSet bool
+ JunkPacketCount int
+ JunkPacketMinSize int
+ JunkPacketMaxSize int
+ InitHeaderJunkSize int
+ ResponseHeaderJunkSize int
+ CookieReplyHeaderJunkSize int
+ TransportHeaderJunkSize int
+ InitPacketMagicHeader uint32
+ ResponsePacketMagicHeader uint32
+ UnderloadPacketMagicHeader uint32
+ TransportPacketMagicHeader uint32
+ // InitPacketMagicHeader Limit
+ // ResponsePacketMagicHeader Limit
+ // UnderloadPacketMagicHeader Limit
+ // TransportPacketMagicHeader Limit
+}
+
+type Limit struct {
+ Min uint32
+ Max uint32
+ HeaderType uint32
+}
+
+func NewLimit(min, max, headerType uint32) (Limit, error) {
+ if min > max {
+ return Limit{}, fmt.Errorf("min (%d) cannot be greater than max (%d)", min, max)
+ }
+
+ return Limit{
+ Min: min,
+ Max: max,
+ HeaderType: headerType,
+ }, nil
+}
+
+func ParseMagicHeader(key, value string, defaultHeaderType uint32) (Limit, error) {
+ // tempAwg.ASecCfg.InitPacketMagicHeader, err = awg.NewLimit(uint32(initPacketMagicHeaderMin), uint32(initPacketMagicHeaderMax), DNewLimit(min, max, headerType)efaultMessageInitiationType)
+ // var min, max, headerType uint32
+ // _, err := fmt.Sscanf(value, "%d-%d:%d", &min, &max, &headerType)
+ // if err != nil {
+ // return Limit{}, fmt.Errorf("invalid magic header format: %s", value)
+ // }
+
+ limits := strings.Split(value, "-")
+ if len(limits) != 2 {
+ return Limit{}, fmt.Errorf("invalid format for key: %s; %s", key, value)
+ }
+
+ min, err := strconv.ParseUint(limits[0], 10, 32)
+ if err != nil {
+ return Limit{}, fmt.Errorf("parse min key: %s; value: ; %w", key, limits[0], err)
+ }
+
+ max, err := strconv.ParseUint(limits[1], 10, 32)
+ if err != nil {
+ return Limit{}, fmt.Errorf("parse max key: %s; value: ; %w", key, limits[0], err)
+ }
+
+ limit, err := NewLimit(uint32(min), uint32(max), defaultHeaderType)
+ if err != nil {
+ return Limit{}, fmt.Errorf("new lmit key: %s; value: ; %w", key, limits[0], err)
+ }
+
+ return limit, nil
+}
+
+type Limits []Limit
+
+func NewLimits(limits []Limit) Limits {
+ slices.SortFunc(limits, func(a, b Limit) int {
+ if a.Min < b.Min {
+ return -1
+ } else if a.Min > b.Min {
+ return 1
+ }
+ return 0
+ })
+
+ return Limits(limits)
+}
+
+type Protocol struct {
+ IsASecOn abool.AtomicBool
+ // TODO: revision the need of the mutex
+ ASecMux sync.RWMutex
+ ASecCfg aSecCfgType
+ JunkCreator junkCreator
+
+ HandshakeHandler SpecialHandshakeHandler
+}
+
+func (protocol *Protocol) CreateInitHeaderJunk() ([]byte, error) {
+ return protocol.createHeaderJunk(protocol.ASecCfg.InitHeaderJunkSize)
+}
+
+func (protocol *Protocol) CreateResponseHeaderJunk() ([]byte, error) {
+ return protocol.createHeaderJunk(protocol.ASecCfg.ResponseHeaderJunkSize)
+}
+
+func (protocol *Protocol) CreateCookieReplyHeaderJunk() ([]byte, error) {
+ return protocol.createHeaderJunk(protocol.ASecCfg.CookieReplyHeaderJunkSize)
+}
+
+func (protocol *Protocol) CreateTransportHeaderJunk(packetSize int) ([]byte, error) {
+ return protocol.createHeaderJunk(protocol.ASecCfg.TransportHeaderJunkSize, packetSize)
+}
+
+func (protocol *Protocol) createHeaderJunk(junkSize int, optExtraSize ...int) ([]byte, error) {
+ extraSize := 0
+ if len(optExtraSize) == 1 {
+ extraSize = optExtraSize[0]
+ }
+
+ var junk []byte
+ protocol.ASecMux.RLock()
+ if junkSize != 0 {
+ buf := make([]byte, 0, junkSize+extraSize)
+ writer := bytes.NewBuffer(buf[:0])
+ err := protocol.JunkCreator.AppendJunk(writer, junkSize)
+ if err != nil {
+ protocol.ASecMux.RUnlock()
+ return nil, err
+ }
+ junk = writer.Bytes()
+ }
+ protocol.ASecMux.RUnlock()
+
+ return junk, nil
+}
diff --git a/device/awg/internal/mock.go b/device/awg/internal/mock.go
new file mode 100644
index 0000000..a2e1c95
--- /dev/null
+++ b/device/awg/internal/mock.go
@@ -0,0 +1,37 @@
+package internal
+
+type mockGenerator struct {
+ size int
+}
+
+func NewMockGenerator(size int) mockGenerator {
+ return mockGenerator{size: size}
+}
+
+func (m mockGenerator) Generate() []byte {
+ return make([]byte, m.size)
+}
+
+func (m mockGenerator) Size() int {
+ return m.size
+}
+
+func (m mockGenerator) Name() string {
+ return "mock"
+}
+
+type mockByteGenerator struct {
+ data []byte
+}
+
+func NewMockByteGenerator(data []byte) mockByteGenerator {
+ return mockByteGenerator{data: data}
+}
+
+func (bg mockByteGenerator) Generate() []byte {
+ return bg.data
+}
+
+func (bg mockByteGenerator) Size() int {
+ return len(bg.data)
+}
diff --git a/device/junk_creator.go b/device/awg/junk_creator.go
similarity index 52%
rename from device/junk_creator.go
rename to device/awg/junk_creator.go
index 3a2d3b4..91fd253 100644
--- a/device/junk_creator.go
+++ b/device/awg/junk_creator.go
@@ -1,4 +1,4 @@
-package device
+package awg
import (
"bytes"
@@ -8,61 +8,62 @@ import (
)
type junkCreator struct {
- device *Device
+ aSecCfg aSecCfgType
cha8Rand *v2.ChaCha8
}
-func NewJunkCreator(d *Device) (junkCreator, error) {
+// TODO: refactor param to only pass the junk related params
+func NewJunkCreator(aSecCfg aSecCfgType) (junkCreator, error) {
buf := make([]byte, 32)
_, err := crand.Read(buf)
if err != nil {
return junkCreator{}, err
}
- return junkCreator{device: d, cha8Rand: v2.NewChaCha8([32]byte(buf))}, nil
+ return junkCreator{aSecCfg: aSecCfg, cha8Rand: v2.NewChaCha8([32]byte(buf))}, nil
}
// Should be called with aSecMux RLocked
-func (jc *junkCreator) createJunkPackets() ([][]byte, error) {
- if jc.device.aSecCfg.junkPacketCount == 0 {
- return nil, nil
+func (jc *junkCreator) CreateJunkPackets(junks *[][]byte) error {
+ if jc.aSecCfg.JunkPacketCount == 0 {
+ return nil
}
- junks := make([][]byte, 0, jc.device.aSecCfg.junkPacketCount)
- for i := 0; i < jc.device.aSecCfg.junkPacketCount; i++ {
+ for range jc.aSecCfg.JunkPacketCount {
packetSize := jc.randomPacketSize()
junk, err := jc.randomJunkWithSize(packetSize)
if err != nil {
- return nil, fmt.Errorf("Failed to create junk packet: %v", err)
+ return fmt.Errorf("create junk packet: %v", err)
}
- junks = append(junks, junk)
+ *junks = append(*junks, junk)
}
- return junks, nil
+ return nil
}
// Should be called with aSecMux RLocked
func (jc *junkCreator) randomPacketSize() int {
return int(
jc.cha8Rand.Uint64()%uint64(
- jc.device.aSecCfg.junkPacketMaxSize-jc.device.aSecCfg.junkPacketMinSize,
+ jc.aSecCfg.JunkPacketMaxSize-jc.aSecCfg.JunkPacketMinSize,
),
- ) + jc.device.aSecCfg.junkPacketMinSize
+ ) + jc.aSecCfg.JunkPacketMinSize
}
// Should be called with aSecMux RLocked
-func (jc *junkCreator) appendJunk(writer *bytes.Buffer, size int) error {
+func (jc *junkCreator) AppendJunk(writer *bytes.Buffer, size int) error {
headerJunk, err := jc.randomJunkWithSize(size)
if err != nil {
- return fmt.Errorf("failed to create header junk: %v", err)
+ return fmt.Errorf("create header junk: %v", err)
}
_, err = writer.Write(headerJunk)
if err != nil {
- return fmt.Errorf("failed to write header junk: %v", err)
+ return fmt.Errorf("write header junk: %v", err)
}
return nil
}
// Should be called with aSecMux RLocked
func (jc *junkCreator) randomJunkWithSize(size int) ([]byte, error) {
+ // TODO: use a memory pool to allocate
junk := make([]byte, size)
_, err := jc.cha8Rand.Read(junk)
return junk, err
diff --git a/device/junk_creator_test.go b/device/awg/junk_creator_test.go
similarity index 61%
rename from device/junk_creator_test.go
rename to device/awg/junk_creator_test.go
index d3cf2b3..424f104 100644
--- a/device/junk_creator_test.go
+++ b/device/awg/junk_creator_test.go
@@ -1,36 +1,27 @@
-package device
+package awg
import (
"bytes"
"fmt"
"testing"
-
- "github.com/amnezia-vpn/amneziawg-go/conn/bindtest"
- "github.com/amnezia-vpn/amneziawg-go/tun/tuntest"
)
func setUpJunkCreator(t *testing.T) (junkCreator, error) {
- cfg, _ := genASecurityConfigs(t)
- tun := tuntest.NewChannelTUN()
- binds := bindtest.NewChannelBinds()
- level := LogLevelVerbose
- dev := NewDevice(
- tun.TUN(),
- binds[0],
- NewLogger(level, ""),
- )
-
- if err := dev.IpcSet(cfg[0]); err != nil {
- t.Errorf("failed to configure device %v", err)
- dev.Close()
- return junkCreator{}, err
- }
-
- jc, err := NewJunkCreator(dev)
+ jc, err := NewJunkCreator(aSecCfgType{
+ IsSet: true,
+ JunkPacketCount: 5,
+ JunkPacketMinSize: 500,
+ JunkPacketMaxSize: 1000,
+ InitHeaderJunkSize: 30,
+ ResponseHeaderJunkSize: 40,
+ InitPacketMagicHeader: 123456,
+ ResponsePacketMagicHeader: 67543,
+ UnderloadPacketMagicHeader: 32345,
+ TransportPacketMagicHeader: 123123,
+ })
if err != nil {
t.Errorf("failed to create junk creator %v", err)
- dev.Close()
return junkCreator{}, err
}
@@ -42,8 +33,9 @@ func Test_junkCreator_createJunkPackets(t *testing.T) {
if err != nil {
return
}
- t.Run("", func(t *testing.T) {
- got, err := jc.createJunkPackets()
+ t.Run("valid", func(t *testing.T) {
+ got := make([][]byte, 0, jc.aSecCfg.JunkPacketCount)
+ err := jc.CreateJunkPackets(&got)
if err != nil {
t.Errorf(
"junkCreator.createJunkPackets() = %v; failed",
@@ -68,7 +60,7 @@ func Test_junkCreator_createJunkPackets(t *testing.T) {
}
func Test_junkCreator_randomJunkWithSize(t *testing.T) {
- t.Run("", func(t *testing.T) {
+ t.Run("valid", func(t *testing.T) {
jc, err := setUpJunkCreator(t)
if err != nil {
return
@@ -78,7 +70,6 @@ func Test_junkCreator_randomJunkWithSize(t *testing.T) {
fmt.Printf("%v\n%v\n", r1, r2)
if bytes.Equal(r1, r2) {
t.Errorf("same junks %v", err)
- jc.device.Close()
return
}
})
@@ -90,14 +81,14 @@ func Test_junkCreator_randomPacketSize(t *testing.T) {
return
}
for range [30]struct{}{} {
- t.Run("", func(t *testing.T) {
- if got := jc.randomPacketSize(); jc.device.aSecCfg.junkPacketMinSize > got ||
- got > jc.device.aSecCfg.junkPacketMaxSize {
+ t.Run("valid", func(t *testing.T) {
+ if got := jc.randomPacketSize(); jc.aSecCfg.JunkPacketMinSize > got ||
+ got > jc.aSecCfg.JunkPacketMaxSize {
t.Errorf(
"junkCreator.randomPacketSize() = %v, not between range [%v,%v]",
got,
- jc.device.aSecCfg.junkPacketMinSize,
- jc.device.aSecCfg.junkPacketMaxSize,
+ jc.aSecCfg.JunkPacketMinSize,
+ jc.aSecCfg.JunkPacketMaxSize,
)
}
})
@@ -109,13 +100,13 @@ func Test_junkCreator_appendJunk(t *testing.T) {
if err != nil {
return
}
- t.Run("", func(t *testing.T) {
+ t.Run("valid", func(t *testing.T) {
s := "apple"
buffer := bytes.NewBuffer([]byte(s))
- err := jc.appendJunk(buffer, 30)
+ err := jc.AppendJunk(buffer, 30)
if err != nil &&
buffer.Len() != len(s)+30 {
- t.Errorf("appendWithJunk() size don't match")
+ t.Error("appendWithJunk() size don't match")
}
read := make([]byte, 50)
buffer.Read(read)
diff --git a/device/awg/special_handshake_handler.go b/device/awg/special_handshake_handler.go
new file mode 100644
index 0000000..e582d97
--- /dev/null
+++ b/device/awg/special_handshake_handler.go
@@ -0,0 +1,73 @@
+package awg
+
+import (
+ "errors"
+ "time"
+
+ "github.com/tevino/abool"
+ "go.uber.org/atomic"
+)
+
+// TODO: atomic?/ and better way to use this
+var PacketCounter *atomic.Uint64 = atomic.NewUint64(0)
+
+// TODO
+var WaitResponse = struct {
+ Channel chan struct{}
+ ShouldWait *abool.AtomicBool
+}{
+ make(chan struct{}, 1),
+ abool.New(),
+}
+
+type SpecialHandshakeHandler struct {
+ isFirstDone bool
+ SpecialJunk TagJunkPacketGenerators
+ ControlledJunk TagJunkPacketGenerators
+
+ nextItime time.Time
+ ITimeout time.Duration // seconds
+
+ IsSet bool
+}
+
+func (handler *SpecialHandshakeHandler) Validate() error {
+ var errs []error
+ if err := handler.SpecialJunk.Validate(); err != nil {
+ errs = append(errs, err)
+ }
+ if err := handler.ControlledJunk.Validate(); err != nil {
+ errs = append(errs, err)
+ }
+ return errors.Join(errs...)
+}
+
+func (handler *SpecialHandshakeHandler) GenerateSpecialJunk() [][]byte {
+ if !handler.SpecialJunk.IsDefined() {
+ return nil
+ }
+
+ // TODO: create tests
+ if !handler.isFirstDone {
+ handler.isFirstDone = true
+ } else if !handler.isTimeToSendSpecial() {
+ return nil
+ }
+
+ rv := handler.SpecialJunk.GeneratePackets()
+ handler.nextItime = time.Now().Add(handler.ITimeout)
+
+ return rv
+}
+
+func (handler *SpecialHandshakeHandler) isTimeToSendSpecial() bool {
+ return time.Now().After(handler.nextItime)
+}
+
+func (handler *SpecialHandshakeHandler) GenerateControlledJunk() [][]byte {
+ if !handler.ControlledJunk.IsDefined() {
+ return nil
+ }
+
+ return handler.ControlledJunk.GeneratePackets()
+}
diff --git a/device/awg/tag_generator.go b/device/awg/tag_generator.go
new file mode 100644
index 0000000..65d8004
--- /dev/null
+++ b/device/awg/tag_generator.go
@@ -0,0 +1,190 @@
+package awg
+
+import (
+ crand "crypto/rand"
+ "encoding/binary"
+ "encoding/hex"
+ "fmt"
+ "strconv"
+ "strings"
+ "time"
+
+ v2 "math/rand/v2"
+ // "go.uber.org/atomic"
+)
+
+type Generator interface {
+ Generate() []byte
+ Size() int
+}
+
+type newGenerator func(string) (Generator, error)
+
+type BytesGenerator struct {
+ value []byte
+ size int
+}
+
+func (bg *BytesGenerator) Generate() []byte {
+ return bg.value
+}
+
+func (bg *BytesGenerator) Size() int {
+ return bg.size
+}
+
+func newBytesGenerator(param string) (Generator, error) {
+ hasPrefix := strings.HasPrefix(param, "0x") || strings.HasPrefix(param, "0X")
+ if !hasPrefix {
+ return nil, fmt.Errorf("not correct hex: %s", param)
+ }
+
+ hex, err := hexToBytes(param)
+ if err != nil {
+ return nil, fmt.Errorf("hexToBytes: %w", err)
+ }
+
+ return &BytesGenerator{value: hex, size: len(hex)}, nil
+}
+
+func hexToBytes(hexStr string) ([]byte, error) {
+ hexStr = strings.TrimPrefix(hexStr, "0x")
+ hexStr = strings.TrimPrefix(hexStr, "0X")
+
+ // Ensure even length (pad with leading zero if needed)
+ if len(hexStr)%2 != 0 {
+ hexStr = "0" + hexStr
+ }
+
+ return hex.DecodeString(hexStr)
+}
+
+type RandomPacketGenerator struct {
+ cha8Rand *v2.ChaCha8
+ size int
+}
+
+func (rpg *RandomPacketGenerator) Generate() []byte {
+ junk := make([]byte, rpg.size)
+ rpg.cha8Rand.Read(junk)
+ return junk
+}
+
+func (rpg *RandomPacketGenerator) Size() int {
+ return rpg.size
+}
+
+func newRandomPacketGenerator(param string) (Generator, error) {
+ size, err := strconv.Atoi(param)
+ if err != nil {
+ return nil, fmt.Errorf("random packet parse int: %w", err)
+ }
+
+ if size > 1000 {
+ return nil, fmt.Errorf("random packet size must be less than 1000")
+ }
+
+ buf := make([]byte, 32)
+ _, err = crand.Read(buf)
+ if err != nil {
+ return nil, fmt.Errorf("random packet crand read: %w", err)
+ }
+
+ return &RandomPacketGenerator{
+ cha8Rand: v2.NewChaCha8([32]byte(buf)),
+ size: size,
+ }, nil
+}
+
+type TimestampGenerator struct {
+}
+
+func (tg *TimestampGenerator) Generate() []byte {
+ buf := make([]byte, 8)
+ binary.BigEndian.PutUint64(buf, uint64(time.Now().Unix()))
+ return buf
+}
+
+func (tg *TimestampGenerator) Size() int {
+ return 8
+}
+
+func newTimestampGenerator(param string) (Generator, error) {
+ if len(param) != 0 {
+ return nil, fmt.Errorf("timestamp param needs to be empty: %s", param)
+ }
+
+ return &TimestampGenerator{}, nil
+}
+
+type WaitTimeoutGenerator struct {
+ waitTimeout time.Duration
+}
+
+func (wtg *WaitTimeoutGenerator) Generate() []byte {
+ time.Sleep(wtg.waitTimeout)
+ return []byte{}
+}
+
+func (wtg *WaitTimeoutGenerator) Size() int {
+ return 0
+}
+
+func newWaitTimeoutGenerator(param string) (Generator, error) {
+ timeout, err := strconv.Atoi(param)
+ if err != nil {
+ return nil, fmt.Errorf("timeout parse int: %w", err)
+ }
+
+ if timeout > 5000 {
+ return nil, fmt.Errorf("timeout must be less than 5000ms")
+ }
+
+ return &WaitTimeoutGenerator{
+ waitTimeout: time.Duration(timeout) * time.Millisecond,
+ }, nil
+}
+
+type PacketCounterGenerator struct {
+}
+
+func (c *PacketCounterGenerator) Generate() []byte {
+ buf := make([]byte, 8)
+ // TODO: better way to handle counter tag
+ binary.BigEndian.PutUint64(buf, PacketCounter.Load())
+ return buf
+}
+
+func (c *PacketCounterGenerator) Size() int {
+ return 8
+}
+
+func newPacketCounterGenerator(param string) (Generator, error) {
+ if len(param) != 0 {
+ return nil, fmt.Errorf("packet counter param needs to be empty: %s", param)
+ }
+
+ return &PacketCounterGenerator{}, nil
+}
+
+type WaitResponseGenerator struct {
+}
+
+func (c *WaitResponseGenerator) Generate() []byte {
+ WaitResponse.ShouldWait.Set()
+ <-WaitResponse.Channel
+ WaitResponse.ShouldWait.UnSet()
+ return []byte{}
+}
+
+func (c *WaitResponseGenerator) Size() int {
+ return 0
+}
+
+func newWaitResponseGenerator(param string) (Generator, error) {
+ if len(param) != 0 {
+ return nil, fmt.Errorf("wait response param needs to be empty: %s", param)
+ }
+
+ return &WaitResponseGenerator{}, nil
+}
diff --git a/device/awg/tag_generator_test.go b/device/awg/tag_generator_test.go
new file mode 100644
index 0000000..4950b33
--- /dev/null
+++ b/device/awg/tag_generator_test.go
@@ -0,0 +1,189 @@
+package awg
+
+import (
+ "encoding/binary"
+ "fmt"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func Test_newBytesGenerator(t *testing.T) {
+ type args struct {
+ param string
+ }
+ tests := []struct {
+ name string
+ args args
+ want []byte
+ wantErr error
+ }{
+ {
+ name: "empty",
+ args: args{
+ param: "",
+ },
+ wantErr: fmt.Errorf("not correct hex"),
+ },
+ {
+ name: "wrong start",
+ args: args{
+ param: "123456",
+ },
+ wantErr: fmt.Errorf("not correct hex"),
+ },
+ {
+ name: "not only hex value with X",
+ args: args{
+ param: "0X12345q",
+ },
+ wantErr: fmt.Errorf("not correct hex"),
+ },
+ {
+ name: "not only hex value with x",
+ args: args{
+ param: "0x12345q",
+ },
+ wantErr: fmt.Errorf("not correct hex"),
+ },
+ {
+ name: "valid hex",
+ args: args{
+ param: "0xf6ab3267fa",
+ },
+ want: []byte{0xf6, 0xab, 0x32, 0x67, 0xfa},
+ },
+ {
+ name: "valid hex with odd length",
+ args: args{
+ param: "0xfab3267fa",
+ },
+ want: []byte{0xf, 0xab, 0x32, 0x67, 0xfa},
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, err := newBytesGenerator(tt.args.param)
+
+ if tt.wantErr != nil {
+ require.ErrorAs(t, err, &tt.wantErr)
+ require.Nil(t, got)
+ return
+ }
+
+ require.Nil(t, err)
+ require.NotNil(t, got)
+
+ gotValues := got.Generate()
+ require.Equal(t, tt.want, gotValues)
+ })
+ }
+}
+
+func Test_newRandomPacketGenerator(t *testing.T) {
+ type args struct {
+ param string
+ }
+ tests := []struct {
+ name string
+ args args
+ wantErr error
+ }{
+ {
+ name: "empty",
+ args: args{
+ param: "",
+ },
+ wantErr: fmt.Errorf("parse int"),
+ },
+ {
+ name: "not an int",
+ args: args{
+ param: "x",
+ },
+ wantErr: fmt.Errorf("parse int"),
+ },
+ {
+ name: "too large",
+ args: args{
+ param: "1001",
+ },
+ wantErr: fmt.Errorf("random packet size must be less than 1000"),
+ },
+ {
+ name: "valid",
+ args: args{
+ param: "12",
+ },
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ got, err := newRandomPacketGenerator(tt.args.param)
+ if tt.wantErr != nil {
+ require.ErrorAs(t, err, &tt.wantErr)
+ require.Nil(t, got)
+ return
+ }
+
+ require.Nil(t, err)
+ require.NotNil(t, got)
+ first := got.Generate()
+
+ second := got.Generate()
+ require.NotEqual(t, first, second)
+ })
+ }
+}
+
+func TestPacketCounterGenerator(t *testing.T) {
+ tests := []struct {
+ name string
+ param string
+ wantErr bool
+ }{
+ {
+ name: "Valid empty param",
+ param: "",
+ wantErr: false,
+ },
+ {
+ name: "Invalid non-empty param",
+ param: "anything",
+ wantErr: true,
+ },
+ }
+
+ for _, tc := range tests {
+ tc := tc // capture range variable
+ t.Run(tc.name, func(t *testing.T) {
+ t.Parallel()
+
+ gen, err := newPacketCounterGenerator(tc.param)
+ if tc.wantErr {
+ require.Error(t, err)
+ return
+ }
+
+ require.NoError(t, err)
+ require.Equal(t, 8, gen.Size())
+
+ // Reset counter to known value for test
+ initialCount := uint64(42)
+ PacketCounter.Store(initialCount)
+
+ output := gen.Generate()
+ require.Equal(t, 8, len(output))
+
+ // Verify counter value in output
+ counterValue := binary.BigEndian.Uint64(output)
+ require.Equal(t, initialCount, counterValue)
+
+ // Increment counter and verify change
+ PacketCounter.Add(1)
+ output = gen.Generate()
+ counterValue = binary.BigEndian.Uint64(output)
+ require.Equal(t, initialCount+1, counterValue)
+ })
+ }
+}
diff --git a/device/awg/tag_junk_packet_generator.go b/device/awg/tag_junk_packet_generator.go
new file mode 100644
index 0000000..fdbebc8
--- /dev/null
+++ b/device/awg/tag_junk_packet_generator.go
@@ -0,0 +1,59 @@
+package awg
+
+import (
+ "fmt"
+ "strconv"
+)
+
+type TagJunkPacketGenerator struct {
+ name string
+ tagValue string
+
+ packetSize int
+ generators []Generator
+}
+
+func newTagJunkPacketGenerator(name, tagValue string, size int) TagJunkPacketGenerator {
+ return TagJunkPacketGenerator{
+ name: name,
+ tagValue: tagValue,
+ generators: make([]Generator, 0, size),
+ }
+}
+
+func (tg *TagJunkPacketGenerator) append(generator Generator) {
+ tg.generators = append(tg.generators, generator)
+ tg.packetSize += generator.Size()
+}
+
+func (tg *TagJunkPacketGenerator) generatePacket() []byte {
+ packet := make([]byte, 0, tg.packetSize)
+ for _, generator := range tg.generators {
+ packet = append(packet, generator.Generate()...)
+ }
+
+ return packet
+}
+
+func (tg *TagJunkPacketGenerator) Name() string {
+ return tg.name
+}
+
+func (tg *TagJunkPacketGenerator) nameIndex() (int, error) {
+ if len(tg.name) != 2 {
+ return 0, fmt.Errorf("name must be 2 character long: %s", tg.name)
+ }
+
+ index, err := strconv.Atoi(tg.name[1:2])
+ if err != nil {
+ return 0, fmt.Errorf("name 2 char should be an int %w", err)
+ }
+ return index, nil
+}
+
+func (tg *TagJunkPacketGenerator) IpcGetFields() IpcFields {
+ return IpcFields{
+ Key: tg.name,
+ Value: tg.tagValue,
+ }
+}
diff --git a/device/awg/tag_junk_packet_generator_test.go b/device/awg/tag_junk_packet_generator_test.go
new file mode 100644
index 0000000..309d425
--- /dev/null
+++ b/device/awg/tag_junk_packet_generator_test.go
@@ -0,0 +1,210 @@
+package awg
+
+import (
+ "testing"
+
+ "github.com/amnezia-vpn/amneziawg-go/device/awg/internal"
+ "github.com/stretchr/testify/require"
+)
+
+func TestNewTagJunkGenerator(t *testing.T) {
+ t.Parallel()
+
+ testCases := []struct {
+ name string
+ genName string
+ size int
+ expected TagJunkPacketGenerator
+ }{
+ {
+ name: "Create new generator with empty name",
+ genName: "",
+ size: 0,
+ expected: TagJunkPacketGenerator{
+ name: "",
+ packetSize: 0,
+ generators: make([]Generator, 0),
+ },
+ },
+ {
+ name: "Create new generator with valid name",
+ genName: "T1",
+ size: 0,
+ expected: TagJunkPacketGenerator{
+ name: "T1",
+ packetSize: 0,
+ generators: make([]Generator, 0),
+ },
+ },
+ {
+ name: "Create new generator with non-zero size",
+ genName: "T2",
+ size: 5,
+ expected: TagJunkPacketGenerator{
+ name: "T2",
+ packetSize: 0,
+ generators: make([]Generator, 5),
+ },
+ },
+ }
+
+ for _, tc := range testCases {
+ tc := tc // capture range variable
+ t.Run(tc.name, func(t *testing.T) {
+ t.Parallel()
+ result := newTagJunkPacketGenerator(tc.genName, "", tc.size)
+ require.Equal(t, tc.expected.name, result.name)
+ require.Equal(t, tc.expected.packetSize, result.packetSize)
+ require.Equal(t, cap(result.generators), len(tc.expected.generators))
+ })
+ }
+}
+
+func TestTagJunkGeneratorAppend(t *testing.T) {
+ t.Parallel()
+
+ testCases := []struct {
+ name string
+ initialState TagJunkPacketGenerator
+ mockSize int
+ expectedLength int
+ expectedSize int
+ }{
+ {
+ name: "Append to empty generator",
+ initialState: newTagJunkPacketGenerator("T1", "", 0),
+ mockSize: 5,
+ expectedLength: 1,
+ expectedSize: 5,
+ },
+ {
+ name: "Append to non-empty generator",
+ initialState: TagJunkPacketGenerator{
+ name: "T2",
+ packetSize: 10,
+ generators: make([]Generator, 2),
+ },
+ mockSize: 7,
+ expectedLength: 3, // 2 existing + 1 new
+ expectedSize: 17, // 10 + 7
+ },
+ }
+
+ for _, tc := range testCases {
+ tc := tc // capture range variable
+ t.Run(tc.name, func(t *testing.T) {
+ t.Parallel()
+
+ tg := tc.initialState
+ mockGen := internal.NewMockGenerator(tc.mockSize)
+
+ tg.append(mockGen)
+
+ require.Equal(t, tc.expectedLength, len(tg.generators))
+ require.Equal(t, tc.expectedSize, tg.packetSize)
+ })
+ }
+}
+
+func TestTagJunkGeneratorGenerate(t *testing.T) {
+ t.Parallel()
+
+ // Create mock generators for testing
+ mockGen1 := internal.NewMockByteGenerator([]byte{0x01, 0x02})
+ mockGen2 := internal.NewMockByteGenerator([]byte{0x03, 0x04, 0x05})
+
+ testCases := []struct {
+ name string
+ setupGenerator func() TagJunkPacketGenerator
+ expected []byte
+ }{
+ {
+ name: "Generate with empty generators",
+ setupGenerator: func() TagJunkPacketGenerator {
+ return newTagJunkPacketGenerator("T1", "", 0)
+ },
+ expected: []byte{},
+ },
+ {
+ name: "Generate with single generator",
+ setupGenerator: func() TagJunkPacketGenerator {
+ tg := newTagJunkPacketGenerator("T2", "", 0)
+ tg.append(mockGen1)
+ return tg
+ },
+ expected: []byte{0x01, 0x02},
+ },
+ {
+ name: "Generate with multiple generators",
+ setupGenerator: func() TagJunkPacketGenerator {
+ tg := newTagJunkPacketGenerator("T3", "", 0)
+ tg.append(mockGen1)
+ tg.append(mockGen2)
+ return tg
+ },
+ expected: []byte{0x01, 0x02, 0x03, 0x04, 0x05},
+ },
+ }
+
+ for _, tc := range testCases {
+ tc := tc // capture range variable
+ t.Run(tc.name, func(t *testing.T) {
+ t.Parallel()
+
+ tg := tc.setupGenerator()
+ result := tg.generatePacket()
+
+ require.Equal(t, tc.expected, result)
+ })
+ }
+}
+
+func TestTagJunkGeneratorNameIndex(t *testing.T) {
+ t.Parallel()
+
+ testCases := []struct {
+ name string
+ generatorName string
+ expectedIndex int
+ expectError bool
+ }{
+ {
+ name: "Valid name with digit",
+ generatorName: "T5",
+ expectedIndex: 5,
+ expectError: false,
+ },
+ {
+ name: "Invalid name - too short",
+ generatorName: "T",
+ expectError: true,
+ },
+ {
+ name: "Invalid name - too long",
+ generatorName: "T55",
+ expectError: true,
+ },
+ {
+ name: "Invalid name - non-digit second character",
+ generatorName: "TX",
+ expectError: true,
+ },
+ }
+
+ for _, tc := range testCases {
+ tc := tc // capture range variable
+ t.Run(tc.name, func(t *testing.T) {
+ t.Parallel()
+
+ tg := TagJunkPacketGenerator{name: tc.generatorName}
+ index, err := tg.nameIndex()
+
+ if tc.expectError {
+ require.Error(t, err)
+ } else {
+ require.NoError(t, err)
+ require.Equal(t, tc.expectedIndex, index)
+ }
+ })
+ }
+}
diff --git a/device/awg/tag_junk_packet_generators.go b/device/awg/tag_junk_packet_generators.go
new file mode 100644
index 0000000..9921eb0
--- /dev/null
+++ b/device/awg/tag_junk_packet_generators.go
@@ -0,0 +1,66 @@
+package awg
+
+import "fmt"
+
+type TagJunkPacketGenerators struct {
+ tagGenerators []TagJunkPacketGenerator
+ length int
+ DefaultJunkCount int // Jc
+}
+
+func (generators *TagJunkPacketGenerators) AppendGenerator(
+ generator TagJunkPacketGenerator,
+) {
+ generators.tagGenerators = append(generators.tagGenerators, generator)
+ generators.length++
+}
+
+func (generators *TagJunkPacketGenerators) IsDefined() bool {
+ return len(generators.tagGenerators) > 0
+}
+
+// validate that packets were defined consecutively
+func (generators *TagJunkPacketGenerators) Validate() error {
+ seen := make([]bool, len(generators.tagGenerators))
+ for _, generator := range generators.tagGenerators {
+ index, err := generator.nameIndex()
+ if index > len(generators.tagGenerators) {
+ return fmt.Errorf("junk packet index should be consecutive")
+ }
+ if err != nil {
+ return fmt.Errorf("name index: %w", err)
+ } else {
+ seen[index-1] = true
+ }
+ }
+
+ for _, found := range seen {
+ if !found {
+ return fmt.Errorf("junk packet index should be consecutive")
+ }
+ }
+
+ return nil
+}
+
+func (generators *TagJunkPacketGenerators) GeneratePackets() [][]byte {
+ var rv = make([][]byte, 0, generators.length+generators.DefaultJunkCount)
+
+ for i, tagGenerator := range generators.tagGenerators {
+ rv = append(rv, make([]byte, tagGenerator.packetSize))
+ copy(rv[i], tagGenerator.generatePacket())
+ PacketCounter.Inc()
+ }
+ PacketCounter.Add(uint64(generators.DefaultJunkCount))
+
+ return rv
+}
+
+func (tg *TagJunkPacketGenerators) IpcGetFields() []IpcFields {
+ rv := make([]IpcFields, 0, len(tg.tagGenerators))
+ for _, generator := range tg.tagGenerators {
+ rv = append(rv, generator.IpcGetFields())
+ }
+
+ return rv
+}
diff --git a/device/awg/tag_junk_packet_generators_test.go b/device/awg/tag_junk_packet_generators_test.go
new file mode 100644
index 0000000..6b1fd47
--- /dev/null
+++ b/device/awg/tag_junk_packet_generators_test.go
@@ -0,0 +1,149 @@
+package awg
+
+import (
+ "testing"
+
+ "github.com/amnezia-vpn/amneziawg-go/device/awg/internal"
+ "github.com/stretchr/testify/require"
+)
+
+func TestTagJunkGeneratorHandlerAppendGenerator(t *testing.T) {
+ tests := []struct {
+ name string
+ generator TagJunkPacketGenerator
+ }{
+ {
+ name: "append single generator",
+ generator: newTagJunkPacketGenerator("t1", "", 10),
+ },
+ }
+
+ for _, tt := range tests {
+ tt := tt
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ generators := &TagJunkPacketGenerators{}
+
+ // Initial length should be 0
+ require.Equal(t, 0, generators.length)
+ require.Empty(t, generators.tagGenerators)
+
+ // After append, length should be 1 and generator should be added
+ generators.AppendGenerator(tt.generator)
+ require.Equal(t, 1, generators.length)
+ require.Len(t, generators.tagGenerators, 1)
+ require.Equal(t, tt.generator, generators.tagGenerators[0])
+ })
+ }
+}
+
+func TestTagJunkGeneratorHandlerValidate(t *testing.T) {
+ tests := []struct {
+ name string
+ generators []TagJunkPacketGenerator
+ wantErr bool
+ errMsg string
+ }{
+ {
+ name: "bad start",
+ generators: []TagJunkPacketGenerator{
+ newTagJunkPacketGenerator("t3", "", 10),
+ newTagJunkPacketGenerator("t4", "", 10),
+ },
+ wantErr: true,
+ errMsg: "junk packet index should be consecutive",
+ },
+ {
+ name: "non-consecutive indices",
+ generators: []TagJunkPacketGenerator{
+ newTagJunkPacketGenerator("t1", "", 10),
+ newTagJunkPacketGenerator("t3", "", 10), // Missing t2
+ },
+ wantErr: true,
+ errMsg: "junk packet index should be consecutive",
+ },
+ {
+ name: "consecutive indices",
+ generators: []TagJunkPacketGenerator{
+ newTagJunkPacketGenerator("t1", "", 10),
+ newTagJunkPacketGenerator("t2", "", 10),
+ newTagJunkPacketGenerator("t3", "", 10),
+ newTagJunkPacketGenerator("t4", "", 10),
+ newTagJunkPacketGenerator("t5", "", 10),
+ },
+ },
+ {
+ name: "nameIndex error",
+ generators: []TagJunkPacketGenerator{
+ newTagJunkPacketGenerator("error", "", 10),
+ },
+ wantErr: true,
+ errMsg: "name must be 2 character long",
+ },
+ }
+
+ for _, tt := range tests {
+ tt := tt
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ generators := &TagJunkPacketGenerators{}
+ for _, gen := range tt.generators {
+ generators.AppendGenerator(gen)
+ }
+
+ err := generators.Validate()
+ if tt.wantErr {
+ require.Error(t, err)
+ require.Contains(t, err.Error(), tt.errMsg)
+ return
+ }
+ require.NoError(t, err)
+ })
+ }
+}
+
+func TestTagJunkGeneratorHandlerGenerate(t *testing.T) {
+ mockByte1 := []byte{0x01, 0x02}
+ mockByte2 := []byte{0x03, 0x04, 0x05}
+ mockGen1 := internal.NewMockByteGenerator(mockByte1)
+ mockGen2 := internal.NewMockByteGenerator(mockByte2)
+
+ tests := []struct {
+ name string
+ setupGenerator func() []TagJunkPacketGenerator
+ expected [][]byte
+ }{
+ {
+ name: "generate with no default junk",
+ setupGenerator: func() []TagJunkPacketGenerator {
+ tg1 := newTagJunkPacketGenerator("t1", "", 0)
+ tg1.append(mockGen1)
+ tg1.append(mockGen2)
+ tg2 := newTagJunkPacketGenerator("t2", "", 0)
+ tg2.append(mockGen2)
+ tg2.append(mockGen1)
+
+ return []TagJunkPacketGenerator{tg1, tg2}
+ },
+ expected: [][]byte{
+ append(mockByte1, mockByte2...),
+ append(mockByte2, mockByte1...),
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ tt := tt
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ generators := &TagJunkPacketGenerators{}
+ tagGenerators := tt.setupGenerator()
+ for _, gen := range tagGenerators {
+ generators.AppendGenerator(gen)
+ }
+
+ result := generators.GeneratePackets()
+ require.Equal(t, result, tt.expected)
+ })
+ }
+}
diff --git a/device/awg/tag_parser.go b/device/awg/tag_parser.go
new file mode 100644
index 0000000..2b09226
--- /dev/null
+++ b/device/awg/tag_parser.go
@@ -0,0 +1,112 @@
+package awg
+
+import (
+ "fmt"
+ "maps"
+ "regexp"
+ "strings"
+)
+
+type IpcFields struct{ Key, Value string }
+
+type EnumTag string
+
+const (
+ BytesEnumTag EnumTag = "b"
+ CounterEnumTag EnumTag = "c"
+ TimestampEnumTag EnumTag = "t"
+ RandomBytesEnumTag EnumTag = "r"
+ WaitTimeoutEnumTag EnumTag = "wt"
+ WaitResponseEnumTag EnumTag = "wr"
+)
+
+var generatorCreator = map[EnumTag]newGenerator{
+ BytesEnumTag: newBytesGenerator,
+ CounterEnumTag: newPacketCounterGenerator,
+ TimestampEnumTag: newTimestampGenerator,
+ RandomBytesEnumTag: newRandomPacketGenerator,
+ WaitTimeoutEnumTag: newWaitTimeoutGenerator,
+ // WaitResponseEnumTag: newWaitResponseGenerator,
+}
+
+// helper map to determine enumTags are unique
+var uniqueTags = map[EnumTag]bool{
+ CounterEnumTag: false,
+ TimestampEnumTag: false,
+}
+
+type Tag struct {
+ Name EnumTag
+ Param string
+}
+
+func parseTag(input string) (Tag, error) {
+ // Regular expression to match
+ re := regexp.MustCompile(`([a-zA-Z]+)(?:\s+([^>]+))?>`)
+
+ match := re.FindStringSubmatch(input)
+ tag := Tag{
+ Name: EnumTag(match[1]),
+ }
+ if len(match) > 2 && match[2] != "" {
+ tag.Param = strings.TrimSpace(match[2])
+ }
+
+ return tag, nil
+}
+
+func Parse(name, input string) (TagJunkPacketGenerator, error) {
+ inputSlice := strings.Split(input, "<")
+ if len(inputSlice) <= 1 {
+ return TagJunkPacketGenerator{}, fmt.Errorf("empty input: %s", input)
+ }
+
+ uniqueTagCheck := make(map[EnumTag]bool, len(uniqueTags))
+ maps.Copy(uniqueTagCheck, uniqueTags)
+
+ // skip byproduct of split
+ inputSlice = inputSlice[1:]
+ rv := newTagJunkPacketGenerator(name, input, len(inputSlice))
+ for _, inputParam := range inputSlice {
+ if len(inputParam) <= 1 {
+ return TagJunkPacketGenerator{}, fmt.Errorf(
+ "empty tag in input: %s",
+ inputSlice,
+ )
+ } else if strings.Count(inputParam, ">") != 1 {
+ return TagJunkPacketGenerator{}, fmt.Errorf("ill formated input: %s", input)
+ }
+
+ tag, _ := parseTag(inputParam)
+ creator, ok := generatorCreator[tag.Name]
+ if !ok {
+ return TagJunkPacketGenerator{}, fmt.Errorf("invalid tag: %s", tag.Name)
+ }
+ if present, ok := uniqueTagCheck[tag.Name]; ok {
+ if present {
+ return TagJunkPacketGenerator{}, fmt.Errorf(
+ "tag %s needs to be unique",
+ tag.Name,
+ )
+ }
+ uniqueTagCheck[tag.Name] = true
+ }
+ generator, err := creator(tag.Param)
+ if err != nil {
+ return TagJunkPacketGenerator{}, fmt.Errorf("gen: %w", err)
+ }
+
+ // TODO: handle counter tag
+ // if tag.Name == CounterEnumTag {
+ // packetCounter, ok := generator.(*PacketCounterGenerator)
+ // if !ok {
+ // log.Fatalf("packet counter generator expected, got %T", generator)
+ // }
+ // PacketCounter = packetCounter.counter
+ // }
+
+ rv.append(generator)
+ }
+
+ return rv, nil
+}
diff --git a/device/awg/tag_parser_test.go b/device/awg/tag_parser_test.go
new file mode 100644
index 0000000..8f828ec
--- /dev/null
+++ b/device/awg/tag_parser_test.go
@@ -0,0 +1,77 @@
+package awg
+
+import (
+ "fmt"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestParse(t *testing.T) {
+ type args struct {
+ name string
+ input string
+ }
+ tests := []struct {
+ name string
+ args args
+ wantErr error
+ }{
+ {
+ name: "invalid name",
+ args: args{name: "apple", input: ""},
+ wantErr: fmt.Errorf("ill formated input"),
+ },
+ {
+ name: "empty",
+ args: args{name: "i1", input: ""},
+ wantErr: fmt.Errorf("ill formated input"),
+ },
+ {
+ name: "extra >",
+ args: args{name: "i1", input: ">"},
+ wantErr: fmt.Errorf("ill formated input"),
+ },
+ {
+ name: "extra <",
+ args: args{name: "i1", input: "<"},
+ wantErr: fmt.Errorf("empty tag in input"),
+ },
+ {
+ name: "empty <>",
+ args: args{name: "i1", input: "<>"},
+ wantErr: fmt.Errorf("empty tag in input"),
+ },
+ {
+ name: "invalid tag",
+ args: args{name: "i1", input: ""},
+ wantErr: fmt.Errorf("invalid tag"),
+ },
+ {
+ name: "counter uniqueness violation",
+ args: args{name: "i1", input: ""},
+ wantErr: fmt.Errorf("parse tag needs to be unique"),
+ },
+ {
+ name: "timestamp uniqueness violation",
+ args: args{name: "i1", input: ""},
+ wantErr: fmt.Errorf("parse tag needs to be unique"),
+ },
+ {
+ name: "valid",
+ args: args{input: ""},
+ },
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ _, err := Parse(tt.args.name, tt.args.input)
+
+ // TODO: ErrorAs doesn't work as you think
+ if tt.wantErr != nil {
+ require.ErrorAs(t, err, &tt.wantErr)
+ return
+ }
+ require.Nil(t, err)
+ })
+ }
+}
diff --git a/device/device.go b/device/device.go
index 124b74e..1829352 100644
--- a/device/device.go
+++ b/device/device.go
@@ -6,19 +6,55 @@
package device
import (
+ "errors"
"runtime"
"sync"
"sync/atomic"
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device/awg"
"github.com/amnezia-vpn/amneziawg-go/ipc"
"github.com/amnezia-vpn/amneziawg-go/ratelimiter"
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/tun"
- "github.com/tevino/abool/v2"
)
+type Version uint8
+
+const (
+ VersionDefault Version = iota
+ VersionAwg
+ VersionAwgSpecialHandshake
+)
+
+// TODO:
+type AtomicVersion struct {
+ value atomic.Uint32
+}
+
+func NewAtomicVersion(v Version) *AtomicVersion {
+ av := &AtomicVersion{}
+ av.Store(v)
+ return av
+}
+
+func (av *AtomicVersion) Load() Version {
+ return Version(av.value.Load())
+}
+
+func (av *AtomicVersion) Store(v Version) {
+ av.value.Store(uint32(v))
+}
+
+func (av *AtomicVersion) CompareAndSwap(old, new Version) bool {
+ return av.value.CompareAndSwap(uint32(old), uint32(new))
+}
+
+func (av *AtomicVersion) Swap(new Version) Version {
+ return Version(av.value.Swap(uint32(new)))
+}
+
type Device struct {
state struct {
// state holds the device's state. It is accessed atomically.
@@ -92,23 +128,8 @@ type Device struct {
closed chan struct{}
log *Logger
- isASecOn abool.AtomicBool
- aSecMux sync.RWMutex
- aSecCfg aSecCfgType
- junkCreator junkCreator
-}
-
-type aSecCfgType struct {
- isSet bool
- junkPacketCount int
- junkPacketMinSize int
- junkPacketMaxSize int
- initPacketJunkSize int
- responsePacketJunkSize int
- initPacketMagicHeader uint32
- responsePacketMagicHeader uint32
- underloadPacketMagicHeader uint32
- transportPacketMagicHeader uint32
+ version Version
+ awg awg.Protocol
}
// deviceState represents the state of a Device.
@@ -557,251 +578,261 @@ func (device *Device) BindClose() error {
device.net.Unlock()
return err
}
-func (device *Device) isAdvancedSecurityOn() bool {
- return device.isASecOn.IsSet()
+func (device *Device) isAWG() bool {
+ return device.version >= VersionAwg
}
func (device *Device) resetProtocol() {
// restore default message type values
- MessageInitiationType = 1
- MessageResponseType = 2
- MessageCookieReplyType = 3
- MessageTransportType = 4
+ MessageInitiationType = DefaultMessageInitiationType
+ MessageResponseType = DefaultMessageResponseType
+ MessageCookieReplyType = DefaultMessageCookieReplyType
+ MessageTransportType = DefaultMessageTransportType
}
-func (device *Device) handlePostConfig(tempASecCfg *aSecCfgType) (err error) {
-
- if !tempASecCfg.isSet {
- return err
+func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
+ if !tempAwg.ASecCfg.IsSet && !tempAwg.HandshakeHandler.IsSet {
+ return nil
}
+ var errs []error
+
isASecOn := false
- device.aSecMux.Lock()
- if tempASecCfg.junkPacketCount < 0 {
- err = ipcErrorf(
+ device.awg.ASecMux.Lock()
+ if tempAwg.ASecCfg.JunkPacketCount < 0 {
+ errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
"JunkPacketCount should be non negative",
+ ),
)
}
- device.aSecCfg.junkPacketCount = tempASecCfg.junkPacketCount
- if tempASecCfg.junkPacketCount != 0 {
+ device.awg.ASecCfg.JunkPacketCount = tempAwg.ASecCfg.JunkPacketCount
+ if tempAwg.ASecCfg.JunkPacketCount != 0 {
isASecOn = true
}
- device.aSecCfg.junkPacketMinSize = tempASecCfg.junkPacketMinSize
- if tempASecCfg.junkPacketMinSize != 0 {
+ device.awg.ASecCfg.JunkPacketMinSize = tempAwg.ASecCfg.JunkPacketMinSize
+ if tempAwg.ASecCfg.JunkPacketMinSize != 0 {
isASecOn = true
}
- if device.aSecCfg.junkPacketCount > 0 &&
- tempASecCfg.junkPacketMaxSize == tempASecCfg.junkPacketMinSize {
+ if device.awg.ASecCfg.JunkPacketCount > 0 &&
+ tempAwg.ASecCfg.JunkPacketMaxSize == tempAwg.ASecCfg.JunkPacketMinSize {
- tempASecCfg.junkPacketMaxSize++ // to make rand gen work
+ tempAwg.ASecCfg.JunkPacketMaxSize++ // to make rand gen work
}
- if tempASecCfg.junkPacketMaxSize >= MaxSegmentSize {
- device.aSecCfg.junkPacketMinSize = 0
- device.aSecCfg.junkPacketMaxSize = 1
- if err != nil {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- "JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d; %w",
- tempASecCfg.junkPacketMaxSize,
- MaxSegmentSize,
- err,
- )
- } else {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- "JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d",
- tempASecCfg.junkPacketMaxSize,
- MaxSegmentSize,
- )
- }
- } else if tempASecCfg.junkPacketMaxSize < tempASecCfg.junkPacketMinSize {
- if err != nil {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- "maxSize: %d; should be greater than minSize: %d; %w",
- tempASecCfg.junkPacketMaxSize,
- tempASecCfg.junkPacketMinSize,
- err,
- )
- } else {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- "maxSize: %d; should be greater than minSize: %d",
- tempASecCfg.junkPacketMaxSize,
- tempASecCfg.junkPacketMinSize,
- )
- }
+ if tempAwg.ASecCfg.JunkPacketMaxSize >= MaxSegmentSize {
+ device.awg.ASecCfg.JunkPacketMinSize = 0
+ device.awg.ASecCfg.JunkPacketMaxSize = 1
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d",
+ tempAwg.ASecCfg.JunkPacketMaxSize,
+ MaxSegmentSize,
+ ))
+ } else if tempAwg.ASecCfg.JunkPacketMaxSize < tempAwg.ASecCfg.JunkPacketMinSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "maxSize: %d; should be greater than minSize: %d",
+ tempAwg.ASecCfg.JunkPacketMaxSize,
+ tempAwg.ASecCfg.JunkPacketMinSize,
+ ))
} else {
- device.aSecCfg.junkPacketMaxSize = tempASecCfg.junkPacketMaxSize
+ device.awg.ASecCfg.JunkPacketMaxSize = tempAwg.ASecCfg.JunkPacketMaxSize
}
- if tempASecCfg.junkPacketMaxSize != 0 {
+ if tempAwg.ASecCfg.JunkPacketMaxSize != 0 {
isASecOn = true
}
- if MessageInitiationSize+tempASecCfg.initPacketJunkSize >= MaxSegmentSize {
- if err != nil {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d; %w`,
- tempASecCfg.initPacketJunkSize,
- MaxSegmentSize,
- err,
- )
- } else {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempASecCfg.initPacketJunkSize,
- MaxSegmentSize,
- )
- }
+ newInitSize := MessageInitiationSize + tempAwg.ASecCfg.InitHeaderJunkSize
+
+ if newInitSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.ASecCfg.InitHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
} else {
- device.aSecCfg.initPacketJunkSize = tempASecCfg.initPacketJunkSize
+ device.awg.ASecCfg.InitHeaderJunkSize = tempAwg.ASecCfg.InitHeaderJunkSize
}
- if tempASecCfg.initPacketJunkSize != 0 {
+ if tempAwg.ASecCfg.InitHeaderJunkSize != 0 {
isASecOn = true
}
- if MessageResponseSize+tempASecCfg.responsePacketJunkSize >= MaxSegmentSize {
- if err != nil {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d; %w`,
- tempASecCfg.responsePacketJunkSize,
- MaxSegmentSize,
- err,
- )
- } else {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempASecCfg.responsePacketJunkSize,
- MaxSegmentSize,
- )
- }
+ newResponseSize := MessageResponseSize + tempAwg.ASecCfg.ResponseHeaderJunkSize
+
+ if newResponseSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.ASecCfg.ResponseHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
} else {
- device.aSecCfg.responsePacketJunkSize = tempASecCfg.responsePacketJunkSize
+ device.awg.ASecCfg.ResponseHeaderJunkSize = tempAwg.ASecCfg.ResponseHeaderJunkSize
}
- if tempASecCfg.responsePacketJunkSize != 0 {
+ if tempAwg.ASecCfg.ResponseHeaderJunkSize != 0 {
isASecOn = true
}
- if tempASecCfg.initPacketMagicHeader > 4 {
+ newCookieSize := MessageCookieReplySize + tempAwg.ASecCfg.CookieReplyHeaderJunkSize
+
+ if newCookieSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `cookie reply size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.ASecCfg.CookieReplyHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
+ } else {
+ device.awg.ASecCfg.CookieReplyHeaderJunkSize = tempAwg.ASecCfg.CookieReplyHeaderJunkSize
+ }
+
+ if tempAwg.ASecCfg.CookieReplyHeaderJunkSize != 0 {
+ isASecOn = true
+ }
+
+ newTransportSize := MessageTransportSize + tempAwg.ASecCfg.TransportHeaderJunkSize
+
+ if newTransportSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `transport size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.ASecCfg.TransportHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
+ } else {
+ device.awg.ASecCfg.TransportHeaderJunkSize = tempAwg.ASecCfg.TransportHeaderJunkSize
+ }
+
+ if tempAwg.ASecCfg.TransportHeaderJunkSize != 0 {
+ isASecOn = true
+ }
+
+ if tempAwg.ASecCfg.InitPacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating init_packet_magic_header")
- device.aSecCfg.initPacketMagicHeader = tempASecCfg.initPacketMagicHeader
- MessageInitiationType = device.aSecCfg.initPacketMagicHeader
+ device.awg.ASecCfg.InitPacketMagicHeader = tempAwg.ASecCfg.InitPacketMagicHeader
+ MessageInitiationType = device.awg.ASecCfg.InitPacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default init type")
- MessageInitiationType = 1
+ MessageInitiationType = DefaultMessageInitiationType
}
- if tempASecCfg.responsePacketMagicHeader > 4 {
+ if tempAwg.ASecCfg.ResponsePacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating response_packet_magic_header")
- device.aSecCfg.responsePacketMagicHeader = tempASecCfg.responsePacketMagicHeader
- MessageResponseType = device.aSecCfg.responsePacketMagicHeader
+ device.awg.ASecCfg.ResponsePacketMagicHeader = tempAwg.ASecCfg.ResponsePacketMagicHeader
+ MessageResponseType = device.awg.ASecCfg.ResponsePacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default response type")
- MessageResponseType = 2
+ MessageResponseType = DefaultMessageResponseType
}
- if tempASecCfg.underloadPacketMagicHeader > 4 {
+ if tempAwg.ASecCfg.UnderloadPacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating underload_packet_magic_header")
- device.aSecCfg.underloadPacketMagicHeader = tempASecCfg.underloadPacketMagicHeader
- MessageCookieReplyType = device.aSecCfg.underloadPacketMagicHeader
+ device.awg.ASecCfg.UnderloadPacketMagicHeader = tempAwg.ASecCfg.UnderloadPacketMagicHeader
+ MessageCookieReplyType = device.awg.ASecCfg.UnderloadPacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default underload type")
- MessageCookieReplyType = 3
+ MessageCookieReplyType = DefaultMessageCookieReplyType
}
- if tempASecCfg.transportPacketMagicHeader > 4 {
+ if tempAwg.ASecCfg.TransportPacketMagicHeader > 4 {
isASecOn = true
device.log.Verbosef("UAPI: Updating transport_packet_magic_header")
- device.aSecCfg.transportPacketMagicHeader = tempASecCfg.transportPacketMagicHeader
- MessageTransportType = device.aSecCfg.transportPacketMagicHeader
+ device.awg.ASecCfg.TransportPacketMagicHeader = tempAwg.ASecCfg.TransportPacketMagicHeader
+ MessageTransportType = device.awg.ASecCfg.TransportPacketMagicHeader
} else {
device.log.Verbosef("UAPI: Using default transport type")
- MessageTransportType = 4
+ MessageTransportType = DefaultMessageTransportType
}
- isSameMap := map[uint32]bool{}
- isSameMap[MessageInitiationType] = true
- isSameMap[MessageResponseType] = true
- isSameMap[MessageCookieReplyType] = true
- isSameMap[MessageTransportType] = true
+ isSameHeaderMap := map[uint32]struct{}{
+ MessageInitiationType: {},
+ MessageResponseType: {},
+ MessageCookieReplyType: {},
+ MessageTransportType: {},
+ }
// size will be different if same values
- if len(isSameMap) != 4 {
- if err != nil {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `magic headers should differ; got: init:%d; recv:%d; unde:%d; tran:%d; %w`,
- MessageInitiationType,
- MessageResponseType,
- MessageCookieReplyType,
- MessageTransportType,
- err,
- )
- } else {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `magic headers should differ; got: init:%d; recv:%d; unde:%d; tran:%d`,
- MessageInitiationType,
- MessageResponseType,
- MessageCookieReplyType,
- MessageTransportType,
- )
+ if len(isSameHeaderMap) != 4 {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `magic headers should differ; got: init:%d; recv:%d; unde:%d; tran:%d`,
+ MessageInitiationType,
+ MessageResponseType,
+ MessageCookieReplyType,
+ MessageTransportType,
+ ),
+ )
+ }
+
+ isSameSizeMap := map[int]struct{}{
+ newInitSize: {},
+ newResponseSize: {},
+ newCookieSize: {},
+ newTransportSize: {},
+ }
+
+ if len(isSameSizeMap) != 4 {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `new sizes should differ; init: %d; response: %d; cookie: %d; trans: %d`,
+ newInitSize,
+ newResponseSize,
+ newCookieSize,
+ newTransportSize,
+ ),
+ )
+ } else {
+ msgTypeToJunkSize = map[uint32]int{
+ MessageInitiationType: device.awg.ASecCfg.InitHeaderJunkSize,
+ MessageResponseType: device.awg.ASecCfg.ResponseHeaderJunkSize,
+ MessageCookieReplyType: device.awg.ASecCfg.CookieReplyHeaderJunkSize,
+ MessageTransportType: device.awg.ASecCfg.TransportHeaderJunkSize,
+ }
+
+ packetSizeToMsgType = map[int]uint32{
+ newInitSize: MessageInitiationType,
+ newResponseSize: MessageResponseType,
+ newCookieSize: MessageCookieReplyType,
+ newTransportSize: MessageTransportType,
}
}
- newInitSize := MessageInitiationSize + device.aSecCfg.initPacketJunkSize
- newResponseSize := MessageResponseSize + device.aSecCfg.responsePacketJunkSize
+ device.awg.IsASecOn.SetTo(isASecOn)
+ var err error
+ device.awg.JunkCreator, err = awg.NewJunkCreator(device.awg.ASecCfg)
+ if err != nil {
+ errs = append(errs, err)
+ }
- if newInitSize == newResponseSize {
- if err != nil {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `new init size:%d; and new response size:%d; should differ; %w`,
- newInitSize,
- newResponseSize,
- err,
- )
+ if tempAwg.HandshakeHandler.IsSet {
+ if err := tempAwg.HandshakeHandler.Validate(); err != nil {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid, "handshake handler validate: %w", err))
} else {
- err = ipcErrorf(
- ipc.IpcErrorInvalid,
- `new init size:%d; and new response size:%d; should differ`,
- newInitSize,
- newResponseSize,
- )
+ device.awg.HandshakeHandler = tempAwg.HandshakeHandler
+ device.awg.HandshakeHandler.ControlledJunk.DefaultJunkCount = tempAwg.ASecCfg.JunkPacketCount
+ device.awg.HandshakeHandler.SpecialJunk.DefaultJunkCount = tempAwg.ASecCfg.JunkPacketCount
+ device.version = VersionAwgSpecialHandshake
}
} else {
- packetSizeToMsgType = map[int]uint32{
- newInitSize: MessageInitiationType,
- newResponseSize: MessageResponseType,
- MessageCookieReplySize: MessageCookieReplyType,
- MessageTransportSize: MessageTransportType,
- }
-
- msgTypeToJunkSize = map[uint32]int{
- MessageInitiationType: device.aSecCfg.initPacketJunkSize,
- MessageResponseType: device.aSecCfg.responsePacketJunkSize,
- MessageCookieReplyType: 0,
- MessageTransportType: 0,
- }
+ device.version = VersionAwg
}
- device.isASecOn.SetTo(isASecOn)
- device.junkCreator, err = NewJunkCreator(device)
- device.aSecMux.Unlock()
+ device.awg.ASecMux.Unlock()
- return err
+ return errors.Join(errs...)
}
diff --git a/device/device_test.go b/device/device_test.go
index f66d326..5824cf9 100644
--- a/device/device_test.go
+++ b/device/device_test.go
@@ -7,19 +7,22 @@ package device
import (
"bytes"
+ "context"
"encoding/hex"
"fmt"
"io"
"math/rand"
"net/netip"
"os"
+ "os/signal"
"runtime"
"runtime/pprof"
"sync"
- "sync/atomic"
"testing"
"time"
+ "go.uber.org/atomic"
+
"github.com/amnezia-vpn/amneziawg-go/conn"
"github.com/amnezia-vpn/amneziawg-go/conn/bindtest"
"github.com/amnezia-vpn/amneziawg-go/tun"
@@ -50,7 +53,7 @@ func uapiCfg(cfg ...string) string {
// genConfigs generates a pair of configs that connect to each other.
// The configs use distinct, probably-usable ports.
-func genConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
+func genConfigs(tb testing.TB, cfg ...string) (cfgs, endpointCfgs [2]string) {
var key1, key2 NoisePrivateKey
_, err := rand.Read(key1[:])
if err != nil {
@@ -62,7 +65,8 @@ func genConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
}
pub1, pub2 := key1.publicKey(), key2.publicKey()
- cfgs[0] = uapiCfg(
+ args0 := append([]string(nil), cfg...)
+ args0 = append(args0, []string{
"private_key", hex.EncodeToString(key1[:]),
"listen_port", "0",
"replace_peers", "true",
@@ -70,12 +74,16 @@ func genConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
"protocol_version", "1",
"replace_allowed_ips", "true",
"allowed_ip", "1.0.0.2/32",
- )
+ }...)
+ cfgs[0] = uapiCfg(args0...)
+
endpointCfgs[0] = uapiCfg(
"public_key", hex.EncodeToString(pub2[:]),
"endpoint", "127.0.0.1:%d",
)
- cfgs[1] = uapiCfg(
+
+ args1 := append([]string(nil), cfg...)
+ args1 = append(args1, []string{
"private_key", hex.EncodeToString(key2[:]),
"listen_port", "0",
"replace_peers", "true",
@@ -83,66 +91,9 @@ func genConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
"protocol_version", "1",
"replace_allowed_ips", "true",
"allowed_ip", "1.0.0.1/32",
- )
- endpointCfgs[1] = uapiCfg(
- "public_key", hex.EncodeToString(pub1[:]),
- "endpoint", "127.0.0.1:%d",
- )
- return
-}
+ }...)
-func genASecurityConfigs(tb testing.TB) (cfgs, endpointCfgs [2]string) {
- var key1, key2 NoisePrivateKey
- _, err := rand.Read(key1[:])
- if err != nil {
- tb.Errorf("unable to generate private key random bytes: %v", err)
- }
- _, err = rand.Read(key2[:])
- if err != nil {
- tb.Errorf("unable to generate private key random bytes: %v", err)
- }
- pub1, pub2 := key1.publicKey(), key2.publicKey()
-
- cfgs[0] = uapiCfg(
- "private_key", hex.EncodeToString(key1[:]),
- "listen_port", "0",
- "replace_peers", "true",
- "jc", "5",
- "jmin", "500",
- "jmax", "1000",
- "s1", "30",
- "s2", "40",
- "h1", "123456",
- "h2", "67543",
- "h4", "32345",
- "h3", "123123",
- "public_key", hex.EncodeToString(pub2[:]),
- "protocol_version", "1",
- "replace_allowed_ips", "true",
- "allowed_ip", "1.0.0.2/32",
- )
- endpointCfgs[0] = uapiCfg(
- "public_key", hex.EncodeToString(pub2[:]),
- "endpoint", "127.0.0.1:%d",
- )
- cfgs[1] = uapiCfg(
- "private_key", hex.EncodeToString(key2[:]),
- "listen_port", "0",
- "replace_peers", "true",
- "jc", "5",
- "jmin", "500",
- "jmax", "1000",
- "s1", "30",
- "s2", "40",
- "h1", "123456",
- "h2", "67543",
- "h4", "32345",
- "h3", "123123",
- "public_key", hex.EncodeToString(pub1[:]),
- "protocol_version", "1",
- "replace_allowed_ips", "true",
- "allowed_ip", "1.0.0.1/32",
- )
+ cfgs[1] = uapiCfg(args1...)
endpointCfgs[1] = uapiCfg(
"public_key", hex.EncodeToString(pub1[:]),
"endpoint", "127.0.0.1:%d",
@@ -185,9 +136,10 @@ func (pair *testPair) Send(
// pong is the new ping
p0, p1 = p1, p0
}
+
msg := tuntest.Ping(p0.ip, p1.ip)
p1.tun.Outbound <- msg
- timer := time.NewTimer(5 * time.Second)
+ timer := time.NewTimer(6 * time.Second)
defer timer.Stop()
var err error
select {
@@ -214,14 +166,12 @@ func (pair *testPair) Send(
// genTestPair creates a testPair.
func genTestPair(
tb testing.TB,
- realSocket, withASecurity bool,
+ realSocket bool,
+ extraCfg ...string,
) (pair testPair) {
var cfg, endpointCfg [2]string
- if withASecurity {
- cfg, endpointCfg = genASecurityConfigs(tb)
- } else {
- cfg, endpointCfg = genConfigs(tb)
- }
+ cfg, endpointCfg = genConfigs(tb, extraCfg...)
+
var binds [2]conn.Bind
if realSocket {
binds[0], binds[1] = conn.NewDefaultBind(), conn.NewDefaultBind()
@@ -265,7 +215,7 @@ func genTestPair(
func TestTwoDevicePing(t *testing.T) {
goroutineLeakCheck(t)
- pair := genTestPair(t, true, false)
+ pair := genTestPair(t, true)
t.Run("ping 1.0.0.1", func(t *testing.T) {
pair.Send(t, Ping, nil)
})
@@ -274,9 +224,23 @@ func TestTwoDevicePing(t *testing.T) {
})
}
-func TestASecurityTwoDevicePing(t *testing.T) {
+// Run test with -race=false to avoid the race for setting the default msgTypes 2 times
+func TestAWGDevicePing(t *testing.T) {
goroutineLeakCheck(t)
- pair := genTestPair(t, true, true)
+
+ pair := genTestPair(t, true,
+ "jc", "5",
+ "jmin", "500",
+ "jmax", "1000",
+ "s1", "30",
+ "s2", "40",
+ "s3", "50",
+ "s4", "5",
+ "h1", "123456",
+ "h2", "67543",
+ "h3", "123123",
+ "h4", "32345",
+ )
t.Run("ping 1.0.0.1", func(t *testing.T) {
pair.Send(t, Ping, nil)
})
@@ -285,13 +249,58 @@ func TestASecurityTwoDevicePing(t *testing.T) {
})
}
+// Needs to be stopped with Ctrl-C
+func TestAWGHandshakeDevicePing(t *testing.T) {
+ t.Skip("This test is intended to be run manually, not as part of the test suite.")
+
+ signalContext, cancel := signal.NotifyContext(context.Background(), os.Interrupt)
+ defer cancel()
+ isRunning := atomic.NewBool(true)
+ go func() {
+ <-signalContext.Done()
+ fmt.Println("Waiting to finish")
+ isRunning.Store(false)
+ }()
+
+ goroutineLeakCheck(t)
+ pair := genTestPair(t, true,
+ "i1", "",
+ "i2", "",
+ "j1", "",
+ "j2", "",
+ "j3", "",
+ "itime", "60",
+ // "jc", "1",
+ // "jmin", "500",
+ // "jmax", "1000",
+ // "s1", "30",
+ // "s2", "40",
+ // "h1", "123456",
+ // "h2", "67543",
+ // "h4", "32345",
+ // "h3", "123123",
+ )
+ t.Run("ping 1.0.0.1", func(t *testing.T) {
+ for isRunning.Load() {
+ pair.Send(t, Ping, nil)
+ time.Sleep(2 * time.Second)
+ }
+ })
+ t.Run("ping 1.0.0.2", func(t *testing.T) {
+ for isRunning.Load() {
+ pair.Send(t, Pong, nil)
+ time.Sleep(2 * time.Second)
+ }
+ })
+}
+
func TestUpDown(t *testing.T) {
goroutineLeakCheck(t)
const itrials = 50
const otrials = 10
for n := 0; n < otrials; n++ {
- pair := genTestPair(t, false, false)
+ pair := genTestPair(t, false)
for i := range pair {
for k := range pair[i].dev.peers.keyMap {
pair[i].dev.IpcSet(fmt.Sprintf("public_key=%s\npersistent_keepalive_interval=1\n", hex.EncodeToString(k[:])))
@@ -325,7 +334,7 @@ func TestUpDown(t *testing.T) {
// TestConcurrencySafety does other things concurrently with tunnel use.
// It is intended to be used with the race detector to catch data races.
func TestConcurrencySafety(t *testing.T) {
- pair := genTestPair(t, true, false)
+ pair := genTestPair(t, true)
done := make(chan struct{})
const warmupIters = 10
@@ -406,7 +415,7 @@ func TestConcurrencySafety(t *testing.T) {
}
func BenchmarkLatency(b *testing.B) {
- pair := genTestPair(b, true, false)
+ pair := genTestPair(b, true)
// Establish a connection.
pair.Send(b, Ping, nil)
@@ -420,7 +429,7 @@ func BenchmarkLatency(b *testing.B) {
}
func BenchmarkThroughput(b *testing.B) {
- pair := genTestPair(b, true, false)
+ pair := genTestPair(b, true)
// Establish a connection.
pair.Send(b, Ping, nil)
@@ -464,7 +473,7 @@ func BenchmarkThroughput(b *testing.B) {
}
func BenchmarkUAPIGet(b *testing.B) {
- pair := genTestPair(b, true, false)
+ pair := genTestPair(b, true)
pair.Send(b, Ping, nil)
pair.Send(b, Pong, nil)
b.ReportAllocs()
diff --git a/device/noise-protocol.go b/device/noise-protocol.go
index 789eb16..f637b24 100644
--- a/device/noise-protocol.go
+++ b/device/noise-protocol.go
@@ -52,11 +52,18 @@ const (
WGLabelCookie = "cookie--"
)
+const (
+ DefaultMessageInitiationType uint32 = 1
+ DefaultMessageResponseType uint32 = 2
+ DefaultMessageCookieReplyType uint32 = 3
+ DefaultMessageTransportType uint32 = 4
+)
+
var (
- MessageInitiationType uint32 = 1
- MessageResponseType uint32 = 2
- MessageCookieReplyType uint32 = 3
- MessageTransportType uint32 = 4
+ MessageInitiationType uint32 = DefaultMessageInitiationType
+ MessageResponseType uint32 = DefaultMessageResponseType
+ MessageCookieReplyType uint32 = DefaultMessageCookieReplyType
+ MessageTransportType uint32 = DefaultMessageTransportType
)
const (
@@ -75,9 +82,10 @@ const (
MessageTransportOffsetContent = 16
)
-var packetSizeToMsgType map[int]uint32
-
-var msgTypeToJunkSize map[uint32]int
+var (
+ packetSizeToMsgType map[int]uint32
+ msgTypeToJunkSize map[uint32]int
+)
/* Type is an 8-bit field, followed by 3 nul bytes,
* by marshalling the messages in little-endian byteorder
@@ -197,12 +205,12 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e
handshake.mixHash(handshake.remoteStatic[:])
- device.aSecMux.RLock()
+ device.awg.ASecMux.RLock()
msg := MessageInitiation{
Type: MessageInitiationType,
Ephemeral: handshake.localEphemeral.publicKey(),
}
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
handshake.mixKey(msg.Ephemeral[:])
handshake.mixHash(msg.Ephemeral[:])
@@ -256,12 +264,12 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer {
chainKey [blake2s.Size]byte
)
- device.aSecMux.RLock()
+ device.awg.ASecMux.RLock()
if msg.Type != MessageInitiationType {
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
return nil
}
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
device.staticIdentity.RLock()
defer device.staticIdentity.RUnlock()
@@ -376,9 +384,9 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
var msg MessageResponse
- device.aSecMux.RLock()
+ device.awg.ASecMux.RLock()
msg.Type = MessageResponseType
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
msg.Sender = handshake.localIndex
msg.Receiver = handshake.remoteIndex
@@ -428,12 +436,12 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
- device.aSecMux.RLock()
+ device.awg.ASecMux.RLock()
if msg.Type != MessageResponseType {
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
return nil
}
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
// lookup handshake by receiver
diff --git a/device/peer.go b/device/peer.go
index 8f88b2a..e8a5168 100644
--- a/device/peer.go
+++ b/device/peer.go
@@ -13,6 +13,7 @@ import (
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device/awg"
)
type Peer struct {
@@ -113,6 +114,16 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
return peer, nil
}
+func (peer *Peer) SendAndCountBuffers(buffers [][]byte) error {
+ err := peer.SendBuffers(buffers)
+ if err == nil {
+ awg.PacketCounter.Add(uint64(len(buffers)))
+ return nil
+ }
+
+ return err
+}
+
func (peer *Peer) SendBuffers(buffers [][]byte) error {
peer.device.net.RLock()
defer peer.device.net.RUnlock()
diff --git a/device/receive.go b/device/receive.go
index 0a4910a..6daba0d 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -129,7 +129,7 @@ func (device *Device) RoutineReceiveIncoming(
}
deathSpiral = 0
- device.aSecMux.RLock()
+ device.awg.ASecMux.RLock()
// handle each packet in the batch
for i, size := range sizes[:count] {
if size < MinMessageSize {
@@ -137,10 +137,14 @@ func (device *Device) RoutineReceiveIncoming(
}
// check size of packet
-
packet := bufsArrs[i][:size]
var msgType uint32
- if device.isAdvancedSecurityOn() {
+ if device.isAWG() {
+ // TODO:
+ // if awg.WaitResponse.ShouldWait.IsSet() {
+ // awg.WaitResponse.Channel <- struct{}{}
+ // }
+
if assumedMsgType, ok := packetSizeToMsgType[size]; ok {
junkSize := msgTypeToJunkSize[assumedMsgType]
// transport size can align with other header types;
@@ -149,19 +153,29 @@ func (device *Device) RoutineReceiveIncoming(
if msgType == assumedMsgType {
packet = packet[junkSize:]
} else {
- device.log.Verbosef("Transport packet lined up with another msg type")
+ device.log.Verbosef("transport packet lined up with another msg type")
msgType = binary.LittleEndian.Uint32(packet[:4])
}
} else {
- msgType = binary.LittleEndian.Uint32(packet[:4])
+ transportJunkSize := device.awg.ASecCfg.TransportHeaderJunkSize
+ msgType = binary.LittleEndian.Uint32(packet[transportJunkSize : transportJunkSize+4])
if msgType != MessageTransportType {
- device.log.Verbosef("ASec: Received message with unknown type")
+ // probably a junk packet
+ device.log.Verbosef("aSec: Received message with unknown type: %d", msgType)
continue
}
+
+ // remove junk from bufsArrs by shifting the packet
+ // this buffer is also used for decryption, so it needs to be corrected
+ copy(bufsArrs[i][:size], packet[transportJunkSize:])
+ size -= transportJunkSize
+ // need to reinitialize packet as well
+ packet = packet[:size]
}
} else {
msgType = binary.LittleEndian.Uint32(packet[:4])
}
+
switch msgType {
// check if transport
@@ -245,7 +259,7 @@ func (device *Device) RoutineReceiveIncoming(
default:
}
}
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
for peer, elemsContainer := range elemsByPeer {
if peer.isRunning.Load() {
peer.queue.inbound.c <- elemsContainer
@@ -304,7 +318,7 @@ func (device *Device) RoutineHandshake(id int) {
for elem := range device.queue.handshake.c {
- device.aSecMux.RLock()
+ device.awg.ASecMux.RLock()
// handle cookie fields and ratelimiting
@@ -456,7 +470,7 @@ func (device *Device) RoutineHandshake(id int) {
peer.SendKeepalive()
}
skip:
- device.aSecMux.RUnlock()
+ device.awg.ASecMux.RUnlock()
device.PutMessageBuffer(elem.buffer)
}
}
diff --git a/device/send.go b/device/send.go
index 7f0faa3..04ca2ad 100644
--- a/device/send.go
+++ b/device/send.go
@@ -124,12 +124,30 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
return err
}
var sendBuffer [][]byte
+
// so only packet processed for cookie generation
var junkedHeader []byte
- if peer.device.isAdvancedSecurityOn() {
- peer.device.aSecMux.RLock()
- junks, err := peer.device.junkCreator.createJunkPackets()
- peer.device.aSecMux.RUnlock()
+ if peer.device.version >= VersionAwg {
+ var junks [][]byte
+ if peer.device.version == VersionAwgSpecialHandshake {
+ peer.device.awg.ASecMux.RLock()
+ // set junks depending on packet type
+ junks = peer.device.awg.HandshakeHandler.GenerateSpecialJunk()
+ if junks == nil {
+ junks = peer.device.awg.HandshakeHandler.GenerateControlledJunk()
+ if junks != nil {
+ peer.device.log.Verbosef("%v - Controlled junks sent", peer)
+ }
+ } else {
+ peer.device.log.Verbosef("%v - Special junks sent", peer)
+ }
+ peer.device.awg.ASecMux.RUnlock()
+ } else {
+ junks = make([][]byte, 0, peer.device.awg.ASecCfg.JunkPacketCount)
+ }
+ peer.device.awg.ASecMux.RLock()
+ err := peer.device.awg.JunkCreator.CreateJunkPackets(&junks)
+ peer.device.awg.ASecMux.RUnlock()
if err != nil {
peer.device.log.Errorf("%v - %v", peer, err)
@@ -145,19 +163,11 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
}
}
- peer.device.aSecMux.RLock()
- if peer.device.aSecCfg.initPacketJunkSize != 0 {
- buf := make([]byte, 0, peer.device.aSecCfg.initPacketJunkSize)
- writer := bytes.NewBuffer(buf[:0])
- err = peer.device.junkCreator.appendJunk(writer, peer.device.aSecCfg.initPacketJunkSize)
- if err != nil {
- peer.device.log.Errorf("%v - %v", peer, err)
- peer.device.aSecMux.RUnlock()
- return err
- }
- junkedHeader = writer.Bytes()
+ junkedHeader, err = peer.device.awg.CreateInitHeaderJunk()
+ if err != nil {
+ peer.device.log.Errorf("%v - %v", peer, err)
+ return err
}
- peer.device.aSecMux.RUnlock()
}
var buf [MessageInitiationSize]byte
@@ -172,7 +182,7 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
sendBuffer = append(sendBuffer, junkedHeader)
- err = peer.SendBuffers(sendBuffer)
+ err = peer.SendAndCountBuffers(sendBuffer)
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err)
}
@@ -193,22 +203,13 @@ func (peer *Peer) SendHandshakeResponse() error {
peer.device.log.Errorf("%v - Failed to create response message: %v", peer, err)
return err
}
- var junkedHeader []byte
- if peer.device.isAdvancedSecurityOn() {
- peer.device.aSecMux.RLock()
- if peer.device.aSecCfg.responsePacketJunkSize != 0 {
- buf := make([]byte, 0, peer.device.aSecCfg.responsePacketJunkSize)
- writer := bytes.NewBuffer(buf[:0])
- err = peer.device.junkCreator.appendJunk(writer, peer.device.aSecCfg.responsePacketJunkSize)
- if err != nil {
- peer.device.aSecMux.RUnlock()
- peer.device.log.Errorf("%v - %v", peer, err)
- return err
- }
- junkedHeader = writer.Bytes()
- }
- peer.device.aSecMux.RUnlock()
+
+ junkedHeader, err := peer.device.awg.CreateResponseHeaderJunk()
+ if err != nil {
+ peer.device.log.Errorf("%v - %v", peer, err)
+ return err
}
+
var buf [MessageResponseSize]byte
writer := bytes.NewBuffer(buf[:0])
@@ -228,7 +229,7 @@ func (peer *Peer) SendHandshakeResponse() error {
peer.timersAnyAuthenticatedPacketSent()
// TODO: allocation could be avoided
- err = peer.SendBuffers([][]byte{junkedHeader})
+ err = peer.SendAndCountBuffers([][]byte{junkedHeader})
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake response: %v", peer, err)
}
@@ -251,11 +252,19 @@ func (device *Device) SendHandshakeCookie(
return err
}
+ junkedHeader, err := device.awg.CreateCookieReplyHeaderJunk()
+ if err != nil {
+ device.log.Errorf("%v - %v", device, err)
+ return err
+ }
+
var buf [MessageCookieReplySize]byte
writer := bytes.NewBuffer(buf[:0])
binary.Write(writer, binary.LittleEndian, reply)
+
+ junkedHeader = append(junkedHeader, writer.Bytes()...)
// TODO: allocation could be avoided
- device.net.bind.Send([][]byte{writer.Bytes()}, initiatingElem.endpoint)
+ device.net.bind.Send([][]byte{junkedHeader}, initiatingElem.endpoint)
return nil
}
@@ -576,6 +585,14 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
for _, elem := range elemsContainer.elems {
if len(elem.packet) != MessageKeepaliveSize {
dataSent = true
+
+ junkedHeader, err := device.awg.CreateTransportHeaderJunk(len(elem.packet))
+ if err != nil {
+ device.log.Errorf("%v - %v", device, err)
+ continue
+ }
+
+ elem.packet = append(junkedHeader, elem.packet...)
}
bufs = append(bufs, elem.packet)
}
@@ -583,10 +600,11 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
- err := peer.SendBuffers(bufs)
+ err := peer.SendAndCountBuffers(bufs)
if dataSent {
peer.timersDataSent()
}
+
for _, elem := range elemsContainer.elems {
device.PutMessageBuffer(elem.buffer)
device.PutOutboundElement(elem)
diff --git a/device/uapi.go b/device/uapi.go
index 870bddc..e9f962a 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -18,6 +18,7 @@ import (
"sync"
"time"
+ "github.com/amnezia-vpn/amneziawg-go/device/awg"
"github.com/amnezia-vpn/amneziawg-go/ipc"
)
@@ -97,33 +98,51 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("fwmark=%d", device.net.fwmark)
}
- if device.isAdvancedSecurityOn() {
- if device.aSecCfg.junkPacketCount != 0 {
- sendf("jc=%d", device.aSecCfg.junkPacketCount)
+ if device.isAWG() {
+ if device.awg.ASecCfg.JunkPacketCount != 0 {
+ sendf("jc=%d", device.awg.ASecCfg.JunkPacketCount)
}
- if device.aSecCfg.junkPacketMinSize != 0 {
- sendf("jmin=%d", device.aSecCfg.junkPacketMinSize)
+ if device.awg.ASecCfg.JunkPacketMinSize != 0 {
+ sendf("jmin=%d", device.awg.ASecCfg.JunkPacketMinSize)
}
- if device.aSecCfg.junkPacketMaxSize != 0 {
- sendf("jmax=%d", device.aSecCfg.junkPacketMaxSize)
+ if device.awg.ASecCfg.JunkPacketMaxSize != 0 {
+ sendf("jmax=%d", device.awg.ASecCfg.JunkPacketMaxSize)
}
- if device.aSecCfg.initPacketJunkSize != 0 {
- sendf("s1=%d", device.aSecCfg.initPacketJunkSize)
+ if device.awg.ASecCfg.InitHeaderJunkSize != 0 {
+ sendf("s1=%d", device.awg.ASecCfg.InitHeaderJunkSize)
}
- if device.aSecCfg.responsePacketJunkSize != 0 {
- sendf("s2=%d", device.aSecCfg.responsePacketJunkSize)
+ if device.awg.ASecCfg.ResponseHeaderJunkSize != 0 {
+ sendf("s2=%d", device.awg.ASecCfg.ResponseHeaderJunkSize)
}
- if device.aSecCfg.initPacketMagicHeader != 0 {
- sendf("h1=%d", device.aSecCfg.initPacketMagicHeader)
+ if device.awg.ASecCfg.CookieReplyHeaderJunkSize != 0 {
+ sendf("s3=%d", device.awg.ASecCfg.CookieReplyHeaderJunkSize)
}
- if device.aSecCfg.responsePacketMagicHeader != 0 {
- sendf("h2=%d", device.aSecCfg.responsePacketMagicHeader)
+ if device.awg.ASecCfg.TransportHeaderJunkSize != 0 {
+ sendf("s4=%d", device.awg.ASecCfg.TransportHeaderJunkSize)
}
- if device.aSecCfg.underloadPacketMagicHeader != 0 {
- sendf("h3=%d", device.aSecCfg.underloadPacketMagicHeader)
+ if device.awg.ASecCfg.InitPacketMagicHeader != 0 {
+ sendf("h1=%d", device.awg.ASecCfg.InitPacketMagicHeader)
}
- if device.aSecCfg.transportPacketMagicHeader != 0 {
- sendf("h4=%d", device.aSecCfg.transportPacketMagicHeader)
+ if device.awg.ASecCfg.ResponsePacketMagicHeader != 0 {
+ sendf("h2=%d", device.awg.ASecCfg.ResponsePacketMagicHeader)
+ }
+ if device.awg.ASecCfg.UnderloadPacketMagicHeader != 0 {
+ sendf("h3=%d", device.awg.ASecCfg.UnderloadPacketMagicHeader)
+ }
+ if device.awg.ASecCfg.TransportPacketMagicHeader != 0 {
+ sendf("h4=%d", device.awg.ASecCfg.TransportPacketMagicHeader)
+ }
+
+ specialJunkIpcFields := device.awg.HandshakeHandler.SpecialJunk.IpcGetFields()
+ for _, field := range specialJunkIpcFields {
+ sendf("%s=%s", field.Key, field.Value)
+ }
+ controlledJunkIpcFields := device.awg.HandshakeHandler.ControlledJunk.IpcGetFields()
+ for _, field := range controlledJunkIpcFields {
+ sendf("%s=%s", field.Key, field.Value)
+ }
+ if device.awg.HandshakeHandler.ITimeout != 0 {
+ sendf("itime=%d", device.awg.HandshakeHandler.ITimeout/time.Second)
}
}
@@ -180,13 +199,13 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
peer := new(ipcSetPeer)
deviceConfig := true
- tempASecCfg := aSecCfgType{}
+ tempAwg := awg.Protocol{}
scanner := bufio.NewScanner(r)
for scanner.Scan() {
line := scanner.Text()
if line == "" {
// Blank line means terminate operation.
- err := device.handlePostConfig(&tempASecCfg)
+ err := device.handlePostConfig(&tempAwg)
if err != nil {
return err
}
@@ -217,7 +236,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
var err error
if deviceConfig {
- err = device.handleDeviceLine(key, value, &tempASecCfg)
+ err = device.handleDeviceLine(key, value, &tempAwg)
} else {
err = device.handlePeerLine(peer, key, value)
}
@@ -225,7 +244,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return err
}
}
- err = device.handlePostConfig(&tempASecCfg)
+ err = device.handlePostConfig(&tempAwg)
if err != nil {
return err
}
@@ -237,7 +256,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return nil
}
-func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgType) error {
+func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol) error {
switch key {
case "private_key":
var sk NoisePrivateKey
@@ -278,7 +297,11 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
case "replace_peers":
if value != "true" {
- return ipcErrorf(ipc.IpcErrorInvalid, "failed to set replace_peers, invalid value: %v", value)
+ return ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "failed to set replace_peers, invalid value: %v",
+ value,
+ )
}
device.log.Verbosef("UAPI: Removing all peers")
device.RemoveAllPeers()
@@ -286,80 +309,138 @@ func (device *Device) handleDeviceLine(key, value string, tempASecCfg *aSecCfgTy
case "jc":
junkPacketCount, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse junk_packet_count %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_count %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_count")
- tempASecCfg.junkPacketCount = junkPacketCount
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.JunkPacketCount = junkPacketCount
+ tempAwg.ASecCfg.IsSet = true
case "jmin":
junkPacketMinSize, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse junk_packet_min_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_min_size %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_min_size")
- tempASecCfg.junkPacketMinSize = junkPacketMinSize
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.JunkPacketMinSize = junkPacketMinSize
+ tempAwg.ASecCfg.IsSet = true
case "jmax":
junkPacketMaxSize, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse junk_packet_max_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_max_size %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_max_size")
- tempASecCfg.junkPacketMaxSize = junkPacketMaxSize
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.JunkPacketMaxSize = junkPacketMaxSize
+ tempAwg.ASecCfg.IsSet = true
case "s1":
initPacketJunkSize, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse init_packet_junk_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse init_packet_junk_size %w", err)
}
device.log.Verbosef("UAPI: Updating init_packet_junk_size")
- tempASecCfg.initPacketJunkSize = initPacketJunkSize
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.InitHeaderJunkSize = initPacketJunkSize
+ tempAwg.ASecCfg.IsSet = true
case "s2":
responsePacketJunkSize, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse response_packet_junk_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse response_packet_junk_size %w", err)
}
device.log.Verbosef("UAPI: Updating response_packet_junk_size")
- tempASecCfg.responsePacketJunkSize = responsePacketJunkSize
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.ResponseHeaderJunkSize = responsePacketJunkSize
+ tempAwg.ASecCfg.IsSet = true
+
+ case "s3":
+ cookieReplyPacketJunkSize, err := strconv.Atoi(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse cookie_reply_packet_junk_size %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating cookie_reply_packet_junk_size")
+ tempAwg.ASecCfg.CookieReplyHeaderJunkSize = cookieReplyPacketJunkSize
+ tempAwg.ASecCfg.IsSet = true
+
+ case "s4":
+ transportPacketJunkSize, err := strconv.Atoi(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse transport_packet_junk_size %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating transport_packet_junk_size")
+ tempAwg.ASecCfg.TransportHeaderJunkSize = transportPacketJunkSize
+ tempAwg.ASecCfg.IsSet = true
case "h1":
initPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse init_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse init_packet_magic_header %w", err)
}
- tempASecCfg.initPacketMagicHeader = uint32(initPacketMagicHeader)
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.InitPacketMagicHeader = uint32(initPacketMagicHeader)
+ tempAwg.ASecCfg.IsSet = true
case "h2":
responsePacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse response_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse response_packet_magic_header %w", err)
}
- tempASecCfg.responsePacketMagicHeader = uint32(responsePacketMagicHeader)
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.ResponsePacketMagicHeader = uint32(responsePacketMagicHeader)
+ tempAwg.ASecCfg.IsSet = true
case "h3":
underloadPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse underload_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse underload_packet_magic_header %w", err)
}
- tempASecCfg.underloadPacketMagicHeader = uint32(underloadPacketMagicHeader)
- tempASecCfg.isSet = true
+ tempAwg.ASecCfg.UnderloadPacketMagicHeader = uint32(underloadPacketMagicHeader)
+ tempAwg.ASecCfg.IsSet = true
case "h4":
transportPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "faield to parse transport_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse transport_packet_magic_header %w", err)
+ }
+ tempAwg.ASecCfg.TransportPacketMagicHeader = uint32(transportPacketMagicHeader)
+ tempAwg.ASecCfg.IsSet = true
+ case "i1", "i2", "i3", "i4", "i5":
+ if len(value) == 0 {
+ device.log.Verbosef("UAPI: received empty %s", key)
+ return nil
}
- tempASecCfg.transportPacketMagicHeader = uint32(transportPacketMagicHeader)
- tempASecCfg.isSet = true
+ generators, err := awg.Parse(key, value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "invalid %s: %w", key, err)
+ }
+ device.log.Verbosef("UAPI: Updating %s", key)
+ tempAwg.HandshakeHandler.SpecialJunk.AppendGenerator(generators)
+ tempAwg.HandshakeHandler.IsSet = true
+ case "j1", "j2", "j3":
+ if len(value) == 0 {
+ device.log.Verbosef("UAPI: received empty %s", key)
+ return nil
+ }
+
+ generators, err := awg.Parse(key, value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "invalid %s: %w", key, err)
+ }
+ device.log.Verbosef("UAPI: Updating %s", key)
+
+ tempAwg.HandshakeHandler.ControlledJunk.AppendGenerator(generators)
+ tempAwg.HandshakeHandler.IsSet = true
+ case "itime":
+ if len(value) == 0 {
+ device.log.Verbosef("UAPI: received empty itime")
+ return nil
+ }
+
+ itime, err := strconv.ParseInt(value, 10, 64)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "parse itime %w", err)
+ }
+ device.log.Verbosef("UAPI: Updating itime")
+
+ tempAwg.HandshakeHandler.ITimeout = time.Duration(itime) * time.Second
+ tempAwg.HandshakeHandler.IsSet = true
default:
return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
}
@@ -432,7 +513,11 @@ func (device *Device) handlePeerLine(
case "update_only":
// allow disabling of creation
if value != "true" {
- return ipcErrorf(ipc.IpcErrorInvalid, "failed to set update only, invalid value: %v", value)
+ return ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "failed to set update only, invalid value: %v",
+ value,
+ )
}
if peer.created && !peer.dummy {
device.RemovePeer(peer.handshake.remoteStatic)
@@ -478,7 +563,11 @@ func (device *Device) handlePeerLine(
secs, err := strconv.ParseUint(value, 10, 16)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "failed to set persistent keepalive interval: %w", err)
+ return ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "failed to set persistent keepalive interval: %w",
+ err,
+ )
}
old := peer.persistentKeepaliveInterval.Swap(uint32(secs))
@@ -489,7 +578,11 @@ func (device *Device) handlePeerLine(
case "replace_allowed_ips":
device.log.Verbosef("%v - UAPI: Removing all allowedips", peer.Peer)
if value != "true" {
- return ipcErrorf(ipc.IpcErrorInvalid, "failed to replace allowedips, invalid value: %v", value)
+ return ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "failed to replace allowedips, invalid value: %v",
+ value,
+ )
}
if peer.dummy {
return nil
@@ -568,7 +661,11 @@ func (device *Device) IpcHandle(socket net.Conn) {
return
}
if nextByte != '\n' {
- err = ipcErrorf(ipc.IpcErrorInvalid, "trailing character in UAPI get: %q", nextByte)
+ err = ipcErrorf(
+ ipc.IpcErrorInvalid,
+ "trailing character in UAPI get: %q",
+ nextByte,
+ )
break
}
err = device.IpcGetOperation(buffered.Writer)
diff --git a/go.mod b/go.mod
index 99569f3..5e5f34d 100644
--- a/go.mod
+++ b/go.mod
@@ -1,17 +1,23 @@
module github.com/amnezia-vpn/amneziawg-go
-go 1.24
+go 1.24.4
require (
+ github.com/stretchr/testify v1.10.0
+ github.com/tevino/abool v1.2.0
github.com/tevino/abool/v2 v2.1.0
- golang.org/x/crypto v0.37.0
- golang.org/x/net v0.39.0
- golang.org/x/sys v0.32.0
+ go.uber.org/atomic v1.11.0
+ golang.org/x/crypto v0.39.0
+ golang.org/x/net v0.41.0
+ golang.org/x/sys v0.33.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
- gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c
+ gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489
)
require (
+ github.com/davecgh/go-spew v1.1.1 // indirect
github.com/google/btree v1.1.3 // indirect
+ github.com/pmezard/go-difflib v1.0.0 // indirect
golang.org/x/time v0.9.0 // indirect
+ gopkg.in/yaml.v3 v3.0.1 // indirect
)
diff --git a/go.sum b/go.sum
index b8ac0bd..6b8f36b 100644
--- a/go.sum
+++ b/go.sum
@@ -1,16 +1,40 @@
+github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
+github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
+github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38=
+github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
+github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
+github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
+github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
+github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
+github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
+github.com/tevino/abool v1.2.0 h1:heAkClL8H6w+mK5md9dzsuohKeXHUpY7Vw0ZCKW+huA=
+github.com/tevino/abool v1.2.0/go.mod h1:qc66Pna1RiIsPa7O4Egxxs9OqkuxDX55zznh9K07Tzg=
github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
-golang.org/x/crypto v0.37.0 h1:kJNSjF/Xp7kU0iB2Z+9viTPMW4EqqsrywMXLJOOsXSE=
-golang.org/x/crypto v0.37.0/go.mod h1:vg+k43peMZ0pUMhYmVAWysMK35e6ioLh3wB8ZCAfbVc=
-golang.org/x/net v0.39.0 h1:ZCu7HMWDxpXpaiKdhzIfaltL9Lp31x/3fCP11bc6/fY=
-golang.org/x/net v0.39.0/go.mod h1:X7NRbYVEA+ewNkCNyJ513WmMdQ3BineSwVtN2zD/d+E=
-golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20=
-golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
+go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
+go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
+golang.org/x/crypto v0.39.0 h1:SHs+kF4LP+f+p14esP5jAoDpHU8Gu/v9lFRK6IT5imM=
+golang.org/x/crypto v0.39.0/go.mod h1:L+Xg3Wf6HoL4Bn4238Z6ft6KfEpN0tJGo53AAPC632U=
+golang.org/x/mod v0.13.0 h1:I/DsJXRlw/8l/0c24sM9yb0T4z9liZTduXvdAWYiysY=
+golang.org/x/mod v0.21.0 h1:vvrHzRwRfVKSiLrG+d4FMl/Qi4ukBCE6kZlTUkDYRT0=
+golang.org/x/mod v0.21.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
+golang.org/x/net v0.41.0 h1:vBTly1HeNPEn3wtREYfy4GZ/NECgw2Cnl+nK6Nz3uvw=
+golang.org/x/net v0.41.0/go.mod h1:B/K4NNqkfmg07DQYrbwvSluqCJOOXwUjeb/5lOisjbA=
+golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
+golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/time v0.9.0 h1:EsRrnYcQiGH+5FfbgvV4AP7qEZstoyrHB0DzarOQ4ZY=
golang.org/x/time v0.9.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
-gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI=
-gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
+gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
+gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
+gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
+gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
+gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489 h1:ze1vwAdliUAr68RQ5NtufWaXaOg8WUO2OACzEV+TNdE=
+gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489/go.mod h1:10sU+Uh5KKNv1+2x2A0Gvzt8FjD3ASIhorV3YsauXhk=
+gvisor.dev/gvisor v0.0.0-20250428193742-2d800c3129d5 h1:sfK5nHuG7lRFZ2FdTT3RimOqWBg8IrVm+/Vko1FVOsk=
+gvisor.dev/gvisor v0.0.0-20250428193742-2d800c3129d5/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
+gvisor.dev/gvisor v0.0.0-20250606233247-e3c4c4cad86f h1:zmc4cHEcCudRt2O8VsCW7nYLfAsbVY2i910/DAop1TM=
+gvisor.dev/gvisor v0.0.0-20250606233247-e3c4c4cad86f/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
From 3f19f1c657d4a338f61eb2495eb4a2a8a6ac4843 Mon Sep 17 00:00:00 2001
From: Yaroslav Gurov
Date: Mon, 7 Jul 2025 15:15:29 +0200
Subject: [PATCH 67/75] fix: restore Dockerfile
---
Dockerfile | 17 ++---------------
1 file changed, 2 insertions(+), 15 deletions(-)
diff --git a/Dockerfile b/Dockerfile
index 6d60440..f165899 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -8,23 +8,10 @@ RUN go mod download && \
FROM alpine:3.19
ARG AWGTOOLS_RELEASE="1.0.20241018"
-RUN apk add linux-headers build-base
-COPY awg-tools /awg-tools
-RUN pwd && ls -la / && ls -la /awg-tools
-WORKDIR /awg-tools/src
-# RUN ls -la && pwd && ls awg-tools
-RUN make
-RUN mkdir -p build && \
- cp wg ./build/awg && \
- cp wg-quick/linux.bash ./build/awg-quick
-
-RUN cp build/awg /usr/bin/awg
-RUN cp build/awg-quick /usr/bin/awg-quick
-
RUN apk --no-cache add iproute2 iptables bash && \
cd /usr/bin/ && \
- # wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
- # unzip -j alpine-3.19-amneziawg-tools.zip && \
+ wget https://github.com/amnezia-vpn/amneziawg-tools/releases/download/v${AWGTOOLS_RELEASE}/alpine-3.19-amneziawg-tools.zip && \
+ unzip -j alpine-3.19-amneziawg-tools.zip && \
chmod +x /usr/bin/awg /usr/bin/awg-quick && \
ln -s /usr/bin/awg /usr/bin/wg && \
ln -s /usr/bin/awg-quick /usr/bin/wg-quick
From f6542209f40f3f8f9e3dc9403d331ad2881fd7e3 Mon Sep 17 00:00:00 2001
From: Mark Puha
Date: Mon, 1 Sep 2025 14:04:52 +0200
Subject: [PATCH 68/75] feat: awg 2.0 (#91)
* feat: ranged H1-H4
* feat: S3, S4 support
* chore: updated awg-tools version
---------
Co-authored-by: Yaroslav Gurov
---
Dockerfile | 2 +-
device/awg/awg.go | 160 +++-----
device/awg/junk_creator.go | 64 ++--
device/awg/junk_creator_test.go | 82 ++--
device/awg/magic_header.go | 97 +++++
device/awg/magic_header_test.go | 488 ++++++++++++++++++++++++
device/awg/prng.go | 50 +++
device/awg/special_handshake_handler.go | 43 +--
device/awg/tag_generator.go | 127 +++---
device/awg/tag_generator_test.go | 140 ++++++-
device/awg/tag_parser.go | 20 +-
device/awg/tag_parser_test.go | 2 +-
device/cookie.go | 3 +-
device/cookie_test.go | 2 +-
device/device.go | 337 ++++++++++------
device/device_test.go | 26 +-
device/noise-protocol.go | 46 ++-
device/receive.go | 48 +--
device/send.go | 53 ++-
device/uapi.go | 149 +++-----
go.mod | 2 +-
go.sum | 14 +-
22 files changed, 1352 insertions(+), 603 deletions(-)
create mode 100644 device/awg/magic_header.go
create mode 100644 device/awg/magic_header_test.go
create mode 100644 device/awg/prng.go
diff --git a/Dockerfile b/Dockerfile
index f165899..98a7e9e 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -6,7 +6,7 @@ RUN go mod download && \
go build -ldflags '-linkmode external -extldflags "-fno-PIC -static"' -v -o /usr/bin
FROM alpine:3.19
-ARG AWGTOOLS_RELEASE="1.0.20241018"
+ARG AWGTOOLS_RELEASE="1.0.20250901"
RUN apk --no-cache add iproute2 iptables bash && \
cd /usr/bin/ && \
diff --git a/device/awg/awg.go b/device/awg/awg.go
index fd5a96d..888a42e 100644
--- a/device/awg/awg.go
+++ b/device/awg/awg.go
@@ -3,142 +3,88 @@ package awg
import (
"bytes"
"fmt"
- "slices"
- "strconv"
- "strings"
"sync"
"github.com/tevino/abool"
)
-type aSecCfgType struct {
- IsSet bool
- JunkPacketCount int
- JunkPacketMinSize int
- JunkPacketMaxSize int
- InitHeaderJunkSize int
- ResponseHeaderJunkSize int
- CookieReplyHeaderJunkSize int
- TransportHeaderJunkSize int
- InitPacketMagicHeader uint32
- ResponsePacketMagicHeader uint32
- UnderloadPacketMagicHeader uint32
- TransportPacketMagicHeader uint32
- // InitPacketMagicHeader Limit
- // ResponsePacketMagicHeader Limit
- // UnderloadPacketMagicHeader Limit
- // TransportPacketMagicHeader Limit
-}
+type Cfg struct {
+ IsSet bool
+ JunkPacketCount int
+ JunkPacketMinSize int
+ JunkPacketMaxSize int
+ InitHeaderJunkSize int
+ ResponseHeaderJunkSize int
+ CookieReplyHeaderJunkSize int
+ TransportHeaderJunkSize int
-type Limit struct {
- Min uint32
- Max uint32
- HeaderType uint32
-}
-
-func NewLimit(min, max, headerType uint32) (Limit, error) {
- if min > max {
- return Limit{}, fmt.Errorf("min (%d) cannot be greater than max (%d)", min, max)
- }
-
- return Limit{
- Min: min,
- Max: max,
- HeaderType: headerType,
- }, nil
-}
-
-func ParseMagicHeader(key, value string, defaultHeaderType uint32) (Limit, error) {
- // tempAwg.ASecCfg.InitPacketMagicHeader, err = awg.NewLimit(uint32(initPacketMagicHeaderMin), uint32(initPacketMagicHeaderMax), DNewLimit(min, max, headerType)efaultMessageInitiationType)
- // var min, max, headerType uint32
- // _, err := fmt.Sscanf(value, "%d-%d:%d", &min, &max, &headerType)
- // if err != nil {
- // return Limit{}, fmt.Errorf("invalid magic header format: %s", value)
- // }
-
- limits := strings.Split(value, "-")
- if len(limits) != 2 {
- return Limit{}, fmt.Errorf("invalid format for key: %s; %s", key, value)
- }
-
- min, err := strconv.ParseUint(limits[0], 10, 32)
- if err != nil {
- return Limit{}, fmt.Errorf("parse min key: %s; value: ; %w", key, limits[0], err)
- }
-
- max, err := strconv.ParseUint(limits[1], 10, 32)
- if err != nil {
- return Limit{}, fmt.Errorf("parse max key: %s; value: ; %w", key, limits[0], err)
- }
-
- limit, err := NewLimit(uint32(min), uint32(max), defaultHeaderType)
- if err != nil {
- return Limit{}, fmt.Errorf("new lmit key: %s; value: ; %w", key, limits[0], err)
- }
-
- return limit, nil
-}
-
-type Limits []Limit
-
-func NewLimits(limits []Limit) Limits {
- slices.SortFunc(limits, func(a, b Limit) int {
- if a.Min < b.Min {
- return -1
- } else if a.Min > b.Min {
- return 1
- }
- return 0
- })
-
- return Limits(limits)
+ MagicHeaders MagicHeaders
}
type Protocol struct {
- IsASecOn abool.AtomicBool
+ IsOn abool.AtomicBool
// TODO: revision the need of the mutex
- ASecMux sync.RWMutex
- ASecCfg aSecCfgType
- JunkCreator junkCreator
+ Mux sync.RWMutex
+ Cfg Cfg
+ JunkCreator JunkCreator
HandshakeHandler SpecialHandshakeHandler
}
func (protocol *Protocol) CreateInitHeaderJunk() ([]byte, error) {
- return protocol.createHeaderJunk(protocol.ASecCfg.InitHeaderJunkSize)
+ protocol.Mux.RLock()
+ defer protocol.Mux.RUnlock()
+
+ return protocol.createHeaderJunk(protocol.Cfg.InitHeaderJunkSize, 0)
}
func (protocol *Protocol) CreateResponseHeaderJunk() ([]byte, error) {
- return protocol.createHeaderJunk(protocol.ASecCfg.ResponseHeaderJunkSize)
+ protocol.Mux.RLock()
+ defer protocol.Mux.RUnlock()
+
+ return protocol.createHeaderJunk(protocol.Cfg.ResponseHeaderJunkSize, 0)
}
func (protocol *Protocol) CreateCookieReplyHeaderJunk() ([]byte, error) {
- return protocol.createHeaderJunk(protocol.ASecCfg.CookieReplyHeaderJunkSize)
+ protocol.Mux.RLock()
+ defer protocol.Mux.RUnlock()
+
+ return protocol.createHeaderJunk(protocol.Cfg.CookieReplyHeaderJunkSize, 0)
}
func (protocol *Protocol) CreateTransportHeaderJunk(packetSize int) ([]byte, error) {
- return protocol.createHeaderJunk(protocol.ASecCfg.TransportHeaderJunkSize, packetSize)
+ protocol.Mux.RLock()
+ defer protocol.Mux.RUnlock()
+
+ return protocol.createHeaderJunk(protocol.Cfg.TransportHeaderJunkSize, packetSize)
}
-func (protocol *Protocol) createHeaderJunk(junkSize int, optExtraSize ...int) ([]byte, error) {
- extraSize := 0
- if len(optExtraSize) == 1 {
- extraSize = optExtraSize[0]
+func (protocol *Protocol) createHeaderJunk(junkSize int, extraSize int) ([]byte, error) {
+ if junkSize == 0 {
+ return nil, nil
}
- var junk []byte
- protocol.ASecMux.RLock()
- if junkSize != 0 {
- buf := make([]byte, 0, junkSize+extraSize)
- writer := bytes.NewBuffer(buf[:0])
- err := protocol.JunkCreator.AppendJunk(writer, junkSize)
- if err != nil {
- protocol.ASecMux.RUnlock()
- return nil, err
+ buf := make([]byte, 0, junkSize+extraSize)
+ writer := bytes.NewBuffer(buf[:0])
+
+ err := protocol.JunkCreator.AppendJunk(writer, junkSize)
+ if err != nil {
+ return nil, fmt.Errorf("append junk: %w", err)
+ }
+
+ return writer.Bytes(), nil
+}
+
+func (protocol *Protocol) GetMagicHeaderMinFor(msgType uint32) (uint32, error) {
+ for _, magicHeader := range protocol.Cfg.MagicHeaders.Values {
+ if magicHeader.Min <= msgType && msgType <= magicHeader.Max {
+ return magicHeader.Min, nil
}
- junk = writer.Bytes()
}
- protocol.ASecMux.RUnlock()
- return junk, nil
+ return 0, fmt.Errorf("no header for value: %d", msgType)
+}
+
+func (protocol *Protocol) GetMsgType(defaultMsgType uint32) (uint32, error) {
+ return protocol.Cfg.MagicHeaders.Get(defaultMsgType)
}
diff --git a/device/awg/junk_creator.go b/device/awg/junk_creator.go
index 91fd253..8ba2918 100644
--- a/device/awg/junk_creator.go
+++ b/device/awg/junk_creator.go
@@ -2,69 +2,49 @@ package awg
import (
"bytes"
- crand "crypto/rand"
"fmt"
- v2 "math/rand/v2"
)
-type junkCreator struct {
- aSecCfg aSecCfgType
- cha8Rand *v2.ChaCha8
+type JunkCreator struct {
+ cfg Cfg
+ randomGenerator PRNG[int]
}
// TODO: refactor param to only pass the junk related params
-func NewJunkCreator(aSecCfg aSecCfgType) (junkCreator, error) {
- buf := make([]byte, 32)
- _, err := crand.Read(buf)
- if err != nil {
- return junkCreator{}, err
- }
- return junkCreator{aSecCfg: aSecCfg, cha8Rand: v2.NewChaCha8([32]byte(buf))}, nil
+func NewJunkCreator(cfg Cfg) JunkCreator {
+ return JunkCreator{cfg: cfg, randomGenerator: NewPRNG[int]()}
}
-// Should be called with aSecMux RLocked
-func (jc *junkCreator) CreateJunkPackets(junks *[][]byte) error {
- if jc.aSecCfg.JunkPacketCount == 0 {
- return nil
+// Should be called with awg mux RLocked
+func (jc *JunkCreator) CreateJunkPackets(junks *[][]byte) {
+ if jc.cfg.JunkPacketCount == 0 {
+ return
}
- for range jc.aSecCfg.JunkPacketCount {
+ for range jc.cfg.JunkPacketCount {
packetSize := jc.randomPacketSize()
- junk, err := jc.randomJunkWithSize(packetSize)
- if err != nil {
- return fmt.Errorf("create junk packet: %v", err)
- }
+ junk := jc.randomJunkWithSize(packetSize)
*junks = append(*junks, junk)
}
- return nil
+ return
}
-// Should be called with aSecMux RLocked
-func (jc *junkCreator) randomPacketSize() int {
- return int(
- jc.cha8Rand.Uint64()%uint64(
- jc.aSecCfg.JunkPacketMaxSize-jc.aSecCfg.JunkPacketMinSize,
- ),
- ) + jc.aSecCfg.JunkPacketMinSize
+// Should be called with awg mux RLocked
+func (jc *JunkCreator) randomPacketSize() int {
+ return jc.randomGenerator.RandomSizeInRange(jc.cfg.JunkPacketMinSize, jc.cfg.JunkPacketMaxSize)
}
-// Should be called with aSecMux RLocked
-func (jc *junkCreator) AppendJunk(writer *bytes.Buffer, size int) error {
- headerJunk, err := jc.randomJunkWithSize(size)
- if err != nil {
- return fmt.Errorf("create header junk: %v", err)
- }
- _, err = writer.Write(headerJunk)
+// Should be called with awg mux RLocked
+func (jc *JunkCreator) AppendJunk(writer *bytes.Buffer, size int) error {
+ headerJunk := jc.randomJunkWithSize(size)
+ _, err := writer.Write(headerJunk)
if err != nil {
return fmt.Errorf("write header junk: %v", err)
}
return nil
}
-// Should be called with aSecMux RLocked
-func (jc *junkCreator) randomJunkWithSize(size int) ([]byte, error) {
- // TODO: use a memory pool to allocate
- junk := make([]byte, size)
- _, err := jc.cha8Rand.Read(junk)
- return junk, err
+// Should be called with awg mux RLocked
+func (jc *JunkCreator) randomJunkWithSize(size int) []byte {
+ return jc.randomGenerator.ReadSize(size)
}
diff --git a/device/awg/junk_creator_test.go b/device/awg/junk_creator_test.go
index 424f104..cdf752b 100644
--- a/device/awg/junk_creator_test.go
+++ b/device/awg/junk_creator_test.go
@@ -6,43 +6,34 @@ import (
"testing"
)
-func setUpJunkCreator(t *testing.T) (junkCreator, error) {
- jc, err := NewJunkCreator(aSecCfgType{
- IsSet: true,
- JunkPacketCount: 5,
- JunkPacketMinSize: 500,
- JunkPacketMaxSize: 1000,
- InitHeaderJunkSize: 30,
- ResponseHeaderJunkSize: 40,
- InitPacketMagicHeader: 123456,
- ResponsePacketMagicHeader: 67543,
- UnderloadPacketMagicHeader: 32345,
- TransportPacketMagicHeader: 123123,
+func setUpJunkCreator() JunkCreator {
+ mh, _ := NewMagicHeaders(
+ []MagicHeader{
+ NewMagicHeaderSameValue(123456),
+ NewMagicHeaderSameValue(67543),
+ NewMagicHeaderSameValue(32345),
+ NewMagicHeaderSameValue(123123),
+ },
+ )
+
+ jc := NewJunkCreator(Cfg{
+ IsSet: true,
+ JunkPacketCount: 5,
+ JunkPacketMinSize: 500,
+ JunkPacketMaxSize: 1000,
+ InitHeaderJunkSize: 30,
+ ResponseHeaderJunkSize: 40,
+ MagicHeaders: mh,
})
- if err != nil {
- t.Errorf("failed to create junk creator %v", err)
- return junkCreator{}, err
- }
-
- return jc, nil
+ return jc
}
func Test_junkCreator_createJunkPackets(t *testing.T) {
- jc, err := setUpJunkCreator(t)
- if err != nil {
- return
- }
+ jc := setUpJunkCreator()
t.Run("valid", func(t *testing.T) {
- got := make([][]byte, 0, jc.aSecCfg.JunkPacketCount)
- err := jc.CreateJunkPackets(&got)
- if err != nil {
- t.Errorf(
- "junkCreator.createJunkPackets() = %v; failed",
- err,
- )
- return
- }
+ got := make([][]byte, 0, jc.cfg.JunkPacketCount)
+ jc.CreateJunkPackets(&got)
seen := make(map[string]bool)
for _, junk := range got {
key := string(junk)
@@ -61,34 +52,28 @@ func Test_junkCreator_createJunkPackets(t *testing.T) {
func Test_junkCreator_randomJunkWithSize(t *testing.T) {
t.Run("valid", func(t *testing.T) {
- jc, err := setUpJunkCreator(t)
- if err != nil {
- return
- }
- r1, _ := jc.randomJunkWithSize(10)
- r2, _ := jc.randomJunkWithSize(10)
+ jc := setUpJunkCreator()
+ r1 := jc.randomJunkWithSize(10)
+ r2 := jc.randomJunkWithSize(10)
fmt.Printf("%v\n%v\n", r1, r2)
if bytes.Equal(r1, r2) {
- t.Errorf("same junks %v", err)
+ t.Errorf("same junks")
return
}
})
}
func Test_junkCreator_randomPacketSize(t *testing.T) {
- jc, err := setUpJunkCreator(t)
- if err != nil {
- return
- }
+ jc := setUpJunkCreator()
for range [30]struct{}{} {
t.Run("valid", func(t *testing.T) {
- if got := jc.randomPacketSize(); jc.aSecCfg.JunkPacketMinSize > got ||
- got > jc.aSecCfg.JunkPacketMaxSize {
+ if got := jc.randomPacketSize(); jc.cfg.JunkPacketMinSize > got ||
+ got > jc.cfg.JunkPacketMaxSize {
t.Errorf(
"junkCreator.randomPacketSize() = %v, not between range [%v,%v]",
got,
- jc.aSecCfg.JunkPacketMinSize,
- jc.aSecCfg.JunkPacketMaxSize,
+ jc.cfg.JunkPacketMinSize,
+ jc.cfg.JunkPacketMaxSize,
)
}
})
@@ -96,10 +81,7 @@ func Test_junkCreator_randomPacketSize(t *testing.T) {
}
func Test_junkCreator_appendJunk(t *testing.T) {
- jc, err := setUpJunkCreator(t)
- if err != nil {
- return
- }
+ jc := setUpJunkCreator()
t.Run("valid", func(t *testing.T) {
s := "apple"
buffer := bytes.NewBuffer([]byte(s))
diff --git a/device/awg/magic_header.go b/device/awg/magic_header.go
new file mode 100644
index 0000000..aaf4e97
--- /dev/null
+++ b/device/awg/magic_header.go
@@ -0,0 +1,97 @@
+package awg
+
+import (
+ "cmp"
+ "fmt"
+ "slices"
+ "strconv"
+ "strings"
+)
+
+type MagicHeader struct {
+ Min uint32
+ Max uint32
+}
+
+func NewMagicHeaderSameValue(value uint32) MagicHeader {
+ return MagicHeader{Min: value, Max: value}
+}
+
+func NewMagicHeader(min, max uint32) (MagicHeader, error) {
+ if min > max {
+ return MagicHeader{}, fmt.Errorf("min (%d) cannot be greater than max (%d)", min, max)
+ }
+
+ return MagicHeader{Min: min, Max: max}, nil
+}
+
+func ParseMagicHeader(key, value string) (MagicHeader, error) {
+ hyphenIdx := strings.Index(value, "-")
+ if hyphenIdx == -1 {
+ // if there is no hyphen, we treat it as single magic header value
+ magicHeader, err := strconv.ParseUint(value, 10, 32)
+ if err != nil {
+ return MagicHeader{}, fmt.Errorf("parse key: %s; value: %s; %w", key, value, err)
+ }
+
+ return NewMagicHeader(uint32(magicHeader), uint32(magicHeader))
+ }
+
+ minStr := value[:hyphenIdx]
+ maxStr := value[hyphenIdx+1:]
+ if len(minStr) == 0 || len(maxStr) == 0 {
+ return MagicHeader{}, fmt.Errorf("invalid value for key: %s; value: %s; expected format: min-max", key, value)
+ }
+
+ min, err := strconv.ParseUint(minStr, 10, 32)
+ if err != nil {
+ return MagicHeader{}, fmt.Errorf("parse min key: %s; value: %s; %w", key, minStr, err)
+ }
+
+ max, err := strconv.ParseUint(maxStr, 10, 32)
+ if err != nil {
+ return MagicHeader{}, fmt.Errorf("parse max key: %s; value: %s; %w", key, maxStr, err)
+ }
+
+ magicHeader, err := NewMagicHeader(uint32(min), uint32(max))
+ if err != nil {
+ return MagicHeader{}, fmt.Errorf("new magicHeader key: %s; value: %s-%s; %w", key, minStr, maxStr, err)
+ }
+
+ return magicHeader, nil
+}
+
+type MagicHeaders struct {
+ Values []MagicHeader
+ randomGenerator RandomNumberGenerator[uint32]
+}
+
+func NewMagicHeaders(headerValues []MagicHeader) (MagicHeaders, error) {
+ if len(headerValues) != 4 {
+ return MagicHeaders{}, fmt.Errorf("all header types should be included: %v", headerValues)
+ }
+
+ sortedMagicHeaders := slices.SortedFunc(slices.Values(headerValues), func(lhs MagicHeader, rhs MagicHeader) int {
+ return cmp.Compare(lhs.Min, rhs.Min)
+ })
+
+ for i := range 3 {
+ if sortedMagicHeaders[i].Max >= sortedMagicHeaders[i+1].Min {
+ return MagicHeaders{}, fmt.Errorf(
+ "magic headers shouldn't overlap; %v > %v",
+ sortedMagicHeaders[i].Max,
+ sortedMagicHeaders[i+1].Min,
+ )
+ }
+ }
+
+ return MagicHeaders{Values: headerValues, randomGenerator: NewPRNG[uint32]()}, nil
+}
+
+func (mh *MagicHeaders) Get(defaultMsgType uint32) (uint32, error) {
+ if defaultMsgType == 0 || defaultMsgType > 4 {
+ return 0, fmt.Errorf("invalid msg type: %d", defaultMsgType)
+ }
+
+ return mh.randomGenerator.RandomSizeInRange(mh.Values[defaultMsgType-1].Min, mh.Values[defaultMsgType-1].Max), nil
+}
diff --git a/device/awg/magic_header_test.go b/device/awg/magic_header_test.go
new file mode 100644
index 0000000..72a823e
--- /dev/null
+++ b/device/awg/magic_header_test.go
@@ -0,0 +1,488 @@
+package awg
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestNewMagicHeaderSameValue(t *testing.T) {
+ tests := []struct {
+ name string
+ value uint32
+ expected MagicHeader
+ }{
+ {
+ name: "zero value",
+ value: 0,
+ expected: MagicHeader{Min: 0, Max: 0},
+ },
+ {
+ name: "small value",
+ value: 1,
+ expected: MagicHeader{Min: 1, Max: 1},
+ },
+ {
+ name: "large value",
+ value: 4294967295, // max uint32
+ expected: MagicHeader{Min: 4294967295, Max: 4294967295},
+ },
+ {
+ name: "medium value",
+ value: 1000,
+ expected: MagicHeader{Min: 1000, Max: 1000},
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ result := NewMagicHeaderSameValue(tt.value)
+ require.Equal(t, tt.expected, result)
+ })
+ }
+}
+
+func TestNewMagicHeader(t *testing.T) {
+ tests := []struct {
+ name string
+ min uint32
+ max uint32
+ expected MagicHeader
+ errorMsg string
+ }{
+ {
+ name: "valid range",
+ min: 1,
+ max: 10,
+ expected: MagicHeader{Min: 1, Max: 10},
+ },
+ {
+ name: "equal values",
+ min: 5,
+ max: 5,
+ expected: MagicHeader{Min: 5, Max: 5},
+ },
+ {
+ name: "zero range",
+ min: 0,
+ max: 0,
+ expected: MagicHeader{Min: 0, Max: 0},
+ },
+ {
+ name: "max uint32 range",
+ min: 4294967294,
+ max: 4294967295,
+ expected: MagicHeader{Min: 4294967294, Max: 4294967295},
+ },
+ {
+ name: "min greater than max",
+ min: 10,
+ max: 5,
+ expected: MagicHeader{},
+ errorMsg: "min (10) cannot be greater than max (5)",
+ },
+ {
+ name: "large min greater than max",
+ min: 4294967295,
+ max: 1,
+ expected: MagicHeader{},
+ errorMsg: "min (4294967295) cannot be greater than max (1)",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ result, err := NewMagicHeader(tt.min, tt.max)
+
+ if tt.errorMsg != "" {
+ require.Error(t, err)
+ require.Contains(t, err.Error(), tt.errorMsg)
+ require.Equal(t, MagicHeader{}, result)
+ } else {
+ require.NoError(t, err)
+ require.Equal(t, tt.expected, result)
+ }
+ })
+ }
+}
+
+func TestParseMagicHeader(t *testing.T) {
+ tests := []struct {
+ name string
+ key string
+ value string
+ expected MagicHeader
+ errorMsg string
+ }{
+ {
+ name: "single value",
+ key: "header1",
+ value: "100",
+ expected: MagicHeader{Min: 100, Max: 100},
+ },
+ {
+ name: "valid range",
+ key: "header2",
+ value: "10-20",
+ expected: MagicHeader{Min: 10, Max: 20},
+ },
+ {
+ name: "zero single value",
+ key: "header3",
+ value: "0",
+ expected: MagicHeader{Min: 0, Max: 0},
+ },
+ {
+ name: "zero range",
+ key: "header4",
+ value: "0-0",
+ expected: MagicHeader{Min: 0, Max: 0},
+ },
+ {
+ name: "max uint32 single",
+ key: "header5",
+ value: "4294967295",
+ expected: MagicHeader{Min: 4294967295, Max: 4294967295},
+ },
+ {
+ name: "max uint32 range",
+ key: "header6",
+ value: "4294967294-4294967295",
+ expected: MagicHeader{Min: 4294967294, Max: 4294967295},
+ },
+ {
+ name: "invalid single value - not number",
+ key: "header7",
+ value: "abc",
+ expected: MagicHeader{},
+ errorMsg: "parse key: header7; value: abc;",
+ },
+ {
+ name: "invalid single value - negative",
+ key: "header8",
+ value: "-5",
+ expected: MagicHeader{},
+ errorMsg: "invalid value for key: header8; value: -5;",
+ },
+ {
+ name: "invalid single value - too large",
+ key: "header9",
+ value: "4294967296",
+ expected: MagicHeader{},
+ errorMsg: "parse key: header9; value: 4294967296;",
+ },
+ {
+ name: "invalid range - min not number",
+ key: "header10",
+ value: "abc-10",
+ expected: MagicHeader{},
+ errorMsg: "parse min key: header10; value: abc;",
+ },
+ {
+ name: "invalid range - max not number",
+ key: "header11",
+ value: "10-abc",
+ expected: MagicHeader{},
+ errorMsg: "parse max key: header11; value: abc;",
+ },
+ {
+ name: "invalid range - min greater than max",
+ key: "header12",
+ value: "20-10",
+ expected: MagicHeader{},
+ errorMsg: "new magicHeader key: header12; value: 20-10;",
+ },
+ {
+ name: "invalid range - too many parts",
+ key: "header13",
+ value: "10-20-30",
+ expected: MagicHeader{},
+ errorMsg: "parse key: header13; value: 10-20-30;",
+ },
+ {
+ name: "empty value",
+ key: "header14",
+ value: "",
+ expected: MagicHeader{},
+ errorMsg: "parse key: header14; value: ;",
+ },
+ {
+ name: "hyphen only",
+ key: "header15",
+ value: "-",
+ expected: MagicHeader{},
+ errorMsg: "invalid value for key: header15; value: -;",
+ },
+ {
+ name: "empty min",
+ key: "header16",
+ value: "-10",
+ expected: MagicHeader{},
+ errorMsg: "invalid value for key: header16; value: -10;",
+ },
+ {
+ name: "empty max",
+ key: "header17",
+ value: "10-",
+ expected: MagicHeader{},
+ errorMsg: "invalid value for key: header17; value: 10-;",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ result, err := ParseMagicHeader(tt.key, tt.value)
+
+ if tt.errorMsg != "" {
+ require.Error(t, err)
+ require.Contains(t, err.Error(), tt.errorMsg)
+ require.Equal(t, MagicHeader{}, result)
+ } else {
+ require.NoError(t, err)
+ require.Equal(t, tt.expected, result)
+ }
+ })
+ }
+}
+
+func TestNewMagicHeaders(t *testing.T) {
+ tests := []struct {
+ name string
+ magicHeaders []MagicHeader
+ errorMsg string
+ }{
+ {
+ name: "valid non-overlapping headers",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 10},
+ {Min: 11, Max: 20},
+ {Min: 21, Max: 30},
+ {Min: 31, Max: 40},
+ },
+ },
+ {
+ name: "valid adjacent headers",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 1},
+ {Min: 2, Max: 2},
+ {Min: 3, Max: 3},
+ {Min: 4, Max: 4},
+ },
+ },
+ {
+ name: "valid zero-based headers",
+ magicHeaders: []MagicHeader{
+ {Min: 0, Max: 0},
+ {Min: 1, Max: 1},
+ {Min: 2, Max: 2},
+ {Min: 3, Max: 3},
+ },
+ },
+ {
+ name: "valid large value headers",
+ magicHeaders: []MagicHeader{
+ {Min: 4294967290, Max: 4294967291},
+ {Min: 4294967292, Max: 4294967293},
+ {Min: 4294967294, Max: 4294967294},
+ {Min: 4294967295, Max: 4294967295},
+ },
+ },
+ {
+ name: "too few headers",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 10},
+ {Min: 11, Max: 20},
+ {Min: 21, Max: 30},
+ },
+ errorMsg: "all header types should be included:",
+ },
+ {
+ name: "too many headers",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 10},
+ {Min: 11, Max: 20},
+ {Min: 21, Max: 30},
+ {Min: 31, Max: 40},
+ {Min: 41, Max: 50},
+ },
+ errorMsg: "all header types should be included:",
+ },
+ {
+ name: "empty headers",
+ magicHeaders: []MagicHeader{},
+ errorMsg: "all header types should be included:",
+ },
+ {
+ name: "overlapping headers",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 15},
+ {Min: 10, Max: 20},
+ {Min: 25, Max: 30},
+ {Min: 35, Max: 40},
+ },
+ errorMsg: "magic headers shouldn't overlap;",
+ },
+ {
+ name: "overlapping headers at limit-first",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 10},
+ {Min: 10, Max: 20},
+ {Min: 25, Max: 30},
+ {Min: 35, Max: 40},
+ },
+ errorMsg: "magic headers shouldn't overlap;",
+ },
+ {
+ name: "overlapping headers at limit-second",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 10},
+ {Min: 15, Max: 25},
+ {Min: 25, Max: 30},
+ {Min: 35, Max: 40},
+ },
+ errorMsg: "magic headers shouldn't overlap;",
+ },
+ {
+ name: "overlapping headers at limit-third",
+ magicHeaders: []MagicHeader{
+ {Min: 1, Max: 10},
+ {Min: 15, Max: 25},
+ {Min: 30, Max: 35},
+ {Min: 35, Max: 40},
+ },
+ errorMsg: "magic headers shouldn't overlap;",
+ },
+ {
+ name: "identical ranges",
+ magicHeaders: []MagicHeader{
+ {Min: 10, Max: 20},
+ {Min: 10, Max: 20},
+ {Min: 25, Max: 30},
+ {Min: 35, Max: 40},
+ },
+ errorMsg: "magic headers shouldn't overlap;",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ result, err := NewMagicHeaders(tt.magicHeaders)
+
+ if tt.errorMsg != "" {
+ require.Error(t, err)
+ require.Contains(t, err.Error(), tt.errorMsg)
+ require.Equal(t, MagicHeaders{}, result)
+ } else {
+ require.NoError(t, err)
+ require.Equal(t, tt.magicHeaders, result.Values)
+ require.NotNil(t, result.randomGenerator)
+ }
+ })
+ }
+}
+
+// Mock PRNG for testing
+type mockPRNG struct {
+ returnValue uint32
+}
+
+func (m *mockPRNG) RandomSizeInRange(min, max uint32) uint32 {
+ return m.returnValue
+}
+
+func (m *mockPRNG) Get() uint64 {
+ return 0
+}
+func (m *mockPRNG) ReadSize(size int) []byte {
+ return make([]byte, size)
+}
+
+func TestMagicHeaders_Get(t *testing.T) {
+ // Create test headers
+ headers := []MagicHeader{
+ {Min: 1, Max: 10},
+ {Min: 11, Max: 20},
+ {Min: 21, Max: 30},
+ {Min: 31, Max: 40},
+ }
+
+ tests := []struct {
+ name string
+ defaultMsgType uint32
+ mockValue uint32
+ expectedValue uint32
+ errorMsg string
+ }{
+ {
+ name: "valid type 1",
+ defaultMsgType: 1,
+ mockValue: 5,
+ expectedValue: 5,
+ },
+ {
+ name: "valid type 2",
+ defaultMsgType: 2,
+ mockValue: 15,
+ expectedValue: 15,
+ },
+ {
+ name: "valid type 3",
+ defaultMsgType: 3,
+ mockValue: 25,
+ expectedValue: 25,
+ },
+ {
+ name: "valid type 4",
+ defaultMsgType: 4,
+ mockValue: 35,
+ expectedValue: 35,
+ },
+ {
+ name: "invalid type 0",
+ defaultMsgType: 0,
+ mockValue: 0,
+ expectedValue: 0,
+ errorMsg: "invalid msg type: 0",
+ },
+ {
+ name: "invalid type 5",
+ defaultMsgType: 5,
+ mockValue: 0,
+ expectedValue: 0,
+ errorMsg: "invalid msg type: 5",
+ },
+ {
+ name: "invalid type max uint32",
+ defaultMsgType: 4294967295,
+ mockValue: 0,
+ expectedValue: 0,
+ errorMsg: "invalid msg type: 4294967295",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ // Create a new instance with mock PRNG for each test
+ testMagicHeaders := MagicHeaders{
+ Values: headers,
+ randomGenerator: &mockPRNG{returnValue: tt.mockValue},
+ }
+
+ result, err := testMagicHeaders.Get(tt.defaultMsgType)
+
+ if tt.errorMsg != "" {
+ require.Error(t, err)
+ require.Contains(t, err.Error(), tt.errorMsg)
+ require.Equal(t, uint32(0), result)
+ } else {
+ require.NoError(t, err)
+ require.Equal(t, tt.expectedValue, result)
+ }
+ })
+ }
+}
diff --git a/device/awg/prng.go b/device/awg/prng.go
new file mode 100644
index 0000000..e7661d7
--- /dev/null
+++ b/device/awg/prng.go
@@ -0,0 +1,50 @@
+package awg
+
+import (
+ crand "crypto/rand"
+ v2 "math/rand/v2"
+
+ "golang.org/x/exp/constraints"
+)
+
+type RandomNumberGenerator[T constraints.Integer] interface {
+ RandomSizeInRange(min, max T) T
+ Get() uint64
+ ReadSize(size int) []byte
+}
+
+type PRNG[T constraints.Integer] struct {
+ cha8Rand *v2.ChaCha8
+}
+
+func NewPRNG[T constraints.Integer]() PRNG[T] {
+ buf := make([]byte, 32)
+ _, _ = crand.Read(buf)
+
+ return PRNG[T]{
+ cha8Rand: v2.NewChaCha8([32]byte(buf)),
+ }
+}
+
+func (p PRNG[T]) RandomSizeInRange(min, max T) T {
+ if min > max {
+ panic("min must be less than max")
+ }
+
+ if min == max {
+ return min
+ }
+
+ return T(p.Get()%uint64(max-min)) + min
+}
+
+func (p PRNG[T]) Get() uint64 {
+ return p.cha8Rand.Uint64()
+}
+
+func (p PRNG[T]) ReadSize(size int) []byte {
+ // TODO: use a memory pool to allocate
+ buf := make([]byte, size)
+ _, _ = p.cha8Rand.Read(buf)
+ return buf
+}
diff --git a/device/awg/special_handshake_handler.go b/device/awg/special_handshake_handler.go
index e582d97..d740879 100644
--- a/device/awg/special_handshake_handler.go
+++ b/device/awg/special_handshake_handler.go
@@ -1,9 +1,6 @@
package awg
import (
- "errors"
- "time"
-
"github.com/tevino/abool"
"go.uber.org/atomic"
)
@@ -21,25 +18,13 @@ var WaitResponse = struct {
}
type SpecialHandshakeHandler struct {
- isFirstDone bool
- SpecialJunk TagJunkPacketGenerators
- ControlledJunk TagJunkPacketGenerators
-
- nextItime time.Time
- ITimeout time.Duration // seconds
+ SpecialJunk TagJunkPacketGenerators
IsSet bool
}
func (handler *SpecialHandshakeHandler) Validate() error {
- var errs []error
- if err := handler.SpecialJunk.Validate(); err != nil {
- errs = append(errs, err)
- }
- if err := handler.ControlledJunk.Validate(); err != nil {
- errs = append(errs, err)
- }
- return errors.Join(errs...)
+ return handler.SpecialJunk.Validate()
}
func (handler *SpecialHandshakeHandler) GenerateSpecialJunk() [][]byte {
@@ -47,27 +32,5 @@ func (handler *SpecialHandshakeHandler) GenerateSpecialJunk() [][]byte {
return nil
}
- // TODO: create tests
- if !handler.isFirstDone {
- handler.isFirstDone = true
- } else if !handler.isTimeToSendSpecial() {
- return nil
- }
-
- rv := handler.SpecialJunk.GeneratePackets()
- handler.nextItime = time.Now().Add(handler.ITimeout)
-
- return rv
-}
-
-func (handler *SpecialHandshakeHandler) isTimeToSendSpecial() bool {
- return time.Now().After(handler.nextItime)
-}
-
-func (handler *SpecialHandshakeHandler) GenerateControlledJunk() [][]byte {
- if !handler.ControlledJunk.IsDefined() {
- return nil
- }
-
- return handler.ControlledJunk.GeneratePackets()
+ return handler.SpecialJunk.GeneratePackets()
}
diff --git a/device/awg/tag_generator.go b/device/awg/tag_generator.go
index 65d8004..3a1d497 100644
--- a/device/awg/tag_generator.go
+++ b/device/awg/tag_generator.go
@@ -59,43 +59,110 @@ func hexToBytes(hexStr string) ([]byte, error) {
return hex.DecodeString(hexStr)
}
-type RandomPacketGenerator struct {
+type randomGeneratorBase struct {
cha8Rand *v2.ChaCha8
size int
}
-func (rpg *RandomPacketGenerator) Generate() []byte {
- junk := make([]byte, rpg.size)
- rpg.cha8Rand.Read(junk)
- return junk
-}
-
-func (rpg *RandomPacketGenerator) Size() int {
- return rpg.size
-}
-
-func newRandomPacketGenerator(param string) (Generator, error) {
+func newRandomGeneratorBase(param string) (*randomGeneratorBase, error) {
size, err := strconv.Atoi(param)
if err != nil {
- return nil, fmt.Errorf("random packet parse int: %w", err)
+ return nil, fmt.Errorf("parse int: %w", err)
}
if size > 1000 {
- return nil, fmt.Errorf("random packet size must be less than 1000")
+ return nil, fmt.Errorf("size must be less than 1000")
}
buf := make([]byte, 32)
_, err = crand.Read(buf)
if err != nil {
- return nil, fmt.Errorf("random packet crand read: %w", err)
+ return nil, fmt.Errorf("crand read: %w", err)
}
- return &RandomPacketGenerator{
+ return &randomGeneratorBase{
cha8Rand: v2.NewChaCha8([32]byte(buf)),
size: size,
}, nil
}
+func (rpg *randomGeneratorBase) generate() []byte {
+ junk := make([]byte, rpg.size)
+ rpg.cha8Rand.Read(junk)
+ return junk
+}
+
+func (rpg *randomGeneratorBase) Size() int {
+ return rpg.size
+}
+
+type RandomBytesGenerator struct {
+ *randomGeneratorBase
+}
+
+func newRandomBytesGenerator(param string) (Generator, error) {
+ rpgBase, err := newRandomGeneratorBase(param)
+ if err != nil {
+ return nil, fmt.Errorf("new random bytes generator: %w", err)
+ }
+
+ return &RandomBytesGenerator{randomGeneratorBase: rpgBase}, nil
+}
+
+func (rpg *RandomBytesGenerator) Generate() []byte {
+ return rpg.generate()
+}
+
+const alphanumericChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
+
+type RandomASCIIGenerator struct {
+ *randomGeneratorBase
+}
+
+func newRandomASCIIGenerator(param string) (Generator, error) {
+ rpgBase, err := newRandomGeneratorBase(param)
+ if err != nil {
+ return nil, fmt.Errorf("new random ascii generator: %w", err)
+ }
+
+ return &RandomASCIIGenerator{randomGeneratorBase: rpgBase}, nil
+}
+
+func (rpg *RandomASCIIGenerator) Generate() []byte {
+ junk := rpg.generate()
+
+ result := make([]byte, rpg.size)
+ for i, b := range junk {
+ result[i] = alphanumericChars[b%byte(len(alphanumericChars))]
+ }
+
+ return result
+}
+
+type RandomDigitGenerator struct {
+ *randomGeneratorBase
+}
+
+func newRandomDigitGenerator(param string) (Generator, error) {
+ rpgBase, err := newRandomGeneratorBase(param)
+ if err != nil {
+ return nil, fmt.Errorf("new random digit generator: %w", err)
+ }
+
+ return &RandomDigitGenerator{randomGeneratorBase: rpgBase}, nil
+}
+
+func (rpg *RandomDigitGenerator) Generate() []byte {
+ junk := rpg.generate()
+
+ result := make([]byte, rpg.size)
+ for i, b := range junk {
+ result[i] = '0' + (b % 10) // Convert to digit character
+ }
+
+ return result
+}
+
type TimestampGenerator struct {
}
@@ -117,34 +184,6 @@ func newTimestampGenerator(param string) (Generator, error) {
return &TimestampGenerator{}, nil
}
-type WaitTimeoutGenerator struct {
- waitTimeout time.Duration
-}
-
-func (wtg *WaitTimeoutGenerator) Generate() []byte {
- time.Sleep(wtg.waitTimeout)
- return []byte{}
-}
-
-func (wtg *WaitTimeoutGenerator) Size() int {
- return 0
-}
-
-func newWaitTimeoutGenerator(param string) (Generator, error) {
- timeout, err := strconv.Atoi(param)
- if err != nil {
- return nil, fmt.Errorf("timeout parse int: %w", err)
- }
-
- if timeout > 5000 {
- return nil, fmt.Errorf("timeout must be less than 5000ms")
- }
-
- return &WaitTimeoutGenerator{
- waitTimeout: time.Duration(timeout) * time.Millisecond,
- }, nil
-}
-
type PacketCounterGenerator struct {
}
diff --git a/device/awg/tag_generator_test.go b/device/awg/tag_generator_test.go
index 4950b33..43efa67 100644
--- a/device/awg/tag_generator_test.go
+++ b/device/awg/tag_generator_test.go
@@ -8,7 +8,9 @@ import (
"github.com/stretchr/testify/require"
)
-func Test_newBytesGenerator(t *testing.T) {
+func TestNewBytesGenerator(t *testing.T) {
+ t.Parallel()
+
type args struct {
param string
}
@@ -63,6 +65,8 @@ func Test_newBytesGenerator(t *testing.T) {
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+
got, err := newBytesGenerator(tt.args.param)
if tt.wantErr != nil {
@@ -80,7 +84,9 @@ func Test_newBytesGenerator(t *testing.T) {
}
}
-func Test_newRandomPacketGenerator(t *testing.T) {
+func TestNewRandomBytesGenerator(t *testing.T) {
+ t.Parallel()
+
type args struct {
param string
}
@@ -117,9 +123,134 @@ func Test_newRandomPacketGenerator(t *testing.T) {
},
},
}
+
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- got, err := newRandomPacketGenerator(tt.args.param)
+ t.Parallel()
+
+ got, err := newRandomBytesGenerator(tt.args.param)
+ if tt.wantErr != nil {
+ require.ErrorAs(t, err, &tt.wantErr)
+ require.Nil(t, got)
+ return
+ }
+
+ require.Nil(t, err)
+ require.NotNil(t, got)
+ first := got.Generate()
+
+ second := got.Generate()
+ require.NotEqual(t, first, second)
+ })
+ }
+}
+
+func TestNewRandomASCIIGenerator(t *testing.T) {
+ t.Parallel()
+
+ type args struct {
+ param string
+ }
+ tests := []struct {
+ name string
+ args args
+ wantErr error
+ }{
+ {
+ name: "empty",
+ args: args{
+ param: "",
+ },
+ wantErr: fmt.Errorf("parse int"),
+ },
+ {
+ name: "not an int",
+ args: args{
+ param: "x",
+ },
+ wantErr: fmt.Errorf("parse int"),
+ },
+ {
+ name: "too large",
+ args: args{
+ param: "1001",
+ },
+ wantErr: fmt.Errorf("random packet size must be less than 1000"),
+ },
+ {
+ name: "valid",
+ args: args{
+ param: "12",
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+
+ got, err := newRandomASCIIGenerator(tt.args.param)
+ if tt.wantErr != nil {
+ require.ErrorAs(t, err, &tt.wantErr)
+ require.Nil(t, got)
+ return
+ }
+
+ require.Nil(t, err)
+ require.NotNil(t, got)
+ first := got.Generate()
+
+ second := got.Generate()
+ require.NotEqual(t, first, second)
+ })
+ }
+}
+
+func TestNewRandomDigitGenerator(t *testing.T) {
+ t.Parallel()
+
+ type args struct {
+ param string
+ }
+ tests := []struct {
+ name string
+ args args
+ wantErr error
+ }{
+ {
+ name: "empty",
+ args: args{
+ param: "",
+ },
+ wantErr: fmt.Errorf("parse int"),
+ },
+ {
+ name: "not an int",
+ args: args{
+ param: "x",
+ },
+ wantErr: fmt.Errorf("parse int"),
+ },
+ {
+ name: "too large",
+ args: args{
+ param: "1001",
+ },
+ wantErr: fmt.Errorf("random packet size must be less than 1000"),
+ },
+ {
+ name: "valid",
+ args: args{
+ param: "12",
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+
+ got, err := newRandomDigitGenerator(tt.args.param)
if tt.wantErr != nil {
require.ErrorAs(t, err, &tt.wantErr)
require.Nil(t, got)
@@ -137,6 +268,8 @@ func Test_newRandomPacketGenerator(t *testing.T) {
}
func TestPacketCounterGenerator(t *testing.T) {
+ t.Parallel()
+
tests := []struct {
name string
param string
@@ -155,7 +288,6 @@ func TestPacketCounterGenerator(t *testing.T) {
}
for _, tc := range tests {
- tc := tc // capture range variable
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
diff --git a/device/awg/tag_parser.go b/device/awg/tag_parser.go
index 2b09226..06ba49b 100644
--- a/device/awg/tag_parser.go
+++ b/device/awg/tag_parser.go
@@ -12,21 +12,21 @@ type IpcFields struct{ Key, Value string }
type EnumTag string
const (
- BytesEnumTag EnumTag = "b"
- CounterEnumTag EnumTag = "c"
- TimestampEnumTag EnumTag = "t"
- RandomBytesEnumTag EnumTag = "r"
- WaitTimeoutEnumTag EnumTag = "wt"
- WaitResponseEnumTag EnumTag = "wr"
+ BytesEnumTag EnumTag = "b"
+ CounterEnumTag EnumTag = "c"
+ TimestampEnumTag EnumTag = "t"
+ RandomBytesEnumTag EnumTag = "r"
+ RandomASCIIEnumTag EnumTag = "rc"
+ RandomDigitEnumTag EnumTag = "rd"
)
var generatorCreator = map[EnumTag]newGenerator{
BytesEnumTag: newBytesGenerator,
CounterEnumTag: newPacketCounterGenerator,
TimestampEnumTag: newTimestampGenerator,
- RandomBytesEnumTag: newRandomPacketGenerator,
- WaitTimeoutEnumTag: newWaitTimeoutGenerator,
- // WaitResponseEnumTag: newWaitResponseGenerator,
+ RandomBytesEnumTag: newRandomBytesGenerator,
+ RandomASCIIEnumTag: newRandomASCIIGenerator,
+ RandomDigitEnumTag: newRandomDigitGenerator,
}
// helper map to determine enumTags are unique
@@ -55,7 +55,7 @@ func parseTag(input string) (Tag, error) {
return tag, nil
}
-func Parse(name, input string) (TagJunkPacketGenerator, error) {
+func ParseTagJunkGenerator(name, input string) (TagJunkPacketGenerator, error) {
inputSlice := strings.Split(input, "<")
if len(inputSlice) <= 1 {
return TagJunkPacketGenerator{}, fmt.Errorf("empty input: %s", input)
diff --git a/device/awg/tag_parser_test.go b/device/awg/tag_parser_test.go
index 8f828ec..3229cee 100644
--- a/device/awg/tag_parser_test.go
+++ b/device/awg/tag_parser_test.go
@@ -64,7 +64,7 @@ func TestParse(t *testing.T) {
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- _, err := Parse(tt.args.name, tt.args.input)
+ _, err := ParseTagJunkGenerator(tt.args.name, tt.args.input)
// TODO: ErrorAs doesn't work as you think
if tt.wantErr != nil {
diff --git a/device/cookie.go b/device/cookie.go
index a093c8b..6a0463c 100644
--- a/device/cookie.go
+++ b/device/cookie.go
@@ -118,6 +118,7 @@ func (st *CookieChecker) CreateReply(
msg []byte,
recv uint32,
src []byte,
+ msgType uint32,
) (*MessageCookieReply, error) {
st.RLock()
@@ -153,7 +154,7 @@ func (st *CookieChecker) CreateReply(
smac1 := smac2 - blake2s.Size128
reply := new(MessageCookieReply)
- reply.Type = MessageCookieReplyType
+ reply.Type = msgType
reply.Receiver = recv
_, err := rand.Read(reply.Nonce[:])
diff --git a/device/cookie_test.go b/device/cookie_test.go
index c937290..e5a2bd4 100644
--- a/device/cookie_test.go
+++ b/device/cookie_test.go
@@ -99,7 +99,7 @@ func TestCookieMAC1(t *testing.T) {
0x8c, 0xe1, 0xe8, 0xfa, 0x67, 0x20, 0x80, 0x6d,
}
generator.AddMacs(msg)
- reply, err := checker.CreateReply(msg, 1377, src)
+ reply, err := checker.CreateReply(msg, 1377, src, DefaultMessageCookieReplyType)
if err != nil {
t.Fatal("Failed to create cookie reply:", err)
}
diff --git a/device/device.go b/device/device.go
index 1829352..46cf04e 100644
--- a/device/device.go
+++ b/device/device.go
@@ -6,7 +6,9 @@
package device
import (
+ "encoding/binary"
"errors"
+ "fmt"
"runtime"
"sync"
"sync/atomic"
@@ -578,6 +580,7 @@ func (device *Device) BindClose() error {
device.net.Unlock()
return err
}
+
func (device *Device) isAWG() bool {
return device.version >= VersionAwg
}
@@ -591,171 +594,123 @@ func (device *Device) resetProtocol() {
}
func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
- if !tempAwg.ASecCfg.IsSet && !tempAwg.HandshakeHandler.IsSet {
+ if !tempAwg.Cfg.IsSet && !tempAwg.HandshakeHandler.IsSet {
return nil
}
var errs []error
- isASecOn := false
- device.awg.ASecMux.Lock()
- if tempAwg.ASecCfg.JunkPacketCount < 0 {
+ isAwgOn := false
+ device.awg.Mux.Lock()
+ if tempAwg.Cfg.JunkPacketCount < 0 {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
"JunkPacketCount should be non negative",
),
)
}
- device.awg.ASecCfg.JunkPacketCount = tempAwg.ASecCfg.JunkPacketCount
- if tempAwg.ASecCfg.JunkPacketCount != 0 {
- isASecOn = true
+ device.awg.Cfg.JunkPacketCount = tempAwg.Cfg.JunkPacketCount
+ if tempAwg.Cfg.JunkPacketCount != 0 {
+ isAwgOn = true
}
- device.awg.ASecCfg.JunkPacketMinSize = tempAwg.ASecCfg.JunkPacketMinSize
- if tempAwg.ASecCfg.JunkPacketMinSize != 0 {
- isASecOn = true
+ device.awg.Cfg.JunkPacketMinSize = tempAwg.Cfg.JunkPacketMinSize
+ if tempAwg.Cfg.JunkPacketMinSize != 0 {
+ isAwgOn = true
}
- if device.awg.ASecCfg.JunkPacketCount > 0 &&
- tempAwg.ASecCfg.JunkPacketMaxSize == tempAwg.ASecCfg.JunkPacketMinSize {
+ if device.awg.Cfg.JunkPacketCount > 0 &&
+ tempAwg.Cfg.JunkPacketMaxSize == tempAwg.Cfg.JunkPacketMinSize {
- tempAwg.ASecCfg.JunkPacketMaxSize++ // to make rand gen work
+ tempAwg.Cfg.JunkPacketMaxSize++ // to make rand gen work
}
- if tempAwg.ASecCfg.JunkPacketMaxSize >= MaxSegmentSize {
- device.awg.ASecCfg.JunkPacketMinSize = 0
- device.awg.ASecCfg.JunkPacketMaxSize = 1
+ if tempAwg.Cfg.JunkPacketMaxSize >= MaxSegmentSize {
+ device.awg.Cfg.JunkPacketMinSize = 0
+ device.awg.Cfg.JunkPacketMaxSize = 1
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
"JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d",
- tempAwg.ASecCfg.JunkPacketMaxSize,
+ tempAwg.Cfg.JunkPacketMaxSize,
MaxSegmentSize,
))
- } else if tempAwg.ASecCfg.JunkPacketMaxSize < tempAwg.ASecCfg.JunkPacketMinSize {
+ } else if tempAwg.Cfg.JunkPacketMaxSize < tempAwg.Cfg.JunkPacketMinSize {
errs = append(errs, ipcErrorf(
ipc.IpcErrorInvalid,
"maxSize: %d; should be greater than minSize: %d",
- tempAwg.ASecCfg.JunkPacketMaxSize,
- tempAwg.ASecCfg.JunkPacketMinSize,
+ tempAwg.Cfg.JunkPacketMaxSize,
+ tempAwg.Cfg.JunkPacketMinSize,
))
} else {
- device.awg.ASecCfg.JunkPacketMaxSize = tempAwg.ASecCfg.JunkPacketMaxSize
+ device.awg.Cfg.JunkPacketMaxSize = tempAwg.Cfg.JunkPacketMaxSize
}
- if tempAwg.ASecCfg.JunkPacketMaxSize != 0 {
- isASecOn = true
+ if tempAwg.Cfg.JunkPacketMaxSize != 0 {
+ isAwgOn = true
}
- newInitSize := MessageInitiationSize + tempAwg.ASecCfg.InitHeaderJunkSize
+ magicHeaders := make([]awg.MagicHeader, 4)
- if newInitSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
+ if len(tempAwg.Cfg.MagicHeaders.Values) != 4 {
+ return ipcErrorf(
ipc.IpcErrorInvalid,
- `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.ASecCfg.InitHeaderJunkSize,
- MaxSegmentSize,
- ),
+ "magic headers should have 4 values; got: %d",
+ len(tempAwg.Cfg.MagicHeaders.Values),
)
- } else {
- device.awg.ASecCfg.InitHeaderJunkSize = tempAwg.ASecCfg.InitHeaderJunkSize
}
- if tempAwg.ASecCfg.InitHeaderJunkSize != 0 {
- isASecOn = true
- }
-
- newResponseSize := MessageResponseSize + tempAwg.ASecCfg.ResponseHeaderJunkSize
-
- if newResponseSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.ASecCfg.ResponseHeaderJunkSize,
- MaxSegmentSize,
- ),
- )
- } else {
- device.awg.ASecCfg.ResponseHeaderJunkSize = tempAwg.ASecCfg.ResponseHeaderJunkSize
- }
-
- if tempAwg.ASecCfg.ResponseHeaderJunkSize != 0 {
- isASecOn = true
- }
-
- newCookieSize := MessageCookieReplySize + tempAwg.ASecCfg.CookieReplyHeaderJunkSize
-
- if newCookieSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `cookie reply size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.ASecCfg.CookieReplyHeaderJunkSize,
- MaxSegmentSize,
- ),
- )
- } else {
- device.awg.ASecCfg.CookieReplyHeaderJunkSize = tempAwg.ASecCfg.CookieReplyHeaderJunkSize
- }
-
- if tempAwg.ASecCfg.CookieReplyHeaderJunkSize != 0 {
- isASecOn = true
- }
-
- newTransportSize := MessageTransportSize + tempAwg.ASecCfg.TransportHeaderJunkSize
-
- if newTransportSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `transport size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.ASecCfg.TransportHeaderJunkSize,
- MaxSegmentSize,
- ),
- )
- } else {
- device.awg.ASecCfg.TransportHeaderJunkSize = tempAwg.ASecCfg.TransportHeaderJunkSize
- }
-
- if tempAwg.ASecCfg.TransportHeaderJunkSize != 0 {
- isASecOn = true
- }
-
- if tempAwg.ASecCfg.InitPacketMagicHeader > 4 {
- isASecOn = true
+ if tempAwg.Cfg.MagicHeaders.Values[0].Min > 4 {
+ isAwgOn = true
device.log.Verbosef("UAPI: Updating init_packet_magic_header")
- device.awg.ASecCfg.InitPacketMagicHeader = tempAwg.ASecCfg.InitPacketMagicHeader
- MessageInitiationType = device.awg.ASecCfg.InitPacketMagicHeader
+ magicHeaders[0] = tempAwg.Cfg.MagicHeaders.Values[0]
+
+ MessageInitiationType = magicHeaders[0].Min
} else {
device.log.Verbosef("UAPI: Using default init type")
MessageInitiationType = DefaultMessageInitiationType
+ magicHeaders[0] = awg.NewMagicHeaderSameValue(DefaultMessageInitiationType)
}
- if tempAwg.ASecCfg.ResponsePacketMagicHeader > 4 {
- isASecOn = true
+ if tempAwg.Cfg.MagicHeaders.Values[1].Min > 4 {
+ isAwgOn = true
+
device.log.Verbosef("UAPI: Updating response_packet_magic_header")
- device.awg.ASecCfg.ResponsePacketMagicHeader = tempAwg.ASecCfg.ResponsePacketMagicHeader
- MessageResponseType = device.awg.ASecCfg.ResponsePacketMagicHeader
+ magicHeaders[1] = tempAwg.Cfg.MagicHeaders.Values[1]
+ MessageResponseType = magicHeaders[1].Min
} else {
device.log.Verbosef("UAPI: Using default response type")
MessageResponseType = DefaultMessageResponseType
+ magicHeaders[1] = awg.NewMagicHeaderSameValue(DefaultMessageResponseType)
}
- if tempAwg.ASecCfg.UnderloadPacketMagicHeader > 4 {
- isASecOn = true
+ if tempAwg.Cfg.MagicHeaders.Values[2].Min > 4 {
+ isAwgOn = true
+
device.log.Verbosef("UAPI: Updating underload_packet_magic_header")
- device.awg.ASecCfg.UnderloadPacketMagicHeader = tempAwg.ASecCfg.UnderloadPacketMagicHeader
- MessageCookieReplyType = device.awg.ASecCfg.UnderloadPacketMagicHeader
+ magicHeaders[2] = tempAwg.Cfg.MagicHeaders.Values[2]
+ MessageCookieReplyType = magicHeaders[2].Min
} else {
device.log.Verbosef("UAPI: Using default underload type")
MessageCookieReplyType = DefaultMessageCookieReplyType
+ magicHeaders[2] = awg.NewMagicHeaderSameValue(DefaultMessageCookieReplyType)
}
- if tempAwg.ASecCfg.TransportPacketMagicHeader > 4 {
- isASecOn = true
+ if tempAwg.Cfg.MagicHeaders.Values[3].Min > 4 {
+ isAwgOn = true
+
device.log.Verbosef("UAPI: Updating transport_packet_magic_header")
- device.awg.ASecCfg.TransportPacketMagicHeader = tempAwg.ASecCfg.TransportPacketMagicHeader
- MessageTransportType = device.awg.ASecCfg.TransportPacketMagicHeader
+ magicHeaders[3] = tempAwg.Cfg.MagicHeaders.Values[3]
+ MessageTransportType = magicHeaders[3].Min
} else {
device.log.Verbosef("UAPI: Using default transport type")
MessageTransportType = DefaultMessageTransportType
+ magicHeaders[3] = awg.NewMagicHeaderSameValue(DefaultMessageTransportType)
+ }
+
+ var err error
+ device.awg.Cfg.MagicHeaders, err = awg.NewMagicHeaders(magicHeaders)
+ if err != nil {
+ errs = append(errs, ipcErrorf(ipc.IpcErrorInvalid, "new magic headers: %w", err))
}
isSameHeaderMap := map[uint32]struct{}{
@@ -778,6 +733,78 @@ func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
)
}
+ newInitSize := MessageInitiationSize + tempAwg.Cfg.InitHeaderJunkSize
+
+ if newInitSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.Cfg.InitHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
+ } else {
+ device.awg.Cfg.InitHeaderJunkSize = tempAwg.Cfg.InitHeaderJunkSize
+ }
+
+ if tempAwg.Cfg.InitHeaderJunkSize != 0 {
+ isAwgOn = true
+ }
+
+ newResponseSize := MessageResponseSize + tempAwg.Cfg.ResponseHeaderJunkSize
+
+ if newResponseSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.Cfg.ResponseHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
+ } else {
+ device.awg.Cfg.ResponseHeaderJunkSize = tempAwg.Cfg.ResponseHeaderJunkSize
+ }
+
+ if tempAwg.Cfg.ResponseHeaderJunkSize != 0 {
+ isAwgOn = true
+ }
+
+ newCookieSize := MessageCookieReplySize + tempAwg.Cfg.CookieReplyHeaderJunkSize
+
+ if newCookieSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `cookie reply size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.Cfg.CookieReplyHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
+ } else {
+ device.awg.Cfg.CookieReplyHeaderJunkSize = tempAwg.Cfg.CookieReplyHeaderJunkSize
+ }
+
+ if tempAwg.Cfg.CookieReplyHeaderJunkSize != 0 {
+ isAwgOn = true
+ }
+
+ newTransportSize := MessageTransportSize + tempAwg.Cfg.TransportHeaderJunkSize
+
+ if newTransportSize >= MaxSegmentSize {
+ errs = append(errs, ipcErrorf(
+ ipc.IpcErrorInvalid,
+ `transport size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
+ tempAwg.Cfg.TransportHeaderJunkSize,
+ MaxSegmentSize,
+ ),
+ )
+ } else {
+ device.awg.Cfg.TransportHeaderJunkSize = tempAwg.Cfg.TransportHeaderJunkSize
+ }
+
+ if tempAwg.Cfg.TransportHeaderJunkSize != 0 {
+ isAwgOn = true
+ }
+
isSameSizeMap := map[int]struct{}{
newInitSize: {},
newResponseSize: {},
@@ -797,10 +824,10 @@ func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
)
} else {
msgTypeToJunkSize = map[uint32]int{
- MessageInitiationType: device.awg.ASecCfg.InitHeaderJunkSize,
- MessageResponseType: device.awg.ASecCfg.ResponseHeaderJunkSize,
- MessageCookieReplyType: device.awg.ASecCfg.CookieReplyHeaderJunkSize,
- MessageTransportType: device.awg.ASecCfg.TransportHeaderJunkSize,
+ MessageInitiationType: device.awg.Cfg.InitHeaderJunkSize,
+ MessageResponseType: device.awg.Cfg.ResponseHeaderJunkSize,
+ MessageCookieReplyType: device.awg.Cfg.CookieReplyHeaderJunkSize,
+ MessageTransportType: device.awg.Cfg.TransportHeaderJunkSize,
}
packetSizeToMsgType = map[int]uint32{
@@ -811,12 +838,8 @@ func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
}
}
- device.awg.IsASecOn.SetTo(isASecOn)
- var err error
- device.awg.JunkCreator, err = awg.NewJunkCreator(device.awg.ASecCfg)
- if err != nil {
- errs = append(errs, err)
- }
+ device.awg.IsOn.SetTo(isAwgOn)
+ device.awg.JunkCreator = awg.NewJunkCreator(device.awg.Cfg)
if tempAwg.HandshakeHandler.IsSet {
if err := tempAwg.HandshakeHandler.Validate(); err != nil {
@@ -824,15 +847,91 @@ func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
ipc.IpcErrorInvalid, "handshake handler validate: %w", err))
} else {
device.awg.HandshakeHandler = tempAwg.HandshakeHandler
- device.awg.HandshakeHandler.ControlledJunk.DefaultJunkCount = tempAwg.ASecCfg.JunkPacketCount
- device.awg.HandshakeHandler.SpecialJunk.DefaultJunkCount = tempAwg.ASecCfg.JunkPacketCount
+ device.awg.HandshakeHandler.SpecialJunk.DefaultJunkCount = tempAwg.Cfg.JunkPacketCount
device.version = VersionAwgSpecialHandshake
}
} else {
device.version = VersionAwg
}
- device.awg.ASecMux.Unlock()
+ device.awg.Mux.Unlock()
return errors.Join(errs...)
}
+
+func (device *Device) ProcessAWGPacket(size int, packet *[]byte, buffer *[MaxMessageSize]byte) (uint32, error) {
+ // TODO:
+ // if awg.WaitResponse.ShouldWait.IsSet() {
+ // awg.WaitResponse.Channel <- struct{}{}
+ // }
+
+ expectedMsgType, isKnownSize := packetSizeToMsgType[size]
+ if !isKnownSize {
+ msgType, err := device.handleTransport(size, packet, buffer)
+
+ if err != nil {
+ return 0, fmt.Errorf("handle transport: %w", err)
+ }
+
+ return msgType, nil
+ }
+
+ junkSize := msgTypeToJunkSize[expectedMsgType]
+
+ // transport size can align with other header types;
+ // making sure we have the right actualMsgType
+ actualMsgType, err := device.getMsgType(packet, junkSize)
+ if err != nil {
+ return 0, fmt.Errorf("get msg type: %w", err)
+ }
+
+ if actualMsgType == expectedMsgType {
+ *packet = (*packet)[junkSize:]
+ return actualMsgType, nil
+ }
+
+ device.log.Verbosef("awg: transport packet lined up with another msg type")
+
+ msgType, err := device.handleTransport(size, packet, buffer)
+ if err != nil {
+ return 0, fmt.Errorf("handle transport: %w", err)
+ }
+
+ return msgType, nil
+}
+
+func (device *Device) getMsgType(packet *[]byte, junkSize int) (uint32, error) {
+ msgTypeValue := binary.LittleEndian.Uint32((*packet)[junkSize : junkSize+4])
+ msgType, err := device.awg.GetMagicHeaderMinFor(msgTypeValue)
+
+ if err != nil {
+ return 0, fmt.Errorf("get magic header min: %w", err)
+ }
+
+ return msgType, nil
+}
+
+func (device *Device) handleTransport(size int, packet *[]byte, buffer *[MaxMessageSize]byte) (uint32, error) {
+ junkSize := device.awg.Cfg.TransportHeaderJunkSize
+
+ msgType, err := device.getMsgType(packet, junkSize)
+ if err != nil {
+ return 0, fmt.Errorf("get msg type: %w", err)
+ }
+
+ if msgType != MessageTransportType {
+ // probably a junk packet
+ return 0, fmt.Errorf("Received message with unknown type: %d", msgType)
+ }
+
+ if junkSize > 0 {
+ // remove junk from buffer by shifting the packet
+ // this buffer is also used for decryption, so it needs to be corrected
+ copy((*buffer)[:size], (*packet)[junkSize:])
+ size -= junkSize
+ // need to reinitialize packet as well
+ (*packet) = (*packet)[:size]
+ }
+
+ return msgType, nil
+}
diff --git a/device/device_test.go b/device/device_test.go
index 5824cf9..2f66185 100644
--- a/device/device_test.go
+++ b/device/device_test.go
@@ -232,14 +232,14 @@ func TestAWGDevicePing(t *testing.T) {
"jc", "5",
"jmin", "500",
"jmax", "1000",
- "s1", "30",
- "s2", "40",
- "s3", "50",
- "s4", "5",
- "h1", "123456",
- "h2", "67543",
- "h3", "123123",
- "h4", "32345",
+ "s1", "15",
+ "s2", "18",
+ "s3", "20",
+ "s4", "25",
+ "h1", "123456-123500",
+ "h2", "67543-67550",
+ "h3", "123123-123200",
+ "h4", "32345-32350",
)
t.Run("ping 1.0.0.1", func(t *testing.T) {
pair.Send(t, Ping, nil)
@@ -264,12 +264,10 @@ func TestAWGHandshakeDevicePing(t *testing.T) {
goroutineLeakCheck(t)
pair := genTestPair(t, true,
- "i1", "",
- "i2", "",
- "j1", "",
- "j2", "",
- "j3", "",
- "itime", "60",
+ "i1", "",
+ "i2", "",
+ "i3", "",
+ "i4", "",
// "jc", "1",
// "jmin", "500",
// "jmax", "1000",
diff --git a/device/noise-protocol.go b/device/noise-protocol.go
index f637b24..6e6fe58 100644
--- a/device/noise-protocol.go
+++ b/device/noise-protocol.go
@@ -205,12 +205,22 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e
handshake.mixHash(handshake.remoteStatic[:])
- device.awg.ASecMux.RLock()
+ msgType := DefaultMessageInitiationType
+ if device.isAWG() {
+ device.awg.Mux.RLock()
+ msgType, err = device.awg.GetMsgType(DefaultMessageInitiationType)
+ if err != nil {
+ device.awg.Mux.RUnlock()
+ return nil, fmt.Errorf("get message type: %w", err)
+ }
+
+ device.awg.Mux.RUnlock()
+ }
+
msg := MessageInitiation{
- Type: MessageInitiationType,
+ Type: msgType,
Ephemeral: handshake.localEphemeral.publicKey(),
}
- device.awg.ASecMux.RUnlock()
handshake.mixKey(msg.Ephemeral[:])
handshake.mixHash(msg.Ephemeral[:])
@@ -264,12 +274,13 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer {
chainKey [blake2s.Size]byte
)
- device.awg.ASecMux.RLock()
+ device.awg.Mux.RLock()
+
if msg.Type != MessageInitiationType {
- device.awg.ASecMux.RUnlock()
+ device.awg.Mux.RUnlock()
return nil
}
- device.awg.ASecMux.RUnlock()
+ device.awg.Mux.RUnlock()
device.staticIdentity.RLock()
defer device.staticIdentity.RUnlock()
@@ -384,9 +395,19 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
var msg MessageResponse
- device.awg.ASecMux.RLock()
- msg.Type = MessageResponseType
- device.awg.ASecMux.RUnlock()
+ if device.isAWG() {
+ device.awg.Mux.RLock()
+ msg.Type, err = device.awg.GetMsgType(DefaultMessageResponseType)
+ if err != nil {
+ device.awg.Mux.RUnlock()
+ return nil, fmt.Errorf("get message type: %w", err)
+ }
+
+ device.awg.Mux.RUnlock()
+ } else {
+ msg.Type = DefaultMessageResponseType
+ }
+
msg.Sender = handshake.localIndex
msg.Receiver = handshake.remoteIndex
@@ -436,12 +457,13 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
- device.awg.ASecMux.RLock()
+ device.awg.Mux.RLock()
+
if msg.Type != MessageResponseType {
- device.awg.ASecMux.RUnlock()
+ device.awg.Mux.RUnlock()
return nil
}
- device.awg.ASecMux.RUnlock()
+ device.awg.Mux.RUnlock()
// lookup handshake by receiver
diff --git a/device/receive.go b/device/receive.go
index 6daba0d..4c34799 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -129,7 +129,7 @@ func (device *Device) RoutineReceiveIncoming(
}
deathSpiral = 0
- device.awg.ASecMux.RLock()
+ device.awg.Mux.RLock()
// handle each packet in the batch
for i, size := range sizes[:count] {
if size < MinMessageSize {
@@ -140,37 +140,11 @@ func (device *Device) RoutineReceiveIncoming(
packet := bufsArrs[i][:size]
var msgType uint32
if device.isAWG() {
- // TODO:
- // if awg.WaitResponse.ShouldWait.IsSet() {
- // awg.WaitResponse.Channel <- struct{}{}
- // }
+ msgType, err = device.ProcessAWGPacket(size, &packet, bufsArrs[i])
- if assumedMsgType, ok := packetSizeToMsgType[size]; ok {
- junkSize := msgTypeToJunkSize[assumedMsgType]
- // transport size can align with other header types;
- // making sure we have the right msgType
- msgType = binary.LittleEndian.Uint32(packet[junkSize : junkSize+4])
- if msgType == assumedMsgType {
- packet = packet[junkSize:]
- } else {
- device.log.Verbosef("transport packet lined up with another msg type")
- msgType = binary.LittleEndian.Uint32(packet[:4])
- }
- } else {
- transportJunkSize := device.awg.ASecCfg.TransportHeaderJunkSize
- msgType = binary.LittleEndian.Uint32(packet[transportJunkSize : transportJunkSize+4])
- if msgType != MessageTransportType {
- // probably a junk packet
- device.log.Verbosef("aSec: Received message with unknown type: %d", msgType)
- continue
- }
-
- // remove junk from bufsArrs by shifting the packet
- // this buffer is also used for decryption, so it needs to be corrected
- copy(bufsArrs[i][:size], packet[transportJunkSize:])
- size -= transportJunkSize
- // need to reinitialize packet as well
- packet = packet[:size]
+ if err != nil {
+ device.log.Verbosef("awg: process packet: %v", err)
+ continue
}
} else {
msgType = binary.LittleEndian.Uint32(packet[:4])
@@ -259,7 +233,7 @@ func (device *Device) RoutineReceiveIncoming(
default:
}
}
- device.awg.ASecMux.RUnlock()
+ device.awg.Mux.RUnlock()
for peer, elemsContainer := range elemsByPeer {
if peer.isRunning.Load() {
peer.queue.inbound.c <- elemsContainer
@@ -318,7 +292,7 @@ func (device *Device) RoutineHandshake(id int) {
for elem := range device.queue.handshake.c {
- device.awg.ASecMux.RLock()
+ device.awg.Mux.RLock()
// handle cookie fields and ratelimiting
@@ -405,6 +379,9 @@ func (device *Device) RoutineHandshake(id int) {
goto skip
}
+ // have to reassign msgType for ranged msgType to work
+ msg.Type = elem.msgType
+
// consume initiation
peer := device.ConsumeMessageInitiation(&msg)
if peer == nil {
@@ -437,6 +414,9 @@ func (device *Device) RoutineHandshake(id int) {
goto skip
}
+ // have to reassign msgType for ranged msgType to work
+ msg.Type = elem.msgType
+
// consume response
peer := device.ConsumeMessageResponse(&msg)
@@ -470,7 +450,7 @@ func (device *Device) RoutineHandshake(id int) {
peer.SendKeepalive()
}
skip:
- device.awg.ASecMux.RUnlock()
+ device.awg.Mux.RUnlock()
device.PutMessageBuffer(elem.buffer)
}
}
diff --git a/device/send.go b/device/send.go
index 04ca2ad..0861a04 100644
--- a/device/send.go
+++ b/device/send.go
@@ -130,29 +130,19 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
if peer.device.version >= VersionAwg {
var junks [][]byte
if peer.device.version == VersionAwgSpecialHandshake {
- peer.device.awg.ASecMux.RLock()
+ peer.device.awg.Mux.RLock()
// set junks depending on packet type
junks = peer.device.awg.HandshakeHandler.GenerateSpecialJunk()
- if junks == nil {
- junks = peer.device.awg.HandshakeHandler.GenerateControlledJunk()
- if junks != nil {
- peer.device.log.Verbosef("%v - Controlled junks sent", peer)
- }
- } else {
+ if junks != nil {
peer.device.log.Verbosef("%v - Special junks sent", peer)
}
- peer.device.awg.ASecMux.RUnlock()
+ peer.device.awg.Mux.RUnlock()
} else {
- junks = make([][]byte, 0, peer.device.awg.ASecCfg.JunkPacketCount)
- }
- peer.device.awg.ASecMux.RLock()
- err := peer.device.awg.JunkCreator.CreateJunkPackets(&junks)
- peer.device.awg.ASecMux.RUnlock()
-
- if err != nil {
- peer.device.log.Errorf("%v - %v", peer, err)
- return err
+ junks = make([][]byte, 0, peer.device.awg.Cfg.JunkPacketCount)
}
+ peer.device.awg.Mux.RLock()
+ peer.device.awg.JunkCreator.CreateJunkPackets(&junks)
+ peer.device.awg.Mux.RUnlock()
if len(junks) > 0 {
err = peer.SendBuffers(junks)
@@ -242,10 +232,24 @@ func (device *Device) SendHandshakeCookie(
device.log.Verbosef("Sending cookie response for denied handshake message for %v", initiatingElem.endpoint.DstToString())
sender := binary.LittleEndian.Uint32(initiatingElem.packet[4:8])
+ msgType := DefaultMessageCookieReplyType
+ if device.isAWG() {
+ device.awg.Mux.RLock()
+
+ var err error
+ msgType, err = device.awg.GetMsgType(DefaultMessageCookieReplyType)
+ device.awg.Mux.RUnlock()
+ if err != nil {
+ device.log.Errorf("Get message type for cookie reply: %v", err)
+ return err
+ }
+ }
+
reply, err := device.cookieChecker.CreateReply(
initiatingElem.packet,
sender,
initiatingElem.endpoint.DstToBytes(),
+ msgType,
)
if err != nil {
device.log.Errorf("Failed to create cookie reply: %v", err)
@@ -528,7 +532,20 @@ func (device *Device) RoutineEncryption(id int) {
fieldReceiver := header[4:8]
fieldNonce := header[8:16]
- binary.LittleEndian.PutUint32(fieldType, MessageTransportType)
+ msgType := DefaultMessageTransportType
+ if device.isAWG() {
+ device.awg.Mux.RLock()
+
+ var err error
+ msgType, err = device.awg.GetMsgType(DefaultMessageTransportType)
+ device.awg.Mux.RUnlock()
+ if err != nil {
+ device.log.Errorf("get message type for transport: %v", err)
+ continue
+ }
+ }
+
+ binary.LittleEndian.PutUint32(fieldType, msgType)
binary.LittleEndian.PutUint32(fieldReceiver, elem.keypair.remoteIndex)
binary.LittleEndian.PutUint64(fieldNonce, elem.nonce)
diff --git a/device/uapi.go b/device/uapi.go
index e9f962a..6c4be05 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -99,51 +99,42 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
}
if device.isAWG() {
- if device.awg.ASecCfg.JunkPacketCount != 0 {
- sendf("jc=%d", device.awg.ASecCfg.JunkPacketCount)
+ if device.awg.Cfg.JunkPacketCount != 0 {
+ sendf("jc=%d", device.awg.Cfg.JunkPacketCount)
}
- if device.awg.ASecCfg.JunkPacketMinSize != 0 {
- sendf("jmin=%d", device.awg.ASecCfg.JunkPacketMinSize)
+ if device.awg.Cfg.JunkPacketMinSize != 0 {
+ sendf("jmin=%d", device.awg.Cfg.JunkPacketMinSize)
}
- if device.awg.ASecCfg.JunkPacketMaxSize != 0 {
- sendf("jmax=%d", device.awg.ASecCfg.JunkPacketMaxSize)
+ if device.awg.Cfg.JunkPacketMaxSize != 0 {
+ sendf("jmax=%d", device.awg.Cfg.JunkPacketMaxSize)
}
- if device.awg.ASecCfg.InitHeaderJunkSize != 0 {
- sendf("s1=%d", device.awg.ASecCfg.InitHeaderJunkSize)
+ if device.awg.Cfg.InitHeaderJunkSize != 0 {
+ sendf("s1=%d", device.awg.Cfg.InitHeaderJunkSize)
}
- if device.awg.ASecCfg.ResponseHeaderJunkSize != 0 {
- sendf("s2=%d", device.awg.ASecCfg.ResponseHeaderJunkSize)
+ if device.awg.Cfg.ResponseHeaderJunkSize != 0 {
+ sendf("s2=%d", device.awg.Cfg.ResponseHeaderJunkSize)
}
- if device.awg.ASecCfg.CookieReplyHeaderJunkSize != 0 {
- sendf("s3=%d", device.awg.ASecCfg.CookieReplyHeaderJunkSize)
+ if device.awg.Cfg.CookieReplyHeaderJunkSize != 0 {
+ sendf("s3=%d", device.awg.Cfg.CookieReplyHeaderJunkSize)
}
- if device.awg.ASecCfg.TransportHeaderJunkSize != 0 {
- sendf("s4=%d", device.awg.ASecCfg.TransportHeaderJunkSize)
+ if device.awg.Cfg.TransportHeaderJunkSize != 0 {
+ sendf("s4=%d", device.awg.Cfg.TransportHeaderJunkSize)
}
- if device.awg.ASecCfg.InitPacketMagicHeader != 0 {
- sendf("h1=%d", device.awg.ASecCfg.InitPacketMagicHeader)
- }
- if device.awg.ASecCfg.ResponsePacketMagicHeader != 0 {
- sendf("h2=%d", device.awg.ASecCfg.ResponsePacketMagicHeader)
- }
- if device.awg.ASecCfg.UnderloadPacketMagicHeader != 0 {
- sendf("h3=%d", device.awg.ASecCfg.UnderloadPacketMagicHeader)
- }
- if device.awg.ASecCfg.TransportPacketMagicHeader != 0 {
- sendf("h4=%d", device.awg.ASecCfg.TransportPacketMagicHeader)
+ for i, magicHeader := range device.awg.Cfg.MagicHeaders.Values {
+ if magicHeader.Min > 4 {
+ if magicHeader.Min == magicHeader.Max {
+ sendf("h%d=%d", i+1, magicHeader.Min)
+ continue
+ }
+
+ sendf("h%d=%d-%d", i+1, magicHeader.Min, magicHeader.Max)
+ }
}
specialJunkIpcFields := device.awg.HandshakeHandler.SpecialJunk.IpcGetFields()
for _, field := range specialJunkIpcFields {
sendf("%s=%s", field.Key, field.Value)
}
- controlledJunkIpcFields := device.awg.HandshakeHandler.ControlledJunk.IpcGetFields()
- for _, field := range controlledJunkIpcFields {
- sendf("%s=%s", field.Key, field.Value)
- }
- if device.awg.HandshakeHandler.ITimeout != 0 {
- sendf("itime=%d", device.awg.HandshakeHandler.ITimeout/time.Second)
- }
}
for _, peer := range device.peers.keyMap {
@@ -200,6 +191,8 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
deviceConfig := true
tempAwg := awg.Protocol{}
+ tempAwg.Cfg.MagicHeaders.Values = make([]awg.MagicHeader, 4)
+
scanner := bufio.NewScanner(r)
for scanner.Scan() {
line := scanner.Text()
@@ -312,8 +305,8 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_count %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_count")
- tempAwg.ASecCfg.JunkPacketCount = junkPacketCount
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.JunkPacketCount = junkPacketCount
+ tempAwg.Cfg.IsSet = true
case "jmin":
junkPacketMinSize, err := strconv.Atoi(value)
@@ -321,8 +314,8 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_min_size %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_min_size")
- tempAwg.ASecCfg.JunkPacketMinSize = junkPacketMinSize
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.JunkPacketMinSize = junkPacketMinSize
+ tempAwg.Cfg.IsSet = true
case "jmax":
junkPacketMaxSize, err := strconv.Atoi(value)
@@ -330,8 +323,8 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_max_size %w", err)
}
device.log.Verbosef("UAPI: Updating junk_packet_max_size")
- tempAwg.ASecCfg.JunkPacketMaxSize = junkPacketMaxSize
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.JunkPacketMaxSize = junkPacketMaxSize
+ tempAwg.Cfg.IsSet = true
case "s1":
initPacketJunkSize, err := strconv.Atoi(value)
@@ -339,8 +332,8 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
return ipcErrorf(ipc.IpcErrorInvalid, "parse init_packet_junk_size %w", err)
}
device.log.Verbosef("UAPI: Updating init_packet_junk_size")
- tempAwg.ASecCfg.InitHeaderJunkSize = initPacketJunkSize
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.InitHeaderJunkSize = initPacketJunkSize
+ tempAwg.Cfg.IsSet = true
case "s2":
responsePacketJunkSize, err := strconv.Atoi(value)
@@ -348,8 +341,8 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
return ipcErrorf(ipc.IpcErrorInvalid, "parse response_packet_junk_size %w", err)
}
device.log.Verbosef("UAPI: Updating response_packet_junk_size")
- tempAwg.ASecCfg.ResponseHeaderJunkSize = responsePacketJunkSize
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.ResponseHeaderJunkSize = responsePacketJunkSize
+ tempAwg.Cfg.IsSet = true
case "s3":
cookieReplyPacketJunkSize, err := strconv.Atoi(value)
@@ -357,8 +350,8 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
return ipcErrorf(ipc.IpcErrorInvalid, "parse cookie_reply_packet_junk_size %w", err)
}
device.log.Verbosef("UAPI: Updating cookie_reply_packet_junk_size")
- tempAwg.ASecCfg.CookieReplyHeaderJunkSize = cookieReplyPacketJunkSize
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.CookieReplyHeaderJunkSize = cookieReplyPacketJunkSize
+ tempAwg.Cfg.IsSet = true
case "s4":
transportPacketJunkSize, err := strconv.Atoi(value)
@@ -366,81 +359,53 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
return ipcErrorf(ipc.IpcErrorInvalid, "parse transport_packet_junk_size %w", err)
}
device.log.Verbosef("UAPI: Updating transport_packet_junk_size")
- tempAwg.ASecCfg.TransportHeaderJunkSize = transportPacketJunkSize
- tempAwg.ASecCfg.IsSet = true
-
+ tempAwg.Cfg.TransportHeaderJunkSize = transportPacketJunkSize
+ tempAwg.Cfg.IsSet = true
case "h1":
- initPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ initMagicHeader, err := awg.ParseMagicHeader(key, value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse init_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
}
- tempAwg.ASecCfg.InitPacketMagicHeader = uint32(initPacketMagicHeader)
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.MagicHeaders.Values[0] = initMagicHeader
+ tempAwg.Cfg.IsSet = true
case "h2":
- responsePacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ responseMagicHeader, err := awg.ParseMagicHeader(key, value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse response_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
}
- tempAwg.ASecCfg.ResponsePacketMagicHeader = uint32(responsePacketMagicHeader)
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.MagicHeaders.Values[1] = responseMagicHeader
+ tempAwg.Cfg.IsSet = true
case "h3":
- underloadPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ cookieReplyMagicHeader, err := awg.ParseMagicHeader(key, value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse underload_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
}
- tempAwg.ASecCfg.UnderloadPacketMagicHeader = uint32(underloadPacketMagicHeader)
- tempAwg.ASecCfg.IsSet = true
+ tempAwg.Cfg.MagicHeaders.Values[2] = cookieReplyMagicHeader
+ tempAwg.Cfg.IsSet = true
case "h4":
- transportPacketMagicHeader, err := strconv.ParseUint(value, 10, 32)
+ transportMagicHeader, err := awg.ParseMagicHeader(key, value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse transport_packet_magic_header %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
}
- tempAwg.ASecCfg.TransportPacketMagicHeader = uint32(transportPacketMagicHeader)
- tempAwg.ASecCfg.IsSet = true
+
+ tempAwg.Cfg.MagicHeaders.Values[3] = transportMagicHeader
+ tempAwg.Cfg.IsSet = true
case "i1", "i2", "i3", "i4", "i5":
if len(value) == 0 {
device.log.Verbosef("UAPI: received empty %s", key)
return nil
}
- generators, err := awg.Parse(key, value)
+ generators, err := awg.ParseTagJunkGenerator(key, value)
if err != nil {
return ipcErrorf(ipc.IpcErrorInvalid, "invalid %s: %w", key, err)
}
device.log.Verbosef("UAPI: Updating %s", key)
tempAwg.HandshakeHandler.SpecialJunk.AppendGenerator(generators)
tempAwg.HandshakeHandler.IsSet = true
- case "j1", "j2", "j3":
- if len(value) == 0 {
- device.log.Verbosef("UAPI: received empty %s", key)
- return nil
- }
-
- generators, err := awg.Parse(key, value)
- if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "invalid %s: %w", key, err)
- }
- device.log.Verbosef("UAPI: Updating %s", key)
-
- tempAwg.HandshakeHandler.ControlledJunk.AppendGenerator(generators)
- tempAwg.HandshakeHandler.IsSet = true
- case "itime":
- if len(value) == 0 {
- device.log.Verbosef("UAPI: received empty itime")
- return nil
- }
-
- itime, err := strconv.ParseInt(value, 10, 64)
- if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse itime %w", err)
- }
- device.log.Verbosef("UAPI: Updating itime")
-
- tempAwg.HandshakeHandler.ITimeout = time.Duration(itime) * time.Second
- tempAwg.HandshakeHandler.IsSet = true
default:
return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
}
diff --git a/go.mod b/go.mod
index 5e5f34d..8c4372d 100644
--- a/go.mod
+++ b/go.mod
@@ -5,9 +5,9 @@ go 1.24.4
require (
github.com/stretchr/testify v1.10.0
github.com/tevino/abool v1.2.0
- github.com/tevino/abool/v2 v2.1.0
go.uber.org/atomic v1.11.0
golang.org/x/crypto v0.39.0
+ golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
golang.org/x/net v0.41.0
golang.org/x/sys v0.33.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
diff --git a/go.sum b/go.sum
index 6b8f36b..3d8b3c2 100644
--- a/go.sum
+++ b/go.sum
@@ -2,24 +2,18 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
-github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38=
-github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
-github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/tevino/abool v1.2.0 h1:heAkClL8H6w+mK5md9dzsuohKeXHUpY7Vw0ZCKW+huA=
github.com/tevino/abool v1.2.0/go.mod h1:qc66Pna1RiIsPa7O4Egxxs9OqkuxDX55zznh9K07Tzg=
-github.com/tevino/abool/v2 v2.1.0 h1:7w+Vf9f/5gmKT4m4qkayb33/92M+Um45F2BkHOR+L/c=
-github.com/tevino/abool/v2 v2.1.0/go.mod h1:+Lmlqk6bHDWHqN1cbxqhwEAwMPXgc8I1SDEamtseuXY=
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
golang.org/x/crypto v0.39.0 h1:SHs+kF4LP+f+p14esP5jAoDpHU8Gu/v9lFRK6IT5imM=
golang.org/x/crypto v0.39.0/go.mod h1:L+Xg3Wf6HoL4Bn4238Z6ft6KfEpN0tJGo53AAPC632U=
-golang.org/x/mod v0.13.0 h1:I/DsJXRlw/8l/0c24sM9yb0T4z9liZTduXvdAWYiysY=
-golang.org/x/mod v0.21.0 h1:vvrHzRwRfVKSiLrG+d4FMl/Qi4ukBCE6kZlTUkDYRT0=
-golang.org/x/mod v0.21.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
+golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
+golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
golang.org/x/net v0.41.0 h1:vBTly1HeNPEn3wtREYfy4GZ/NECgw2Cnl+nK6Nz3uvw=
golang.org/x/net v0.41.0/go.mod h1:B/K4NNqkfmg07DQYrbwvSluqCJOOXwUjeb/5lOisjbA=
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
@@ -34,7 +28,3 @@ gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489 h1:ze1vwAdliUAr68RQ5NtufWaXaOg8WUO2OACzEV+TNdE=
gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489/go.mod h1:10sU+Uh5KKNv1+2x2A0Gvzt8FjD3ASIhorV3YsauXhk=
-gvisor.dev/gvisor v0.0.0-20250428193742-2d800c3129d5 h1:sfK5nHuG7lRFZ2FdTT3RimOqWBg8IrVm+/Vko1FVOsk=
-gvisor.dev/gvisor v0.0.0-20250428193742-2d800c3129d5/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
-gvisor.dev/gvisor v0.0.0-20250606233247-e3c4c4cad86f h1:zmc4cHEcCudRt2O8VsCW7nYLfAsbVY2i910/DAop1TM=
-gvisor.dev/gvisor v0.0.0-20250606233247-e3c4c4cad86f/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
From 0361c54dca92071463e9611796c04464386457ac Mon Sep 17 00:00:00 2001
From: Yaroslav Gurov <31506978+ygurov@users.noreply.github.com>
Date: Mon, 1 Dec 2025 13:07:48 +0100
Subject: [PATCH 69/75] fix: refactor processing of junk packets (#103)
- fix the bug that transport packet interprets as init/resp/cookie with the same size
- cleanup error responses
- reduce buffer allocations
---
device/awg/awg.go | 90 ----
device/awg/internal/mock.go | 37 --
device/awg/junk_creator.go | 50 --
device/awg/junk_creator_test.go | 97 ----
device/awg/magic_header.go | 97 ----
device/awg/magic_header_test.go | 488 ------------------
device/awg/prng.go | 50 --
device/awg/special_handshake_handler.go | 36 --
device/awg/tag_generator.go | 229 --------
device/awg/tag_generator_test.go | 321 ------------
device/awg/tag_junk_packet_generator.go | 59 ---
device/awg/tag_junk_packet_generator_test.go | 210 --------
device/awg/tag_junk_packet_generators.go | 66 ---
device/awg/tag_junk_packet_generators_test.go | 149 ------
device/awg/tag_parser.go | 112 ----
device/awg/tag_parser_test.go | 77 ---
device/cookie_test.go | 2 +-
device/device.go | 425 +--------------
device/magic-header.go | 63 +++
device/noise-protocol.go | 55 +-
device/obf.go | 140 +++++
device/obf_bytes.go | 47 ++
device/obf_data.go | 25 +
device/obf_datasize.go | 38 ++
device/obf_datastring.go | 29 ++
device/obf_rand.go | 39 ++
device/obf_randchars.go | 48 ++
device/obf_randdigits.go | 48 ++
device/obf_timestamp.go | 31 ++
device/peer.go | 11 -
device/receive.go | 78 ++-
device/send.go | 138 ++---
device/uapi.go | 299 +++++++----
33 files changed, 852 insertions(+), 2832 deletions(-)
delete mode 100644 device/awg/awg.go
delete mode 100644 device/awg/internal/mock.go
delete mode 100644 device/awg/junk_creator.go
delete mode 100644 device/awg/junk_creator_test.go
delete mode 100644 device/awg/magic_header.go
delete mode 100644 device/awg/magic_header_test.go
delete mode 100644 device/awg/prng.go
delete mode 100644 device/awg/special_handshake_handler.go
delete mode 100644 device/awg/tag_generator.go
delete mode 100644 device/awg/tag_generator_test.go
delete mode 100644 device/awg/tag_junk_packet_generator.go
delete mode 100644 device/awg/tag_junk_packet_generator_test.go
delete mode 100644 device/awg/tag_junk_packet_generators.go
delete mode 100644 device/awg/tag_junk_packet_generators_test.go
delete mode 100644 device/awg/tag_parser.go
delete mode 100644 device/awg/tag_parser_test.go
create mode 100644 device/magic-header.go
create mode 100644 device/obf.go
create mode 100644 device/obf_bytes.go
create mode 100644 device/obf_data.go
create mode 100644 device/obf_datasize.go
create mode 100644 device/obf_datastring.go
create mode 100644 device/obf_rand.go
create mode 100644 device/obf_randchars.go
create mode 100644 device/obf_randdigits.go
create mode 100644 device/obf_timestamp.go
diff --git a/device/awg/awg.go b/device/awg/awg.go
deleted file mode 100644
index 888a42e..0000000
--- a/device/awg/awg.go
+++ /dev/null
@@ -1,90 +0,0 @@
-package awg
-
-import (
- "bytes"
- "fmt"
- "sync"
-
- "github.com/tevino/abool"
-)
-
-type Cfg struct {
- IsSet bool
- JunkPacketCount int
- JunkPacketMinSize int
- JunkPacketMaxSize int
- InitHeaderJunkSize int
- ResponseHeaderJunkSize int
- CookieReplyHeaderJunkSize int
- TransportHeaderJunkSize int
-
- MagicHeaders MagicHeaders
-}
-
-type Protocol struct {
- IsOn abool.AtomicBool
- // TODO: revision the need of the mutex
- Mux sync.RWMutex
- Cfg Cfg
- JunkCreator JunkCreator
-
- HandshakeHandler SpecialHandshakeHandler
-}
-
-func (protocol *Protocol) CreateInitHeaderJunk() ([]byte, error) {
- protocol.Mux.RLock()
- defer protocol.Mux.RUnlock()
-
- return protocol.createHeaderJunk(protocol.Cfg.InitHeaderJunkSize, 0)
-}
-
-func (protocol *Protocol) CreateResponseHeaderJunk() ([]byte, error) {
- protocol.Mux.RLock()
- defer protocol.Mux.RUnlock()
-
- return protocol.createHeaderJunk(protocol.Cfg.ResponseHeaderJunkSize, 0)
-}
-
-func (protocol *Protocol) CreateCookieReplyHeaderJunk() ([]byte, error) {
- protocol.Mux.RLock()
- defer protocol.Mux.RUnlock()
-
- return protocol.createHeaderJunk(protocol.Cfg.CookieReplyHeaderJunkSize, 0)
-}
-
-func (protocol *Protocol) CreateTransportHeaderJunk(packetSize int) ([]byte, error) {
- protocol.Mux.RLock()
- defer protocol.Mux.RUnlock()
-
- return protocol.createHeaderJunk(protocol.Cfg.TransportHeaderJunkSize, packetSize)
-}
-
-func (protocol *Protocol) createHeaderJunk(junkSize int, extraSize int) ([]byte, error) {
- if junkSize == 0 {
- return nil, nil
- }
-
- buf := make([]byte, 0, junkSize+extraSize)
- writer := bytes.NewBuffer(buf[:0])
-
- err := protocol.JunkCreator.AppendJunk(writer, junkSize)
- if err != nil {
- return nil, fmt.Errorf("append junk: %w", err)
- }
-
- return writer.Bytes(), nil
-}
-
-func (protocol *Protocol) GetMagicHeaderMinFor(msgType uint32) (uint32, error) {
- for _, magicHeader := range protocol.Cfg.MagicHeaders.Values {
- if magicHeader.Min <= msgType && msgType <= magicHeader.Max {
- return magicHeader.Min, nil
- }
- }
-
- return 0, fmt.Errorf("no header for value: %d", msgType)
-}
-
-func (protocol *Protocol) GetMsgType(defaultMsgType uint32) (uint32, error) {
- return protocol.Cfg.MagicHeaders.Get(defaultMsgType)
-}
diff --git a/device/awg/internal/mock.go b/device/awg/internal/mock.go
deleted file mode 100644
index a2e1c95..0000000
--- a/device/awg/internal/mock.go
+++ /dev/null
@@ -1,37 +0,0 @@
-package internal
-
-type mockGenerator struct {
- size int
-}
-
-func NewMockGenerator(size int) mockGenerator {
- return mockGenerator{size: size}
-}
-
-func (m mockGenerator) Generate() []byte {
- return make([]byte, m.size)
-}
-
-func (m mockGenerator) Size() int {
- return m.size
-}
-
-func (m mockGenerator) Name() string {
- return "mock"
-}
-
-type mockByteGenerator struct {
- data []byte
-}
-
-func NewMockByteGenerator(data []byte) mockByteGenerator {
- return mockByteGenerator{data: data}
-}
-
-func (bg mockByteGenerator) Generate() []byte {
- return bg.data
-}
-
-func (bg mockByteGenerator) Size() int {
- return len(bg.data)
-}
diff --git a/device/awg/junk_creator.go b/device/awg/junk_creator.go
deleted file mode 100644
index 8ba2918..0000000
--- a/device/awg/junk_creator.go
+++ /dev/null
@@ -1,50 +0,0 @@
-package awg
-
-import (
- "bytes"
- "fmt"
-)
-
-type JunkCreator struct {
- cfg Cfg
- randomGenerator PRNG[int]
-}
-
-// TODO: refactor param to only pass the junk related params
-func NewJunkCreator(cfg Cfg) JunkCreator {
- return JunkCreator{cfg: cfg, randomGenerator: NewPRNG[int]()}
-}
-
-// Should be called with awg mux RLocked
-func (jc *JunkCreator) CreateJunkPackets(junks *[][]byte) {
- if jc.cfg.JunkPacketCount == 0 {
- return
- }
-
- for range jc.cfg.JunkPacketCount {
- packetSize := jc.randomPacketSize()
- junk := jc.randomJunkWithSize(packetSize)
- *junks = append(*junks, junk)
- }
- return
-}
-
-// Should be called with awg mux RLocked
-func (jc *JunkCreator) randomPacketSize() int {
- return jc.randomGenerator.RandomSizeInRange(jc.cfg.JunkPacketMinSize, jc.cfg.JunkPacketMaxSize)
-}
-
-// Should be called with awg mux RLocked
-func (jc *JunkCreator) AppendJunk(writer *bytes.Buffer, size int) error {
- headerJunk := jc.randomJunkWithSize(size)
- _, err := writer.Write(headerJunk)
- if err != nil {
- return fmt.Errorf("write header junk: %v", err)
- }
- return nil
-}
-
-// Should be called with awg mux RLocked
-func (jc *JunkCreator) randomJunkWithSize(size int) []byte {
- return jc.randomGenerator.ReadSize(size)
-}
diff --git a/device/awg/junk_creator_test.go b/device/awg/junk_creator_test.go
deleted file mode 100644
index cdf752b..0000000
--- a/device/awg/junk_creator_test.go
+++ /dev/null
@@ -1,97 +0,0 @@
-package awg
-
-import (
- "bytes"
- "fmt"
- "testing"
-)
-
-func setUpJunkCreator() JunkCreator {
- mh, _ := NewMagicHeaders(
- []MagicHeader{
- NewMagicHeaderSameValue(123456),
- NewMagicHeaderSameValue(67543),
- NewMagicHeaderSameValue(32345),
- NewMagicHeaderSameValue(123123),
- },
- )
-
- jc := NewJunkCreator(Cfg{
- IsSet: true,
- JunkPacketCount: 5,
- JunkPacketMinSize: 500,
- JunkPacketMaxSize: 1000,
- InitHeaderJunkSize: 30,
- ResponseHeaderJunkSize: 40,
- MagicHeaders: mh,
- })
-
- return jc
-}
-
-func Test_junkCreator_createJunkPackets(t *testing.T) {
- jc := setUpJunkCreator()
- t.Run("valid", func(t *testing.T) {
- got := make([][]byte, 0, jc.cfg.JunkPacketCount)
- jc.CreateJunkPackets(&got)
- seen := make(map[string]bool)
- for _, junk := range got {
- key := string(junk)
- if seen[key] {
- t.Errorf(
- "junkCreator.createJunkPackets() = %v, duplicate key: %v",
- got,
- junk,
- )
- return
- }
- seen[key] = true
- }
- })
-}
-
-func Test_junkCreator_randomJunkWithSize(t *testing.T) {
- t.Run("valid", func(t *testing.T) {
- jc := setUpJunkCreator()
- r1 := jc.randomJunkWithSize(10)
- r2 := jc.randomJunkWithSize(10)
- fmt.Printf("%v\n%v\n", r1, r2)
- if bytes.Equal(r1, r2) {
- t.Errorf("same junks")
- return
- }
- })
-}
-
-func Test_junkCreator_randomPacketSize(t *testing.T) {
- jc := setUpJunkCreator()
- for range [30]struct{}{} {
- t.Run("valid", func(t *testing.T) {
- if got := jc.randomPacketSize(); jc.cfg.JunkPacketMinSize > got ||
- got > jc.cfg.JunkPacketMaxSize {
- t.Errorf(
- "junkCreator.randomPacketSize() = %v, not between range [%v,%v]",
- got,
- jc.cfg.JunkPacketMinSize,
- jc.cfg.JunkPacketMaxSize,
- )
- }
- })
- }
-}
-
-func Test_junkCreator_appendJunk(t *testing.T) {
- jc := setUpJunkCreator()
- t.Run("valid", func(t *testing.T) {
- s := "apple"
- buffer := bytes.NewBuffer([]byte(s))
- err := jc.AppendJunk(buffer, 30)
- if err != nil &&
- buffer.Len() != len(s)+30 {
- t.Error("appendWithJunk() size don't match")
- }
- read := make([]byte, 50)
- buffer.Read(read)
- fmt.Println(string(read))
- })
-}
diff --git a/device/awg/magic_header.go b/device/awg/magic_header.go
deleted file mode 100644
index aaf4e97..0000000
--- a/device/awg/magic_header.go
+++ /dev/null
@@ -1,97 +0,0 @@
-package awg
-
-import (
- "cmp"
- "fmt"
- "slices"
- "strconv"
- "strings"
-)
-
-type MagicHeader struct {
- Min uint32
- Max uint32
-}
-
-func NewMagicHeaderSameValue(value uint32) MagicHeader {
- return MagicHeader{Min: value, Max: value}
-}
-
-func NewMagicHeader(min, max uint32) (MagicHeader, error) {
- if min > max {
- return MagicHeader{}, fmt.Errorf("min (%d) cannot be greater than max (%d)", min, max)
- }
-
- return MagicHeader{Min: min, Max: max}, nil
-}
-
-func ParseMagicHeader(key, value string) (MagicHeader, error) {
- hyphenIdx := strings.Index(value, "-")
- if hyphenIdx == -1 {
- // if there is no hyphen, we treat it as single magic header value
- magicHeader, err := strconv.ParseUint(value, 10, 32)
- if err != nil {
- return MagicHeader{}, fmt.Errorf("parse key: %s; value: %s; %w", key, value, err)
- }
-
- return NewMagicHeader(uint32(magicHeader), uint32(magicHeader))
- }
-
- minStr := value[:hyphenIdx]
- maxStr := value[hyphenIdx+1:]
- if len(minStr) == 0 || len(maxStr) == 0 {
- return MagicHeader{}, fmt.Errorf("invalid value for key: %s; value: %s; expected format: min-max", key, value)
- }
-
- min, err := strconv.ParseUint(minStr, 10, 32)
- if err != nil {
- return MagicHeader{}, fmt.Errorf("parse min key: %s; value: %s; %w", key, minStr, err)
- }
-
- max, err := strconv.ParseUint(maxStr, 10, 32)
- if err != nil {
- return MagicHeader{}, fmt.Errorf("parse max key: %s; value: %s; %w", key, maxStr, err)
- }
-
- magicHeader, err := NewMagicHeader(uint32(min), uint32(max))
- if err != nil {
- return MagicHeader{}, fmt.Errorf("new magicHeader key: %s; value: %s-%s; %w", key, minStr, maxStr, err)
- }
-
- return magicHeader, nil
-}
-
-type MagicHeaders struct {
- Values []MagicHeader
- randomGenerator RandomNumberGenerator[uint32]
-}
-
-func NewMagicHeaders(headerValues []MagicHeader) (MagicHeaders, error) {
- if len(headerValues) != 4 {
- return MagicHeaders{}, fmt.Errorf("all header types should be included: %v", headerValues)
- }
-
- sortedMagicHeaders := slices.SortedFunc(slices.Values(headerValues), func(lhs MagicHeader, rhs MagicHeader) int {
- return cmp.Compare(lhs.Min, rhs.Min)
- })
-
- for i := range 3 {
- if sortedMagicHeaders[i].Max >= sortedMagicHeaders[i+1].Min {
- return MagicHeaders{}, fmt.Errorf(
- "magic headers shouldn't overlap; %v > %v",
- sortedMagicHeaders[i].Max,
- sortedMagicHeaders[i+1].Min,
- )
- }
- }
-
- return MagicHeaders{Values: headerValues, randomGenerator: NewPRNG[uint32]()}, nil
-}
-
-func (mh *MagicHeaders) Get(defaultMsgType uint32) (uint32, error) {
- if defaultMsgType == 0 || defaultMsgType > 4 {
- return 0, fmt.Errorf("invalid msg type: %d", defaultMsgType)
- }
-
- return mh.randomGenerator.RandomSizeInRange(mh.Values[defaultMsgType-1].Min, mh.Values[defaultMsgType-1].Max), nil
-}
diff --git a/device/awg/magic_header_test.go b/device/awg/magic_header_test.go
deleted file mode 100644
index 72a823e..0000000
--- a/device/awg/magic_header_test.go
+++ /dev/null
@@ -1,488 +0,0 @@
-package awg
-
-import (
- "testing"
-
- "github.com/stretchr/testify/require"
-)
-
-func TestNewMagicHeaderSameValue(t *testing.T) {
- tests := []struct {
- name string
- value uint32
- expected MagicHeader
- }{
- {
- name: "zero value",
- value: 0,
- expected: MagicHeader{Min: 0, Max: 0},
- },
- {
- name: "small value",
- value: 1,
- expected: MagicHeader{Min: 1, Max: 1},
- },
- {
- name: "large value",
- value: 4294967295, // max uint32
- expected: MagicHeader{Min: 4294967295, Max: 4294967295},
- },
- {
- name: "medium value",
- value: 1000,
- expected: MagicHeader{Min: 1000, Max: 1000},
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- result := NewMagicHeaderSameValue(tt.value)
- require.Equal(t, tt.expected, result)
- })
- }
-}
-
-func TestNewMagicHeader(t *testing.T) {
- tests := []struct {
- name string
- min uint32
- max uint32
- expected MagicHeader
- errorMsg string
- }{
- {
- name: "valid range",
- min: 1,
- max: 10,
- expected: MagicHeader{Min: 1, Max: 10},
- },
- {
- name: "equal values",
- min: 5,
- max: 5,
- expected: MagicHeader{Min: 5, Max: 5},
- },
- {
- name: "zero range",
- min: 0,
- max: 0,
- expected: MagicHeader{Min: 0, Max: 0},
- },
- {
- name: "max uint32 range",
- min: 4294967294,
- max: 4294967295,
- expected: MagicHeader{Min: 4294967294, Max: 4294967295},
- },
- {
- name: "min greater than max",
- min: 10,
- max: 5,
- expected: MagicHeader{},
- errorMsg: "min (10) cannot be greater than max (5)",
- },
- {
- name: "large min greater than max",
- min: 4294967295,
- max: 1,
- expected: MagicHeader{},
- errorMsg: "min (4294967295) cannot be greater than max (1)",
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- result, err := NewMagicHeader(tt.min, tt.max)
-
- if tt.errorMsg != "" {
- require.Error(t, err)
- require.Contains(t, err.Error(), tt.errorMsg)
- require.Equal(t, MagicHeader{}, result)
- } else {
- require.NoError(t, err)
- require.Equal(t, tt.expected, result)
- }
- })
- }
-}
-
-func TestParseMagicHeader(t *testing.T) {
- tests := []struct {
- name string
- key string
- value string
- expected MagicHeader
- errorMsg string
- }{
- {
- name: "single value",
- key: "header1",
- value: "100",
- expected: MagicHeader{Min: 100, Max: 100},
- },
- {
- name: "valid range",
- key: "header2",
- value: "10-20",
- expected: MagicHeader{Min: 10, Max: 20},
- },
- {
- name: "zero single value",
- key: "header3",
- value: "0",
- expected: MagicHeader{Min: 0, Max: 0},
- },
- {
- name: "zero range",
- key: "header4",
- value: "0-0",
- expected: MagicHeader{Min: 0, Max: 0},
- },
- {
- name: "max uint32 single",
- key: "header5",
- value: "4294967295",
- expected: MagicHeader{Min: 4294967295, Max: 4294967295},
- },
- {
- name: "max uint32 range",
- key: "header6",
- value: "4294967294-4294967295",
- expected: MagicHeader{Min: 4294967294, Max: 4294967295},
- },
- {
- name: "invalid single value - not number",
- key: "header7",
- value: "abc",
- expected: MagicHeader{},
- errorMsg: "parse key: header7; value: abc;",
- },
- {
- name: "invalid single value - negative",
- key: "header8",
- value: "-5",
- expected: MagicHeader{},
- errorMsg: "invalid value for key: header8; value: -5;",
- },
- {
- name: "invalid single value - too large",
- key: "header9",
- value: "4294967296",
- expected: MagicHeader{},
- errorMsg: "parse key: header9; value: 4294967296;",
- },
- {
- name: "invalid range - min not number",
- key: "header10",
- value: "abc-10",
- expected: MagicHeader{},
- errorMsg: "parse min key: header10; value: abc;",
- },
- {
- name: "invalid range - max not number",
- key: "header11",
- value: "10-abc",
- expected: MagicHeader{},
- errorMsg: "parse max key: header11; value: abc;",
- },
- {
- name: "invalid range - min greater than max",
- key: "header12",
- value: "20-10",
- expected: MagicHeader{},
- errorMsg: "new magicHeader key: header12; value: 20-10;",
- },
- {
- name: "invalid range - too many parts",
- key: "header13",
- value: "10-20-30",
- expected: MagicHeader{},
- errorMsg: "parse key: header13; value: 10-20-30;",
- },
- {
- name: "empty value",
- key: "header14",
- value: "",
- expected: MagicHeader{},
- errorMsg: "parse key: header14; value: ;",
- },
- {
- name: "hyphen only",
- key: "header15",
- value: "-",
- expected: MagicHeader{},
- errorMsg: "invalid value for key: header15; value: -;",
- },
- {
- name: "empty min",
- key: "header16",
- value: "-10",
- expected: MagicHeader{},
- errorMsg: "invalid value for key: header16; value: -10;",
- },
- {
- name: "empty max",
- key: "header17",
- value: "10-",
- expected: MagicHeader{},
- errorMsg: "invalid value for key: header17; value: 10-;",
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- result, err := ParseMagicHeader(tt.key, tt.value)
-
- if tt.errorMsg != "" {
- require.Error(t, err)
- require.Contains(t, err.Error(), tt.errorMsg)
- require.Equal(t, MagicHeader{}, result)
- } else {
- require.NoError(t, err)
- require.Equal(t, tt.expected, result)
- }
- })
- }
-}
-
-func TestNewMagicHeaders(t *testing.T) {
- tests := []struct {
- name string
- magicHeaders []MagicHeader
- errorMsg string
- }{
- {
- name: "valid non-overlapping headers",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 10},
- {Min: 11, Max: 20},
- {Min: 21, Max: 30},
- {Min: 31, Max: 40},
- },
- },
- {
- name: "valid adjacent headers",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 1},
- {Min: 2, Max: 2},
- {Min: 3, Max: 3},
- {Min: 4, Max: 4},
- },
- },
- {
- name: "valid zero-based headers",
- magicHeaders: []MagicHeader{
- {Min: 0, Max: 0},
- {Min: 1, Max: 1},
- {Min: 2, Max: 2},
- {Min: 3, Max: 3},
- },
- },
- {
- name: "valid large value headers",
- magicHeaders: []MagicHeader{
- {Min: 4294967290, Max: 4294967291},
- {Min: 4294967292, Max: 4294967293},
- {Min: 4294967294, Max: 4294967294},
- {Min: 4294967295, Max: 4294967295},
- },
- },
- {
- name: "too few headers",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 10},
- {Min: 11, Max: 20},
- {Min: 21, Max: 30},
- },
- errorMsg: "all header types should be included:",
- },
- {
- name: "too many headers",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 10},
- {Min: 11, Max: 20},
- {Min: 21, Max: 30},
- {Min: 31, Max: 40},
- {Min: 41, Max: 50},
- },
- errorMsg: "all header types should be included:",
- },
- {
- name: "empty headers",
- magicHeaders: []MagicHeader{},
- errorMsg: "all header types should be included:",
- },
- {
- name: "overlapping headers",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 15},
- {Min: 10, Max: 20},
- {Min: 25, Max: 30},
- {Min: 35, Max: 40},
- },
- errorMsg: "magic headers shouldn't overlap;",
- },
- {
- name: "overlapping headers at limit-first",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 10},
- {Min: 10, Max: 20},
- {Min: 25, Max: 30},
- {Min: 35, Max: 40},
- },
- errorMsg: "magic headers shouldn't overlap;",
- },
- {
- name: "overlapping headers at limit-second",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 10},
- {Min: 15, Max: 25},
- {Min: 25, Max: 30},
- {Min: 35, Max: 40},
- },
- errorMsg: "magic headers shouldn't overlap;",
- },
- {
- name: "overlapping headers at limit-third",
- magicHeaders: []MagicHeader{
- {Min: 1, Max: 10},
- {Min: 15, Max: 25},
- {Min: 30, Max: 35},
- {Min: 35, Max: 40},
- },
- errorMsg: "magic headers shouldn't overlap;",
- },
- {
- name: "identical ranges",
- magicHeaders: []MagicHeader{
- {Min: 10, Max: 20},
- {Min: 10, Max: 20},
- {Min: 25, Max: 30},
- {Min: 35, Max: 40},
- },
- errorMsg: "magic headers shouldn't overlap;",
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- result, err := NewMagicHeaders(tt.magicHeaders)
-
- if tt.errorMsg != "" {
- require.Error(t, err)
- require.Contains(t, err.Error(), tt.errorMsg)
- require.Equal(t, MagicHeaders{}, result)
- } else {
- require.NoError(t, err)
- require.Equal(t, tt.magicHeaders, result.Values)
- require.NotNil(t, result.randomGenerator)
- }
- })
- }
-}
-
-// Mock PRNG for testing
-type mockPRNG struct {
- returnValue uint32
-}
-
-func (m *mockPRNG) RandomSizeInRange(min, max uint32) uint32 {
- return m.returnValue
-}
-
-func (m *mockPRNG) Get() uint64 {
- return 0
-}
-func (m *mockPRNG) ReadSize(size int) []byte {
- return make([]byte, size)
-}
-
-func TestMagicHeaders_Get(t *testing.T) {
- // Create test headers
- headers := []MagicHeader{
- {Min: 1, Max: 10},
- {Min: 11, Max: 20},
- {Min: 21, Max: 30},
- {Min: 31, Max: 40},
- }
-
- tests := []struct {
- name string
- defaultMsgType uint32
- mockValue uint32
- expectedValue uint32
- errorMsg string
- }{
- {
- name: "valid type 1",
- defaultMsgType: 1,
- mockValue: 5,
- expectedValue: 5,
- },
- {
- name: "valid type 2",
- defaultMsgType: 2,
- mockValue: 15,
- expectedValue: 15,
- },
- {
- name: "valid type 3",
- defaultMsgType: 3,
- mockValue: 25,
- expectedValue: 25,
- },
- {
- name: "valid type 4",
- defaultMsgType: 4,
- mockValue: 35,
- expectedValue: 35,
- },
- {
- name: "invalid type 0",
- defaultMsgType: 0,
- mockValue: 0,
- expectedValue: 0,
- errorMsg: "invalid msg type: 0",
- },
- {
- name: "invalid type 5",
- defaultMsgType: 5,
- mockValue: 0,
- expectedValue: 0,
- errorMsg: "invalid msg type: 5",
- },
- {
- name: "invalid type max uint32",
- defaultMsgType: 4294967295,
- mockValue: 0,
- expectedValue: 0,
- errorMsg: "invalid msg type: 4294967295",
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- // Create a new instance with mock PRNG for each test
- testMagicHeaders := MagicHeaders{
- Values: headers,
- randomGenerator: &mockPRNG{returnValue: tt.mockValue},
- }
-
- result, err := testMagicHeaders.Get(tt.defaultMsgType)
-
- if tt.errorMsg != "" {
- require.Error(t, err)
- require.Contains(t, err.Error(), tt.errorMsg)
- require.Equal(t, uint32(0), result)
- } else {
- require.NoError(t, err)
- require.Equal(t, tt.expectedValue, result)
- }
- })
- }
-}
diff --git a/device/awg/prng.go b/device/awg/prng.go
deleted file mode 100644
index e7661d7..0000000
--- a/device/awg/prng.go
+++ /dev/null
@@ -1,50 +0,0 @@
-package awg
-
-import (
- crand "crypto/rand"
- v2 "math/rand/v2"
-
- "golang.org/x/exp/constraints"
-)
-
-type RandomNumberGenerator[T constraints.Integer] interface {
- RandomSizeInRange(min, max T) T
- Get() uint64
- ReadSize(size int) []byte
-}
-
-type PRNG[T constraints.Integer] struct {
- cha8Rand *v2.ChaCha8
-}
-
-func NewPRNG[T constraints.Integer]() PRNG[T] {
- buf := make([]byte, 32)
- _, _ = crand.Read(buf)
-
- return PRNG[T]{
- cha8Rand: v2.NewChaCha8([32]byte(buf)),
- }
-}
-
-func (p PRNG[T]) RandomSizeInRange(min, max T) T {
- if min > max {
- panic("min must be less than max")
- }
-
- if min == max {
- return min
- }
-
- return T(p.Get()%uint64(max-min)) + min
-}
-
-func (p PRNG[T]) Get() uint64 {
- return p.cha8Rand.Uint64()
-}
-
-func (p PRNG[T]) ReadSize(size int) []byte {
- // TODO: use a memory pool to allocate
- buf := make([]byte, size)
- _, _ = p.cha8Rand.Read(buf)
- return buf
-}
diff --git a/device/awg/special_handshake_handler.go b/device/awg/special_handshake_handler.go
deleted file mode 100644
index d740879..0000000
--- a/device/awg/special_handshake_handler.go
+++ /dev/null
@@ -1,36 +0,0 @@
-package awg
-
-import (
- "github.com/tevino/abool"
- "go.uber.org/atomic"
-)
-
-// TODO: atomic?/ and better way to use this
-var PacketCounter *atomic.Uint64 = atomic.NewUint64(0)
-
-// TODO
-var WaitResponse = struct {
- Channel chan struct{}
- ShouldWait *abool.AtomicBool
-}{
- make(chan struct{}, 1),
- abool.New(),
-}
-
-type SpecialHandshakeHandler struct {
- SpecialJunk TagJunkPacketGenerators
-
- IsSet bool
-}
-
-func (handler *SpecialHandshakeHandler) Validate() error {
- return handler.SpecialJunk.Validate()
-}
-
-func (handler *SpecialHandshakeHandler) GenerateSpecialJunk() [][]byte {
- if !handler.SpecialJunk.IsDefined() {
- return nil
- }
-
- return handler.SpecialJunk.GeneratePackets()
-}
diff --git a/device/awg/tag_generator.go b/device/awg/tag_generator.go
deleted file mode 100644
index 3a1d497..0000000
--- a/device/awg/tag_generator.go
+++ /dev/null
@@ -1,229 +0,0 @@
-package awg
-
-import (
- crand "crypto/rand"
- "encoding/binary"
- "encoding/hex"
- "fmt"
- "strconv"
- "strings"
- "time"
-
- v2 "math/rand/v2"
- // "go.uber.org/atomic"
-)
-
-type Generator interface {
- Generate() []byte
- Size() int
-}
-
-type newGenerator func(string) (Generator, error)
-
-type BytesGenerator struct {
- value []byte
- size int
-}
-
-func (bg *BytesGenerator) Generate() []byte {
- return bg.value
-}
-
-func (bg *BytesGenerator) Size() int {
- return bg.size
-}
-
-func newBytesGenerator(param string) (Generator, error) {
- hasPrefix := strings.HasPrefix(param, "0x") || strings.HasPrefix(param, "0X")
- if !hasPrefix {
- return nil, fmt.Errorf("not correct hex: %s", param)
- }
-
- hex, err := hexToBytes(param)
- if err != nil {
- return nil, fmt.Errorf("hexToBytes: %w", err)
- }
-
- return &BytesGenerator{value: hex, size: len(hex)}, nil
-}
-
-func hexToBytes(hexStr string) ([]byte, error) {
- hexStr = strings.TrimPrefix(hexStr, "0x")
- hexStr = strings.TrimPrefix(hexStr, "0X")
-
- // Ensure even length (pad with leading zero if needed)
- if len(hexStr)%2 != 0 {
- hexStr = "0" + hexStr
- }
-
- return hex.DecodeString(hexStr)
-}
-
-type randomGeneratorBase struct {
- cha8Rand *v2.ChaCha8
- size int
-}
-
-func newRandomGeneratorBase(param string) (*randomGeneratorBase, error) {
- size, err := strconv.Atoi(param)
- if err != nil {
- return nil, fmt.Errorf("parse int: %w", err)
- }
-
- if size > 1000 {
- return nil, fmt.Errorf("size must be less than 1000")
- }
-
- buf := make([]byte, 32)
- _, err = crand.Read(buf)
- if err != nil {
- return nil, fmt.Errorf("crand read: %w", err)
- }
-
- return &randomGeneratorBase{
- cha8Rand: v2.NewChaCha8([32]byte(buf)),
- size: size,
- }, nil
-}
-
-func (rpg *randomGeneratorBase) generate() []byte {
- junk := make([]byte, rpg.size)
- rpg.cha8Rand.Read(junk)
- return junk
-}
-
-func (rpg *randomGeneratorBase) Size() int {
- return rpg.size
-}
-
-type RandomBytesGenerator struct {
- *randomGeneratorBase
-}
-
-func newRandomBytesGenerator(param string) (Generator, error) {
- rpgBase, err := newRandomGeneratorBase(param)
- if err != nil {
- return nil, fmt.Errorf("new random bytes generator: %w", err)
- }
-
- return &RandomBytesGenerator{randomGeneratorBase: rpgBase}, nil
-}
-
-func (rpg *RandomBytesGenerator) Generate() []byte {
- return rpg.generate()
-}
-
-const alphanumericChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
-
-type RandomASCIIGenerator struct {
- *randomGeneratorBase
-}
-
-func newRandomASCIIGenerator(param string) (Generator, error) {
- rpgBase, err := newRandomGeneratorBase(param)
- if err != nil {
- return nil, fmt.Errorf("new random ascii generator: %w", err)
- }
-
- return &RandomASCIIGenerator{randomGeneratorBase: rpgBase}, nil
-}
-
-func (rpg *RandomASCIIGenerator) Generate() []byte {
- junk := rpg.generate()
-
- result := make([]byte, rpg.size)
- for i, b := range junk {
- result[i] = alphanumericChars[b%byte(len(alphanumericChars))]
- }
-
- return result
-}
-
-type RandomDigitGenerator struct {
- *randomGeneratorBase
-}
-
-func newRandomDigitGenerator(param string) (Generator, error) {
- rpgBase, err := newRandomGeneratorBase(param)
- if err != nil {
- return nil, fmt.Errorf("new random digit generator: %w", err)
- }
-
- return &RandomDigitGenerator{randomGeneratorBase: rpgBase}, nil
-}
-
-func (rpg *RandomDigitGenerator) Generate() []byte {
- junk := rpg.generate()
-
- result := make([]byte, rpg.size)
- for i, b := range junk {
- result[i] = '0' + (b % 10) // Convert to digit character
- }
-
- return result
-}
-
-type TimestampGenerator struct {
-}
-
-func (tg *TimestampGenerator) Generate() []byte {
- buf := make([]byte, 8)
- binary.BigEndian.PutUint64(buf, uint64(time.Now().Unix()))
- return buf
-}
-
-func (tg *TimestampGenerator) Size() int {
- return 8
-}
-
-func newTimestampGenerator(param string) (Generator, error) {
- if len(param) != 0 {
- return nil, fmt.Errorf("timestamp param needs to be empty: %s", param)
- }
-
- return &TimestampGenerator{}, nil
-}
-
-type PacketCounterGenerator struct {
-}
-
-func (c *PacketCounterGenerator) Generate() []byte {
- buf := make([]byte, 8)
- // TODO: better way to handle counter tag
- binary.BigEndian.PutUint64(buf, PacketCounter.Load())
- return buf
-}
-
-func (c *PacketCounterGenerator) Size() int {
- return 8
-}
-
-func newPacketCounterGenerator(param string) (Generator, error) {
- if len(param) != 0 {
- return nil, fmt.Errorf("packet counter param needs to be empty: %s", param)
- }
-
- return &PacketCounterGenerator{}, nil
-}
-
-type WaitResponseGenerator struct {
-}
-
-func (c *WaitResponseGenerator) Generate() []byte {
- WaitResponse.ShouldWait.Set()
- <-WaitResponse.Channel
- WaitResponse.ShouldWait.UnSet()
- return []byte{}
-}
-
-func (c *WaitResponseGenerator) Size() int {
- return 0
-}
-
-func newWaitResponseGenerator(param string) (Generator, error) {
- if len(param) != 0 {
- return nil, fmt.Errorf("wait response param needs to be empty: %s", param)
- }
-
- return &WaitResponseGenerator{}, nil
-}
diff --git a/device/awg/tag_generator_test.go b/device/awg/tag_generator_test.go
deleted file mode 100644
index 43efa67..0000000
--- a/device/awg/tag_generator_test.go
+++ /dev/null
@@ -1,321 +0,0 @@
-package awg
-
-import (
- "encoding/binary"
- "fmt"
- "testing"
-
- "github.com/stretchr/testify/require"
-)
-
-func TestNewBytesGenerator(t *testing.T) {
- t.Parallel()
-
- type args struct {
- param string
- }
- tests := []struct {
- name string
- args args
- want []byte
- wantErr error
- }{
- {
- name: "empty",
- args: args{
- param: "",
- },
- wantErr: fmt.Errorf("not correct hex"),
- },
- {
- name: "wrong start",
- args: args{
- param: "123456",
- },
- wantErr: fmt.Errorf("not correct hex"),
- },
- {
- name: "not only hex value with X",
- args: args{
- param: "0X12345q",
- },
- wantErr: fmt.Errorf("not correct hex"),
- },
- {
- name: "not only hex value with x",
- args: args{
- param: "0x12345q",
- },
- wantErr: fmt.Errorf("not correct hex"),
- },
- {
- name: "valid hex",
- args: args{
- param: "0xf6ab3267fa",
- },
- want: []byte{0xf6, 0xab, 0x32, 0x67, 0xfa},
- },
- {
- name: "valid hex with odd length",
- args: args{
- param: "0xfab3267fa",
- },
- want: []byte{0xf, 0xab, 0x32, 0x67, 0xfa},
- },
- }
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
-
- got, err := newBytesGenerator(tt.args.param)
-
- if tt.wantErr != nil {
- require.ErrorAs(t, err, &tt.wantErr)
- require.Nil(t, got)
- return
- }
-
- require.Nil(t, err)
- require.NotNil(t, got)
-
- gotValues := got.Generate()
- require.Equal(t, tt.want, gotValues)
- })
- }
-}
-
-func TestNewRandomBytesGenerator(t *testing.T) {
- t.Parallel()
-
- type args struct {
- param string
- }
- tests := []struct {
- name string
- args args
- wantErr error
- }{
- {
- name: "empty",
- args: args{
- param: "",
- },
- wantErr: fmt.Errorf("parse int"),
- },
- {
- name: "not an int",
- args: args{
- param: "x",
- },
- wantErr: fmt.Errorf("parse int"),
- },
- {
- name: "too large",
- args: args{
- param: "1001",
- },
- wantErr: fmt.Errorf("random packet size must be less than 1000"),
- },
- {
- name: "valid",
- args: args{
- param: "12",
- },
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
-
- got, err := newRandomBytesGenerator(tt.args.param)
- if tt.wantErr != nil {
- require.ErrorAs(t, err, &tt.wantErr)
- require.Nil(t, got)
- return
- }
-
- require.Nil(t, err)
- require.NotNil(t, got)
- first := got.Generate()
-
- second := got.Generate()
- require.NotEqual(t, first, second)
- })
- }
-}
-
-func TestNewRandomASCIIGenerator(t *testing.T) {
- t.Parallel()
-
- type args struct {
- param string
- }
- tests := []struct {
- name string
- args args
- wantErr error
- }{
- {
- name: "empty",
- args: args{
- param: "",
- },
- wantErr: fmt.Errorf("parse int"),
- },
- {
- name: "not an int",
- args: args{
- param: "x",
- },
- wantErr: fmt.Errorf("parse int"),
- },
- {
- name: "too large",
- args: args{
- param: "1001",
- },
- wantErr: fmt.Errorf("random packet size must be less than 1000"),
- },
- {
- name: "valid",
- args: args{
- param: "12",
- },
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
-
- got, err := newRandomASCIIGenerator(tt.args.param)
- if tt.wantErr != nil {
- require.ErrorAs(t, err, &tt.wantErr)
- require.Nil(t, got)
- return
- }
-
- require.Nil(t, err)
- require.NotNil(t, got)
- first := got.Generate()
-
- second := got.Generate()
- require.NotEqual(t, first, second)
- })
- }
-}
-
-func TestNewRandomDigitGenerator(t *testing.T) {
- t.Parallel()
-
- type args struct {
- param string
- }
- tests := []struct {
- name string
- args args
- wantErr error
- }{
- {
- name: "empty",
- args: args{
- param: "",
- },
- wantErr: fmt.Errorf("parse int"),
- },
- {
- name: "not an int",
- args: args{
- param: "x",
- },
- wantErr: fmt.Errorf("parse int"),
- },
- {
- name: "too large",
- args: args{
- param: "1001",
- },
- wantErr: fmt.Errorf("random packet size must be less than 1000"),
- },
- {
- name: "valid",
- args: args{
- param: "12",
- },
- },
- }
-
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
-
- got, err := newRandomDigitGenerator(tt.args.param)
- if tt.wantErr != nil {
- require.ErrorAs(t, err, &tt.wantErr)
- require.Nil(t, got)
- return
- }
-
- require.Nil(t, err)
- require.NotNil(t, got)
- first := got.Generate()
-
- second := got.Generate()
- require.NotEqual(t, first, second)
- })
- }
-}
-
-func TestPacketCounterGenerator(t *testing.T) {
- t.Parallel()
-
- tests := []struct {
- name string
- param string
- wantErr bool
- }{
- {
- name: "Valid empty param",
- param: "",
- wantErr: false,
- },
- {
- name: "Invalid non-empty param",
- param: "anything",
- wantErr: true,
- },
- }
-
- for _, tc := range tests {
- t.Run(tc.name, func(t *testing.T) {
- t.Parallel()
-
- gen, err := newPacketCounterGenerator(tc.param)
- if tc.wantErr {
- require.Error(t, err)
- return
- }
-
- require.NoError(t, err)
- require.Equal(t, 8, gen.Size())
-
- // Reset counter to known value for test
- initialCount := uint64(42)
- PacketCounter.Store(initialCount)
-
- output := gen.Generate()
- require.Equal(t, 8, len(output))
-
- // Verify counter value in output
- counterValue := binary.BigEndian.Uint64(output)
- require.Equal(t, initialCount, counterValue)
-
- // Increment counter and verify change
- PacketCounter.Add(1)
- output = gen.Generate()
- counterValue = binary.BigEndian.Uint64(output)
- require.Equal(t, initialCount+1, counterValue)
- })
- }
-}
diff --git a/device/awg/tag_junk_packet_generator.go b/device/awg/tag_junk_packet_generator.go
deleted file mode 100644
index fdbebc8..0000000
--- a/device/awg/tag_junk_packet_generator.go
+++ /dev/null
@@ -1,59 +0,0 @@
-package awg
-
-import (
- "fmt"
- "strconv"
-)
-
-type TagJunkPacketGenerator struct {
- name string
- tagValue string
-
- packetSize int
- generators []Generator
-}
-
-func newTagJunkPacketGenerator(name, tagValue string, size int) TagJunkPacketGenerator {
- return TagJunkPacketGenerator{
- name: name,
- tagValue: tagValue,
- generators: make([]Generator, 0, size),
- }
-}
-
-func (tg *TagJunkPacketGenerator) append(generator Generator) {
- tg.generators = append(tg.generators, generator)
- tg.packetSize += generator.Size()
-}
-
-func (tg *TagJunkPacketGenerator) generatePacket() []byte {
- packet := make([]byte, 0, tg.packetSize)
- for _, generator := range tg.generators {
- packet = append(packet, generator.Generate()...)
- }
-
- return packet
-}
-
-func (tg *TagJunkPacketGenerator) Name() string {
- return tg.name
-}
-
-func (tg *TagJunkPacketGenerator) nameIndex() (int, error) {
- if len(tg.name) != 2 {
- return 0, fmt.Errorf("name must be 2 character long: %s", tg.name)
- }
-
- index, err := strconv.Atoi(tg.name[1:2])
- if err != nil {
- return 0, fmt.Errorf("name 2 char should be an int %w", err)
- }
- return index, nil
-}
-
-func (tg *TagJunkPacketGenerator) IpcGetFields() IpcFields {
- return IpcFields{
- Key: tg.name,
- Value: tg.tagValue,
- }
-}
diff --git a/device/awg/tag_junk_packet_generator_test.go b/device/awg/tag_junk_packet_generator_test.go
deleted file mode 100644
index 309d425..0000000
--- a/device/awg/tag_junk_packet_generator_test.go
+++ /dev/null
@@ -1,210 +0,0 @@
-package awg
-
-import (
- "testing"
-
- "github.com/amnezia-vpn/amneziawg-go/device/awg/internal"
- "github.com/stretchr/testify/require"
-)
-
-func TestNewTagJunkGenerator(t *testing.T) {
- t.Parallel()
-
- testCases := []struct {
- name string
- genName string
- size int
- expected TagJunkPacketGenerator
- }{
- {
- name: "Create new generator with empty name",
- genName: "",
- size: 0,
- expected: TagJunkPacketGenerator{
- name: "",
- packetSize: 0,
- generators: make([]Generator, 0),
- },
- },
- {
- name: "Create new generator with valid name",
- genName: "T1",
- size: 0,
- expected: TagJunkPacketGenerator{
- name: "T1",
- packetSize: 0,
- generators: make([]Generator, 0),
- },
- },
- {
- name: "Create new generator with non-zero size",
- genName: "T2",
- size: 5,
- expected: TagJunkPacketGenerator{
- name: "T2",
- packetSize: 0,
- generators: make([]Generator, 5),
- },
- },
- }
-
- for _, tc := range testCases {
- tc := tc // capture range variable
- t.Run(tc.name, func(t *testing.T) {
- t.Parallel()
- result := newTagJunkPacketGenerator(tc.genName, "", tc.size)
- require.Equal(t, tc.expected.name, result.name)
- require.Equal(t, tc.expected.packetSize, result.packetSize)
- require.Equal(t, cap(result.generators), len(tc.expected.generators))
- })
- }
-}
-
-func TestTagJunkGeneratorAppend(t *testing.T) {
- t.Parallel()
-
- testCases := []struct {
- name string
- initialState TagJunkPacketGenerator
- mockSize int
- expectedLength int
- expectedSize int
- }{
- {
- name: "Append to empty generator",
- initialState: newTagJunkPacketGenerator("T1", "", 0),
- mockSize: 5,
- expectedLength: 1,
- expectedSize: 5,
- },
- {
- name: "Append to non-empty generator",
- initialState: TagJunkPacketGenerator{
- name: "T2",
- packetSize: 10,
- generators: make([]Generator, 2),
- },
- mockSize: 7,
- expectedLength: 3, // 2 existing + 1 new
- expectedSize: 17, // 10 + 7
- },
- }
-
- for _, tc := range testCases {
- tc := tc // capture range variable
- t.Run(tc.name, func(t *testing.T) {
- t.Parallel()
-
- tg := tc.initialState
- mockGen := internal.NewMockGenerator(tc.mockSize)
-
- tg.append(mockGen)
-
- require.Equal(t, tc.expectedLength, len(tg.generators))
- require.Equal(t, tc.expectedSize, tg.packetSize)
- })
- }
-}
-
-func TestTagJunkGeneratorGenerate(t *testing.T) {
- t.Parallel()
-
- // Create mock generators for testing
- mockGen1 := internal.NewMockByteGenerator([]byte{0x01, 0x02})
- mockGen2 := internal.NewMockByteGenerator([]byte{0x03, 0x04, 0x05})
-
- testCases := []struct {
- name string
- setupGenerator func() TagJunkPacketGenerator
- expected []byte
- }{
- {
- name: "Generate with empty generators",
- setupGenerator: func() TagJunkPacketGenerator {
- return newTagJunkPacketGenerator("T1", "", 0)
- },
- expected: []byte{},
- },
- {
- name: "Generate with single generator",
- setupGenerator: func() TagJunkPacketGenerator {
- tg := newTagJunkPacketGenerator("T2", "", 0)
- tg.append(mockGen1)
- return tg
- },
- expected: []byte{0x01, 0x02},
- },
- {
- name: "Generate with multiple generators",
- setupGenerator: func() TagJunkPacketGenerator {
- tg := newTagJunkPacketGenerator("T3", "", 0)
- tg.append(mockGen1)
- tg.append(mockGen2)
- return tg
- },
- expected: []byte{0x01, 0x02, 0x03, 0x04, 0x05},
- },
- }
-
- for _, tc := range testCases {
- tc := tc // capture range variable
- t.Run(tc.name, func(t *testing.T) {
- t.Parallel()
-
- tg := tc.setupGenerator()
- result := tg.generatePacket()
-
- require.Equal(t, tc.expected, result)
- })
- }
-}
-
-func TestTagJunkGeneratorNameIndex(t *testing.T) {
- t.Parallel()
-
- testCases := []struct {
- name string
- generatorName string
- expectedIndex int
- expectError bool
- }{
- {
- name: "Valid name with digit",
- generatorName: "T5",
- expectedIndex: 5,
- expectError: false,
- },
- {
- name: "Invalid name - too short",
- generatorName: "T",
- expectError: true,
- },
- {
- name: "Invalid name - too long",
- generatorName: "T55",
- expectError: true,
- },
- {
- name: "Invalid name - non-digit second character",
- generatorName: "TX",
- expectError: true,
- },
- }
-
- for _, tc := range testCases {
- tc := tc // capture range variable
- t.Run(tc.name, func(t *testing.T) {
- t.Parallel()
-
- tg := TagJunkPacketGenerator{name: tc.generatorName}
- index, err := tg.nameIndex()
-
- if tc.expectError {
- require.Error(t, err)
- } else {
- require.NoError(t, err)
- require.Equal(t, tc.expectedIndex, index)
- }
- })
- }
-}
diff --git a/device/awg/tag_junk_packet_generators.go b/device/awg/tag_junk_packet_generators.go
deleted file mode 100644
index 9921eb0..0000000
--- a/device/awg/tag_junk_packet_generators.go
+++ /dev/null
@@ -1,66 +0,0 @@
-package awg
-
-import "fmt"
-
-type TagJunkPacketGenerators struct {
- tagGenerators []TagJunkPacketGenerator
- length int
- DefaultJunkCount int // Jc
-}
-
-func (generators *TagJunkPacketGenerators) AppendGenerator(
- generator TagJunkPacketGenerator,
-) {
- generators.tagGenerators = append(generators.tagGenerators, generator)
- generators.length++
-}
-
-func (generators *TagJunkPacketGenerators) IsDefined() bool {
- return len(generators.tagGenerators) > 0
-}
-
-// validate that packets were defined consecutively
-func (generators *TagJunkPacketGenerators) Validate() error {
- seen := make([]bool, len(generators.tagGenerators))
- for _, generator := range generators.tagGenerators {
- index, err := generator.nameIndex()
- if index > len(generators.tagGenerators) {
- return fmt.Errorf("junk packet index should be consecutive")
- }
- if err != nil {
- return fmt.Errorf("name index: %w", err)
- } else {
- seen[index-1] = true
- }
- }
-
- for _, found := range seen {
- if !found {
- return fmt.Errorf("junk packet index should be consecutive")
- }
- }
-
- return nil
-}
-
-func (generators *TagJunkPacketGenerators) GeneratePackets() [][]byte {
- var rv = make([][]byte, 0, generators.length+generators.DefaultJunkCount)
-
- for i, tagGenerator := range generators.tagGenerators {
- rv = append(rv, make([]byte, tagGenerator.packetSize))
- copy(rv[i], tagGenerator.generatePacket())
- PacketCounter.Inc()
- }
- PacketCounter.Add(uint64(generators.DefaultJunkCount))
-
- return rv
-}
-
-func (tg *TagJunkPacketGenerators) IpcGetFields() []IpcFields {
- rv := make([]IpcFields, 0, len(tg.tagGenerators))
- for _, generator := range tg.tagGenerators {
- rv = append(rv, generator.IpcGetFields())
- }
-
- return rv
-}
diff --git a/device/awg/tag_junk_packet_generators_test.go b/device/awg/tag_junk_packet_generators_test.go
deleted file mode 100644
index 6b1fd47..0000000
--- a/device/awg/tag_junk_packet_generators_test.go
+++ /dev/null
@@ -1,149 +0,0 @@
-package awg
-
-import (
- "testing"
-
- "github.com/amnezia-vpn/amneziawg-go/device/awg/internal"
- "github.com/stretchr/testify/require"
-)
-
-func TestTagJunkGeneratorHandlerAppendGenerator(t *testing.T) {
- tests := []struct {
- name string
- generator TagJunkPacketGenerator
- }{
- {
- name: "append single generator",
- generator: newTagJunkPacketGenerator("t1", "", 10),
- },
- }
-
- for _, tt := range tests {
- tt := tt
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- generators := &TagJunkPacketGenerators{}
-
- // Initial length should be 0
- require.Equal(t, 0, generators.length)
- require.Empty(t, generators.tagGenerators)
-
- // After append, length should be 1 and generator should be added
- generators.AppendGenerator(tt.generator)
- require.Equal(t, 1, generators.length)
- require.Len(t, generators.tagGenerators, 1)
- require.Equal(t, tt.generator, generators.tagGenerators[0])
- })
- }
-}
-
-func TestTagJunkGeneratorHandlerValidate(t *testing.T) {
- tests := []struct {
- name string
- generators []TagJunkPacketGenerator
- wantErr bool
- errMsg string
- }{
- {
- name: "bad start",
- generators: []TagJunkPacketGenerator{
- newTagJunkPacketGenerator("t3", "", 10),
- newTagJunkPacketGenerator("t4", "", 10),
- },
- wantErr: true,
- errMsg: "junk packet index should be consecutive",
- },
- {
- name: "non-consecutive indices",
- generators: []TagJunkPacketGenerator{
- newTagJunkPacketGenerator("t1", "", 10),
- newTagJunkPacketGenerator("t3", "", 10), // Missing t2
- },
- wantErr: true,
- errMsg: "junk packet index should be consecutive",
- },
- {
- name: "consecutive indices",
- generators: []TagJunkPacketGenerator{
- newTagJunkPacketGenerator("t1", "", 10),
- newTagJunkPacketGenerator("t2", "", 10),
- newTagJunkPacketGenerator("t3", "", 10),
- newTagJunkPacketGenerator("t4", "", 10),
- newTagJunkPacketGenerator("t5", "", 10),
- },
- },
- {
- name: "nameIndex error",
- generators: []TagJunkPacketGenerator{
- newTagJunkPacketGenerator("error", "", 10),
- },
- wantErr: true,
- errMsg: "name must be 2 character long",
- },
- }
-
- for _, tt := range tests {
- tt := tt
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- generators := &TagJunkPacketGenerators{}
- for _, gen := range tt.generators {
- generators.AppendGenerator(gen)
- }
-
- err := generators.Validate()
- if tt.wantErr {
- require.Error(t, err)
- require.Contains(t, err.Error(), tt.errMsg)
- return
- }
- require.NoError(t, err)
- })
- }
-}
-
-func TestTagJunkGeneratorHandlerGenerate(t *testing.T) {
- mockByte1 := []byte{0x01, 0x02}
- mockByte2 := []byte{0x03, 0x04, 0x05}
- mockGen1 := internal.NewMockByteGenerator(mockByte1)
- mockGen2 := internal.NewMockByteGenerator(mockByte2)
-
- tests := []struct {
- name string
- setupGenerator func() []TagJunkPacketGenerator
- expected [][]byte
- }{
- {
- name: "generate with no default junk",
- setupGenerator: func() []TagJunkPacketGenerator {
- tg1 := newTagJunkPacketGenerator("t1", "", 0)
- tg1.append(mockGen1)
- tg1.append(mockGen2)
- tg2 := newTagJunkPacketGenerator("t2", "", 0)
- tg2.append(mockGen2)
- tg2.append(mockGen1)
-
- return []TagJunkPacketGenerator{tg1, tg2}
- },
- expected: [][]byte{
- append(mockByte1, mockByte2...),
- append(mockByte2, mockByte1...),
- },
- },
- }
-
- for _, tt := range tests {
- tt := tt
- t.Run(tt.name, func(t *testing.T) {
- t.Parallel()
- generators := &TagJunkPacketGenerators{}
- tagGenerators := tt.setupGenerator()
- for _, gen := range tagGenerators {
- generators.AppendGenerator(gen)
- }
-
- result := generators.GeneratePackets()
- require.Equal(t, result, tt.expected)
- })
- }
-}
diff --git a/device/awg/tag_parser.go b/device/awg/tag_parser.go
deleted file mode 100644
index 06ba49b..0000000
--- a/device/awg/tag_parser.go
+++ /dev/null
@@ -1,112 +0,0 @@
-package awg
-
-import (
- "fmt"
- "maps"
- "regexp"
- "strings"
-)
-
-type IpcFields struct{ Key, Value string }
-
-type EnumTag string
-
-const (
- BytesEnumTag EnumTag = "b"
- CounterEnumTag EnumTag = "c"
- TimestampEnumTag EnumTag = "t"
- RandomBytesEnumTag EnumTag = "r"
- RandomASCIIEnumTag EnumTag = "rc"
- RandomDigitEnumTag EnumTag = "rd"
-)
-
-var generatorCreator = map[EnumTag]newGenerator{
- BytesEnumTag: newBytesGenerator,
- CounterEnumTag: newPacketCounterGenerator,
- TimestampEnumTag: newTimestampGenerator,
- RandomBytesEnumTag: newRandomBytesGenerator,
- RandomASCIIEnumTag: newRandomASCIIGenerator,
- RandomDigitEnumTag: newRandomDigitGenerator,
-}
-
-// helper map to determine enumTags are unique
-var uniqueTags = map[EnumTag]bool{
- CounterEnumTag: false,
- TimestampEnumTag: false,
-}
-
-type Tag struct {
- Name EnumTag
- Param string
-}
-
-func parseTag(input string) (Tag, error) {
- // Regular expression to match
- re := regexp.MustCompile(`([a-zA-Z]+)(?:\s+([^>]+))?>`)
-
- match := re.FindStringSubmatch(input)
- tag := Tag{
- Name: EnumTag(match[1]),
- }
- if len(match) > 2 && match[2] != "" {
- tag.Param = strings.TrimSpace(match[2])
- }
-
- return tag, nil
-}
-
-func ParseTagJunkGenerator(name, input string) (TagJunkPacketGenerator, error) {
- inputSlice := strings.Split(input, "<")
- if len(inputSlice) <= 1 {
- return TagJunkPacketGenerator{}, fmt.Errorf("empty input: %s", input)
- }
-
- uniqueTagCheck := make(map[EnumTag]bool, len(uniqueTags))
- maps.Copy(uniqueTagCheck, uniqueTags)
-
- // skip byproduct of split
- inputSlice = inputSlice[1:]
- rv := newTagJunkPacketGenerator(name, input, len(inputSlice))
- for _, inputParam := range inputSlice {
- if len(inputParam) <= 1 {
- return TagJunkPacketGenerator{}, fmt.Errorf(
- "empty tag in input: %s",
- inputSlice,
- )
- } else if strings.Count(inputParam, ">") != 1 {
- return TagJunkPacketGenerator{}, fmt.Errorf("ill formated input: %s", input)
- }
-
- tag, _ := parseTag(inputParam)
- creator, ok := generatorCreator[tag.Name]
- if !ok {
- return TagJunkPacketGenerator{}, fmt.Errorf("invalid tag: %s", tag.Name)
- }
- if present, ok := uniqueTagCheck[tag.Name]; ok {
- if present {
- return TagJunkPacketGenerator{}, fmt.Errorf(
- "tag %s needs to be unique",
- tag.Name,
- )
- }
- uniqueTagCheck[tag.Name] = true
- }
- generator, err := creator(tag.Param)
- if err != nil {
- return TagJunkPacketGenerator{}, fmt.Errorf("gen: %w", err)
- }
-
- // TODO: handle counter tag
- // if tag.Name == CounterEnumTag {
- // packetCounter, ok := generator.(*PacketCounterGenerator)
- // if !ok {
- // log.Fatalf("packet counter generator expected, got %T", generator)
- // }
- // PacketCounter = packetCounter.counter
- // }
-
- rv.append(generator)
- }
-
- return rv, nil
-}
diff --git a/device/awg/tag_parser_test.go b/device/awg/tag_parser_test.go
deleted file mode 100644
index 3229cee..0000000
--- a/device/awg/tag_parser_test.go
+++ /dev/null
@@ -1,77 +0,0 @@
-package awg
-
-import (
- "fmt"
- "testing"
-
- "github.com/stretchr/testify/require"
-)
-
-func TestParse(t *testing.T) {
- type args struct {
- name string
- input string
- }
- tests := []struct {
- name string
- args args
- wantErr error
- }{
- {
- name: "invalid name",
- args: args{name: "apple", input: ""},
- wantErr: fmt.Errorf("ill formated input"),
- },
- {
- name: "empty",
- args: args{name: "i1", input: ""},
- wantErr: fmt.Errorf("ill formated input"),
- },
- {
- name: "extra >",
- args: args{name: "i1", input: ">"},
- wantErr: fmt.Errorf("ill formated input"),
- },
- {
- name: "extra <",
- args: args{name: "i1", input: "<"},
- wantErr: fmt.Errorf("empty tag in input"),
- },
- {
- name: "empty <>",
- args: args{name: "i1", input: "<>"},
- wantErr: fmt.Errorf("empty tag in input"),
- },
- {
- name: "invalid tag",
- args: args{name: "i1", input: ""},
- wantErr: fmt.Errorf("invalid tag"),
- },
- {
- name: "counter uniqueness violation",
- args: args{name: "i1", input: ""},
- wantErr: fmt.Errorf("parse tag needs to be unique"),
- },
- {
- name: "timestamp uniqueness violation",
- args: args{name: "i1", input: ""},
- wantErr: fmt.Errorf("parse tag needs to be unique"),
- },
- {
- name: "valid",
- args: args{input: ""},
- },
- }
- for _, tt := range tests {
- t.Run(tt.name, func(t *testing.T) {
- _, err := ParseTagJunkGenerator(tt.args.name, tt.args.input)
-
- // TODO: ErrorAs doesn't work as you think
- if tt.wantErr != nil {
- require.ErrorAs(t, err, &tt.wantErr)
- return
- }
- require.Nil(t, err)
- })
- }
-}
diff --git a/device/cookie_test.go b/device/cookie_test.go
index e5a2bd4..9df3049 100644
--- a/device/cookie_test.go
+++ b/device/cookie_test.go
@@ -99,7 +99,7 @@ func TestCookieMAC1(t *testing.T) {
0x8c, 0xe1, 0xe8, 0xfa, 0x67, 0x20, 0x80, 0x6d,
}
generator.AddMacs(msg)
- reply, err := checker.CreateReply(msg, 1377, src, DefaultMessageCookieReplyType)
+ reply, err := checker.CreateReply(msg, 1377, src, MessageCookieReplyType)
if err != nil {
t.Fatal("Failed to create cookie reply:", err)
}
diff --git a/device/device.go b/device/device.go
index 46cf04e..2fdf85c 100644
--- a/device/device.go
+++ b/device/device.go
@@ -6,57 +6,17 @@
package device
import (
- "encoding/binary"
- "errors"
- "fmt"
"runtime"
"sync"
"sync/atomic"
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
- "github.com/amnezia-vpn/amneziawg-go/device/awg"
- "github.com/amnezia-vpn/amneziawg-go/ipc"
"github.com/amnezia-vpn/amneziawg-go/ratelimiter"
"github.com/amnezia-vpn/amneziawg-go/rwcancel"
"github.com/amnezia-vpn/amneziawg-go/tun"
)
-type Version uint8
-
-const (
- VersionDefault Version = iota
- VersionAwg
- VersionAwgSpecialHandshake
-)
-
-// TODO:
-type AtomicVersion struct {
- value atomic.Uint32
-}
-
-func NewAtomicVersion(v Version) *AtomicVersion {
- av := &AtomicVersion{}
- av.Store(v)
- return av
-}
-
-func (av *AtomicVersion) Load() Version {
- return Version(av.value.Load())
-}
-
-func (av *AtomicVersion) Store(v Version) {
- av.value.Store(uint32(v))
-}
-
-func (av *AtomicVersion) CompareAndSwap(old, new Version) bool {
- return av.value.CompareAndSwap(uint32(old), uint32(new))
-}
-
-func (av *AtomicVersion) Swap(new Version) Version {
- return Version(av.value.Swap(uint32(new)))
-}
-
type Device struct {
state struct {
// state holds the device's state. It is accessed atomically.
@@ -130,8 +90,27 @@ type Device struct {
closed chan struct{}
log *Logger
- version Version
- awg awg.Protocol
+ junk struct {
+ min int
+ max int
+ count int
+ }
+
+ headers struct {
+ init *magicHeader
+ cookie *magicHeader
+ response *magicHeader
+ transport *magicHeader
+ }
+
+ paddings struct {
+ init int
+ response int
+ cookie int
+ transport int
+ }
+
+ ipackets [5]*obfChain
}
// deviceState represents the state of a Device.
@@ -342,6 +321,11 @@ func NewDevice(tunDevice tun.Device, bind conn.Bind, logger *Logger) *Device {
device.rate.limiter.Init()
device.indexTable.Init()
+ device.headers.init = &magicHeader{start: MessageInitiationType, end: MessageInitiationType}
+ device.headers.response = &magicHeader{start: MessageResponseType, end: MessageResponseType}
+ device.headers.cookie = &magicHeader{start: MessageCookieReplyType, end: MessageCookieReplyType}
+ device.headers.transport = &magicHeader{start: MessageTransportType, end: MessageTransportType}
+
device.PopulatePools()
// create queues
@@ -439,8 +423,6 @@ func (device *Device) Close() {
device.rate.limiter.Close()
- device.resetProtocol()
-
device.log.Verbosef("Device closed")
close(device.closed)
}
@@ -580,358 +562,3 @@ func (device *Device) BindClose() error {
device.net.Unlock()
return err
}
-
-func (device *Device) isAWG() bool {
- return device.version >= VersionAwg
-}
-
-func (device *Device) resetProtocol() {
- // restore default message type values
- MessageInitiationType = DefaultMessageInitiationType
- MessageResponseType = DefaultMessageResponseType
- MessageCookieReplyType = DefaultMessageCookieReplyType
- MessageTransportType = DefaultMessageTransportType
-}
-
-func (device *Device) handlePostConfig(tempAwg *awg.Protocol) error {
- if !tempAwg.Cfg.IsSet && !tempAwg.HandshakeHandler.IsSet {
- return nil
- }
-
- var errs []error
-
- isAwgOn := false
- device.awg.Mux.Lock()
- if tempAwg.Cfg.JunkPacketCount < 0 {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- "JunkPacketCount should be non negative",
- ),
- )
- }
- device.awg.Cfg.JunkPacketCount = tempAwg.Cfg.JunkPacketCount
- if tempAwg.Cfg.JunkPacketCount != 0 {
- isAwgOn = true
- }
-
- device.awg.Cfg.JunkPacketMinSize = tempAwg.Cfg.JunkPacketMinSize
- if tempAwg.Cfg.JunkPacketMinSize != 0 {
- isAwgOn = true
- }
-
- if device.awg.Cfg.JunkPacketCount > 0 &&
- tempAwg.Cfg.JunkPacketMaxSize == tempAwg.Cfg.JunkPacketMinSize {
-
- tempAwg.Cfg.JunkPacketMaxSize++ // to make rand gen work
- }
-
- if tempAwg.Cfg.JunkPacketMaxSize >= MaxSegmentSize {
- device.awg.Cfg.JunkPacketMinSize = 0
- device.awg.Cfg.JunkPacketMaxSize = 1
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- "JunkPacketMaxSize: %d; should be smaller than maxSegmentSize: %d",
- tempAwg.Cfg.JunkPacketMaxSize,
- MaxSegmentSize,
- ))
- } else if tempAwg.Cfg.JunkPacketMaxSize < tempAwg.Cfg.JunkPacketMinSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- "maxSize: %d; should be greater than minSize: %d",
- tempAwg.Cfg.JunkPacketMaxSize,
- tempAwg.Cfg.JunkPacketMinSize,
- ))
- } else {
- device.awg.Cfg.JunkPacketMaxSize = tempAwg.Cfg.JunkPacketMaxSize
- }
-
- if tempAwg.Cfg.JunkPacketMaxSize != 0 {
- isAwgOn = true
- }
-
- magicHeaders := make([]awg.MagicHeader, 4)
-
- if len(tempAwg.Cfg.MagicHeaders.Values) != 4 {
- return ipcErrorf(
- ipc.IpcErrorInvalid,
- "magic headers should have 4 values; got: %d",
- len(tempAwg.Cfg.MagicHeaders.Values),
- )
- }
-
- if tempAwg.Cfg.MagicHeaders.Values[0].Min > 4 {
- isAwgOn = true
- device.log.Verbosef("UAPI: Updating init_packet_magic_header")
- magicHeaders[0] = tempAwg.Cfg.MagicHeaders.Values[0]
-
- MessageInitiationType = magicHeaders[0].Min
- } else {
- device.log.Verbosef("UAPI: Using default init type")
- MessageInitiationType = DefaultMessageInitiationType
- magicHeaders[0] = awg.NewMagicHeaderSameValue(DefaultMessageInitiationType)
- }
-
- if tempAwg.Cfg.MagicHeaders.Values[1].Min > 4 {
- isAwgOn = true
-
- device.log.Verbosef("UAPI: Updating response_packet_magic_header")
- magicHeaders[1] = tempAwg.Cfg.MagicHeaders.Values[1]
- MessageResponseType = magicHeaders[1].Min
- } else {
- device.log.Verbosef("UAPI: Using default response type")
- MessageResponseType = DefaultMessageResponseType
- magicHeaders[1] = awg.NewMagicHeaderSameValue(DefaultMessageResponseType)
- }
-
- if tempAwg.Cfg.MagicHeaders.Values[2].Min > 4 {
- isAwgOn = true
-
- device.log.Verbosef("UAPI: Updating underload_packet_magic_header")
- magicHeaders[2] = tempAwg.Cfg.MagicHeaders.Values[2]
- MessageCookieReplyType = magicHeaders[2].Min
- } else {
- device.log.Verbosef("UAPI: Using default underload type")
- MessageCookieReplyType = DefaultMessageCookieReplyType
- magicHeaders[2] = awg.NewMagicHeaderSameValue(DefaultMessageCookieReplyType)
- }
-
- if tempAwg.Cfg.MagicHeaders.Values[3].Min > 4 {
- isAwgOn = true
-
- device.log.Verbosef("UAPI: Updating transport_packet_magic_header")
- magicHeaders[3] = tempAwg.Cfg.MagicHeaders.Values[3]
- MessageTransportType = magicHeaders[3].Min
- } else {
- device.log.Verbosef("UAPI: Using default transport type")
- MessageTransportType = DefaultMessageTransportType
- magicHeaders[3] = awg.NewMagicHeaderSameValue(DefaultMessageTransportType)
- }
-
- var err error
- device.awg.Cfg.MagicHeaders, err = awg.NewMagicHeaders(magicHeaders)
- if err != nil {
- errs = append(errs, ipcErrorf(ipc.IpcErrorInvalid, "new magic headers: %w", err))
- }
-
- isSameHeaderMap := map[uint32]struct{}{
- MessageInitiationType: {},
- MessageResponseType: {},
- MessageCookieReplyType: {},
- MessageTransportType: {},
- }
-
- // size will be different if same values
- if len(isSameHeaderMap) != 4 {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `magic headers should differ; got: init:%d; recv:%d; unde:%d; tran:%d`,
- MessageInitiationType,
- MessageResponseType,
- MessageCookieReplyType,
- MessageTransportType,
- ),
- )
- }
-
- newInitSize := MessageInitiationSize + tempAwg.Cfg.InitHeaderJunkSize
-
- if newInitSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `init header size(148) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.Cfg.InitHeaderJunkSize,
- MaxSegmentSize,
- ),
- )
- } else {
- device.awg.Cfg.InitHeaderJunkSize = tempAwg.Cfg.InitHeaderJunkSize
- }
-
- if tempAwg.Cfg.InitHeaderJunkSize != 0 {
- isAwgOn = true
- }
-
- newResponseSize := MessageResponseSize + tempAwg.Cfg.ResponseHeaderJunkSize
-
- if newResponseSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `response header size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.Cfg.ResponseHeaderJunkSize,
- MaxSegmentSize,
- ),
- )
- } else {
- device.awg.Cfg.ResponseHeaderJunkSize = tempAwg.Cfg.ResponseHeaderJunkSize
- }
-
- if tempAwg.Cfg.ResponseHeaderJunkSize != 0 {
- isAwgOn = true
- }
-
- newCookieSize := MessageCookieReplySize + tempAwg.Cfg.CookieReplyHeaderJunkSize
-
- if newCookieSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `cookie reply size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.Cfg.CookieReplyHeaderJunkSize,
- MaxSegmentSize,
- ),
- )
- } else {
- device.awg.Cfg.CookieReplyHeaderJunkSize = tempAwg.Cfg.CookieReplyHeaderJunkSize
- }
-
- if tempAwg.Cfg.CookieReplyHeaderJunkSize != 0 {
- isAwgOn = true
- }
-
- newTransportSize := MessageTransportSize + tempAwg.Cfg.TransportHeaderJunkSize
-
- if newTransportSize >= MaxSegmentSize {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `transport size(92) + junkSize:%d; should be smaller than maxSegmentSize: %d`,
- tempAwg.Cfg.TransportHeaderJunkSize,
- MaxSegmentSize,
- ),
- )
- } else {
- device.awg.Cfg.TransportHeaderJunkSize = tempAwg.Cfg.TransportHeaderJunkSize
- }
-
- if tempAwg.Cfg.TransportHeaderJunkSize != 0 {
- isAwgOn = true
- }
-
- isSameSizeMap := map[int]struct{}{
- newInitSize: {},
- newResponseSize: {},
- newCookieSize: {},
- newTransportSize: {},
- }
-
- if len(isSameSizeMap) != 4 {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid,
- `new sizes should differ; init: %d; response: %d; cookie: %d; trans: %d`,
- newInitSize,
- newResponseSize,
- newCookieSize,
- newTransportSize,
- ),
- )
- } else {
- msgTypeToJunkSize = map[uint32]int{
- MessageInitiationType: device.awg.Cfg.InitHeaderJunkSize,
- MessageResponseType: device.awg.Cfg.ResponseHeaderJunkSize,
- MessageCookieReplyType: device.awg.Cfg.CookieReplyHeaderJunkSize,
- MessageTransportType: device.awg.Cfg.TransportHeaderJunkSize,
- }
-
- packetSizeToMsgType = map[int]uint32{
- newInitSize: MessageInitiationType,
- newResponseSize: MessageResponseType,
- newCookieSize: MessageCookieReplyType,
- newTransportSize: MessageTransportType,
- }
- }
-
- device.awg.IsOn.SetTo(isAwgOn)
- device.awg.JunkCreator = awg.NewJunkCreator(device.awg.Cfg)
-
- if tempAwg.HandshakeHandler.IsSet {
- if err := tempAwg.HandshakeHandler.Validate(); err != nil {
- errs = append(errs, ipcErrorf(
- ipc.IpcErrorInvalid, "handshake handler validate: %w", err))
- } else {
- device.awg.HandshakeHandler = tempAwg.HandshakeHandler
- device.awg.HandshakeHandler.SpecialJunk.DefaultJunkCount = tempAwg.Cfg.JunkPacketCount
- device.version = VersionAwgSpecialHandshake
- }
- } else {
- device.version = VersionAwg
- }
-
- device.awg.Mux.Unlock()
-
- return errors.Join(errs...)
-}
-
-func (device *Device) ProcessAWGPacket(size int, packet *[]byte, buffer *[MaxMessageSize]byte) (uint32, error) {
- // TODO:
- // if awg.WaitResponse.ShouldWait.IsSet() {
- // awg.WaitResponse.Channel <- struct{}{}
- // }
-
- expectedMsgType, isKnownSize := packetSizeToMsgType[size]
- if !isKnownSize {
- msgType, err := device.handleTransport(size, packet, buffer)
-
- if err != nil {
- return 0, fmt.Errorf("handle transport: %w", err)
- }
-
- return msgType, nil
- }
-
- junkSize := msgTypeToJunkSize[expectedMsgType]
-
- // transport size can align with other header types;
- // making sure we have the right actualMsgType
- actualMsgType, err := device.getMsgType(packet, junkSize)
- if err != nil {
- return 0, fmt.Errorf("get msg type: %w", err)
- }
-
- if actualMsgType == expectedMsgType {
- *packet = (*packet)[junkSize:]
- return actualMsgType, nil
- }
-
- device.log.Verbosef("awg: transport packet lined up with another msg type")
-
- msgType, err := device.handleTransport(size, packet, buffer)
- if err != nil {
- return 0, fmt.Errorf("handle transport: %w", err)
- }
-
- return msgType, nil
-}
-
-func (device *Device) getMsgType(packet *[]byte, junkSize int) (uint32, error) {
- msgTypeValue := binary.LittleEndian.Uint32((*packet)[junkSize : junkSize+4])
- msgType, err := device.awg.GetMagicHeaderMinFor(msgTypeValue)
-
- if err != nil {
- return 0, fmt.Errorf("get magic header min: %w", err)
- }
-
- return msgType, nil
-}
-
-func (device *Device) handleTransport(size int, packet *[]byte, buffer *[MaxMessageSize]byte) (uint32, error) {
- junkSize := device.awg.Cfg.TransportHeaderJunkSize
-
- msgType, err := device.getMsgType(packet, junkSize)
- if err != nil {
- return 0, fmt.Errorf("get msg type: %w", err)
- }
-
- if msgType != MessageTransportType {
- // probably a junk packet
- return 0, fmt.Errorf("Received message with unknown type: %d", msgType)
- }
-
- if junkSize > 0 {
- // remove junk from buffer by shifting the packet
- // this buffer is also used for decryption, so it needs to be corrected
- copy((*buffer)[:size], (*packet)[junkSize:])
- size -= junkSize
- // need to reinitialize packet as well
- (*packet) = (*packet)[:size]
- }
-
- return msgType, nil
-}
diff --git a/device/magic-header.go b/device/magic-header.go
new file mode 100644
index 0000000..78e59d6
--- /dev/null
+++ b/device/magic-header.go
@@ -0,0 +1,63 @@
+package device
+
+import (
+ "crypto/rand"
+ "errors"
+ "fmt"
+ "math/big"
+ "strconv"
+ "strings"
+)
+
+type magicHeader struct {
+ start uint32
+ end uint32
+}
+
+func newMagicHeader(spec string) (*magicHeader, error) {
+ parts := strings.Split(spec, "-")
+ if len(parts) < 1 || len(parts) > 2 {
+ return nil, errors.New("bad format")
+ }
+
+ start, err := strconv.ParseUint(parts[0], 10, 32)
+ if err != nil {
+ return nil, fmt.Errorf("failed to parse %s: %w", parts[0], err)
+ }
+
+ var end uint64
+ if len(parts) > 1 {
+ end, err = strconv.ParseUint(parts[1], 10, 32)
+ if err != nil {
+ return nil, fmt.Errorf("failed to parse %s: %w", parts[1], err)
+ }
+ } else {
+ end = start
+ }
+
+ if end < start {
+ return nil, errors.New("wrong range specified")
+ }
+
+ return &magicHeader{
+ start: uint32(start),
+ end: uint32(end),
+ }, nil
+}
+
+func (h *magicHeader) GenSpec() string {
+ if h.start == h.end {
+ return fmt.Sprintf("%d", h.start)
+ }
+ return fmt.Sprintf("%d-%d", h.start, h.end)
+}
+
+func (h *magicHeader) Validate(val uint32) bool {
+ return h.start <= val && val <= h.end
+}
+
+func (h *magicHeader) Generate() uint32 {
+ high := int64(h.end - h.start + 1)
+ r, _ := rand.Int(rand.Reader, big.NewInt(high))
+ return h.start + uint32(r.Int64())
+}
diff --git a/device/noise-protocol.go b/device/noise-protocol.go
index 6e6fe58..86346ac 100644
--- a/device/noise-protocol.go
+++ b/device/noise-protocol.go
@@ -53,17 +53,11 @@ const (
)
const (
- DefaultMessageInitiationType uint32 = 1
- DefaultMessageResponseType uint32 = 2
- DefaultMessageCookieReplyType uint32 = 3
- DefaultMessageTransportType uint32 = 4
-)
-
-var (
- MessageInitiationType uint32 = DefaultMessageInitiationType
- MessageResponseType uint32 = DefaultMessageResponseType
- MessageCookieReplyType uint32 = DefaultMessageCookieReplyType
- MessageTransportType uint32 = DefaultMessageTransportType
+ MessageUnknownType uint32 = 0
+ MessageInitiationType uint32 = 1
+ MessageResponseType uint32 = 2
+ MessageCookieReplyType uint32 = 3
+ MessageTransportType uint32 = 4
)
const (
@@ -82,11 +76,6 @@ const (
MessageTransportOffsetContent = 16
)
-var (
- packetSizeToMsgType map[int]uint32
- msgTypeToJunkSize map[uint32]int
-)
-
/* Type is an 8-bit field, followed by 3 nul bytes,
* by marshalling the messages in little-endian byteorder
* we can treat these as a 32-bit unsigned int (for now)
@@ -205,17 +194,7 @@ func (device *Device) CreateMessageInitiation(peer *Peer) (*MessageInitiation, e
handshake.mixHash(handshake.remoteStatic[:])
- msgType := DefaultMessageInitiationType
- if device.isAWG() {
- device.awg.Mux.RLock()
- msgType, err = device.awg.GetMsgType(DefaultMessageInitiationType)
- if err != nil {
- device.awg.Mux.RUnlock()
- return nil, fmt.Errorf("get message type: %w", err)
- }
-
- device.awg.Mux.RUnlock()
- }
+ msgType := device.headers.init.Generate()
msg := MessageInitiation{
Type: msgType,
@@ -274,13 +253,9 @@ func (device *Device) ConsumeMessageInitiation(msg *MessageInitiation) *Peer {
chainKey [blake2s.Size]byte
)
- device.awg.Mux.RLock()
-
if msg.Type != MessageInitiationType {
- device.awg.Mux.RUnlock()
return nil
}
- device.awg.Mux.RUnlock()
device.staticIdentity.RLock()
defer device.staticIdentity.RUnlock()
@@ -395,19 +370,7 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
var msg MessageResponse
- if device.isAWG() {
- device.awg.Mux.RLock()
- msg.Type, err = device.awg.GetMsgType(DefaultMessageResponseType)
- if err != nil {
- device.awg.Mux.RUnlock()
- return nil, fmt.Errorf("get message type: %w", err)
- }
-
- device.awg.Mux.RUnlock()
- } else {
- msg.Type = DefaultMessageResponseType
- }
-
+ msg.Type = device.headers.response.Generate()
msg.Sender = handshake.localIndex
msg.Receiver = handshake.remoteIndex
@@ -457,13 +420,9 @@ func (device *Device) CreateMessageResponse(peer *Peer) (*MessageResponse, error
}
func (device *Device) ConsumeMessageResponse(msg *MessageResponse) *Peer {
- device.awg.Mux.RLock()
-
if msg.Type != MessageResponseType {
- device.awg.Mux.RUnlock()
return nil
}
- device.awg.Mux.RUnlock()
// lookup handshake by receiver
diff --git a/device/obf.go b/device/obf.go
new file mode 100644
index 0000000..53c55ff
--- /dev/null
+++ b/device/obf.go
@@ -0,0 +1,140 @@
+package device
+
+import (
+ "errors"
+ "fmt"
+ "strings"
+)
+
+type obfBuilder func(val string) (obf, error)
+
+var obfBuilders = map[string]obfBuilder{
+ "b": newBytesObf,
+ "t": newTimestampObf,
+ "r": newRandObf,
+ "rc": newRandCharObf,
+ "rd": newRandDigitsObf,
+ "d": newDataObf,
+ "ds": newDataStringObf,
+ "dz": newDataSizeObf,
+}
+
+type obf interface {
+ Obfuscate(dst, src []byte)
+ Deobfuscate(dst, src []byte) bool
+ ObfuscatedLen(srcLen int) int
+ DeobfuscatedLen(srcLen int) int
+}
+
+type obfChain struct {
+ Spec string
+ obfs []obf
+}
+
+func newObfChain(spec string) (*obfChain, error) {
+ var (
+ obfs []obf
+ errs []error
+ )
+
+ remaining := spec[:]
+ for {
+ start := strings.IndexByte(remaining, '<')
+ if start == -1 {
+ break
+ }
+
+ end := strings.IndexByte(remaining[start:], '>')
+ if end == -1 {
+ return nil, errors.New("missing enclosing >")
+ }
+ end += start
+
+ tag := remaining[start+1 : end]
+ parts := strings.Fields(tag)
+ if len(parts) == 0 {
+ errs = append(errs, errors.New("empty tag"))
+ remaining = remaining[end+1:]
+ continue
+ }
+
+ key := parts[0]
+ builder, ok := obfBuilders[key]
+ if !ok {
+ errs = append(errs, fmt.Errorf("unknown tag <%s>", key))
+ remaining = remaining[end+1:]
+ continue
+ }
+
+ val := ""
+ if len(parts) > 1 {
+ val = parts[1]
+ }
+
+ o, err := builder(val)
+ if err != nil {
+ errs = append(errs, fmt.Errorf("failed to build <%s>: %w", key, err))
+ remaining = remaining[end+1:]
+ continue
+ }
+
+ obfs = append(obfs, o)
+ remaining = remaining[end+1:]
+ }
+
+ if len(errs) > 0 {
+ return nil, errors.Join(errs...)
+ }
+
+ return &obfChain{
+ Spec: spec,
+ obfs: obfs,
+ }, nil
+}
+
+func (c *obfChain) Obfuscate(dst, src []byte) {
+ written := 0
+ for _, o := range c.obfs {
+ obfLen := o.ObfuscatedLen(len(src))
+ o.Obfuscate(dst[written:written+obfLen], src)
+ written += obfLen
+ }
+}
+
+func (c *obfChain) Deobfuscate(dst, src []byte) bool {
+ dynamicLen := len(src) - c.ObfuscatedLen(0)
+
+ written, read := 0, 0
+
+ for _, o := range c.obfs {
+ deobfLen := o.DeobfuscatedLen(dynamicLen)
+ obfLen := o.ObfuscatedLen(deobfLen)
+
+ if !o.Deobfuscate(dst[written:written+deobfLen], src[read:read+obfLen]) {
+ return false
+ }
+
+ written += deobfLen
+ read += obfLen
+ }
+
+ return true
+}
+
+func (c *obfChain) ObfuscatedLen(n int) int {
+ total := 0
+ for _, o := range c.obfs {
+ total += o.ObfuscatedLen(n)
+ }
+ return total
+}
+
+func (c *obfChain) DeobfuscatedLen(n int) int {
+ dynamicLen := n - c.ObfuscatedLen(0)
+
+ total := 0
+ for _, o := range c.obfs {
+ total += o.DeobfuscatedLen(dynamicLen)
+ }
+ return total
+}
diff --git a/device/obf_bytes.go b/device/obf_bytes.go
new file mode 100644
index 0000000..68d722b
--- /dev/null
+++ b/device/obf_bytes.go
@@ -0,0 +1,47 @@
+package device
+
+import (
+ "bytes"
+ "encoding/hex"
+ "errors"
+ "strings"
+)
+
+func newBytesObf(val string) (obf, error) {
+ val = strings.TrimPrefix(val, "0x")
+
+ if len(val) == 0 {
+ return nil, errors.New("empty argument")
+ }
+
+ if len(val)%2 != 0 {
+ return nil, errors.New("odd amount of symbols")
+ }
+
+ bytes, err := hex.DecodeString(val)
+ if err != nil {
+ return nil, err
+ }
+
+ return &bytesObf{data: bytes}, nil
+}
+
+type bytesObf struct {
+ data []byte
+}
+
+func (o *bytesObf) Obfuscate(dst, src []byte) {
+ copy(dst, o.data)
+}
+
+func (o *bytesObf) Deobfuscate(dst, src []byte) bool {
+ return bytes.Equal(o.data, src[:o.ObfuscatedLen(0)])
+}
+
+func (o *bytesObf) ObfuscatedLen(srcLen int) int {
+ return len(o.data)
+}
+
+func (o *bytesObf) DeobfuscatedLen(srcLen int) int {
+ return 0
+}
diff --git a/device/obf_data.go b/device/obf_data.go
new file mode 100644
index 0000000..42d3f65
--- /dev/null
+++ b/device/obf_data.go
@@ -0,0 +1,25 @@
+package device
+
+func newDataObf(val string) (obf, error) {
+ return &dataObf{}, nil
+}
+
+type dataObf struct {
+}
+
+func (obf *dataObf) Obfuscate(dst, src []byte) {
+ copy(dst, src)
+}
+
+func (obf *dataObf) Deobfuscate(dst, src []byte) bool {
+ copy(dst, src)
+ return true
+}
+
+func (o *dataObf) ObfuscatedLen(n int) int {
+ return n
+}
+
+func (o *dataObf) DeobfuscatedLen(n int) int {
+ return n
+}
diff --git a/device/obf_datasize.go b/device/obf_datasize.go
new file mode 100644
index 0000000..7267e2a
--- /dev/null
+++ b/device/obf_datasize.go
@@ -0,0 +1,38 @@
+package device
+
+import "strconv"
+
+func newDataSizeObf(val string) (obf, error) {
+ length, err := strconv.Atoi(val)
+ if err != nil {
+ return nil, err
+ }
+
+ return &dataSizeObf{
+ length: length,
+ }, nil
+}
+
+type dataSizeObf struct {
+ length int
+}
+
+func (o *dataSizeObf) Obfuscate(dst, src []byte) {
+ srcLen := len(src)
+ for i := o.length - 1; i >= 0; i-- {
+ dst[i] = byte(srcLen & 0xFF)
+ srcLen >>= 8
+ }
+}
+
+func (o *dataSizeObf) Deobfuscate(dst, src []byte) bool {
+ return true
+}
+
+func (o *dataSizeObf) ObfuscatedLen(srcLen int) int {
+ return o.length
+}
+
+func (o *dataSizeObf) DeobfuscatedLen(srcLen int) int {
+ return 0
+}
diff --git a/device/obf_datastring.go b/device/obf_datastring.go
new file mode 100644
index 0000000..2701e95
--- /dev/null
+++ b/device/obf_datastring.go
@@ -0,0 +1,29 @@
+package device
+
+import (
+ "encoding/base64"
+)
+
+func newDataStringObf(val string) (obf, error) {
+ return &dataStringObf{}, nil
+}
+
+type dataStringObf struct {
+}
+
+func (o *dataStringObf) Obfuscate(dst, src []byte) {
+ base64.RawStdEncoding.Encode(dst, src)
+}
+
+func (o *dataStringObf) Deobfuscate(dst, src []byte) bool {
+ base64.RawStdEncoding.Decode(dst, src)
+ return true
+}
+
+func (o *dataStringObf) ObfuscatedLen(n int) int {
+ return base64.RawStdEncoding.EncodedLen(n)
+}
+
+func (o *dataStringObf) DeobfuscatedLen(n int) int {
+ return base64.RawStdEncoding.DecodedLen(n)
+}
diff --git a/device/obf_rand.go b/device/obf_rand.go
new file mode 100644
index 0000000..edf461e
--- /dev/null
+++ b/device/obf_rand.go
@@ -0,0 +1,39 @@
+package device
+
+import (
+ "crypto/rand"
+ "strconv"
+)
+
+func newRandObf(val string) (obf, error) {
+ length, err := strconv.Atoi(val)
+ if err != nil {
+ return nil, err
+ }
+
+ return &randObf{
+ length: length,
+ }, nil
+}
+
+type randObf struct {
+ length int
+}
+
+func (o *randObf) Obfuscate(dst, src []byte) {
+ rand.Read(dst[:o.length])
+}
+
+func (o *randObf) Deobfuscate(dst, src []byte) bool {
+ // there is no way to validate randomness :)
+ // assume that it is always true
+ return true
+}
+
+func (o *randObf) ObfuscatedLen(n int) int {
+ return o.length
+}
+
+func (o *randObf) DeobfuscatedLen(n int) int {
+ return 0
+}
diff --git a/device/obf_randchars.go b/device/obf_randchars.go
new file mode 100644
index 0000000..1d9968c
--- /dev/null
+++ b/device/obf_randchars.go
@@ -0,0 +1,48 @@
+package device
+
+import (
+ "crypto/rand"
+ "strconv"
+ "unicode"
+)
+
+const chars52 = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
+
+func newRandCharObf(val string) (obf, error) {
+ length, err := strconv.Atoi(val)
+ if err != nil {
+ return nil, err
+ }
+
+ return &randCharObf{
+ length: length,
+ }, nil
+}
+
+type randCharObf struct {
+ length int
+}
+
+func (o *randCharObf) Obfuscate(dst, src []byte) {
+ rand.Read(dst[:o.length])
+ for i := range dst[:o.length] {
+ dst[i] = chars52[dst[i]%52]
+ }
+}
+
+func (o *randCharObf) Deobfuscate(dst, src []byte) bool {
+ for _, b := range src[:o.length] {
+ if !unicode.IsLetter(rune(b)) {
+ return false
+ }
+ }
+ return true
+}
+
+func (o *randCharObf) ObfuscatedLen(n int) int {
+ return o.length
+}
+
+func (o *randCharObf) DeobfuscatedLen(n int) int {
+ return 0
+}
diff --git a/device/obf_randdigits.go b/device/obf_randdigits.go
new file mode 100644
index 0000000..4794bb1
--- /dev/null
+++ b/device/obf_randdigits.go
@@ -0,0 +1,48 @@
+package device
+
+import (
+ "crypto/rand"
+ "strconv"
+ "unicode"
+)
+
+const digits10 = "0123456789"
+
+func newRandDigitsObf(val string) (obf, error) {
+ length, err := strconv.Atoi(val)
+ if err != nil {
+ return nil, err
+ }
+
+ return &randDigitObf{
+ length: length,
+ }, nil
+}
+
+type randDigitObf struct {
+ length int
+}
+
+func (o *randDigitObf) Obfuscate(dst, src []byte) {
+ rand.Read(dst[:o.length])
+ for i := range dst[:o.length] {
+ dst[i] = digits10[dst[i]%10]
+ }
+}
+
+func (o *randDigitObf) Deobfuscate(dst, src []byte) bool {
+ for _, b := range src[:o.length] {
+ if !unicode.IsDigit(rune(b)) {
+ return false
+ }
+ }
+ return true
+}
+
+func (o *randDigitObf) ObfuscatedLen(n int) int {
+ return o.length
+}
+
+func (o *randDigitObf) DeobfuscatedLen(n int) int {
+ return 0
+}
diff --git a/device/obf_timestamp.go b/device/obf_timestamp.go
new file mode 100644
index 0000000..0a8180b
--- /dev/null
+++ b/device/obf_timestamp.go
@@ -0,0 +1,31 @@
+package device
+
+import (
+ "encoding/binary"
+ "time"
+)
+
+func newTimestampObf(_ string) (obf, error) {
+ return ×tampObf{}, nil
+}
+
+type timestampObf struct{}
+
+func (o *timestampObf) Obfuscate(dst, src []byte) {
+ t := uint32(time.Now().Unix())
+ binary.BigEndian.PutUint32(dst, t)
+}
+
+func (o *timestampObf) Deobfuscate(dst, src []byte) bool {
+ // replay attack check?
+ // requires time to be always synchronized
+ return true
+}
+
+func (o *timestampObf) ObfuscatedLen(n int) int {
+ return 4
+}
+
+func (o *timestampObf) DeobfuscatedLen(n int) int {
+ return 0
+}
diff --git a/device/peer.go b/device/peer.go
index e8a5168..8f88b2a 100644
--- a/device/peer.go
+++ b/device/peer.go
@@ -13,7 +13,6 @@ import (
"time"
"github.com/amnezia-vpn/amneziawg-go/conn"
- "github.com/amnezia-vpn/amneziawg-go/device/awg"
)
type Peer struct {
@@ -114,16 +113,6 @@ func (device *Device) NewPeer(pk NoisePublicKey) (*Peer, error) {
return peer, nil
}
-func (peer *Peer) SendAndCountBuffers(buffers [][]byte) error {
- err := peer.SendBuffers(buffers)
- if err == nil {
- awg.PacketCounter.Add(uint64(len(buffers)))
- return nil
- }
-
- return err
-}
-
func (peer *Peer) SendBuffers(buffers [][]byte) error {
peer.device.net.RLock()
defer peer.device.net.RUnlock()
diff --git a/device/receive.go b/device/receive.go
index 4c34799..b3b6105 100644
--- a/device/receive.go
+++ b/device/receive.go
@@ -97,13 +97,13 @@ func (device *Device) RoutineReceiveIncoming(
elemsByPeer = make(map[*Peer]*QueueInboundElementsContainer, maxBatchSize)
)
- for i := range bufsArrs {
+ for i := range maxBatchSize {
bufsArrs[i] = device.GetMessageBuffer()
bufs[i] = bufsArrs[i][:]
}
defer func() {
- for i := 0; i < maxBatchSize; i++ {
+ for i := range maxBatchSize {
if bufsArrs[i] != nil {
device.PutMessageBuffer(bufsArrs[i])
}
@@ -129,7 +129,6 @@ func (device *Device) RoutineReceiveIncoming(
}
deathSpiral = 0
- device.awg.Mux.RLock()
// handle each packet in the batch
for i, size := range sizes[:count] {
if size < MinMessageSize {
@@ -138,16 +137,12 @@ func (device *Device) RoutineReceiveIncoming(
// check size of packet
packet := bufsArrs[i][:size]
- var msgType uint32
- if device.isAWG() {
- msgType, err = device.ProcessAWGPacket(size, &packet, bufsArrs[i])
- if err != nil {
- device.log.Verbosef("awg: process packet: %v", err)
- continue
- }
- } else {
- msgType = binary.LittleEndian.Uint32(packet[:4])
+ // get message padding and type based on information from S1-S4 and H1-H4
+ msgType, padding := device.DeterminePacketTypeAndPadding(packet, MessageUnknownType)
+ if padding > 0 {
+ copy(packet, packet[padding:])
+ packet = packet[:len(packet)-padding]
}
switch msgType {
@@ -233,7 +228,6 @@ func (device *Device) RoutineReceiveIncoming(
default:
}
}
- device.awg.Mux.RUnlock()
for peer, elemsContainer := range elemsByPeer {
if peer.isRunning.Load() {
peer.queue.inbound.c <- elemsContainer
@@ -291,9 +285,6 @@ func (device *Device) RoutineHandshake(id int) {
device.log.Verbosef("Routine: handshake worker %d - started", id)
for elem := range device.queue.handshake.c {
-
- device.awg.Mux.RLock()
-
// handle cookie fields and ratelimiting
switch elem.msgType {
@@ -450,7 +441,6 @@ func (device *Device) RoutineHandshake(id int) {
peer.SendKeepalive()
}
skip:
- device.awg.Mux.RUnlock()
device.PutMessageBuffer(elem.buffer)
}
}
@@ -569,3 +559,57 @@ func (peer *Peer) RoutineSequentialReceiver(maxBatchSize int) {
device.PutInboundElementsContainer(elemsContainer)
}
}
+
+func (device *Device) DeterminePacketTypeAndPadding(packet []byte, expectedType uint32) (uint32, int) {
+ size := len(packet)
+
+ if expectedType == MessageUnknownType || expectedType == MessageInitiationType {
+ padding := device.paddings.init
+ header := device.headers.init
+
+ if size == padding+MessageInitiationSize {
+ data := packet[padding:]
+ if header.Validate(binary.LittleEndian.Uint32(data)) {
+ return MessageInitiationType, padding
+ }
+ }
+ }
+
+ if expectedType == MessageUnknownType || expectedType == MessageResponseType {
+ padding := device.paddings.response
+ header := device.headers.response
+
+ if size == padding+MessageResponseSize {
+ data := packet[padding:]
+ if header.Validate(binary.LittleEndian.Uint32(data)) {
+ return MessageResponseType, padding
+ }
+ }
+ }
+
+ if expectedType == MessageUnknownType || expectedType == MessageCookieReplyType {
+ padding := device.paddings.cookie
+ header := device.headers.cookie
+
+ if size == padding+MessageCookieReplySize {
+ data := packet[padding:]
+ if header.Validate(binary.LittleEndian.Uint32(data)) {
+ return MessageCookieReplyType, padding
+ }
+ }
+ }
+
+ if expectedType == MessageUnknownType || expectedType == MessageTransportType {
+ padding := device.paddings.transport
+ header := device.headers.transport
+
+ if size >= padding+MessageTransportHeaderSize {
+ data := packet[padding:]
+ if header.Validate(binary.LittleEndian.Uint32(data)) {
+ return MessageTransportType, padding
+ }
+ }
+ }
+
+ return MessageUnknownType, 0
+}
diff --git a/device/send.go b/device/send.go
index 0861a04..5e5cc1b 100644
--- a/device/send.go
+++ b/device/send.go
@@ -7,8 +7,10 @@ package device
import (
"bytes"
+ "crypto/rand"
"encoding/binary"
"errors"
+ "math/big"
"net"
"os"
"sync"
@@ -123,41 +125,28 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
peer.device.log.Errorf("%v - Failed to create initiation message: %v", peer, err)
return err
}
+
var sendBuffer [][]byte
- // so only packet processed for cookie generation
- var junkedHeader []byte
- if peer.device.version >= VersionAwg {
- var junks [][]byte
- if peer.device.version == VersionAwgSpecialHandshake {
- peer.device.awg.Mux.RLock()
- // set junks depending on packet type
- junks = peer.device.awg.HandshakeHandler.GenerateSpecialJunk()
- if junks != nil {
- peer.device.log.Verbosef("%v - Special junks sent", peer)
- }
- peer.device.awg.Mux.RUnlock()
- } else {
- junks = make([][]byte, 0, peer.device.awg.Cfg.JunkPacketCount)
+ for _, ipacket := range peer.device.ipackets {
+ if ipacket != nil {
+ buf := make([]byte, ipacket.ObfuscatedLen(0))
+ ipacket.Obfuscate(buf, nil)
+ sendBuffer = append(sendBuffer, buf)
}
- peer.device.awg.Mux.RLock()
- peer.device.awg.JunkCreator.CreateJunkPackets(&junks)
- peer.device.awg.Mux.RUnlock()
+ }
- if len(junks) > 0 {
- err = peer.SendBuffers(junks)
+ jc := peer.device.junk.count
+ jmin := peer.device.junk.min
+ jmax := peer.device.junk.max
- if err != nil {
- peer.device.log.Errorf("%v - Failed to send junk packets: %v", peer, err)
- return err
- }
- }
+ for i := 0; i < jc; i++ {
+ nBig, _ := rand.Int(rand.Reader, big.NewInt(int64(jmax-jmin+1)))
+ n := int(nBig.Int64()) + jmin
- junkedHeader, err = peer.device.awg.CreateInitHeaderJunk()
- if err != nil {
- peer.device.log.Errorf("%v - %v", peer, err)
- return err
- }
+ buf := make([]byte, n)
+ rand.Read(buf)
+ sendBuffer = append(sendBuffer, buf)
}
var buf [MessageInitiationSize]byte
@@ -165,14 +154,20 @@ func (peer *Peer) SendHandshakeInitiation(isRetry bool) error {
binary.Write(writer, binary.LittleEndian, msg)
packet := writer.Bytes()
peer.cookieGenerator.AddMacs(packet)
- junkedHeader = append(junkedHeader, packet...)
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
- sendBuffer = append(sendBuffer, junkedHeader)
+ if padding := peer.device.paddings.init; padding > 0 {
+ buf := make([]byte, padding+len(packet))
+ rand.Read(buf[:padding])
+ copy(buf[padding:], packet)
+ packet = buf
+ }
- err = peer.SendAndCountBuffers(sendBuffer)
+ sendBuffer = append(sendBuffer, packet)
+
+ err = peer.SendBuffers(sendBuffer)
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake initiation: %v", peer, err)
}
@@ -194,19 +189,12 @@ func (peer *Peer) SendHandshakeResponse() error {
return err
}
- junkedHeader, err := peer.device.awg.CreateResponseHeaderJunk()
- if err != nil {
- peer.device.log.Errorf("%v - %v", peer, err)
- return err
- }
-
var buf [MessageResponseSize]byte
writer := bytes.NewBuffer(buf[:0])
binary.Write(writer, binary.LittleEndian, response)
packet := writer.Bytes()
peer.cookieGenerator.AddMacs(packet)
- junkedHeader = append(junkedHeader, packet...)
err = peer.BeginSymmetricSession()
if err != nil {
@@ -218,32 +206,26 @@ func (peer *Peer) SendHandshakeResponse() error {
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
+ if padding := peer.device.paddings.response; padding > 0 {
+ buf := make([]byte, padding+len(packet))
+ rand.Read(buf[:padding])
+ copy(buf[padding:], packet)
+ packet = buf
+ }
+
// TODO: allocation could be avoided
- err = peer.SendAndCountBuffers([][]byte{junkedHeader})
+ err = peer.SendBuffers([][]byte{packet})
if err != nil {
peer.device.log.Errorf("%v - Failed to send handshake response: %v", peer, err)
}
return err
}
-func (device *Device) SendHandshakeCookie(
- initiatingElem *QueueHandshakeElement,
-) error {
+func (device *Device) SendHandshakeCookie(initiatingElem *QueueHandshakeElement) error {
device.log.Verbosef("Sending cookie response for denied handshake message for %v", initiatingElem.endpoint.DstToString())
sender := binary.LittleEndian.Uint32(initiatingElem.packet[4:8])
- msgType := DefaultMessageCookieReplyType
- if device.isAWG() {
- device.awg.Mux.RLock()
-
- var err error
- msgType, err = device.awg.GetMsgType(DefaultMessageCookieReplyType)
- device.awg.Mux.RUnlock()
- if err != nil {
- device.log.Errorf("Get message type for cookie reply: %v", err)
- return err
- }
- }
+ msgType := device.headers.cookie.Generate()
reply, err := device.cookieChecker.CreateReply(
initiatingElem.packet,
@@ -256,19 +238,20 @@ func (device *Device) SendHandshakeCookie(
return err
}
- junkedHeader, err := device.awg.CreateCookieReplyHeaderJunk()
- if err != nil {
- device.log.Errorf("%v - %v", device, err)
- return err
- }
-
var buf [MessageCookieReplySize]byte
writer := bytes.NewBuffer(buf[:0])
binary.Write(writer, binary.LittleEndian, reply)
+ packet := writer.Bytes()
+
+ if padding := device.paddings.cookie; padding > 0 {
+ buf := make([]byte, padding+len(packet))
+ rand.Read(buf[:padding])
+ copy(buf[padding:], packet)
+ packet = buf
+ }
- junkedHeader = append(junkedHeader, writer.Bytes()...)
// TODO: allocation could be avoided
- device.net.bind.Send([][]byte{junkedHeader}, initiatingElem.endpoint)
+ device.net.bind.Send([][]byte{packet}, initiatingElem.endpoint)
return nil
}
@@ -532,18 +515,7 @@ func (device *Device) RoutineEncryption(id int) {
fieldReceiver := header[4:8]
fieldNonce := header[8:16]
- msgType := DefaultMessageTransportType
- if device.isAWG() {
- device.awg.Mux.RLock()
-
- var err error
- msgType, err = device.awg.GetMsgType(DefaultMessageTransportType)
- device.awg.Mux.RUnlock()
- if err != nil {
- device.log.Errorf("get message type for transport: %v", err)
- continue
- }
- }
+ msgType := device.headers.transport.Generate()
binary.LittleEndian.PutUint32(fieldType, msgType)
binary.LittleEndian.PutUint32(fieldReceiver, elem.keypair.remoteIndex)
@@ -603,13 +575,15 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
if len(elem.packet) != MessageKeepaliveSize {
dataSent = true
- junkedHeader, err := device.awg.CreateTransportHeaderJunk(len(elem.packet))
- if err != nil {
- device.log.Errorf("%v - %v", device, err)
- continue
+ if padding := device.paddings.transport; padding > 0 {
+ // elem.packet is stored at the start of elem.buffer
+ // with zero padding
+ for i := len(elem.packet) - 1; i >= 0; i-- {
+ elem.buffer[i+padding] = elem.buffer[i]
+ }
+ rand.Read(elem.buffer[:padding])
+ elem.packet = elem.buffer[:padding+len(elem.packet)]
}
-
- elem.packet = append(junkedHeader, elem.packet...)
}
bufs = append(bufs, elem.packet)
}
@@ -617,7 +591,7 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
peer.timersAnyAuthenticatedPacketTraversal()
peer.timersAnyAuthenticatedPacketSent()
- err := peer.SendAndCountBuffers(bufs)
+ err := peer.SendBuffers(bufs)
if dataSent {
peer.timersDataSent()
}
diff --git a/device/uapi.go b/device/uapi.go
index 6c4be05..cff247e 100644
--- a/device/uapi.go
+++ b/device/uapi.go
@@ -18,7 +18,6 @@ import (
"sync"
"time"
- "github.com/amnezia-vpn/amneziawg-go/device/awg"
"github.com/amnezia-vpn/amneziawg-go/ipc"
)
@@ -98,42 +97,53 @@ func (device *Device) IpcGetOperation(w io.Writer) error {
sendf("fwmark=%d", device.net.fwmark)
}
- if device.isAWG() {
- if device.awg.Cfg.JunkPacketCount != 0 {
- sendf("jc=%d", device.awg.Cfg.JunkPacketCount)
- }
- if device.awg.Cfg.JunkPacketMinSize != 0 {
- sendf("jmin=%d", device.awg.Cfg.JunkPacketMinSize)
- }
- if device.awg.Cfg.JunkPacketMaxSize != 0 {
- sendf("jmax=%d", device.awg.Cfg.JunkPacketMaxSize)
- }
- if device.awg.Cfg.InitHeaderJunkSize != 0 {
- sendf("s1=%d", device.awg.Cfg.InitHeaderJunkSize)
- }
- if device.awg.Cfg.ResponseHeaderJunkSize != 0 {
- sendf("s2=%d", device.awg.Cfg.ResponseHeaderJunkSize)
- }
- if device.awg.Cfg.CookieReplyHeaderJunkSize != 0 {
- sendf("s3=%d", device.awg.Cfg.CookieReplyHeaderJunkSize)
- }
- if device.awg.Cfg.TransportHeaderJunkSize != 0 {
- sendf("s4=%d", device.awg.Cfg.TransportHeaderJunkSize)
- }
- for i, magicHeader := range device.awg.Cfg.MagicHeaders.Values {
- if magicHeader.Min > 4 {
- if magicHeader.Min == magicHeader.Max {
- sendf("h%d=%d", i+1, magicHeader.Min)
- continue
- }
+ if device.junk.count != 0 {
+ sendf("jc=%d", device.junk.count)
+ }
- sendf("h%d=%d-%d", i+1, magicHeader.Min, magicHeader.Max)
- }
- }
+ if device.junk.min != 0 {
+ sendf("jmin=%d", device.junk.min)
+ }
- specialJunkIpcFields := device.awg.HandshakeHandler.SpecialJunk.IpcGetFields()
- for _, field := range specialJunkIpcFields {
- sendf("%s=%s", field.Key, field.Value)
+ if device.junk.max != 0 {
+ sendf("jmax=%d", device.junk.max)
+ }
+
+ if device.paddings.init != 0 {
+ sendf("s1=%d", device.paddings.init)
+ }
+
+ if device.paddings.response != 0 {
+ sendf("s2=%d", device.paddings.response)
+ }
+
+ if device.paddings.cookie != 0 {
+ sendf("s3=%d", device.paddings.cookie)
+ }
+
+ if device.paddings.transport != 0 {
+ sendf("s4=%d", device.paddings.transport)
+ }
+
+ if device.headers.init != nil {
+ sendf("h1=%s", device.headers.init.GenSpec())
+ }
+
+ if device.headers.response != nil {
+ sendf("h2=%s", device.headers.response.GenSpec())
+ }
+
+ if device.headers.cookie != nil {
+ sendf("h3=%s", device.headers.cookie.GenSpec())
+ }
+
+ if device.headers.transport != nil {
+ sendf("h4=%s", device.headers.transport.GenSpec())
+ }
+
+ for i, ipacket := range device.ipackets {
+ if ipacket != nil {
+ sendf("i%d=%s", i+1, ipacket.Spec)
}
}
@@ -187,20 +197,18 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
}
}()
+ ipcDev := new(ipcSetDevice)
peer := new(ipcSetPeer)
deviceConfig := true
- tempAwg := awg.Protocol{}
- tempAwg.Cfg.MagicHeaders.Values = make([]awg.MagicHeader, 4)
-
scanner := bufio.NewScanner(r)
for scanner.Scan() {
line := scanner.Text()
if line == "" {
// Blank line means terminate operation.
- err := device.handlePostConfig(&tempAwg)
+ err := ipcDev.mergeWithDevice(device)
if err != nil {
- return err
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to merge with device: %w", err)
}
peer.handlePostConfig()
return nil
@@ -229,7 +237,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
var err error
if deviceConfig {
- err = device.handleDeviceLine(key, value, &tempAwg)
+ err = device.handleDeviceLine(key, value)
} else {
err = device.handlePeerLine(peer, key, value)
}
@@ -237,9 +245,9 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return err
}
}
- err = device.handlePostConfig(&tempAwg)
+ err = ipcDev.mergeWithDevice(device)
if err != nil {
- return err
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to merge with device: %w", err)
}
peer.handlePostConfig()
@@ -249,7 +257,7 @@ func (device *Device) IpcSetOperation(r io.Reader) (err error) {
return nil
}
-func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol) error {
+func (device *Device) handleDeviceLine(key, value string) error {
switch key {
case "private_key":
var sk NoisePrivateKey
@@ -300,112 +308,145 @@ func (device *Device) handleDeviceLine(key, value string, tempAwg *awg.Protocol)
device.RemoveAllPeers()
case "jc":
- junkPacketCount, err := strconv.Atoi(value)
+ jc, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_count %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jc: %w", err)
}
- device.log.Verbosef("UAPI: Updating junk_packet_count")
- tempAwg.Cfg.JunkPacketCount = junkPacketCount
- tempAwg.Cfg.IsSet = true
+ if jc <= 0 {
+ return ipcErrorf(ipc.IpcErrorInvalid, "jc must be a positive value")
+ }
+ device.log.Verbosef("UAPI: Updating junk count")
+ device.junk.count = jc
case "jmin":
- junkPacketMinSize, err := strconv.Atoi(value)
+ jmin, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_min_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmin: %w", err)
}
- device.log.Verbosef("UAPI: Updating junk_packet_min_size")
- tempAwg.Cfg.JunkPacketMinSize = junkPacketMinSize
- tempAwg.Cfg.IsSet = true
+ if jmin <= 0 {
+ return ipcErrorf(ipc.IpcErrorInvalid, "jmin must be a positive value")
+ }
+ device.log.Verbosef("UAPI: Updating junk min")
+ device.junk.min = jmin
case "jmax":
- junkPacketMaxSize, err := strconv.Atoi(value)
+ jmax, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse junk_packet_max_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse jmax: %w", err)
}
- device.log.Verbosef("UAPI: Updating junk_packet_max_size")
- tempAwg.Cfg.JunkPacketMaxSize = junkPacketMaxSize
- tempAwg.Cfg.IsSet = true
+ if jmax <= 0 {
+ return ipcErrorf(ipc.IpcErrorInvalid, "jmax must be a positive value")
+ }
+ device.log.Verbosef("UAPI: Updating junk max")
+ device.junk.max = jmax
case "s1":
- initPacketJunkSize, err := strconv.Atoi(value)
+ padding, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse init_packet_junk_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s1: %w", err)
}
- device.log.Verbosef("UAPI: Updating init_packet_junk_size")
- tempAwg.Cfg.InitHeaderJunkSize = initPacketJunkSize
- tempAwg.Cfg.IsSet = true
+ if padding < 0 {
+ return ipcErrorf(ipc.IpcErrorInvalid, "s1 must be non-negative")
+ }
+ device.log.Verbosef("UAPI: Updating s1 padding")
+ device.paddings.init = padding
case "s2":
- responsePacketJunkSize, err := strconv.Atoi(value)
+ padding, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse response_packet_junk_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s2: %w", err)
}
- device.log.Verbosef("UAPI: Updating response_packet_junk_size")
- tempAwg.Cfg.ResponseHeaderJunkSize = responsePacketJunkSize
- tempAwg.Cfg.IsSet = true
+ if padding < 0 {
+ return ipcErrorf(ipc.IpcErrorInvalid, "s2 must be non-negative")
+ }
+ device.log.Verbosef("UAPI: Updating s2 padding")
+ device.paddings.response = padding
case "s3":
- cookieReplyPacketJunkSize, err := strconv.Atoi(value)
+ padding, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse cookie_reply_packet_junk_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s3: %w", err)
}
- device.log.Verbosef("UAPI: Updating cookie_reply_packet_junk_size")
- tempAwg.Cfg.CookieReplyHeaderJunkSize = cookieReplyPacketJunkSize
- tempAwg.Cfg.IsSet = true
+ if padding < 0 {
+ return ipcErrorf(ipc.IpcErrorInvalid, "s3 must be non-negative")
+ }
+ device.log.Verbosef("UAPI: Updating s3 padding")
+ device.paddings.cookie = padding
case "s4":
- transportPacketJunkSize, err := strconv.Atoi(value)
+ padding, err := strconv.Atoi(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "parse transport_packet_junk_size %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse s4: %w", err)
}
- device.log.Verbosef("UAPI: Updating transport_packet_junk_size")
- tempAwg.Cfg.TransportHeaderJunkSize = transportPacketJunkSize
- tempAwg.Cfg.IsSet = true
+ if padding < 0 {
+ return ipcErrorf(ipc.IpcErrorInvalid, "s4 must be non-negative")
+ }
+ device.log.Verbosef("UAPI: Updating s4 padding")
+ device.paddings.transport = padding
+
case "h1":
- initMagicHeader, err := awg.ParseMagicHeader(key, value)
+ header, err := newMagicHeader(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H1: %w", err)
}
+ device.headers.init = header
- tempAwg.Cfg.MagicHeaders.Values[0] = initMagicHeader
- tempAwg.Cfg.IsSet = true
case "h2":
- responseMagicHeader, err := awg.ParseMagicHeader(key, value)
+ header, err := newMagicHeader(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H2: %w", err)
}
+ device.headers.response = header
- tempAwg.Cfg.MagicHeaders.Values[1] = responseMagicHeader
- tempAwg.Cfg.IsSet = true
case "h3":
- cookieReplyMagicHeader, err := awg.ParseMagicHeader(key, value)
+ header, err := newMagicHeader(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H3: %w", err)
}
+ device.headers.cookie = header
- tempAwg.Cfg.MagicHeaders.Values[2] = cookieReplyMagicHeader
- tempAwg.Cfg.IsSet = true
case "h4":
- transportMagicHeader, err := awg.ParseMagicHeader(key, value)
+ header, err := newMagicHeader(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "uapi: %w", err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse H4: %w", err)
}
+ device.headers.transport = header
- tempAwg.Cfg.MagicHeaders.Values[3] = transportMagicHeader
- tempAwg.Cfg.IsSet = true
- case "i1", "i2", "i3", "i4", "i5":
- if len(value) == 0 {
- device.log.Verbosef("UAPI: received empty %s", key)
- return nil
- }
-
- generators, err := awg.ParseTagJunkGenerator(key, value)
+ case "i1":
+ chain, err := newObfChain(value)
if err != nil {
- return ipcErrorf(ipc.IpcErrorInvalid, "invalid %s: %w", key, err)
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I1: %w", err)
}
- device.log.Verbosef("UAPI: Updating %s", key)
- tempAwg.HandshakeHandler.SpecialJunk.AppendGenerator(generators)
- tempAwg.HandshakeHandler.IsSet = true
+ device.ipackets[0] = chain
+
+ case "i2":
+ chain, err := newObfChain(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I2: %w", err)
+ }
+ device.ipackets[1] = chain
+
+ case "i3":
+ chain, err := newObfChain(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I3: %w", err)
+ }
+ device.ipackets[2] = chain
+
+ case "i4":
+ chain, err := newObfChain(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I4: %w", err)
+ }
+ device.ipackets[3] = chain
+
+ case "i5":
+ chain, err := newObfChain(value)
+ if err != nil {
+ return ipcErrorf(ipc.IpcErrorInvalid, "failed to parse I5: %w", err)
+ }
+ device.ipackets[4] = chain
+
default:
return ipcErrorf(ipc.IpcErrorInvalid, "invalid UAPI device key: %v", key)
}
@@ -654,3 +695,49 @@ func (device *Device) IpcHandle(socket net.Conn) {
buffered.Flush()
}
}
+
+type ipcSetDevice struct {
+ headers struct {
+ init *magicHeader
+ response *magicHeader
+ cookie *magicHeader
+ transport *magicHeader
+ }
+}
+
+func (d *ipcSetDevice) mergeWithDevice(device *Device) error {
+ if d.headers.init == nil {
+ d.headers.init = device.headers.init
+ }
+
+ if d.headers.response == nil {
+ d.headers.response = device.headers.response
+ }
+
+ if d.headers.cookie == nil {
+ d.headers.cookie = device.headers.cookie
+ }
+
+ if d.headers.transport == nil {
+ d.headers.transport = device.headers.transport
+ }
+
+ headers := []*magicHeader{d.headers.init, d.headers.response, d.headers.cookie, d.headers.transport}
+ for i := 0; i < len(headers); i++ {
+ for j := i + 1; j < len(headers); j++ {
+ left := headers[i]
+ right := headers[j]
+
+ if left.start <= right.end && right.start <= left.end {
+ return errors.New("headers must not overlap")
+ }
+ }
+ }
+
+ device.headers.init = d.headers.init
+ device.headers.response = d.headers.response
+ device.headers.cookie = d.headers.cookie
+ device.headers.transport = d.headers.transport
+
+ return nil
+}
From 730d6c39d0c4e348a3d080bebe496664215e5c99 Mon Sep 17 00:00:00 2001
From: Yaroslav Gurov
Date: Sun, 30 Nov 2025 16:14:47 +0100
Subject: [PATCH 70/75] chore: add docs for the params from awg2
---
README.md | 63 ++++++++++++++++++++++++++++++++++++++++++++++++++++++-
1 file changed, 62 insertions(+), 1 deletion(-)
diff --git a/README.md b/README.md
index 428b752..f98db43 100644
--- a/README.md
+++ b/README.md
@@ -50,4 +50,65 @@ $ git clone https://github.com/amnezia-vpn/amneziawg-go
$ cd amneziawg-go
$ make
```
-
+
+## Configuration
+
+> [!NOTE]
+> If there is no value specified (for any param), AWG treats it as 0
+
+### Junk packets
+
+The amount of junk packets specified in `Jc` with a random size between `Jmin` and `Jmax` would be generated and sent prior every handshake
+
+- `Jc: int`, recommended range is 4-12
+- `Jmin: int` <= `Jmax:int`
+
+> [!TIP]
+> Junk packets do not carry any actual data, so there is no need to specify it on both sides. General recommendation is to use it on the client side only
+
+> [!IMPORTANT]
+> If Jmax >= system MTU (not the one specified in AWG), then the system can fracture this packet into fragments, which looks suspicious from the censor side
+
+### Message paddings
+
+- `S1: int` - padding of handshake initial message
+- `S2: int` - padding of handshake response message
+- `S3: int` - padding of handshake cookie message
+- `S4: int` - padding of transport messages
+
+### Message headers
+
+Every message in wireguard has `int32` type at the beginning of the packet. This field could be controlled by specifying the params below:
+
+- `H1: string` - header range of handshake initial message
+- `H2: string` - header range of handshake initial message
+- `H3: string` - header range of handshake cookie message
+- `H4: string` - header range of transport message
+
+Values could be specified as:
+- range: `x-y`, x <= y; e.g. `123-456`
+- single value `1234`
+
+### Custom signature packets
+
+These packets are being send prior to every handshake, in the same way as Junk packets do. The sending order is `I1`, `I2`, `I3`, `I4`, `I5`. If there is no value specified, the packet is skipped.
+
+- `I1: string`
+- `I2: string`
+- `I3: string`
+- `I4: string`
+- `I5: string`
+
+Value is a sequence of tags specified below:
+- `` - static bytes tag. Dumps `[seq]` as-is to the packet. `[seq]` is hex-encoded sequence which represents bytes sequence (2 hex numbers per byte) and is always even-sized
+- `` - random bytes tag. Dumps `[size]` amount of randomly-generated bytes to the packet
+- `` - random digits tag. Dumps `[size]` amount of randomly-generated bytes from `[0-9]` set to the packet
+- `` - random chars tag. Dumps `[size]` amount of randomly-generated bytes from `[a-zA-Z] set to the packet
+- `` - timestamp tag. Dumps 4-bytes long current system time in UNIX format
+- `` - packet counter tag. Dumps 4-bytes long amount of packets sent by AWG
+
+> [!TIP]
+> Custom signature packets does not carry any actual data, so there is no need to specify it on both sides. General recommendation is to use it on the client side only
+
+> [!IMPORTANT]
+> If the final size of any packet exceeds system MTU, it would be fractured into fragments, which looks suspicious
\ No newline at end of file
From e796d477d89e6851b2bb4871bf75f1e621f94ace Mon Sep 17 00:00:00 2001
From: vkamn
Date: Thu, 11 Dec 2025 18:56:42 +0800
Subject: [PATCH 71/75] chore: update license (#105)
Signed-off-by: vkamn
---
LICENSE | 2 ++
1 file changed, 2 insertions(+)
diff --git a/LICENSE b/LICENSE
index f85e365..ab45fb3 100644
--- a/LICENSE
+++ b/LICENSE
@@ -1,3 +1,5 @@
+Copyright (C) 2017-2025 WireGuard LLC. All Rights Reserved.
+
Permission is hereby granted, free of charge, to any person obtaining a copy of
this software and associated documentation files (the "Software"), to deal in
the Software without restriction, including without limitation the rights to
From 449d7cffd4adf86971bd679d0be5384b443e8be5 Mon Sep 17 00:00:00 2001
From: Yaroslav Gurov <31506978+ygurov@users.noreply.github.com>
Date: Fri, 19 Dec 2025 03:14:48 +0100
Subject: [PATCH 72/75] Feature/outline glue (#106)
* feat: added outline integration layer
* chore: make the function used in RegisterFallbackParser a standalone one
* fix: check if domain has a dot prior trimming it
* fix: use net.JoinHostPort instead of plain concat
---
go.mod | 17 ++-
go.sum | 34 ++++++
outline/dialer.go | 77 ++++++++++++++
outline/fallback.go | 224 +++++++++++++++++++++++++++++++++++++++
outline/fallback_test.go | 52 +++++++++
5 files changed, 400 insertions(+), 4 deletions(-)
create mode 100644 outline/dialer.go
create mode 100644 outline/fallback.go
create mode 100644 outline/fallback_test.go
diff --git a/go.mod b/go.mod
index 8c4372d..a5f6548 100644
--- a/go.mod
+++ b/go.mod
@@ -6,18 +6,27 @@ require (
github.com/stretchr/testify v1.10.0
github.com/tevino/abool v1.2.0
go.uber.org/atomic v1.11.0
- golang.org/x/crypto v0.39.0
- golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
- golang.org/x/net v0.41.0
- golang.org/x/sys v0.33.0
+ golang.org/x/crypto v0.42.0
+ golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
+ golang.org/x/net v0.44.0
+ golang.org/x/sys v0.36.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
gvisor.dev/gvisor v0.0.0-20231202080848-1f7806d17489
)
require (
+ github.com/Jigsaw-Code/outline-sdk v0.0.20 // indirect
+ github.com/Jigsaw-Code/outline-sdk/x v0.0.8 // indirect
github.com/davecgh/go-spew v1.1.1 // indirect
+ github.com/goccy/go-yaml v1.17.1 // indirect
github.com/google/btree v1.1.3 // indirect
+ github.com/gorilla/websocket v1.5.3 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect
+ github.com/shadowsocks/go-shadowsocks2 v0.1.5 // indirect
+ golang.org/x/mobile v0.0.0-20240520174638-fa72addaaa1b // indirect
+ golang.org/x/mod v0.28.0 // indirect
+ golang.org/x/sync v0.17.0 // indirect
golang.org/x/time v0.9.0 // indirect
+ golang.org/x/tools v0.37.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
)
diff --git a/go.sum b/go.sum
index 3d8b3c2..8dde512 100644
--- a/go.sum
+++ b/go.sum
@@ -1,25 +1,59 @@
+github.com/Jigsaw-Code/outline-sdk v0.0.20 h1:4ep7MK9lFmcyPIRIbn4xrP1VKdJNsqR6+iJEOHDKnNg=
+github.com/Jigsaw-Code/outline-sdk v0.0.20/go.mod h1:CFDKyGZA4zatKE4vMLe8TyQpZCyINOeRFbMAmYHxodw=
+github.com/Jigsaw-Code/outline-sdk/x v0.0.8 h1:fFHFXW7CKhRiegyNSdP25S/WIiVrRnMKysHDoO/N2Xg=
+github.com/Jigsaw-Code/outline-sdk/x v0.0.8/go.mod h1:zqSH7yEYIQ0pYOhrr4QnodATVb5X/eZXV4AjUp9zhvs=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
+github.com/goccy/go-yaml v1.17.1 h1:LI34wktB2xEE3ONG/2Ar54+/HJVBriAGJ55PHls4YuY=
+github.com/goccy/go-yaml v1.17.1/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA=
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
+github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
+github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
+github.com/riobard/go-bloom v0.0.0-20200614022211-cdc8013cb5b3/go.mod h1:HgjTstvQsPGkxUsCd2KWxErBblirPizecHcpD3ffK+s=
+github.com/shadowsocks/go-shadowsocks2 v0.1.5 h1:PDSQv9y2S85Fl7VBeOMF9StzeXZyK1HakRm86CUbr28=
+github.com/shadowsocks/go-shadowsocks2 v0.1.5/go.mod h1:AGGpIoek4HRno4xzyFiAtLHkOpcoznZEkAccaI/rplM=
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/tevino/abool v1.2.0 h1:heAkClL8H6w+mK5md9dzsuohKeXHUpY7Vw0ZCKW+huA=
github.com/tevino/abool v1.2.0/go.mod h1:qc66Pna1RiIsPa7O4Egxxs9OqkuxDX55zznh9K07Tzg=
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
+golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
+golang.org/x/crypto v0.0.0-20210220033148-5ea612d1eb83/go.mod h1:jdWPYTVW3xRLrWPugEBEK3UY2ZEsg3UU495nc5E+M+I=
golang.org/x/crypto v0.39.0 h1:SHs+kF4LP+f+p14esP5jAoDpHU8Gu/v9lFRK6IT5imM=
golang.org/x/crypto v0.39.0/go.mod h1:L+Xg3Wf6HoL4Bn4238Z6ft6KfEpN0tJGo53AAPC632U=
+golang.org/x/crypto v0.42.0 h1:chiH31gIWm57EkTXpwnqf8qeuMUi0yekh6mT2AvFlqI=
+golang.org/x/crypto v0.42.0/go.mod h1:4+rDnOTJhQCx2q7/j6rAN5XDw8kPjeaXEUR2eL94ix8=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
+golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
+golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
+golang.org/x/mobile v0.0.0-20240520174638-fa72addaaa1b h1:WX7nnnLfCEXg+FmdYZPai2XuP3VqCP1HZVMST0n9DF0=
+golang.org/x/mobile v0.0.0-20240520174638-fa72addaaa1b/go.mod h1:EiXZlVfUTaAyySFVJb9rsODuiO+WXu8HrUuySb7nYFw=
+golang.org/x/mod v0.28.0 h1:gQBtGhjxykdjY9YhZpSlZIsbnaE2+PgjfLWUQTnoZ1U=
+golang.org/x/mod v0.28.0/go.mod h1:yfB/L0NOf/kmEbXjzCPOx1iK1fRutOydrCMsqRhEBxI=
+golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.41.0 h1:vBTly1HeNPEn3wtREYfy4GZ/NECgw2Cnl+nK6Nz3uvw=
golang.org/x/net v0.41.0/go.mod h1:B/K4NNqkfmg07DQYrbwvSluqCJOOXwUjeb/5lOisjbA=
+golang.org/x/net v0.44.0 h1:evd8IRDyfNBMBTTY5XRF1vaZlD+EmWx6x8PkhR04H/I=
+golang.org/x/net v0.44.0/go.mod h1:ECOoLqd5U3Lhyeyo/QDCEVQ4sNgYsqvCZ722XogGieY=
+golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
+golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
+golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
+golang.org/x/sys v0.0.0-20191026070338-33540a1f6037/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
+golang.org/x/sys v0.36.0 h1:KVRy2GtZBrk1cBYA7MKu5bEZFxQk4NIDV6RLVcC8o0k=
+golang.org/x/sys v0.36.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
+golang.org/x/term v0.0.0-20201117132131-f5c789dd3221/go.mod h1:Nr5EML6q2oocZ2LXRh80K7BxOlk5/8JxuGnuhpl+muw=
+golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/time v0.9.0 h1:EsRrnYcQiGH+5FfbgvV4AP7qEZstoyrHB0DzarOQ4ZY=
golang.org/x/time v0.9.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
+golang.org/x/tools v0.37.0 h1:DVSRzp7FwePZW356yEAChSdNcQo6Nsp+fex1SUW09lE=
+golang.org/x/tools v0.37.0/go.mod h1:MBN5QPQtLMHVdvsbtarmTNukZDdgwdwlO5qGacAzF0w=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
diff --git a/outline/dialer.go b/outline/dialer.go
new file mode 100644
index 0000000..86aaff5
--- /dev/null
+++ b/outline/dialer.go
@@ -0,0 +1,77 @@
+package outline
+
+import (
+ "context"
+ "fmt"
+ "net"
+ "net/netip"
+
+ "github.com/Jigsaw-Code/outline-sdk/transport"
+ "github.com/amnezia-vpn/amneziawg-go/conn"
+ "github.com/amnezia-vpn/amneziawg-go/device"
+ "github.com/amnezia-vpn/amneziawg-go/tun/netstack"
+ "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
+)
+
+type DialerOptions struct {
+ Ipc string
+ Prefixes []netip.Prefix
+ Mtu int
+ Dns []netip.Addr
+}
+
+func NewStreamDialer(opts DialerOptions) (*StreamDialer, error) {
+ var localAddresses []netip.Addr
+ for _, prefix := range opts.Prefixes {
+ localAddresses = append(localAddresses, prefix.Addr())
+ }
+
+ tun, tnet, err := netstack.CreateNetTUN(localAddresses, opts.Dns, opts.Mtu)
+ if err != nil {
+ return nil, fmt.Errorf("failed to create network tun: %v", err)
+ }
+
+ awgLogger := device.Logger{
+ Verbosef: func(format string, args ...any) {
+ },
+ Errorf: func(format string, args ...any) {
+ },
+ }
+
+ dev := device.NewDevice(tun, conn.NewDefaultBind(), &awgLogger)
+ if err := dev.IpcSet(opts.Ipc); err != nil {
+ return nil, fmt.Errorf("failed to configure device: %v", err)
+ }
+
+ if err := dev.Up(); err != nil {
+ return nil, fmt.Errorf("failed to start awg device: %v", err)
+ }
+
+ return &StreamDialer{
+ tnet: tnet,
+ }, nil
+}
+
+var _ transport.StreamDialer = (*StreamDialer)(nil)
+
+type StreamDialer struct {
+ tnet *netstack.Net
+}
+
+func (d *StreamDialer) DialStream(ctx context.Context, raddr string) (transport.StreamConn, error) {
+ host, port, err := net.SplitHostPort(raddr)
+ if err != nil {
+ return nil, fmt.Errorf("failed to parse raddr: %v", err)
+ }
+ if l := len(host); l > 0 && host[l-1] == '.' {
+ host = host[:l-1]
+ raddr = net.JoinHostPort(host, port)
+ }
+
+ conn, err := d.tnet.DialContext(ctx, "tcp", raddr)
+ if err != nil {
+ return nil, err
+ }
+
+ return conn.(*gonet.TCPConn), nil
+}
diff --git a/outline/fallback.go b/outline/fallback.go
new file mode 100644
index 0000000..0f75dc9
--- /dev/null
+++ b/outline/fallback.go
@@ -0,0 +1,224 @@
+package outline
+
+import (
+ "context"
+ "encoding/base64"
+ "encoding/hex"
+ "fmt"
+ "net/netip"
+ "strconv"
+ "strings"
+
+ "github.com/Jigsaw-Code/outline-sdk/transport"
+ "github.com/Jigsaw-Code/outline-sdk/x/mobileproxy"
+ "github.com/Jigsaw-Code/outline-sdk/x/smart"
+ "github.com/goccy/go-yaml"
+)
+
+type DeviceConfig struct {
+ PrivateKey string `yaml:"private_key"`
+ Address []string `yaml:"address"`
+ Dns []string `yaml:"dns"`
+ Mtu int `yaml:"mtu,omitempty"`
+ Jc int `yaml:"jc,omitempty"`
+ Jmin int `yaml:"jmin,omitempty"`
+ Jmax int `yaml:"jmax,omitempty"`
+ S1 int `yaml:"s1,omitempty"`
+ S2 int `yaml:"s2,omitempty"`
+ S3 int `yaml:"s3,omitempty"`
+ S4 int `yaml:"s4,omitempty"`
+ H1 string `yaml:"h1,omitempty"`
+ H2 string `yaml:"h2,omitempty"`
+ H3 string `yaml:"h3,omitempty"`
+ H4 string `yaml:"h4,omitempty"`
+ I1 string `yaml:"i1,omitempty"`
+ I2 string `yaml:"i2,omitempty"`
+ I3 string `yaml:"i3,omitempty"`
+ I4 string `yaml:"i4,omitempty"`
+ I5 string `yaml:"i5,omitempty"`
+ Peers []PeerConfig `yaml:"peers,omitempty"`
+}
+
+type PeerConfig struct {
+ PublicKey string `yaml:"public_key"`
+ PresharedKey string `yaml:"preshared_key,omitempty"`
+ Endpoint string `yaml:"endpoint"`
+ AllowedIPs []string `yaml:"allowed_ips"`
+ PersistentKeepaliveInterval uint16 `yaml:"persistent_keepalive_interval,omitempty"`
+}
+
+func mapYamlToConfig(y smart.YAMLNode) (*DeviceConfig, error) {
+ bytes, err := yaml.Marshal(y)
+ if err != nil {
+ return nil, fmt.Errorf("failed to marshal yaml: %v", err)
+ }
+
+ var cfg DeviceConfig
+ if err = yaml.Unmarshal(bytes, &cfg); err != nil {
+ return nil, fmt.Errorf("failed to unmarshal yaml: %v", err)
+ }
+
+ return &cfg, nil
+}
+
+func genIpcString(cfg *DeviceConfig) (string, error) {
+ privateKeyBytes, err := base64.StdEncoding.DecodeString(cfg.PrivateKey)
+ if err != nil {
+ return "", fmt.Errorf("failed to decode private key: %v", err)
+ }
+
+ var b strings.Builder
+
+ b.WriteString("private_key=")
+ b.WriteString(hex.EncodeToString(privateKeyBytes))
+
+ if cfg.Jc != 0 {
+ b.WriteString("\njc=")
+ b.WriteString(strconv.Itoa(cfg.Jc))
+ }
+ if cfg.Jmin != 0 {
+ b.WriteString("\njmin=")
+ b.WriteString(strconv.Itoa(cfg.Jmin))
+ }
+ if cfg.Jmax != 0 {
+ b.WriteString("\njmax=")
+ b.WriteString(strconv.Itoa(cfg.Jmax))
+ }
+ if cfg.S1 != 0 {
+ b.WriteString("\ns1=")
+ b.WriteString(strconv.Itoa(cfg.S1))
+ }
+ if cfg.S2 != 0 {
+ b.WriteString("\ns2=")
+ b.WriteString(strconv.Itoa(cfg.S2))
+ }
+ if cfg.S3 != 0 {
+ b.WriteString("\ns3=")
+ b.WriteString(strconv.Itoa(cfg.S3))
+ }
+ if cfg.S4 != 0 {
+ b.WriteString("\ns4=")
+ b.WriteString(strconv.Itoa(cfg.S4))
+ }
+ if cfg.H1 != "" {
+ b.WriteString("\nh1=")
+ b.WriteString(cfg.H1)
+ }
+ if cfg.H2 != "" {
+ b.WriteString("\nh2=")
+ b.WriteString(cfg.H2)
+ }
+ if cfg.H3 != "" {
+ b.WriteString("\nh3=")
+ b.WriteString(cfg.H3)
+ }
+ if cfg.H4 != "" {
+ b.WriteString("\nh4=")
+ b.WriteString(cfg.H4)
+ }
+ if cfg.I1 != "" {
+ b.WriteString("\ni1=")
+ b.WriteString(cfg.I1)
+ }
+ if cfg.I2 != "" {
+ b.WriteString("\ni2=")
+ b.WriteString(cfg.I2)
+ }
+ if cfg.I3 != "" {
+ b.WriteString("\ni3=")
+ b.WriteString(cfg.I3)
+ }
+ if cfg.I4 != "" {
+ b.WriteString("\ni4=")
+ b.WriteString(cfg.I4)
+ }
+ if cfg.I5 != "" {
+ b.WriteString("\ni5=")
+ b.WriteString(cfg.I5)
+ }
+
+ for _, peer := range cfg.Peers {
+ publicKeyBytes, err := base64.StdEncoding.DecodeString(peer.PublicKey)
+ if err != nil {
+ return "", fmt.Errorf("failed to decode public key: %v", err)
+ }
+
+ b.WriteString("\npublic_key=")
+ b.WriteString(hex.EncodeToString(publicKeyBytes))
+
+ b.WriteString("\nendpoint=")
+ b.WriteString(peer.Endpoint)
+
+ for _, allowedIp := range peer.AllowedIPs {
+ b.WriteString("\nallowed_ip=")
+ b.WriteString(allowedIp)
+ }
+
+ if peer.PresharedKey != "" {
+ presharedKeyBytes, err := base64.StdEncoding.DecodeString(peer.PresharedKey)
+ if err != nil {
+ return "", fmt.Errorf("failed to decode preshared key: %v", err)
+ }
+
+ b.WriteString("\npreshared_key=")
+ b.WriteString(hex.EncodeToString(presharedKeyBytes))
+ }
+
+ if peer.PersistentKeepaliveInterval != 0 {
+ b.WriteString("\npersistent_keepalive_interval=")
+ b.WriteString(strconv.Itoa(int(peer.PersistentKeepaliveInterval)))
+ }
+ }
+
+ return b.String(), nil
+}
+
+func FallbackParser(ctx context.Context, y smart.YAMLNode) (transport.StreamDialer, string, error) {
+ cfg, err := mapYamlToConfig(y)
+ if err != nil {
+ return nil, "", fmt.Errorf("failed to map yaml to config: %v", err)
+ }
+
+ ipc, err := genIpcString(cfg)
+ if err != nil {
+ return nil, "", fmt.Errorf("faield to generate ipc config: %v", err)
+ }
+
+ var prefixes []netip.Prefix
+ for _, address := range cfg.Address {
+ prefix, err := netip.ParsePrefix(address)
+ if err != nil {
+ return nil, "", fmt.Errorf("failed to parse address: %v", err)
+ }
+ prefixes = append(prefixes, prefix)
+ }
+
+ var dns []netip.Addr
+ for _, saddr := range cfg.Dns {
+ addr, err := netip.ParseAddr(saddr)
+ if err != nil {
+ return nil, "", fmt.Errorf("failed to parse dns: %v", err)
+ }
+ dns = append(dns, addr)
+ }
+
+ if cfg.Mtu == 0 {
+ cfg.Mtu = 1408
+ }
+
+ dialer, err := NewStreamDialer(DialerOptions{
+ Ipc: ipc,
+ Prefixes: prefixes,
+ Mtu: cfg.Mtu,
+ Dns: dns,
+ })
+ if err != nil {
+ return nil, "", fmt.Errorf("failed to create dialer: %v", err)
+ }
+
+ return dialer, ipc, nil
+}
+
+func RegisterFallbackParser(opt *mobileproxy.SmartDialerOptions, name string) {
+ opt.RegisterFallbackParser(name, FallbackParser)
+}
diff --git a/outline/fallback_test.go b/outline/fallback_test.go
new file mode 100644
index 0000000..496f949
--- /dev/null
+++ b/outline/fallback_test.go
@@ -0,0 +1,52 @@
+package outline_test
+
+import (
+ "testing"
+
+ "github.com/Jigsaw-Code/outline-sdk/x/mobileproxy"
+ awg "github.com/amnezia-vpn/amneziawg-go/outline"
+)
+
+const cfg = `
+dns:
+ - {system: {}}
+tls:
+ - ""
+fallback:
+ - awg:
+ address: [10.0.0.0/32]
+ dns: [8.8.8.8, 8.8.4.4]
+ private_key: +CdqlYvjqZ3OUr4mLWvGJo1h67CWpQwMIxA5OpyiJUM=
+ jc: 4
+ jmin: 50
+ jmax: 100
+ s1: 87
+ s2: 65
+ s3: 43
+ s4: 21
+ h1: 1000000000-1000000001
+ h2: 2000000000-2000000002
+ h3: 3000000000-3000000003
+ h4: 4000000000-4000000004
+ peers:
+ - public_key: EGxNYihRLKQ9nvdOE5j5aZ7rtw3ttzJS1xxaJpgYYHI=
+ preshared_key: 2OiSh6rP3t/g39jgJNGK70B+nize821yIFNtUqi8/XU=
+ endpoint: 123.123.123.123:51820
+ allowed_ips: [0.0.0.0/0, ::/0]
+ persistent_keepalive_interval: 25
+`
+
+var testDomains = mobileproxy.NewListFromLines("example.com")
+
+func Test_outlineIntegration(t *testing.T) {
+ opts := mobileproxy.NewSmartDialerOptions(testDomains, cfg)
+ opts.SetLogWriter(mobileproxy.NewStderrLogWriter())
+ awg.RegisterFallbackParser(opts, "awg")
+ dialer, err := opts.NewStreamDialer()
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, err = mobileproxy.RunProxy("", dialer); err != nil {
+ t.Fatal(err)
+ }
+}
From e7ef4339e718641fc7bc1b0ea41b538108de77cc Mon Sep 17 00:00:00 2001
From: Yaroslav Gurov
Date: Mon, 23 Mar 2026 11:01:42 +0000
Subject: [PATCH 73/75] readme: remove tag from tag reference
---
README.md | 1 -
1 file changed, 1 deletion(-)
diff --git a/README.md b/README.md
index f98db43..301b5bf 100644
--- a/README.md
+++ b/README.md
@@ -105,7 +105,6 @@ Value is a sequence of tags specified below:
- `` - random digits tag. Dumps `[size]` amount of randomly-generated bytes from `[0-9]` set to the packet
- `` - random chars tag. Dumps `[size]` amount of randomly-generated bytes from `[a-zA-Z] set to the packet
- `` - timestamp tag. Dumps 4-bytes long current system time in UNIX format
-- `` - packet counter tag. Dumps 4-bytes long amount of packets sent by AWG
> [!TIP]
> Custom signature packets does not carry any actual data, so there is no need to specify it on both sides. General recommendation is to use it on the client side only
From 12a012205e3c444be02aba91a840455f74c127e1 Mon Sep 17 00:00:00 2001
From: Yaroslav Gurov
Date: Tue, 31 Mar 2026 15:48:57 +0000
Subject: [PATCH 74/75] readme: actualize type for H1-H4
---
README.md | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/README.md b/README.md
index 301b5bf..e1c7309 100644
--- a/README.md
+++ b/README.md
@@ -78,7 +78,7 @@ The amount of junk packets specified in `Jc` with a random size between `Jmin` a
### Message headers
-Every message in wireguard has `int32` type at the beginning of the packet. This field could be controlled by specifying the params below:
+Every message in wireguard has `uint32` type at the beginning of the packet. This field could be controlled by specifying the params below:
- `H1: string` - header range of handshake initial message
- `H2: string` - header range of handshake initial message
From f4f4c999267437c3eb909e8d0e5278fb4596d9a7 Mon Sep 17 00:00:00 2001
From: admin
Date: Tue, 31 Mar 2026 16:37:57 +0300
Subject: [PATCH 75/75] fix: apply S4 transport padding to keepalive packets
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
Keepalive packets were excluded from S4 padding because the padding
logic was nested inside the dataSent guard. The receiving side
(DeterminePacketTypeAndPadding) expects S4 padding on all transport
packets, so unpadded keepalives fail H4 header validation and are
silently dropped.
This prevents the responder from completing key confirmation —
lastHandshakeNano stays 0 until real data flows through the tunnel.
---
device/send.go | 17 ++++++++---------
1 file changed, 8 insertions(+), 9 deletions(-)
diff --git a/device/send.go b/device/send.go
index 5e5cc1b..0cc57da 100644
--- a/device/send.go
+++ b/device/send.go
@@ -574,16 +574,15 @@ func (peer *Peer) RoutineSequentialSender(maxBatchSize int) {
for _, elem := range elemsContainer.elems {
if len(elem.packet) != MessageKeepaliveSize {
dataSent = true
-
- if padding := device.paddings.transport; padding > 0 {
- // elem.packet is stored at the start of elem.buffer
- // with zero padding
- for i := len(elem.packet) - 1; i >= 0; i-- {
- elem.buffer[i+padding] = elem.buffer[i]
- }
- rand.Read(elem.buffer[:padding])
- elem.packet = elem.buffer[:padding+len(elem.packet)]
+ }
+ if padding := device.paddings.transport; padding > 0 {
+ // elem.packet is stored at the start of elem.buffer
+ // with zero padding
+ for i := len(elem.packet) - 1; i >= 0; i-- {
+ elem.buffer[i+padding] = elem.buffer[i]
}
+ rand.Read(elem.buffer[:padding])
+ elem.packet = elem.buffer[:padding+len(elem.packet)]
}
bufs = append(bufs, elem.packet)
}