mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
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:
committed by
gVisor bot
parent
381a17d923
commit
f2a57c9dac
@@ -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",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user