Fix ipv6 header view ownership.

During ipv6 header processing, the packet can be released and
switched to a different fragment, but we still hold
the same ipv6 header bytes from the released packet.
This causes a use-after-free.

PiperOrigin-RevId: 484627475
This commit is contained in:
Lucas Manning
2022-10-28 15:00:57 -07:00
committed by gVisor bot
parent a1468b7a62
commit 6b3b5493d0
3 changed files with 155 additions and 28 deletions
+1
View File
@@ -42,6 +42,7 @@ go_test(
"//pkg/bufferv2",
"//pkg/refs",
"//pkg/refsvfs2",
"//pkg/sync",
"//pkg/tcpip",
"//pkg/tcpip/checker",
"//pkg/tcpip/checksum",
+13 -6
View File
@@ -1046,11 +1046,13 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) {
return
}
h, ok := e.protocol.parseAndValidate(pkt)
hView, ok := e.protocol.parseAndValidate(pkt)
if !ok {
stats.MalformedPacketsReceived.Increment()
return
}
defer hView.Release()
h := header.IPv6(hView.AsSlice())
if !e.nic.IsLoopback() {
if !e.protocol.options.AllowExternalLoopbackTraffic {
@@ -1101,11 +1103,13 @@ func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksu
defer pkt.DecRef()
pkt.RXTransportChecksumValidated = canSkipRXChecksum
h, ok := e.protocol.parseAndValidate(pkt)
hView, ok := e.protocol.parseAndValidate(pkt)
if !ok {
stats.MalformedPacketsReceived.Increment()
return
}
defer hView.Release()
h := header.IPv6(hView.AsSlice())
e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */)
}
@@ -2537,19 +2541,22 @@ func (p *protocol) forwardPendingMulticastPacket(pkt stack.PacketBufferPtr, inst
func (*protocol) Wait() {}
// parseAndValidate parses the packet (including its transport layer header) and
// returns the parsed IP header.
// returns a view containing the parsed IP header. The caller is responsible
// for releasing the returned View.
//
// Returns true if the IP header was successfully parsed.
func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (header.IPv6, bool) {
func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (*bufferv2.View, bool) {
transProtoNum, hasTransportHdr, ok := p.Parse(pkt)
if !ok {
return nil, false
}
h := header.IPv6(pkt.NetworkHeader().Slice())
hView := pkt.NetworkHeader().View()
h := header.IPv6(hView.AsSlice())
// Do not include the link header's size when calculating the size of the IP
// packet.
if !h.IsValid(pkt.Size() - len(pkt.LinkHeader().Slice())) {
hView.Release()
return nil, false
}
@@ -2557,7 +2564,7 @@ func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (header.IPv6, boo
p.parseTransport(pkt, transProtoNum)
}
return h, true
return hView, true
}
func (p *protocol) parseTransport(pkt stack.PacketBufferPtr, transProtoNum tcpip.TransportProtocolNumber) {
+141 -22
View File
@@ -26,6 +26,7 @@ import (
"github.com/google/go-cmp/cmp"
"gvisor.dev/gvisor/pkg/bufferv2"
"gvisor.dev/gvisor/pkg/sync"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/checker"
"gvisor.dev/gvisor/pkg/tcpip/checksum"
@@ -1109,6 +1110,28 @@ type fragmentData struct {
data []byte
}
func udpGen(payload []byte, multiplier uint8, src, dst tcpip.Address) []byte {
payloadLen := len(payload)
for i := 0; i < payloadLen; i++ {
payload[i] = uint8(i) * multiplier
}
udpLength := header.UDPMinimumSize + payloadLen
hdr := prependable.New(udpLength)
u := header.UDP(hdr.Prepend(udpLength))
u.Encode(&header.UDPFields{
SrcPort: 5555,
DstPort: 80,
Length: uint16(udpLength),
})
copy(u.Payload(), payload)
sum := header.PseudoHeaderChecksum(udp.ProtocolNumber, src, dst, uint16(udpLength))
sum = checksum.Checksum(payload, sum)
u.SetChecksum(^u.CalculateChecksum(sum))
return hdr.View()
}
func TestReceiveIPv6Fragments(t *testing.T) {
const (
udpPayload1Length = 256
@@ -1124,28 +1147,6 @@ func TestReceiveIPv6Fragments(t *testing.T) {
routingExtHdrLen = 8
)
udpGen := func(payload []byte, multiplier uint8, src, dst tcpip.Address) []byte {
payloadLen := len(payload)
for i := 0; i < payloadLen; i++ {
payload[i] = uint8(i) * multiplier
}
udpLength := header.UDPMinimumSize + payloadLen
hdr := prependable.New(udpLength)
u := header.UDP(hdr.Prepend(udpLength))
u.Encode(&header.UDPFields{
SrcPort: 5555,
DstPort: 80,
Length: uint16(udpLength),
})
copy(u.Payload(), payload)
sum := header.PseudoHeaderChecksum(udp.ProtocolNumber, src, dst, uint16(udpLength))
sum = checksum.Checksum(payload, sum)
u.SetChecksum(^u.CalculateChecksum(sum))
return hdr.View()
}
var udpPayload1Addr1ToAddr2Buf [udpPayload1Length]byte
udpPayload1Addr1ToAddr2 := udpPayload1Addr1ToAddr2Buf[:]
ipv6Payload1Addr1ToAddr2 := udpGen(udpPayload1Addr1ToAddr2, 1, addr1, addr2)
@@ -1936,6 +1937,124 @@ func TestReceiveIPv6Fragments(t *testing.T) {
}
}
func TestConcurrentFragmentWrites(t *testing.T) {
const udpPayload1Length = 256
const udpPayload2Length = 128
var udpPayload1Addr1ToAddr2Buf [udpPayload1Length]byte
udpPayload1Addr1ToAddr2 := udpPayload1Addr1ToAddr2Buf[:]
ipv6Payload1Addr1ToAddr2 := udpGen(udpPayload1Addr1ToAddr2, 1, addr1, addr2)
var udpPayload2Addr1ToAddr2Buf [udpPayload2Length]byte
udpPayload2Addr1ToAddr2 := udpPayload2Addr1ToAddr2Buf[:]
ipv6Payload2Addr1ToAddr2 := udpGen(udpPayload2Addr1ToAddr2, 2, addr1, addr2)
fragments := []fragmentData{
{
srcAddr: addr1,
dstAddr: addr2,
nextHdr: fragmentExtHdrID,
data: append(
// Fragment extension header.
//
// Fragment offset = 0, More = true, ID = 1
[]byte{uint8(header.UDPProtocolNumber), 0, 0, 1, 0, 0, 0, 1},
ipv6Payload1Addr1ToAddr2[:64]...,
),
},
{
srcAddr: addr1,
dstAddr: addr2,
nextHdr: fragmentExtHdrID,
data: append(
// Fragment extension header.
//
// Fragment offset = 0, More = true, ID = 2
[]byte{uint8(header.UDPProtocolNumber), 0, 0, 1, 0, 0, 0, 2},
ipv6Payload2Addr1ToAddr2[:32]...,
),
},
{
srcAddr: addr1,
dstAddr: addr2,
nextHdr: fragmentExtHdrID,
data: append(
// Fragment extension header.
//
// Fragment offset = 8, More = false, ID = 1
[]byte{uint8(header.UDPProtocolNumber), 0, 0, 64, 0, 0, 0, 1},
ipv6Payload1Addr1ToAddr2[64:]...,
),
},
}
c := newTestContext()
defer c.cleanup()
s := c.s
e := channel.New(0, header.IPv6MinimumMTU, linkAddr1)
defer e.Close()
if err := s.CreateNIC(nicID, e); err != nil {
t.Fatalf("CreateNIC(%d, _) = %s", nicID, err)
}
protocolAddr := tcpip.ProtocolAddress{
Protocol: ProtocolNumber,
AddressWithPrefix: addr2.WithPrefix(),
}
if err := s.AddProtocolAddress(nicID, protocolAddr, stack.AddressProperties{}); err != nil {
t.Fatalf("AddProtocolAddress(%d, %+v, {}): %s", nicID, protocolAddr, err)
}
wq := waiter.Queue{}
we, ch := waiter.NewChannelEntry(waiter.ReadableEvents)
wq.EventRegister(&we)
defer wq.EventUnregister(&we)
defer close(ch)
ep, err := s.NewEndpoint(udp.ProtocolNumber, ProtocolNumber, &wq)
if err != nil {
t.Fatalf("NewEndpoint(%d, %d, _): %s", udp.ProtocolNumber, ProtocolNumber, err)
}
defer ep.Close()
bindAddr := tcpip.FullAddress{Addr: addr2, Port: 80}
if err := ep.Bind(bindAddr); err != nil {
t.Fatalf("Bind(%+v): %s", bindAddr, err)
}
var wg sync.WaitGroup
defer wg.Wait()
for i := 0; i < 3; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for i := 0; i < 10; i++ {
for _, f := range fragments {
hdr := prependable.New(header.IPv6MinimumSize)
// Serialize IPv6 fixed header.
ip := header.IPv6(hdr.Prepend(header.IPv6MinimumSize))
ip.Encode(&header.IPv6Fields{
PayloadLength: uint16(len(f.data)),
// We're lying about transport protocol here so that we can generate
// raw extension headers for the tests.
TransportProtocol: tcpip.TransportProtocolNumber(f.nextHdr),
HopLimit: 255,
SrcAddr: f.srcAddr,
DstAddr: f.dstAddr,
})
buf := bufferv2.MakeWithData(hdr.View())
buf.Append(bufferv2.NewViewWithData(f.data))
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
Payload: buf,
})
e.InjectInbound(ProtocolNumber, pkt)
pkt.DecRef()
}
}
}()
}
}
func TestInvalidIPv6Fragments(t *testing.T) {
const (
addr1 = tcpip.Address("\x0a\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x01")