Files
Lucas ManningandgVisor bot afa323bd30 Replace most instances of IncRef with Clone.
Incrementing the reference count of a packet as a means of granting ownership
is unsafe when the packet is shared across gorountines. The underlying buffer's
reference count is unchanged since it "technically" has the same owning
PacketBuffer, which means different goroutines operating on the underlying
buffer (and packet itself) race.

Clones are roughly as fast as IncRefs because the PacketBuffers allocate from
a pool and the underlying buffers are cloned with copy-on-write
semantics.

I've left IncRef in places where the original packet in obviously going out of
scope at the end of the function or in some tests.

Reported-by: syzbot+e026046f4bf8ad09ae1f@syzkaller.appspotmail.com
Reported-by: syzbot+559365d6050db4b30e0f@syzkaller.appspotmail.com
Reported-by: syzbot+63c78a2c88a5744c636b@syzkaller.appspotmail.com
PiperOrigin-RevId: 705676806
2024-12-12 17:09:40 -08:00

729 lines
19 KiB
Go

// Copyright 2018 The gVisor Authors.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//go:build linux
// +build linux
package fdbased
import (
"bytes"
"fmt"
"os"
"slices"
"testing"
"time"
"unsafe"
"github.com/google/go-cmp/cmp"
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/buffer"
"gvisor.dev/gvisor/pkg/rand"
"gvisor.dev/gvisor/pkg/refs"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header"
"gvisor.dev/gvisor/pkg/tcpip/stack"
)
const (
mtu = 1500
laddr = tcpip.LinkAddress("\x11\x22\x33\x44\x55\x66")
raddr = tcpip.LinkAddress("\x77\x88\x99\xaa\xbb\xcc")
proto = 10
csumOffset = 48
gsoMSS = 500
)
type packetInfo struct {
Proto tcpip.NetworkProtocolNumber
Contents *stack.PacketBuffer
}
type packetContents struct {
LinkHeader []byte
NetworkHeader []byte
TransportHeader []byte
Data []byte
}
func checkPacketInfoEqual(t *testing.T, got, want packetInfo) {
t.Helper()
if diff := cmp.Diff(
want, got,
cmp.Transformer("ExtractPacketBuffer", func(pk *stack.PacketBuffer) *packetContents {
if pk == nil {
return nil
}
return &packetContents{
LinkHeader: pk.LinkHeader().Slice(),
NetworkHeader: pk.NetworkHeader().Slice(),
TransportHeader: pk.TransportHeader().Slice(),
Data: pk.Data().AsRange().ToSlice(),
}
}),
); diff != "" {
t.Errorf("unexpected packetInfo (-want +got):\n%s", diff)
}
}
type testContext struct {
t *testing.T
readFDs []int
writeFDs []int
ep stack.LinkEndpoint
ch chan packetInfo
done chan struct{}
}
func newContext(t *testing.T, opt *Options) *testContext {
firstFDPair, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_SEQPACKET, 0)
if err != nil {
t.Fatalf("Socketpair failed: %v", err)
}
secondFDPair, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_SEQPACKET, 0)
if err != nil {
t.Fatalf("Socketpair failed: %v", err)
}
done := make(chan struct{}, 2)
opt.ClosedFunc = func(tcpip.Error) {
done <- struct{}{}
}
opt.FDs = []int{firstFDPair[1], secondFDPair[1]}
ep, err := New(opt)
if err != nil {
t.Fatalf("Failed to create FD endpoint: %v", err)
}
c := &testContext{
t: t,
readFDs: []int{firstFDPair[0], secondFDPair[0]},
writeFDs: opt.FDs,
ep: ep,
ch: make(chan packetInfo, 100),
done: done,
}
ep.Attach(c)
return c
}
func (c *testContext) cleanup() {
for _, fd := range c.readFDs {
unix.Close(fd)
}
<-c.done
<-c.done
for _, fd := range c.writeFDs {
unix.Close(fd)
}
}
func (c *testContext) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) {
c.ch <- packetInfo{protocol, pkt.Clone()}
}
func (c *testContext) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer) {
c.t.Fatal("DeliverLinkPacket not implemented")
}
func TestNoEthernetProperties(t *testing.T) {
c := newContext(t, &Options{MTU: mtu})
defer c.cleanup()
if want, v := uint16(0), c.ep.MaxHeaderLength(); want != v {
t.Fatalf("MaxHeaderLength() = %v, want %v", v, want)
}
if want, v := uint32(mtu), c.ep.MTU(); want != v {
t.Fatalf("MTU() = %v, want %v", v, want)
}
}
func TestEthernetProperties(t *testing.T) {
c := newContext(t, &Options{EthernetHeader: true, MTU: mtu})
defer c.cleanup()
if want, v := uint16(header.EthernetMinimumSize), c.ep.MaxHeaderLength(); want != v {
t.Fatalf("MaxHeaderLength() = %v, want %v", v, want)
}
if want, v := uint32(mtu), c.ep.MTU(); want != v {
t.Fatalf("MTU() = %v, want %v", v, want)
}
}
func TestAddress(t *testing.T) {
addrs := []tcpip.LinkAddress{"", "abc", "def"}
for _, a := range addrs {
t.Run(fmt.Sprintf("Address: %q", a), func(t *testing.T) {
c := newContext(t, &Options{Address: a, MTU: mtu})
defer c.cleanup()
if want, v := a, c.ep.LinkAddress(); want != v {
t.Fatalf("LinkAddress() = %v, want %v", v, want)
}
})
}
}
func TestSetAddress(t *testing.T) {
addrs := []tcpip.LinkAddress{"abc", "def"}
c := newContext(t, &Options{Address: laddr, MTU: mtu})
defer c.cleanup()
for _, addr := range addrs {
c.ep.SetLinkAddress(addr)
if want, v := addr, c.ep.LinkAddress(); want != v {
t.Errorf("LinkAddress() = %v, want %v", v, want)
}
}
}
func TestMTU(t *testing.T) {
mtus := []uint32{200, 300}
c := newContext(t, &Options{MTU: mtu})
defer c.cleanup()
for _, m := range mtus {
c.ep.SetMTU(m)
if want, v := m, c.ep.MTU(); want != v {
t.Errorf("MTU() = %v, want %v", v, want)
}
}
}
func testWritePacket(t *testing.T, plen int, eth bool, gsoMaxSize uint32, hash uint32) {
c := newContext(t, &Options{Address: laddr, MTU: mtu, EthernetHeader: eth, GSOMaxSize: gsoMaxSize})
defer c.cleanup()
// Build payload.
payload := make([]byte, plen)
if _, err := rand.Read(payload); err != nil {
t.Fatalf("rand.Read(payload): %s", err)
}
// Build packet buffer.
const netHdrLen = 100
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
ReserveHeaderBytes: int(c.ep.MaxHeaderLength()) + netHdrLen,
Payload: buffer.MakeWithData(payload),
})
defer pkt.DecRef()
pkt.Hash = hash
// Every PacketBuffer must have these set:
// See nic.writePacket.
pkt.EgressRoute.LocalLinkAddress = laddr
pkt.EgressRoute.RemoteLinkAddress = raddr
pkt.NetworkProtocolNumber = proto
// Build header.
b := pkt.NetworkHeader().Push(netHdrLen)
if _, err := rand.Read(b); err != nil {
t.Fatalf("rand.Read(b): %s", err)
}
// Write.
want := append(append([]byte{}, b...), payload...)
const l3HdrLen = header.IPv6MinimumSize
if gsoMaxSize != 0 {
pkt.GSOOptions = stack.GSO{
Type: stack.GSOTCPv6,
NeedsCsum: true,
CsumOffset: csumOffset,
MSS: gsoMSS,
L3HdrLen: l3HdrLen,
}
}
c.ep.AddHeader(pkt)
var pkts stack.PacketBufferList
pkts.PushBack(pkt)
if _, err := c.ep.WritePackets(pkts); err != nil {
t.Fatalf("WritePackets failed: %s", err)
}
// Read from the corresponding FD, then compare with what we wrote.
b = make([]byte, mtu)
fd := c.readFDs[hash%uint32(len(c.readFDs))]
n, err := unix.Read(fd, b)
if err != nil {
t.Fatalf("Read failed: %v", err)
}
b = b[:n]
if gsoMaxSize != 0 {
vnetHdr := *(*virtioNetHdr)(unsafe.Pointer(&b[0]))
if vnetHdr.flags&_VIRTIO_NET_HDR_F_NEEDS_CSUM == 0 {
t.Fatalf("virtioNetHdr.flags %v doesn't contain %v", vnetHdr.flags, _VIRTIO_NET_HDR_F_NEEDS_CSUM)
}
const csumStart = header.EthernetMinimumSize + l3HdrLen
if vnetHdr.csumStart != csumStart {
t.Fatalf("vnetHdr.csumStart = %v, want %v", vnetHdr.csumStart, csumStart)
}
if vnetHdr.csumOffset != csumOffset {
t.Fatalf("vnetHdr.csumOffset = %v, want %v", vnetHdr.csumOffset, csumOffset)
}
gsoType := uint8(0)
if plen > gsoMSS {
gsoType = _VIRTIO_NET_HDR_GSO_TCPV6
}
if vnetHdr.gsoType != gsoType {
t.Fatalf("vnetHdr.gsoType = %v, want %v", vnetHdr.gsoType, gsoType)
}
b = b[virtioNetHdrSize:]
}
if eth {
h := header.Ethernet(b)
b = b[header.EthernetMinimumSize:]
if a := h.SourceAddress(); a != laddr {
t.Fatalf("SourceAddress() = %v, want %v", a, laddr)
}
if a := h.DestinationAddress(); a != raddr {
t.Fatalf("DestinationAddress() = %v, want %v", a, raddr)
}
if et := h.Type(); et != proto {
t.Fatalf("Type() = %v, want %v", et, proto)
}
}
if len(b) != len(want) {
t.Fatalf("Read returned %v bytes, want %v", len(b), len(want))
}
if !bytes.Equal(b, want) {
t.Fatalf("Read returned %x, want %x", b, want)
}
}
func TestWritePacket(t *testing.T) {
lengths := []int{0, 100, 1000}
eths := []bool{true, false}
gsos := []uint32{0, 32768}
for _, eth := range eths {
for _, plen := range lengths {
for _, gso := range gsos {
t.Run(
fmt.Sprintf("Eth=%v,PayloadLen=%v,GSOMaxSize=%v", eth, plen, gso),
func(t *testing.T) {
testWritePacket(t, plen, eth, gso, 0)
},
)
}
}
}
}
func TestHashedWritePacket(t *testing.T) {
lengths := []int{0, 100, 1000}
eths := []bool{true, false}
gsos := []uint32{0, 32768}
hashes := []uint32{0, 1}
for _, eth := range eths {
for _, plen := range lengths {
for _, gso := range gsos {
for _, hash := range hashes {
t.Run(
fmt.Sprintf("Eth=%v,PayloadLen=%v,GSOMaxSize=%v,Hash=%d", eth, plen, gso, hash),
func(t *testing.T) {
testWritePacket(t, plen, eth, gso, hash)
},
)
}
}
}
}
}
func TestPreserveSrcAddress(t *testing.T) {
baddr := tcpip.LinkAddress("\xcc\xbb\xaa\x77\x88\x99")
c := newContext(t, &Options{Address: laddr, MTU: mtu, EthernetHeader: true})
defer c.cleanup()
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
// WritePacket panics given a prependable with anything less than
// the minimum size of the ethernet header.
// TODO(b/153685824): Figure out if this should use c.ep.MaxHeaderLength().
ReserveHeaderBytes: header.EthernetMinimumSize,
})
defer pkt.DecRef()
// Every PacketBuffer must have these set:
// See nic.writePacket.
pkt.NetworkProtocolNumber = proto
// Set LocalLinkAddress in route to the value of the bridged address.
pkt.EgressRoute.LocalLinkAddress = baddr
pkt.EgressRoute.RemoteLinkAddress = raddr
c.ep.AddHeader(pkt)
var pkts stack.PacketBufferList
pkts.PushBack(pkt)
if _, err := c.ep.WritePackets(pkts); err != nil {
t.Fatalf("WritePackets failed: %s", err)
}
// Read from the FD, then compare with what we wrote.
b := make([]byte, mtu)
n, err := unix.Read(c.readFDs[0], b)
if err != nil {
t.Fatalf("Read failed: %v", err)
}
b = b[:n]
h := header.Ethernet(b)
if a := h.SourceAddress(); a != baddr {
t.Fatalf("SourceAddress() = %v, want %v", a, baddr)
}
}
func TestDeliverPacket(t *testing.T) {
lengths := []int{100, 1000}
eths := []bool{true, false}
for _, eth := range eths {
for _, plen := range lengths {
t.Run(fmt.Sprintf("Eth=%v,PayloadLen=%v", eth, plen), func(t *testing.T) {
c := newContext(t, &Options{Address: laddr, MTU: mtu, EthernetHeader: eth})
defer c.cleanup()
// Build packet.
all := make([]byte, plen)
if _, err := rand.Read(all); err != nil {
t.Fatalf("rand.Read(all): %s", err)
}
// Make it look like an IPv4 packet.
all[0] = 0x40
wantPkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
ReserveHeaderBytes: header.EthernetMinimumSize,
Payload: buffer.MakeWithData(all),
})
defer wantPkt.DecRef()
if eth {
hdr := header.Ethernet(wantPkt.LinkHeader().Push(header.EthernetMinimumSize))
hdr.Encode(&header.EthernetFields{
SrcAddr: raddr,
DstAddr: laddr,
Type: proto,
})
all = append(hdr, all...)
}
// Write packet via the file descriptor.
if _, err := unix.Write(c.readFDs[0], all); err != nil {
t.Fatalf("Write failed: %v", err)
}
// Receive packet through the endpoint.
select {
case pi := <-c.ch:
defer pi.Contents.DecRef()
want := packetInfo{
Proto: proto,
Contents: wantPkt,
}
if !eth {
want.Proto = header.IPv4ProtocolNumber
}
checkPacketInfoEqual(t, pi, want)
case <-time.After(10 * time.Second):
t.Fatalf("Timed out waiting for packet")
}
})
}
}
}
func TestBufConfigMaxLength(t *testing.T) {
got := 0
for _, i := range BufConfig {
got += i
}
want := header.MaxIPPacketSize // maximum TCP packet size
if got < want {
t.Errorf("total buffer size is invalid: got %d, want >= %d", got, want)
}
}
func TestBufConfigFirst(t *testing.T) {
// The stack assumes that the TCP/IP header is entirely contained in the first view.
// Therefore, the first view needs to be large enough to contain the maximum TCP/IP
// header, which is 120 bytes (60 bytes for IP + 60 bytes for TCP).
want := 120
got := BufConfig[0]
if got < want {
t.Errorf("first view has an invalid size: got %d, want >= %d", got, want)
}
}
var capLengthTestCases = []struct {
comment string
config []int
n int
wantUsed int
wantLengths []int
}{
{
comment: "Single slice",
config: []int{256},
n: 128,
wantUsed: 1,
wantLengths: []int{128},
},
{
comment: "Multiple slices",
config: []int{128, 256},
n: 256,
wantUsed: 2,
wantLengths: []int{128, 128},
},
{
comment: "Entire buffer",
config: []int{128, 256},
n: 384,
wantUsed: 2,
wantLengths: []int{128, 256},
},
{
comment: "Entire buffer but not on the last slice",
config: []int{128, 256, 512},
n: 384,
wantUsed: 2,
wantLengths: []int{128, 256},
},
}
func TestIovecBuffer(t *testing.T) {
for _, c := range capLengthTestCases {
t.Run(c.comment, func(t *testing.T) {
b := newIovecBuffer(c.config, false /* skipsVnetHdr */)
defer b.release()
// Test initial allocation.
iovecs := b.nextIovecs()
if got, want := len(iovecs), len(c.config); got != want {
t.Fatalf("len(iovecs) = %d, want %d", got, want)
}
// Make a copy as iovecs points to internal slice. We will need this state
// later.
oldIovecs := append([]unix.Iovec(nil), iovecs...)
// Test the buffer that get pulled.
buf := b.pullBuffer(c.n)
defer buf.Release()
var lengths []int
buf.Apply(func(v *buffer.View) {
lengths = append(lengths, v.Size())
})
if !slices.Equal(lengths, c.wantLengths) {
t.Errorf("Pulled view lengths = %v, want %v", lengths, c.wantLengths)
}
// Test that new views get reallocated.
for i, newIov := range b.nextIovecs() {
if i < c.wantUsed {
if newIov.Base == oldIovecs[i].Base {
t.Errorf("b.views[%d] should have been reallocated", i)
}
} else {
if newIov.Base != oldIovecs[i].Base {
t.Errorf("b.views[%d] should not have been reallocated", i)
}
}
}
})
}
}
func TestIovecBufferSkipVnetHdr(t *testing.T) {
for _, test := range []struct {
desc string
readN int
wantLen int
}{
{
desc: "nothing read",
readN: 0,
wantLen: 0,
},
{
desc: "smaller than vnet header",
readN: virtioNetHdrSize - 1,
wantLen: 0,
},
{
desc: "header skipped",
readN: virtioNetHdrSize + 512,
wantLen: 512,
},
} {
t.Run(test.desc, func(t *testing.T) {
b := newIovecBuffer([]int{128, 256, 512, 1024}, true)
defer b.release()
// Pretend a read happened.
b.nextIovecs()
buf := b.pullBuffer(test.readN)
defer buf.Release()
if got, want := int(buf.Size()), test.wantLen; got != want {
t.Errorf("b.pullView(%d).Size() = %d; want %d", test.readN, got, want)
}
if got, want := len(buf.Flatten()), test.wantLen; got != want {
t.Errorf("b.pullView(%d).ToOwnedView() has length %d; want %d", test.readN, got, want)
}
})
}
}
// fakeNetworkDispatcher delivers packets to pkts.
type fakeNetworkDispatcher struct {
pkts []*stack.PacketBuffer
}
func (d *fakeNetworkDispatcher) DeliverNetworkPacket(_ tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) {
d.pkts = append(d.pkts, pkt.Clone())
}
func (*fakeNetworkDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer) {
panic("not implemented")
}
func TestDispatchPacketFormat(t *testing.T) {
for _, test := range []struct {
name string
newDispatcher func(fd int, e *endpoint, opts *Options) (linkDispatcher, error)
ethHdr []byte
netHdr []byte
}{
{
name: "readVDispatcher",
newDispatcher: newReadVDispatcher,
ethHdr: []byte{
1, 2, 3, 4, 5, 60,
1, 2, 3, 4, 5, 61,
8, 0,
},
netHdr: []byte{0x40, 0, 0, 0},
},
{
name: "recvMMsgDispatcher",
newDispatcher: newRecvMMsgDispatcher,
ethHdr: []byte{
1, 2, 3, 4, 5, 60,
1, 2, 3, 4, 5, 61,
8, 0,
},
netHdr: []byte{0x40, 0, 0, 0},
},
{
name: "readVDispatcherNoEth",
newDispatcher: newReadVDispatcher,
netHdr: []byte{0x40, 0, 0, 0},
},
{
name: "readVDispatcherNoEthIPv6",
newDispatcher: newReadVDispatcher,
netHdr: []byte{
0x60, 0, 0, 0, // IPv6 Preamble
0, 0x04, 0, 0, // IPv6 Preamble
byte(header.IPv6HopByHopOptionsExtHdrIdentifier), // Next header
0xff, // Hop limit
0, 0, 0, 0, 0, 0, 0, 0, // Src addr
0, 0, 0, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 0, 0, // Dst addr
0, 0, 0, 0, 0, 0, 0, 0,
// Hop by hop extension header.
byte(header.IPv6DestinationOptionsExtHdrIdentifier),
0x8, // Length.
0, 0, 0, 0, 0, 0, // Padding.
// Destination options extension header.
0, // Next header (none)
0x8, // Length.
0, 0, 0, 0, 0, 0, // Padding.
// payload data.
1, 2, 3, 4,
},
},
} {
t.Run(test.name, func(t *testing.T) {
// Create a socket pair to send/recv.
fds, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_DGRAM, 0)
if err != nil {
t.Fatal(err)
}
data := append(test.ethHdr, test.netHdr...)
err = unix.Sendmsg(fds[1], data, nil, nil, 0)
if err != nil {
t.Fatal(err)
}
// Create and run dispatcher once.
sink := &fakeNetworkDispatcher{}
var addr tcpip.LinkAddress
if len(test.ethHdr) > 0 {
addr = tcpip.LinkAddress(test.ethHdr[:header.EthernetAddressSize])
}
d, err := test.newDispatcher(fds[0], &endpoint{
addr: addr,
hdrSize: len(test.ethHdr),
dispatcher: sink,
}, &Options{ProcessorsPerChannel: 1})
if err != nil {
t.Fatal(err)
}
defer d.release()
if ok, err := d.dispatch(); !ok || err != nil {
t.Fatalf("d.dispatch() = %v, %v", ok, err)
}
// Verify packet.
if got, want := len(sink.pkts), 1; got != want {
t.Fatalf("len(sink.pkts) = %d, want %d", got, want)
}
pkt := sink.pkts[0]
defer pkt.DecRef()
if len(test.ethHdr) > 0 {
if got, want := len(pkt.LinkHeader().Slice()), header.EthernetMinimumSize; got != want {
t.Errorf("pkt.LinkHeader().View().Size() = %d, want %d", got, want)
}
}
if got, want := pkt.Data().Size(), len(test.netHdr); got != want {
t.Errorf("pkt.Data().Size() = %d, want %d", got, want)
}
wantProto := header.IPv4ProtocolNumber
if header.IPVersion(test.netHdr[:]) == header.IPv6Version {
wantProto = header.IPv6ProtocolNumber
}
if pkt.NetworkProtocolNumber != wantProto {
t.Errorf("pkt.NetworkProtocolNumber = %d, want %d", pkt.NetworkProtocolNumber, wantProto)
}
})
}
}
func TestMain(m *testing.M) {
refs.SetLeakMode(refs.LeaksPanic)
code := m.Run()
refs.DoLeakCheck()
os.Exit(code)
}