mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
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
729 lines
19 KiB
Go
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)
|
|
}
|