From f2a57c9dac27bd6dcd1c53156d94bdcbdb7a542e Mon Sep 17 00:00:00 2001 From: Bhasker Hariharan Date: Thu, 6 Jan 2022 17:30:26 -0800 Subject: [PATCH] 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 --- pkg/tcpip/link/sharedmem/BUILD | 3 + pkg/tcpip/link/sharedmem/pipe/pipe_test.go | 2 +- pkg/tcpip/link/sharedmem/pipe/rx.go | 5 ++ pkg/tcpip/link/sharedmem/queuepair.go | 6 ++ pkg/tcpip/link/sharedmem/server_tx.go | 55 ++++++++---- .../link/sharedmem/sharedmem_server_test.go | 87 +++++++++++++++++-- pkg/tcpip/link/sharedmem/tx.go | 4 +- pkg/tcpip/link/sniffer/sniffer.go | 2 +- 8 files changed, 133 insertions(+), 31 deletions(-) diff --git a/pkg/tcpip/link/sharedmem/BUILD b/pkg/tcpip/link/sharedmem/BUILD index fe3ba2ed8..598719fe9 100644 --- a/pkg/tcpip/link/sharedmem/BUILD +++ b/pkg/tcpip/link/sharedmem/BUILD @@ -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", ], ) diff --git a/pkg/tcpip/link/sharedmem/pipe/pipe_test.go b/pkg/tcpip/link/sharedmem/pipe/pipe_test.go index 2777f1411..c91671981 100644 --- a/pkg/tcpip/link/sharedmem/pipe/pipe_test.go +++ b/pkg/tcpip/link/sharedmem/pipe/pipe_test.go @@ -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) diff --git a/pkg/tcpip/link/sharedmem/pipe/rx.go b/pkg/tcpip/link/sharedmem/pipe/rx.go index f22e533ac..1f44f5f14 100644 --- a/pkg/tcpip/link/sharedmem/pipe/rx.go +++ b/pkg/tcpip/link/sharedmem/pipe/rx.go @@ -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 diff --git a/pkg/tcpip/link/sharedmem/queuepair.go b/pkg/tcpip/link/sharedmem/queuepair.go index b12647fdd..c289bb873 100644 --- a/pkg/tcpip/link/sharedmem/queuepair.go +++ b/pkg/tcpip/link/sharedmem/queuepair.go @@ -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. diff --git a/pkg/tcpip/link/sharedmem/server_tx.go b/pkg/tcpip/link/sharedmem/server_tx.go index 79a9e382b..be0ebf931 100644 --- a/pkg/tcpip/link/sharedmem/server_tx.go +++ b/pkg/tcpip/link/sharedmem/server_tx.go @@ -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 diff --git a/pkg/tcpip/link/sharedmem/sharedmem_server_test.go b/pkg/tcpip/link/sharedmem/sharedmem_server_test.go index 1bc58614e..8ebb9591e 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_server_test.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_server_test.go @@ -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) + } +} diff --git a/pkg/tcpip/link/sharedmem/tx.go b/pkg/tcpip/link/sharedmem/tx.go index a74fc012b..ab7d47e8b 100644 --- a/pkg/tcpip/link/sharedmem/tx.go +++ b/pkg/tcpip/link/sharedmem/tx.go @@ -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 diff --git a/pkg/tcpip/link/sniffer/sniffer.go b/pkg/tcpip/link/sniffer/sniffer.go index fabe6411c..339992dce 100644 --- a/pkg/tcpip/link/sniffer/sniffer.go +++ b/pkg/tcpip/link/sniffer/sniffer.go @@ -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) }