Skip to content

Commit f324ed8

Browse files
committed
vpn: fix remote crash and CPU-exhaustion from malformed tunnel packets
Parse now validates IPv4 IHL against the packet length, and RecalculateChecksum guards its header/transport slicing, so a peer can no longer crash the process with a short packet claiming a large IHL or TCP/UDP protocol. ReadFrom refuses to overflow Packet.Buffer instead of spinning forever on zero-length reads, and ReadPacketHeader rejects sizes above the real buffer capacity (new vpn.MaxPacketBodySize) rather than the 1 MiB protocol bound that was ~300x too large.
1 parent b914715 commit f324ed8

9 files changed

Lines changed: 250 additions & 39 deletions

File tree

application_test.go

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -405,6 +405,16 @@ func TestSOCKS5ProxyFallbackToOldProtocol(t *testing.T) {
405405
// Remove new protocol handler from peer2 to simulate old peer
406406
peer2.app.P2p.Host().RemoveStreamHandler(protocol.Socks5NoAuthMethod)
407407

408+
// Wait until peer1's peerstore reflects the removal (identify-push). Otherwise
409+
// NewStreamMulti optimistically negotiates the stale /socks5-noauth/ against a
410+
// peer that no longer handles it, and the dial fails with EOF.
411+
ts.Eventually(func() bool {
412+
supported, err := peer1.app.P2p.Host().Peerstore().
413+
SupportsProtocols(peer2.app.P2p.PeerID(), protocol.Socks5NoAuthMethod)
414+
ts.NoError(err)
415+
return len(supported) == 0
416+
}, 15*time.Second, 100*time.Millisecond)
417+
408418
// Allow peer1 to use peer2 as exit node
409419
peer1Config, err := peer2.api.KnownPeerConfig(peer1.PeerID())
410420
ts.NoError(err)

protocol/protocol_test.go

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -95,13 +95,22 @@ func TestReadPacketHeader_RejectsBothFlags(t *testing.T) {
9595

9696
func TestReadPacketHeader_RejectsHugeSize(t *testing.T) {
9797
var header [8]byte
98-
binary.BigEndian.PutUint64(header[:], tunnelMaxLength+1)
98+
binary.BigEndian.PutUint64(header[:], uint64(vpn.MaxPacketBodySize)+1)
9999

100100
_, _, err := ReadPacketHeader(bytes.NewReader(header[:]))
101101
require.Error(t, err)
102102
require.Contains(t, err.Error(), "exceeds max")
103103
}
104104

105+
func TestReadPacketHeader_AcceptsMaxBodySize(t *testing.T) {
106+
var header [8]byte
107+
binary.BigEndian.PutUint64(header[:], uint64(vpn.MaxPacketBodySize))
108+
109+
size, _, err := ReadPacketHeader(bytes.NewReader(header[:]))
110+
require.NoError(t, err)
111+
require.Equal(t, uint64(vpn.MaxPacketBodySize), size)
112+
}
113+
105114
func TestReadPacketHeader_TruncatedStream(t *testing.T) {
106115
_, _, err := ReadPacketHeader(bytes.NewReader([]byte{1, 2, 3}))
107116
require.Error(t, err)

protocol/tunnel.go

Lines changed: 2 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -25,11 +25,6 @@ const (
2525
tunnelFlagGatewayForward uint64 = 1 << 63
2626
tunnelFlagGatewayReturn uint64 = 1 << 62
2727
tunnelLengthMask uint64 = ^(tunnelFlagGatewayForward | tunnelFlagGatewayReturn)
28-
// tunnelMaxLength bounds the per-packet size we'll accept on read. The
29-
// real cap is vpn.maxContentSize (MTU + overhead); this is a generous
30-
// upper bound to reject obvious garbage early without crossing package
31-
// dependencies.
32-
tunnelMaxLength uint64 = 1 << 20
3328
)
3429

3530
// ReadPacketHeader reads an 8-byte tunnel packet header from the stream and
@@ -53,8 +48,8 @@ func ReadPacketHeader(stream io.Reader) (size uint64, dir vpn.GatewayDir, err er
5348
dir = vpn.GatewayDirReturn
5449
}
5550
size = v & tunnelLengthMask
56-
if size > tunnelMaxLength {
57-
return 0, 0, fmt.Errorf("invalid tunnel header: size %d exceeds max %d", size, tunnelMaxLength)
51+
if size > uint64(vpn.MaxPacketBodySize) {
52+
return 0, 0, fmt.Errorf("invalid tunnel header: size %d exceeds max %d", size, vpn.MaxPacketBodySize)
5853
}
5954
return size, dir, nil
6055
}

vpn/iface_darwin.go

Lines changed: 14 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ import (
77
"fmt"
88
"net"
99
"os/exec"
10+
"strings"
1011

1112
"golang.zx2c4.com/wireguard/tun"
1213
)
@@ -21,26 +22,33 @@ func newTUN(ifname string, mtu int, localIP net.IP, ipMask net.IPMask) (tun.Devi
2122
if err != nil {
2223
return nil, fmt.Errorf("create tun: %v", err)
2324
}
25+
// Close the freshly created device if any later setup step fails, otherwise
26+
// the TUN interface leaks.
27+
success := false
28+
defer func() {
29+
if !success {
30+
_ = tunDevice.Close()
31+
}
32+
}()
2433
// Interface name must be utun[0-9]*
2534
realIfname, err := tunDevice.Name()
2635
if err != nil {
2736
return nil, fmt.Errorf("get interface name: %v", err)
2837
}
2938

30-
err = exec.Command("ifconfig", realIfname, "inet", ipNet.String(), localIP.String()).Run()
31-
if err != nil {
32-
return nil, fmt.Errorf("unable to setup interface mask: %v", err)
39+
if out, err := exec.Command("ifconfig", realIfname, "inet", ipNet.String(), localIP.String()).CombinedOutput(); err != nil {
40+
return nil, fmt.Errorf("unable to setup interface mask: %v: %s", err, strings.TrimSpace(string(out)))
3341
}
3442

3543
ipNetMasked := &net.IPNet{
3644
IP: localIP.Mask(ipMask),
3745
Mask: ipMask,
3846
}
39-
err = exec.Command("route", "-q", "-n", "add", "-inet", ipNetMasked.String(), "-iface", realIfname).Run()
40-
if err != nil {
41-
return nil, fmt.Errorf("unable to setup interface route: %v", err)
47+
if out, err := exec.Command("route", "-q", "-n", "add", "-inet", ipNetMasked.String(), "-iface", realIfname).CombinedOutput(); err != nil {
48+
return nil, fmt.Errorf("unable to setup interface route: %v: %s", err, strings.TrimSpace(string(out)))
4249
}
4350

51+
success = true
4452
return tunDevice, nil
4553
}
4654

vpn/iface_linux.go

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,14 @@ func newTUN(ifname string, mtu int, localIP net.IP, ipMask net.IPMask) (tun.Devi
1616
if err != nil {
1717
return nil, fmt.Errorf("create tun: %v", err)
1818
}
19+
// Close the freshly created device if any later setup step fails, otherwise
20+
// the TUN interface leaks.
21+
success := false
22+
defer func() {
23+
if !success {
24+
_ = tunDevice.Close()
25+
}
26+
}()
1927

2028
link, err := netlink.LinkByName(ifname)
2129
if err != nil {
@@ -36,6 +44,7 @@ func newTUN(ifname string, mtu int, localIP net.IP, ipMask net.IPMask) (tun.Devi
3644
return nil, fmt.Errorf("unable to UP interface: %v", err)
3745
}
3846

47+
success = true
3948
return tunDevice, nil
4049
}
4150

vpn/iface_windows.go

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,14 @@ func newTUN(ifname string, mtu int, localIP net.IP, ipMask net.IPMask) (tun.Devi
4444
if err != nil {
4545
return nil, fmt.Errorf("do as system: %v", err)
4646
}
47+
// Close the freshly created adapter if any later setup step fails, otherwise
48+
// the Wintun adapter leaks.
49+
success := false
50+
defer func() {
51+
if !success {
52+
_ = tunDevice.Close()
53+
}
54+
}()
4755

4856
nativeTunDevice := tunDevice.(*tun.NativeTun)
4957
luid := winipcfg.LUID(nativeTunDevice.LUID())
@@ -54,7 +62,6 @@ func newTUN(ifname string, mtu int, localIP net.IP, ipMask net.IPMask) (tun.Devi
5462
// force NLMTU ourselves via winipcfg, otherwise local apps see MTU=65535,
5563
// send oversized packets, and the size check in ReadTUNPackets silently drops them.
5664
if err := setInterfaceMTU(logger, luid, winipcfg.AddressFamily(windows.AF_INET), uint32(mtu)); err != nil {
57-
tunDevice.Close()
5865
return nil, fmt.Errorf("set IPv4 MTU on tun: %v", err)
5966
}
6067
// TODO: support ipv6. Forwarding still ignores IPv6 packets (see Device.WritePacket),
@@ -74,6 +81,7 @@ func newTUN(ifname string, mtu int, localIP net.IP, ipMask net.IPMask) (tun.Devi
7481
return nil, fmt.Errorf("unable to setup interface IP: %v", err)
7582
}
7683

84+
success = true
7785
return tunDevice, nil
7886
}
7987

vpn/packet.go

Lines changed: 64 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package vpn
22

33
import (
44
"encoding/binary"
5+
"fmt"
56
"io"
67
"net"
78

@@ -61,12 +62,17 @@ func (data *Packet) clear() {
6162

6263
func (data *Packet) CopyTo(copyPacket *Packet) {
6364
*copyPacket = *data
64-
// set Packet reference to a new buffer
65+
// Buffer is an array copied by value above, but Packet/Src/Dst are slice
66+
// headers still aliasing data.Buffer. Re-derive them as sub-slices of the
67+
// copy's own buffer.
6568
copyPacket.Packet = copyPacket.Buffer[tunPacketOffset : len(data.Packet)+tunPacketOffset]
69+
if data.Src != nil {
70+
copyPacket.setAddrs()
71+
}
6672
}
6773

6874
func (data *Packet) ReadFrom(stream io.Reader) (int64, error) {
69-
var totalRead = tunPacketOffset
75+
totalRead := tunPacketOffset
7076
for {
7177
n, err := stream.Read(data.Buffer[totalRead:])
7278
totalRead += n
@@ -76,6 +82,13 @@ func (data *Packet) ReadFrom(stream io.Reader) (int64, error) {
7682
} else if err != nil {
7783
return int64(totalRead - tunPacketOffset), err
7884
}
85+
// The buffer is sized with slack above the MTU (maxContentSize), so a
86+
// valid single-packet body never fills it. A full buffer means an
87+
// oversized frame; reject it instead of reading into the empty tail
88+
// (which would spin on (0, nil) reads).
89+
if totalRead == len(data.Buffer) {
90+
return int64(totalRead - tunPacketOffset), fmt.Errorf("packet body exceeds max size %d bytes", MaxPacketBodySize)
91+
}
7992
}
8093
}
8194

@@ -90,19 +103,25 @@ func (data *Packet) Parse() bool {
90103
if len(packet) < ipv4.HeaderLen {
91104
return false
92105
}
106+
// Validate the header length against the actual packet so downstream
107+
// slicing (RecalculateChecksum) can't run past the packet and panic.
108+
// Transport-length is not enforced here (it would drop legit non-first IP
109+
// fragments); RecalculateChecksum guards its own transport slicing.
110+
ipHeaderLen := int(packet[0]&0x0f) << 2
111+
if ipHeaderLen < ipv4.HeaderLen || ipHeaderLen > len(packet) {
112+
return false
113+
}
114+
data.IPProtocol = packet[9]
93115

94-
data.Src = packet[device.IPv4offsetSrc : device.IPv4offsetSrc+net.IPv4len]
95-
data.Dst = packet[device.IPv4offsetDst : device.IPv4offsetDst+net.IPv4len]
96116
data.IsIPv6 = false
97-
data.IPProtocol = data.Packet[9]
117+
data.setAddrs()
98118
case ipv6.Version:
99119
if len(packet) < ipv6.HeaderLen {
100120
return false
101121
}
102122

103-
data.Src = packet[device.IPv6offsetSrc : device.IPv6offsetSrc+net.IPv6len]
104-
data.Dst = packet[device.IPv6offsetDst : device.IPv6offsetDst+net.IPv6len]
105123
data.IsIPv6 = true
124+
data.setAddrs()
106125
// TODO: set data.IPProtocol
107126
default:
108127
return false
@@ -114,24 +133,45 @@ func (data *Packet) Parse() bool {
114133
func (data *Packet) RecalculateChecksum() {
115134
if data.IsIPv6 {
116135
// TODO
117-
} else {
118-
ipHeaderLen := int(data.Packet[0]&0x0f) << 2
119-
copy(data.Packet[ipv4offsetChecksum:], []byte{0, 0})
120-
ipChecksum := checksumIPv4Header(data.Packet[:ipHeaderLen])
121-
binary.BigEndian.PutUint16(data.Packet[ipv4offsetChecksum:], ipChecksum)
122-
123-
switch protocol := data.Packet[9]; protocol {
124-
case IPProtocolTCP:
125-
tcpOffsetChecksum := ipHeaderLen + 16
126-
copy(data.Packet[tcpOffsetChecksum:], []byte{0, 0})
127-
checksum := checksumIPv4TCPUDP(data.Packet[ipHeaderLen:], uint32(protocol), data.Src, data.Dst)
128-
binary.BigEndian.PutUint16(data.Packet[tcpOffsetChecksum:], checksum)
129-
case IPProtocolUDP:
130-
udpOffsetChecksum := ipHeaderLen + 6
131-
copy(data.Packet[udpOffsetChecksum:], []byte{0, 0})
132-
checksum := checksumIPv4TCPUDP(data.Packet[ipHeaderLen:], uint32(protocol), data.Src, data.Dst)
133-
binary.BigEndian.PutUint16(data.Packet[udpOffsetChecksum:], checksum)
136+
return
137+
}
138+
// Guard against malformed lengths so a bad packet can't slice past bounds and
139+
// panic, even if RecalculateChecksum is called without a prior valid Parse.
140+
ipHeaderLen := int(data.Packet[0]&0x0f) << 2
141+
if ipHeaderLen < ipv4.HeaderLen || ipHeaderLen > len(data.Packet) {
142+
return
143+
}
144+
copy(data.Packet[ipv4offsetChecksum:], []byte{0, 0})
145+
ipChecksum := checksumIPv4Header(data.Packet[:ipHeaderLen])
146+
binary.BigEndian.PutUint16(data.Packet[ipv4offsetChecksum:], ipChecksum)
147+
148+
switch protocol := data.Packet[9]; protocol {
149+
case IPProtocolTCP:
150+
if len(data.Packet) < ipHeaderLen+18 {
151+
return
134152
}
153+
tcpOffsetChecksum := ipHeaderLen + 16
154+
copy(data.Packet[tcpOffsetChecksum:], []byte{0, 0})
155+
checksum := checksumIPv4TCPUDP(data.Packet[ipHeaderLen:], uint32(protocol), data.Src, data.Dst)
156+
binary.BigEndian.PutUint16(data.Packet[tcpOffsetChecksum:], checksum)
157+
case IPProtocolUDP:
158+
if len(data.Packet) < ipHeaderLen+8 {
159+
return
160+
}
161+
udpOffsetChecksum := ipHeaderLen + 6
162+
copy(data.Packet[udpOffsetChecksum:], []byte{0, 0})
163+
checksum := checksumIPv4TCPUDP(data.Packet[ipHeaderLen:], uint32(protocol), data.Src, data.Dst)
164+
binary.BigEndian.PutUint16(data.Packet[udpOffsetChecksum:], checksum)
165+
}
166+
}
167+
168+
func (data *Packet) setAddrs() {
169+
if data.IsIPv6 {
170+
data.Src = data.Packet[device.IPv6offsetSrc : device.IPv6offsetSrc+net.IPv6len]
171+
data.Dst = data.Packet[device.IPv6offsetDst : device.IPv6offsetDst+net.IPv6len]
172+
} else {
173+
data.Src = data.Packet[device.IPv4offsetSrc : device.IPv4offsetSrc+net.IPv4len]
174+
data.Dst = data.Packet[device.IPv4offsetDst : device.IPv4offsetDst+net.IPv4len]
135175
}
136176
}
137177

0 commit comments

Comments
 (0)