diff --git a/pkg/sentry/socket/netstack/stack.go b/pkg/sentry/socket/netstack/stack.go index 7e4502a0c..c579f5a2f 100644 --- a/pkg/sentry/socket/netstack/stack.go +++ b/pkg/sentry/socket/netstack/stack.go @@ -298,7 +298,7 @@ func (s *Stack) newVeth(ctx context.Context, linkAttrs map[uint16]nlmsg.BytesVie } } } - ep, peerEP := veth.NewPair(defaultMTU) + ep, peerEP := veth.NewPair(defaultMTU, veth.DefaultBacklogSize) id := s.Stack.NextNICID() peerID := peerStack.Stack.NextNICID() if ifname == "" { diff --git a/pkg/tcpip/link/veth/BUILD b/pkg/tcpip/link/veth/BUILD index c0dcf97c5..3980a6a08 100644 --- a/pkg/tcpip/link/veth/BUILD +++ b/pkg/tcpip/link/veth/BUILD @@ -46,6 +46,7 @@ go_test( deps = [ "//pkg/buffer", "//pkg/refs", + "//pkg/sync", "//pkg/tcpip", "//pkg/tcpip/header", "//pkg/tcpip/link/ethernet", diff --git a/pkg/tcpip/link/veth/veth.go b/pkg/tcpip/link/veth/veth.go index e9e35b510..c22f8e960 100644 --- a/pkg/tcpip/link/veth/veth.go +++ b/pkg/tcpip/link/veth/veth.go @@ -21,6 +21,9 @@ import ( "gvisor.dev/gvisor/pkg/tcpip/stack" ) +// DefaultBacklogSize is the default size of a veth device's buffer. +const DefaultBacklogSize = 1000 + var _ stack.LinkEndpoint = (*Endpoint)(nil) var _ stack.GSOEndpoint = (*Endpoint)(nil) @@ -62,8 +65,6 @@ type vethPacket struct { pkt *stack.PacketBuffer } -const backlogQueueSize = 64 - // Endpoint is link layer endpoint that redirects packets to a pair veth endpoint. // // +stateify savable @@ -84,7 +85,7 @@ type Endpoint struct { } // NewPair creates a new veth pair. -func NewPair(mtu uint32) (*Endpoint, *Endpoint) { +func NewPair(mtu, backlogQueueSize uint32) (*Endpoint, *Endpoint) { veth := veth{ backlogQueue: make(chan vethPacket, backlogQueueSize), mtu: mtu, @@ -219,14 +220,18 @@ func (e *Endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) Payload: payload.DeepClone(), }) payload.Release() - (e.veth.backlogQueue) <- vethPacket{ + select { + case (e.veth.backlogQueue) <- vethPacket{ e: e.peer, protocol: pkt.NetworkProtocolNumber, pkt: newPkt, + }: + n++ + default: + newPkt.DecRef() + return n, &tcpip.ErrNoBufferSpace{} } - n++ } - return n, nil } diff --git a/pkg/tcpip/link/veth/veth_test.go b/pkg/tcpip/link/veth/veth_test.go index e32502be5..985c9e8d2 100644 --- a/pkg/tcpip/link/veth/veth_test.go +++ b/pkg/tcpip/link/veth/veth_test.go @@ -21,6 +21,7 @@ import ( "gvisor.dev/gvisor/pkg/buffer" "gvisor.dev/gvisor/pkg/refs" + "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/header" "gvisor.dev/gvisor/pkg/tcpip/link/ethernet" @@ -30,7 +31,7 @@ import ( func TestSetLinkAddress(t *testing.T) { addrs := []tcpip.LinkAddress{"abc", "def"} - e, e2 := veth.NewPair(1500) + e, e2 := veth.NewPair(1500, veth.DefaultBacklogSize) defer e.Close() defer e2.Close() for _, addr := range addrs { @@ -44,11 +45,13 @@ func TestSetLinkAddress(t *testing.T) { type testNetworkDispatcher struct { ch chan *stack.PacketBuffer + wg *sync.WaitGroup } func (d *testNetworkDispatcher) DeliverNetworkPacket(_ tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { pkt.IncRef() d.ch <- pkt + d.wg.Wait() } func (*testNetworkDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer) { @@ -64,7 +67,7 @@ func TestWritePacket(t *testing.T) { nicID = 5 ) - veth1, veth2 := veth.NewPair(1500) + veth1, veth2 := veth.NewPair(1500, veth.DefaultBacklogSize) veth1.SetLinkAddress(localLinkAddr) veth2.SetLinkAddress(remoteLinkAddr) @@ -73,7 +76,8 @@ func TestWritePacket(t *testing.T) { t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) } - sink := &testNetworkDispatcher{ch: make(chan *stack.PacketBuffer, 1)} + var wg sync.WaitGroup + sink := &testNetworkDispatcher{ch: make(chan *stack.PacketBuffer, 1), wg: &wg} veth2Ethernet := ethernet.New(veth2) veth2Ethernet.Attach(sink) @@ -98,13 +102,63 @@ func TestWritePacket(t *testing.T) { } } +func TestVethOverflows(t *testing.T) { + const ( + localLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x06") + remoteLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x07") + + netProto = 55 + nicID = 5 + ) + + backlogSize := uint32(1) + veth1, veth2 := veth.NewPair(1500, backlogSize) + veth1.SetLinkAddress(localLinkAddr) + veth2.SetLinkAddress(remoteLinkAddr) + + s := stack.New(stack.Options{}) + if err := s.CreateNIC(nicID, ethernet.New(veth1)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) + } + + var wg sync.WaitGroup + wg.Add(1) + // Use a unbuffered channel in the dispatcher so that the received packet won't be processed. + sink := &testNetworkDispatcher{ch: make(chan *stack.PacketBuffer), wg: &wg} + veth2Ethernet := ethernet.New(veth2) + veth2Ethernet.Attach(sink) + + // Send 3 packets, the first packet will be blocked at the sink's channel, the + // second packet will be waiting at the veth device's backlog queue, the third + // packet will be rejected since it overflows the backlog queue. + if err := s.WritePacketToRemote(nicID, remoteLinkAddr, netProto, buffer.Buffer{}); err != nil { + t.Fatalf("s.WritePacketToRemote(%d, %s, _): %s", nicID, remoteLinkAddr, err) + } + select { + case pkt := <-sink.ch: + defer pkt.DecRef() + if err := s.WritePacketToRemote(nicID, remoteLinkAddr, netProto, buffer.Buffer{}); err != nil { + t.Fatalf("s.WritePacketToRemote(%d, %s, _): %s", nicID, remoteLinkAddr, err) + } + if err := s.WritePacketToRemote(nicID, remoteLinkAddr, netProto, buffer.Buffer{}); err == nil { + t.Fatalf("s.WritePacketToRemote(%d, %s, _) got: %v, want: %v", nicID, remoteLinkAddr, err, &tcpip.ErrNoBufferSpace{}) + } + wg.Done() + } + pkt := <-sink.ch + if pkt == nil { + t.Fatal("expected to read a packet") + } + pkt.DecRef() +} + func TestDestroyDevices(t *testing.T) { const ( vethFirstID = 5 vethSecondID = 6 ) - veth1, veth2 := veth.NewPair(1500) + veth1, veth2 := veth.NewPair(1500, veth.DefaultBacklogSize) s1 := stack.New(stack.Options{}) if err := s1.CreateNIC(vethFirstID, ethernet.New(veth1)); err != nil { @@ -129,7 +183,7 @@ func TestDestroyDevices(t *testing.T) { func TestMTU(t *testing.T) { mtus := []uint32{100, 200} - e, e2 := veth.NewPair(1500) + e, e2 := veth.NewPair(1500, veth.DefaultBacklogSize) defer e.Close() defer e2.Close() for _, mtu := range mtus { diff --git a/pkg/tcpip/stack/bridge_test.go b/pkg/tcpip/stack/bridge_test.go index c066abd0b..3bbf5bce5 100644 --- a/pkg/tcpip/stack/bridge_test.go +++ b/pkg/tcpip/stack/bridge_test.go @@ -116,7 +116,7 @@ func TestWritePacketBetweenDevices(t *testing.T) { ) // Creates a pair of veth devices which will be attached to different // network stacks. - veth1, veth2 := veth.NewPair(1500) + veth1, veth2 := veth.NewPair(1500, veth.DefaultBacklogSize) veth1.SetLinkAddress(vethLinkAddr1) veth2.SetLinkAddress(vethLinkAddr2) ch1 := channel.New(1, header.EthernetMinimumSize, channelLinkAddr1) @@ -199,7 +199,7 @@ func TestBridgeFDB(t *testing.T) { vethID2 = 8 bridgeID = 9 ) - veth1, veth2 := veth.NewPair(1500) + veth1, veth2 := veth.NewPair(1500, veth.DefaultBacklogSize) ch := channel.New(1, header.EthernetMinimumSize, channelLinkAddr) bridgeEndpoint := stack.NewBridgeEndpoint(1500) bridgeEndpoint.SetLinkAddress(bridgeLinkAddr)