Fixes multiple bugs in server_rx implementation.

fillPacket was incorrectly setting the buffer sizes and
causing large packets to be egressed incorrectly resulting
in packet drops when the MTU was > buffer size.

PiperOrigin-RevId: 420178122
This commit is contained in:
Bhasker Hariharan
2022-01-06 17:33:31 -08:00
committed by gVisor bot
parent 381a17d923
commit f2a57c9dac
8 changed files with 133 additions and 31 deletions
+3
View File
@@ -56,15 +56,18 @@ go_test(
srcs = ["sharedmem_server_test.go"],
deps = [
":sharedmem",
"//pkg/log",
"//pkg/tcpip",
"//pkg/tcpip/adapters/gonet",
"//pkg/tcpip/header",
"//pkg/tcpip/link/qdisc/fifo",
"//pkg/tcpip/link/sniffer",
"//pkg/tcpip/network/ipv4",
"//pkg/tcpip/network/ipv6",
"//pkg/tcpip/stack",
"//pkg/tcpip/transport/tcp",
"//pkg/tcpip/transport/udp",
"@org_golang_x_sync//errgroup:go_default_library",
"@org_golang_x_sys//unix:go_default_library",
],
)
+1 -1
View File
@@ -461,7 +461,7 @@ func TestConcurrentReaderWriter(t *testing.T) {
tr := rand.New(rand.NewSource(99))
rr := rand.New(rand.NewSource(99))
b := make([]byte, 100)
b := make([]byte, 4096)
var tx Tx
tx.Init(b)
+5
View File
@@ -87,6 +87,11 @@ func (r *Rx) Flush() {
r.tail = r.head
}
// Abort unpulls any pulled buffers.
func (r *Rx) Abort() {
r.head = r.tail
}
// Bytes returns the byte slice on which the pipe operates.
func (r *Rx) Bytes() []byte {
return r.p.buffer
+6
View File
@@ -50,6 +50,12 @@ const (
// defaultSharedDataSize is the size of the sharedData region used to
// enable/disable notifications.
defaultSharedDataSize = 4 << 10 // 4KiB
// DefaultBufferSize is the size of each individual buffer that the data
// region is broken down into to hold packet data. Should be larger than
// 1500 + 14 (Ethernet header) + 10 (VirtIO header) to fit each packet
// in a single buffer.
DefaultBufferSize = 2048
)
// A QueuePair represents a pair of TX/RX queues.
+37 -18
View File
@@ -113,6 +113,27 @@ func (s *serverTx) cleanup() {
s.eventFD.Close()
}
// acquireBuffers acquires enough buffers to hold all the data in views or
// returns nil if not enough buffers are currently available.
func (s *serverTx) acquireBuffers(views []buffer.View, buffers []queue.RxBuffer) (acquiredBuffers []queue.RxBuffer) {
acquiredBuffers = buffers[:0]
wantBytes := 0
for i := range views {
wantBytes += len(views[i])
}
for wantBytes > 0 {
var b []byte
if b = s.fillPipe.Pull(); b == nil {
s.fillPipe.Abort()
return nil
}
rxBuffer := queue.DecodeRxBufferHeader(b)
acquiredBuffers = append(acquiredBuffers, rxBuffer)
wantBytes -= int(rxBuffer.Size)
}
return acquiredBuffers
}
// fillPacket copies the data in the provided views into buffers pulled from the
// fillPipe and returns a slice of RxBuffers that contain the copied data as
// well as the total number of bytes copied.
@@ -120,50 +141,48 @@ func (s *serverTx) cleanup() {
// To avoid allocations the filledBuffers are appended to the buffers slice
// which will be grown as required.
func (s *serverTx) fillPacket(views []buffer.View, buffers []queue.RxBuffer) (filledBuffers []queue.RxBuffer, totalCopied uint32) {
filledBuffers = buffers[:0]
// fillBuffer copies as much of the views as possible into the provided buffer
// and returns any left over views (if any).
fillBuffer := func(buffer *queue.RxBuffer, views []buffer.View) (left []buffer.View) {
if len(views) == 0 {
return nil
}
availBytes := buffer.Size
copied := uint64(0)
availBytes := buffer.Size
for availBytes > 0 && len(views) > 0 {
n := copy(s.data[buffer.Offset+copied:][:uint64(buffer.Size)-copied], views[0])
copied += uint64(n)
availBytes -= uint32(n)
views[0].TrimFront(n)
if !views[0].IsEmpty() {
break
}
views = views[1:]
copied += uint64(n)
availBytes -= uint32(n)
}
buffer.Size = uint32(copied)
return views
}
for len(views) > 0 {
var b []byte
// Spin till we get a free buffer reposted by the peer.
for {
if b = s.fillPipe.Pull(); b != nil {
break
}
}
rxBuffer := queue.DecodeRxBufferHeader(b)
bufs := s.acquireBuffers(views, buffers)
if bufs == nil {
return nil, 0
}
for i := 0; len(views) > 0 && i < len(bufs); i++ {
// Copy the packet into the posted buffer.
views = fillBuffer(&rxBuffer, views)
totalCopied += rxBuffer.Size
filledBuffers = append(filledBuffers, rxBuffer)
views = fillBuffer(&bufs[i], views)
totalCopied += bufs[i].Size
}
return filledBuffers, totalCopied
return bufs, totalCopied
}
func (s *serverTx) transmit(views []buffer.View) bool {
buffers := make([]queue.RxBuffer, 8)
buffers, totalCopied := s.fillPacket(views, buffers)
if totalCopied == 0 {
// drop the packet as not enough buffers were probably available
// to send.
return false
}
b := s.completionPipe.Push(queue.RxCompletionSize(len(buffers)))
if b == nil {
return false
@@ -22,13 +22,17 @@ import (
"io"
"net"
"net/http"
"strings"
"syscall"
"testing"
"golang.org/x/sync/errgroup"
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/log"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
"gvisor.dev/gvisor/pkg/tcpip/header"
"gvisor.dev/gvisor/pkg/tcpip/link/qdisc/fifo"
"gvisor.dev/gvisor/pkg/tcpip/link/sharedmem"
"gvisor.dev/gvisor/pkg/tcpip/link/sniffer"
"gvisor.dev/gvisor/pkg/tcpip/network/ipv4"
@@ -45,13 +49,18 @@ const (
remoteIPv4Address = tcpip.Address("\x0a\x00\x00\x02")
serverPort = 10001
defaultMTU = 1500
defaultMTU = 65536
defaultBufferSize = 1500
// qDisc options
numQueues = 1
queueLen = 1000
)
type stackOptions struct {
ep stack.LinkEndpoint
addr tcpip.Address
ep stack.LinkEndpoint
addr tcpip.Address
enablePacketLogs bool
}
func newStackWithOptions(stackOpts stackOptions) (*stack.Stack, error) {
@@ -67,9 +76,16 @@ func newStackWithOptions(stackOpts stackOptions) (*stack.Stack, error) {
TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol},
})
nicID := tcpip.NICID(1)
sniffEP := sniffer.New(stackOpts.ep)
opts := stack.NICOptions{Name: "eth0"}
if err := st.CreateNICWithOptions(nicID, sniffEP, opts); err != nil {
ep := stackOpts.ep
if stackOpts.enablePacketLogs {
ep = sniffer.New(stackOpts.ep)
}
qDisc := fifo.New(ep, int(numQueues), int(queueLen))
opts := stack.NICOptions{
Name: "eth0",
QDisc: qDisc,
}
if err := st.CreateNICWithOptions(nicID, ep, opts); err != nil {
return nil, fmt.Errorf("method CreateNICWithOptions(%d, _, %v) failed: %s", nicID, opts, err)
}
@@ -106,7 +122,7 @@ func newClientStack(t *testing.T, qPair *sharedmem.QueuePair, peerFD int) (*stac
if err != nil {
return nil, fmt.Errorf("failed to create sharedmem endpoint: %s", err)
}
st, err := newStackWithOptions(stackOptions{ep: ep, addr: localIPv4Address})
st, err := newStackWithOptions(stackOptions{ep: ep, addr: localIPv4Address, enablePacketLogs: true})
if err != nil {
return nil, fmt.Errorf("failed to create client stack: %s", err)
}
@@ -125,7 +141,7 @@ func newServerStack(t *testing.T, qPair *sharedmem.QueuePair, peerFD int) (*stac
if err != nil {
return nil, fmt.Errorf("failed to create sharedmem endpoint: %s", err)
}
st, err := newStackWithOptions(stackOptions{ep: ep, addr: remoteIPv4Address})
st, err := newStackWithOptions(stackOptions{ep: ep, addr: remoteIPv4Address, enablePacketLogs: true})
if err != nil {
return nil, fmt.Errorf("failed to create client stack: %s", err)
}
@@ -185,7 +201,7 @@ func TestServerRoundTrip(t *testing.T) {
t.Fatalf("failed to start TCP Listener: %s", err)
}
defer l.Close()
var responseString = "response"
var responseString = strings.Repeat("response", 8<<10)
go func() {
http.Serve(l, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Write([]byte(responseString))
@@ -218,3 +234,56 @@ func TestServerRoundTrip(t *testing.T) {
t.Fatalf("unexpected response got: %s, want: %s", got, want)
}
}
func TestServerRoundTripStress(t *testing.T) {
ctx := newTestContext(t)
defer ctx.cleanup()
listenAddr := tcpip.FullAddress{Addr: remoteIPv4Address, Port: serverPort}
l, err := gonet.ListenTCP(ctx.serverStk, listenAddr, ipv4.ProtocolNumber)
if err != nil {
t.Fatalf("failed to start TCP Listener: %s", err)
}
defer l.Close()
var responseString = strings.Repeat("response", 8<<10)
go func() {
http.Serve(l, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Write([]byte(responseString))
}))
}()
dialFunc := func(address, protocol string) (net.Conn, error) {
return gonet.DialTCP(ctx.clientStk, listenAddr, ipv4.ProtocolNumber)
}
serverURL := fmt.Sprintf("http://[%s]:%d/", net.IP(remoteIPv4Address), serverPort)
var errs errgroup.Group
for i := 0; i < 1000; i++ {
errs.Go(func() error {
httpClient := &http.Client{
Transport: &http.Transport{
Dial: dialFunc,
},
}
response, err := httpClient.Get(serverURL)
if err != nil {
return fmt.Errorf("httpClient.Get(\"/\") failed: %s", err)
}
if got, want := response.StatusCode, http.StatusOK; got != want {
return fmt.Errorf("unexpected status code got: %d, want: %d", got, want)
}
body, err := io.ReadAll(response.Body)
if err != nil {
return fmt.Errorf("io.ReadAll(response.Body) failed: %s", err)
}
response.Body.Close()
if got, want := string(body), responseString; got != want {
return fmt.Errorf("unexpected response got: %s, want: %s", got, want)
}
log.Infof("worker: %d read %d bytes", len(body))
return nil
})
}
if err := errs.Wait(); err != nil {
t.Fatalf("request failed: %s", err)
}
}
+2 -2
View File
@@ -43,7 +43,7 @@ type tx struct {
//
// The caller always retains ownership of all file descriptors passed in. The
// queue implementation will duplicate any that it may need in the future.
func (t *tx) init(mtu uint32, c *QueueConfig) error {
func (t *tx) init(bufferSize uint32, c *QueueConfig) error {
// Map in all buffers.
txPipe, err := getBuffer(c.TxPipeFD)
if err != nil {
@@ -73,7 +73,7 @@ func (t *tx) init(mtu uint32, c *QueueConfig) error {
// Initialize state based on buffers.
t.q.Init(txPipe, rxPipe, sharedDataPointer(sharedData))
t.ids.init()
t.bufs.init(0, len(data), int(mtu))
t.bufs.init(0, len(data), int(bufferSize))
t.data = data
t.eventFD = c.EventFD
t.sharedDataFD = c.SharedDataFD
+1 -1
View File
@@ -166,7 +166,7 @@ func (e *endpoint) dumpPacket(dir direction, protocol tcpip.NetworkProtocolNumbe
// forwards the request to the lower endpoint.
func (e *endpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) {
for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() {
e.dumpPacket(directionSend, protocol, pkt)
e.dumpPacket(directionSend, pkt.NetworkProtocolNumber, pkt)
}
return e.Endpoint.WritePackets(r, pkts, protocol)
}