lx: gate reserved-byte clear on receive so AmneziaWG magic survives

The Cloudflare "reserved" bytes (1-3) were zeroed unconditionally on
every received datagram across all StdNetBind/WinRingBind receive paths.
AmneziaWG reads its magic header as LittleEndian.Uint32(packet[padding:])
where padding is s1/s2/s4; with small padding (0-3) the magic overlaps
bytes 1-3, so clearing them collapses it out of the ranged h1-h4 window
and every packet is dropped (handshake included) — the AWG endpoint
never comes up. Plain WG (types 1-4, bytes 1-3 already zero) and large
padding are unaffected, which is why it went unnoticed.

Gate all five receive clears (bind_std receiveIP, msgx_darwin
receiveSingle + makeReceiveMsgX, bind_windows receiveIPv4/v6) behind a
new hasReserved() so bytes 1-3 are only touched when a WARP reserved
value is actually configured. Send paths already gate on a per-endpoint
loaded/non-zero check, so they are left unchanged. The reserved map is
populated before the receive goroutines start and never mutated after,
so the lock-free read is safe.

Tests: awg_stdnetbind_reserved_lx_test.go brings up two Devices over
StdNetBind with zero padding (magic in bytes 0-3) and asserts delivery
(red before the fix, green after); reserved_gate_lx_test.go pins the
hasReserved() gate.
This commit is contained in:
Leadaxe 2026-07-17 01:32:40 +02:00
parent ee7ff1b77f
commit 1e787bb3e0
7 changed files with 283 additions and 11 deletions

View file

@ -359,7 +359,7 @@ func (s *StdNetBind) receiveIP(
if sizes[i] == 0 {
continue
}
if msg.N > 3 {
if msg.N > 3 && s.hasReserved() { // lx: only strip reserved bytes for WARP (see hasReserved)
common.ClearArray(bufs[i][1:4])
}
ep := &StdNetEndpoint{AddrPort: M.AddrPortFromNet(msg.Addr)} // TODO: remove allocation
@ -539,6 +539,20 @@ func (s *StdNetBind) SetReservedForEndpoint(destination netip.AddrPort, reserved
s.reservedForEndpoint[destination] = reserved
}
// lx: hasReserved reports whether any Cloudflare "reserved" value is set. The
// receive path must only zero bytes 1-3 when a reserved value exists (WARP);
// otherwise an AmneziaWG magic header that lands in bytes 1-3 (small s1/s2/s4
// padding) would be corrupted and the packet dropped. The send path already
// gates its stamp on a per-endpoint `loaded` check, so no change is needed there.
func (s *StdNetBind) hasReserved() bool {
for _, reserved := range s.reservedForEndpoint {
if reserved != [3]uint8{} {
return true
}
}
return false
}
func (s *StdNetBind) send(conn *net.UDPConn, pc batchWriter, msgs []ipv6.Message) error {
var (
n int

View file

@ -461,7 +461,7 @@ func (bind *WinRingBind) receiveIPv4(bufs [][]byte, sizes []int, eps []Endpoint)
bind.mu.RLock()
defer bind.mu.RUnlock()
n, ep, err := bind.v4.Receive(bufs[0], &bind.isOpen)
if n > 3 {
if n > 3 && bind.hasReserved() { // lx: only strip reserved bytes for WARP (see hasReserved)
common.ClearArray(bufs[0][1:4])
}
sizes[0] = n
@ -473,7 +473,7 @@ func (bind *WinRingBind) receiveIPv6(bufs [][]byte, sizes []int, eps []Endpoint)
bind.mu.RLock()
defer bind.mu.RUnlock()
n, ep, err := bind.v6.Receive(bufs[0], &bind.isOpen)
if n > 3 {
if n > 3 && bind.hasReserved() { // lx: only strip reserved bytes for WARP (see hasReserved)
common.ClearArray(bufs[0][1:4])
}
sizes[0] = n
@ -576,6 +576,18 @@ func (bind *WinRingBind) SetReservedForEndpoint(destination netip.AddrPort, rese
bind.reservedForEndpoint[*endpoint.(*WinRingEndpoint)] = reserved
}
// lx: hasReserved reports whether any Cloudflare "reserved" value is set. See
// the StdNetBind.hasReserved comment — the unconditional receive clear would
// corrupt an AmneziaWG magic header sitting in bytes 1-3 (small padding).
func (bind *WinRingBind) hasReserved() bool {
for _, reserved := range bind.reservedForEndpoint {
if reserved != [3]uint8{} {
return true
}
}
return false
}
func (s *StdNetBind) BindSocketToInterface4(interfaceIndex uint32, blackhole bool) error {
s.mu.Lock()
defer s.mu.Unlock()

View file

@ -234,7 +234,7 @@ func (s *StdNetBind) receiveSingle(conn *net.UDPConn, bufs [][]byte, sizes []int
return 0, err
}
sizes[0] = n
if n > 3 {
if n > 3 && s.hasReserved() { // lx: only strip reserved bytes for WARP (see hasReserved)
bufs[0][1] = 0
bufs[0][2] = 0
bufs[0][3] = 0
@ -299,7 +299,7 @@ func (s *StdNetBind) makeReceiveMsgX(conn *net.UDPConn, isV6 bool) (ReceiveFunc,
numMsgs := int(n)
for i := 0; i < numMsgs; i++ {
sizes[i] = int(state.hdrs[i].DataLen)
if sizes[i] > 3 {
if sizes[i] > 3 && s.hasReserved() { // lx: only strip reserved bytes for WARP (see hasReserved)
bufs[i][1] = 0
bufs[i][2] = 0
bufs[i][3] = 0

View file

@ -0,0 +1,56 @@
/* SPDX-License-Identifier: MIT
*
* lx: unit coverage for the StdNetBind.hasReserved() gate that guards the
* receive-side reserved-clear. receiveIP zeroes bytes 1-3 (Cloudflare WARP
* "reserved") only when a non-zero reserved value is set for some endpoint;
* otherwise an AmneziaWG magic header landing in bytes 1-3 (small s1/s2/s4
* padding) would be corrupted and the packet dropped. This test pins the gate
* itself; the end-to-end handshake proof lives in the device package.
*/
package conn
import (
"net/netip"
"testing"
)
func stdNetBindForTest(t *testing.T) *StdNetBind {
t.Helper()
b, ok := NewStdNetBind(nil).(*StdNetBind)
if !ok {
t.Fatalf("NewStdNetBind did not return *StdNetBind")
}
return b
}
func TestStdNetBindHasReserved(t *testing.T) {
b := stdNetBindForTest(t)
if b.hasReserved() {
t.Fatal("fresh bind must report no reserved value")
}
ep := netip.MustParseAddrPort("127.0.0.1:51820")
// An all-zero reserved value is indistinguishable from "unset" and must
// not arm the clear.
b.SetReservedForEndpoint(ep, [3]byte{0, 0, 0})
if b.hasReserved() {
t.Fatal("all-zero reserved must not count as reserved")
}
// Any non-zero byte (WARP anycast tag) arms the clear.
b.SetReservedForEndpoint(ep, [3]byte{0, 0, 1})
if !b.hasReserved() {
t.Fatal("non-zero reserved (byte 3) must count as reserved")
}
// A second endpoint's non-zero value must also be seen.
b2 := stdNetBindForTest(t)
ep2 := netip.MustParseAddrPort("192.0.2.1:2408")
b2.SetReservedForEndpoint(ep, [3]byte{0, 0, 0})
b2.SetReservedForEndpoint(ep2, [3]byte{0xAB, 0, 0})
if !b2.hasReserved() {
t.Fatal("non-zero reserved on any endpoint must count as reserved")
}
}

View file

@ -0,0 +1,190 @@
/* SPDX-License-Identifier: MIT
*
* lx: e2e regression for the reserved-clear vs AWG magic-header collision,
* exercised over the StdNetBind path (no detour) with real loopback UDP.
*
* Bug model. On receive, StdNetBind.receiveIP unconditionally zeroed bytes
* 1-3 of every datagram >3 bytes (the Cloudflare WARP "reserved" field).
* AmneziaWG reads its magic header as LittleEndian.Uint32(packet[padding:]),
* where padding is s1/s2 (handshake) or s4 (transport). With small padding
* (0..3) the 4-byte magic overlaps bytes 1-3, so the unconditional clear
* corrupts it: the value falls outside the ranged h1-h4 window, the packet is
* classified MessageUnknownType and dropped. WARP was never configured on
* these binds (no SetReservedForEndpoint), so the clear was pure collateral.
*
* The fix gates the clear behind StdNetBind.hasReserved(): bytes 1-3 are only
* zeroed when a non-zero reserved value is actually set for some endpoint.
* With no reserved value the magic survives and the handshake completes.
*
* This test provokes the worst case: padding = 0 (no s1/s2/s4 at all), so the
* h1 initiation magic sits in bytes [0..3] and its high bytes (1-3) are the
* ones the clear would destroy. The h1-h4 ranges are chosen entirely above
* 0x10000000, so after zeroing bytes 1-3 the surviving value is <= 255 and can
* never land back inside any range -> guaranteed drop on the buggy tree.
*
* GREEN on the fixed tree. To see RED, temporarily restore the unconditional
* clear in conn/bind_std.go receiveIP:
* if msg.N > 3 {
* common.ClearArray(bufs[i][1:4])
* }
* and the handshake times out (init magic zeroed in bytes 1-3).
*/
package device
import (
"bufio"
"bytes"
"context"
"encoding/hex"
"fmt"
"strconv"
"strings"
"testing"
"time"
"github.com/sagernet/wireguard-go/conn"
)
// magic header ranges kept entirely above 0x10000000 (268435456). Any value
// the sender picks therefore has a non-zero byte among positions 1-3; zeroing
// those bytes collapses the value to <= 0xFF, which is below every range start,
// so a corrupted magic can never validate. Distinct windows per message type.
const (
lxH1Lo, lxH1Hi = 268500000, 268600000 // init
lxH2Lo, lxH2Hi = 300000000, 300100000 // response
lxH3Lo, lxH3Hi = 400000000, 400100000 // cookie
lxH4Lo, lxH4Hi = 500000000, 500100000 // transport
)
// lxReadListenPort parses listen_port=<n> out of a device's IpcGet dump.
func lxReadListenPort(t *testing.T, dev *Device) uint16 {
t.Helper()
dump, err := dev.IpcGet()
if err != nil {
t.Fatalf("IpcGet: %v", err)
}
scanner := bufio.NewScanner(strings.NewReader(dump))
for scanner.Scan() {
line := scanner.Text()
if v, ok := strings.CutPrefix(line, "listen_port="); ok {
p, err := strconv.Atoi(v)
if err != nil {
t.Fatalf("parse listen_port %q: %v", v, err)
}
return uint16(p)
}
}
t.Fatalf("listen_port not found in dump:\n%s", dump)
return 0
}
// newStdNetPaddedPair builds two Up()'d Devices peered over real loopback UDP
// (NewStdNetBind), configured with ranged h1-h4 + junk and *no* s1/s2/s4
// (padding = 0). Endpoints are wired after Up, once the ephemeral ports are
// known. No reserved value is ever set, so hasReserved() is false.
func newStdNetPaddedPair(t *testing.T) (devA, devB *Device, tunA, tunB *chanTun) {
t.Helper()
skA, err := newPrivateKey()
if err != nil {
t.Fatalf("newPrivateKey A: %v", err)
}
skB, err := newPrivateKey()
if err != nil {
t.Fatalf("newPrivateKey B: %v", err)
}
pkA := skA.publicKey()
pkB := skB.publicKey()
tunA = newChanTun()
tunB = newChanTun()
devA = NewDevice(context.Background(), tunA, conn.NewStdNetBind(nil), NewLogger(LogLevelError, "devA: "), 1)
devB = NewDevice(context.Background(), tunB, conn.NewStdNetBind(nil), NewLogger(LogLevelError, "devB: "), 1)
t.Cleanup(devA.Close)
t.Cleanup(devB.Close)
// obfuscation shared by both ends. Ranged magic headers, junk packets,
// and deliberately no s1/s2/s4 so padding stays 0 for every message type.
obf := fmt.Sprintf(
"jc=3\njmin=8\njmax=16\n"+
"h1=%d-%d\nh2=%d-%d\nh3=%d-%d\nh4=%d-%d\n",
lxH1Lo, lxH1Hi, lxH2Lo, lxH2Hi, lxH3Lo, lxH3Hi, lxH4Lo, lxH4Hi)
// Bring both up on an ephemeral port (listen_port=0), no endpoint yet.
cfgA := fmt.Sprintf("private_key=%s\nlisten_port=0\n%sreplace_peers=true\npublic_key=%s\nallowed_ip=%s/32\n",
hex.EncodeToString(skA[:]), obf, hex.EncodeToString(pkB[:]), testIPB)
cfgB := fmt.Sprintf("private_key=%s\nlisten_port=0\n%sreplace_peers=true\npublic_key=%s\nallowed_ip=%s/32\n",
hex.EncodeToString(skB[:]), obf, hex.EncodeToString(pkA[:]), testIPA)
if err := devA.IpcSet(cfgA); err != nil {
t.Fatalf("IpcSet A: %v", err)
}
if err := devB.IpcSet(cfgB); err != nil {
t.Fatalf("IpcSet B: %v", err)
}
if devA.paddings.init != 0 || devA.paddings.response != 0 || devA.paddings.transport != 0 {
t.Fatalf("padding must be 0 for this test: init=%d resp=%d transport=%d",
devA.paddings.init, devA.paddings.response, devA.paddings.transport)
}
if err := devA.Up(); err != nil {
t.Fatalf("Up A: %v", err)
}
if err := devB.Up(); err != nil {
t.Fatalf("Up B: %v", err)
}
portA := lxReadListenPort(t, devA)
portB := lxReadListenPort(t, devB)
if portA == 0 || portB == 0 {
t.Fatalf("ephemeral ports not assigned: A=%d B=%d", portA, portB)
}
// Now that ports are known, point each peer at the other over loopback.
if err := devA.IpcSet(fmt.Sprintf("public_key=%s\nupdate_only=true\nendpoint=127.0.0.1:%d\n",
hex.EncodeToString(pkB[:]), portB)); err != nil {
t.Fatalf("set endpoint A->B: %v", err)
}
if err := devB.IpcSet(fmt.Sprintf("public_key=%s\nupdate_only=true\nendpoint=127.0.0.1:%d\n",
hex.EncodeToString(pkA[:]), portA)); err != nil {
t.Fatalf("set endpoint B->A: %v", err)
}
return devA, devB, tunA, tunB
}
// TestStdNetBindReservedClearVsMagic_ZeroPadding drives a real handshake and a
// data packet A->B over loopback UDP through StdNetBind, with padding=0 so the
// h1 magic overlaps the reserved bytes 1-3. It passes only when receive does
// not blindly clear those bytes (the fix).
func TestStdNetBindReservedClearVsMagic_ZeroPadding(t *testing.T) {
devA, _, tunA, tunB := newStdNetPaddedPair(t)
_ = devA
pkt := buildIPv4Packet(testIPA, testIPB, 8)
// Re-inject periodically: the first packet triggers the handshake and may
// be dropped until keys are established.
send := func() { tunA.toDevice <- pkt }
send()
deadline := time.After(15 * time.Second)
retry := time.NewTicker(500 * time.Millisecond)
defer retry.Stop()
for {
select {
case got := <-tunB.fromDevice:
if bytes.Equal(got, pkt) {
return // delivered end to end: magic survived, handshake ok
}
t.Logf("ignoring unexpected packet len=%d", len(got))
case <-retry.C:
send()
case <-deadline:
t.Fatal("timed out waiting for packet on peer tun " +
"(handshake never completed: reserved-clear likely corrupted the h1 magic)")
}
}
}

View file

@ -23,8 +23,8 @@ func TestParseObfLen(t *testing.T) {
{"0", 0, false},
{"100", 100, false},
{fmt.Sprintf("%d", MaxMessageSize), MaxMessageSize, false},
{"-1", 0, true}, // would panic slice bounds in Obfuscate
{fmt.Sprintf("%d", MaxMessageSize + 1), 0, true}, // would OOM the handshake make
{"-1", 0, true}, // would panic slice bounds in Obfuscate
{fmt.Sprintf("%d", MaxMessageSize+1), 0, true}, // would OOM the handshake make
{"2000000000", 0, true},
{"abc", 0, true},
}

View file

@ -188,10 +188,10 @@ func (t *chanTun) Write(bufs [][]byte, offset int) (int, error) {
return len(bufs), nil
}
func (t *chanTun) MTU() (int, error) { return DefaultMTU, nil }
func (t *chanTun) Name() (string, error) { return "chantun", nil }
func (t *chanTun) Events() <-chan tun.Event { return t.events }
func (t *chanTun) BatchSize() int { return 1 }
func (t *chanTun) MTU() (int, error) { return DefaultMTU, nil }
func (t *chanTun) Name() (string, error) { return "chantun", nil }
func (t *chanTun) Events() <-chan tun.Event { return t.events }
func (t *chanTun) BatchSize() int { return 1 }
func (t *chanTun) Close() error {
t.closeOnce.Do(func() {