DecRef ptr niling

PiperOrigin-RevId: 480518221
This commit is contained in:
Kevin Krakauer
2022-10-11 20:27:23 -07:00
committed by gVisor bot
parent 7bb273341e
commit 607fdc536d
14 changed files with 60 additions and 50 deletions
-4
View File
@@ -218,7 +218,3 @@ analyzers:
suppress:
- "comment on exported type Translation" # Intentional.
- "comment on exported type PinnedRange" # Intentional.
ST1016: # CheckReceiverNamesIdentical
internal:
exclude:
- pkg/tcpip/stack/packet_buffer.go # TODO(b/233086175): Remove.
+2 -2
View File
@@ -63,7 +63,7 @@ func (q *queue) Read() stack.PacketBufferPtr {
case p := <-q.c:
return p
default:
return nil
return stack.PacketBufferPtr{}
}
}
@@ -72,7 +72,7 @@ func (q *queue) ReadContext(ctx context.Context) stack.PacketBufferPtr {
case pkt := <-q.c:
return pkt
case <-ctx.Done():
return nil
return stack.PacketBufferPtr{}
}
}
@@ -71,10 +71,10 @@ func (pl *packetBufferCircularList) pushBack(pb stack.PacketBufferPtr) {
//go:nosplit
func (pl *packetBufferCircularList) removeFront() stack.PacketBufferPtr {
if pl.isEmpty() {
return nil
return stack.PacketBufferPtr{}
}
ret := pl.pbs[pl.head]
pl.pbs[pl.head] = nil
pl.pbs[pl.head] = stack.PacketBufferPtr{}
pl.head = (pl.head + 1) % len(pl.pbs)
pl.size--
return ret
+1 -1
View File
@@ -84,7 +84,7 @@ func TestWriteRefusedAfterClosed(t *testing.T) {
linkEp := fifo.New(nil, 1, 2)
linkEp.Close()
err := linkEp.WritePacket(nil)
err := linkEp.WritePacket(stack.PacketBufferPtr{})
_, ok := err.(*tcpip.ErrClosedForSend)
if !ok {
t.Errorf("got err = %s, want %s", err, &tcpip.ErrClosedForSend{})
@@ -158,25 +158,25 @@ func (f *Fragmentation) Process(
id FragmentID, first, last uint16, more bool, proto uint8, pkt stack.PacketBufferPtr) (
stack.PacketBufferPtr, uint8, bool, error) {
if first > last {
return nil, 0, false, fmt.Errorf("first=%d is greater than last=%d: %w", first, last, ErrInvalidArgs)
return stack.PacketBufferPtr{}, 0, false, fmt.Errorf("first=%d is greater than last=%d: %w", first, last, ErrInvalidArgs)
}
if first%f.blockSize != 0 {
return nil, 0, false, fmt.Errorf("first=%d is not a multiple of block size=%d: %w", first, f.blockSize, ErrInvalidArgs)
return stack.PacketBufferPtr{}, 0, false, fmt.Errorf("first=%d is not a multiple of block size=%d: %w", first, f.blockSize, ErrInvalidArgs)
}
fragmentSize := last - first + 1
if more && fragmentSize%f.blockSize != 0 {
return nil, 0, false, fmt.Errorf("fragment size=%d bytes is not a multiple of block size=%d on non-final fragment: %w", fragmentSize, f.blockSize, ErrInvalidArgs)
return stack.PacketBufferPtr{}, 0, false, fmt.Errorf("fragment size=%d bytes is not a multiple of block size=%d on non-final fragment: %w", fragmentSize, f.blockSize, ErrInvalidArgs)
}
if l := pkt.Data().Size(); l != int(fragmentSize) {
return nil, 0, false, fmt.Errorf("got fragment size=%d bytes not equal to the expected fragment size=%d bytes (first=%d last=%d): %w", l, fragmentSize, first, last, ErrInvalidArgs)
return stack.PacketBufferPtr{}, 0, false, fmt.Errorf("got fragment size=%d bytes not equal to the expected fragment size=%d bytes (first=%d last=%d): %w", l, fragmentSize, first, last, ErrInvalidArgs)
}
f.mu.Lock()
if f.reassemblers == nil {
return nil, 0, false, fmt.Errorf("Release() called before fragmentation processing could finish")
return stack.PacketBufferPtr{}, 0, false, fmt.Errorf("Release() called before fragmentation processing could finish")
}
r, ok := f.reassemblers[id]
@@ -201,7 +201,7 @@ func (f *Fragmentation) Process(
f.mu.Lock()
f.release(r, false /* timedOut */)
f.mu.Unlock()
return nil, 0, false, fmt.Errorf("fragmentation processing error: %w", err)
return stack.PacketBufferPtr{}, 0, false, fmt.Errorf("fragmentation processing error: %w", err)
}
f.mu.Lock()
f.memSize += memConsumed
@@ -253,12 +253,12 @@ func (f *Fragmentation) release(r *reassembler, timedOut bool) {
}
if !r.pkt.IsNil() {
r.pkt.DecRef()
r.pkt = nil
r.pkt = stack.PacketBufferPtr{}
}
for _, h := range r.holes {
if !h.pkt.IsNil() {
h.pkt.DecRef()
h.pkt = nil
h.pkt = stack.PacketBufferPtr{}
}
}
r.holes = nil
@@ -113,6 +113,7 @@ func TestFragmentationProcess(t *testing.T) {
f := NewFragmentation(minBlockSize, 2048, 512, reassembleTimeout, &faketime.NullClock{}, nil)
firstFragmentProto := c.in[0].proto
for i, in := range c.in {
in := in
defer in.pkt.DecRef()
defer c.out[i].buf.Release()
resPkt, proto, done, err := f.Process(in.id, in.first, in.last, in.more, in.proto, in.pkt)
@@ -631,7 +632,7 @@ func TestTimeoutHandler(t *testing.T) {
},
},
wantError: false,
wantPkt: nil,
wantPkt: stack.PacketBufferPtr{},
},
{
name: "second pkt is ignored",
@@ -663,7 +664,7 @@ func TestTimeoutHandler(t *testing.T) {
},
},
wantError: true,
wantPkt: nil,
wantPkt: stack.PacketBufferPtr{},
},
}
@@ -671,7 +672,7 @@ func TestTimeoutHandler(t *testing.T) {
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
handler := &testTimeoutHandler{pkt: nil}
handler := &testTimeoutHandler{pkt: stack.PacketBufferPtr{}}
f := NewFragmentation(minBlockSize, HighFragThreshold, LowFragThreshold, reassembleTimeout, &faketime.NullClock{}, handler)
@@ -702,7 +703,7 @@ func TestTimeoutHandler(t *testing.T) {
}
func TestFragmentSurvivesReleaseJob(t *testing.T) {
handler := &testTimeoutHandler{pkt: nil}
handler := &testTimeoutHandler{pkt: stack.PacketBufferPtr{}}
c := faketime.NewManualClock()
f := NewFragmentation(minBlockSize, HighFragThreshold, LowFragThreshold, reassembleTimeout, c, handler)
pkt := pkt(2, "01")
@@ -67,7 +67,7 @@ func (r *reassembler) process(first, last uint16, more bool, proto uint8, pkt st
// A concurrent goroutine might have already reassembled
// the packet and emptied the heap while this goroutine
// was waiting on the mutex. We don't have to do anything in this case.
return nil, 0, false, 0, nil
return stack.PacketBufferPtr{}, 0, false, 0, nil
}
var holeFound bool
@@ -91,12 +91,12 @@ func (r *reassembler) process(first, last uint16, more bool, proto uint8, pkt st
// https://github.com/torvalds/linux/blob/38525c6/net/ipv4/inet_fragment.c#L349
if first < currentHole.first || currentHole.last < last {
// Incoming fragment only partially fits in the free hole.
return nil, 0, false, 0, ErrFragmentOverlap
return stack.PacketBufferPtr{}, 0, false, 0, ErrFragmentOverlap
}
if !more {
if !currentHole.final || currentHole.filled && currentHole.last != last {
// We have another final fragment, which does not perfectly overlap.
return nil, 0, false, 0, ErrFragmentConflict
return stack.PacketBufferPtr{}, 0, false, 0, ErrFragmentConflict
}
}
@@ -155,12 +155,12 @@ func (r *reassembler) process(first, last uint16, more bool, proto uint8, pkt st
}
if !holeFound {
// Incoming fragment is beyond end.
return nil, 0, false, 0, ErrFragmentConflict
return stack.PacketBufferPtr{}, 0, false, 0, ErrFragmentConflict
}
// Check if all the holes have been filled and we are ready to reassemble.
if r.filled < len(r.holes) {
return nil, 0, false, memConsumed, nil
return stack.PacketBufferPtr{}, 0, false, memConsumed, nil
}
sort.Slice(r.holes, func(i, j int) bool {
+2 -1
View File
@@ -256,7 +256,8 @@ func (e *endpoint) handleICMP(pkt stack.PacketBufferPtr) {
// It's possible that a raw socket expects to receive this.
e.dispatcher.DeliverTransportPacket(header.ICMPv4ProtocolNumber, pkt)
pkt = nil
pkt = stack.PacketBufferPtr{}
_ = pkt // Suppress unused variable warning.
sent := e.stats.icmp.packetsSent
if !e.protocol.allowICMPReply(header.ICMPv4EchoReply, header.ICMPv4UnusedCode) {
+1 -1
View File
@@ -34,7 +34,7 @@ go_template_instance(
prefix = "packetBuffer",
template = "//pkg/refsvfs2:refs_template",
types = {
"T": "PacketBuffer",
"T": "packetBuffer",
},
)
+30 -16
View File
@@ -35,7 +35,7 @@ const (
var pkPool = sync.Pool{
New: func() interface{} {
return &PacketBuffer{}
return &packetBuffer{}
},
}
@@ -59,9 +59,14 @@ type PacketBufferOptions struct {
}
// PacketBufferPtr is a pointer to a PacketBuffer.
type PacketBufferPtr = *PacketBuffer
//
// +stateify savable
type PacketBufferPtr struct {
// packetBuffer is the underlying packet buffer.
*packetBuffer
}
// A PacketBuffer contains all the data of a network packet.
// A packetBuffer contains all the data of a network packet.
//
// As a PacketBuffer traverses up the stack, it may be necessary to pass it to
// multiple endpoints.
@@ -103,7 +108,7 @@ type PacketBufferPtr = *PacketBuffer
// starting offset of each header in `buf`.
//
// +stateify savable
type PacketBuffer struct {
type packetBuffer struct {
_ sync.NoCopy
packetBufferRefs
@@ -173,7 +178,7 @@ type PacketBuffer struct {
// NewPacketBuffer creates a new PacketBuffer with opts.
func NewPacketBuffer(opts PacketBufferOptions) PacketBufferPtr {
pk := pkPool.Get().(PacketBufferPtr)
pk := pkPool.Get().(*packetBuffer)
pk.reset()
if opts.ReserveHeaderBytes != 0 {
v := bufferv2.NewViewSize(opts.ReserveHeaderBytes)
@@ -186,31 +191,36 @@ func NewPacketBuffer(opts PacketBufferOptions) PacketBufferPtr {
pk.NetworkPacketInfo.IsForwardedPacket = opts.IsForwardedPacket
pk.onRelease = opts.OnRelease
pk.InitRefs()
return pk
return PacketBufferPtr{
packetBuffer: pk,
}
}
// IncRef increments the PacketBuffer's refcount.
func (pk PacketBufferPtr) IncRef() PacketBufferPtr {
pk.packetBufferRefs.IncRef()
return pk
return PacketBufferPtr{
packetBuffer: pk.packetBuffer,
}
}
// DecRef decrements the PacketBuffer's refcount. If the refcount is
// decremented to zero, the PacketBuffer is returned to the PacketBuffer
// pool.
func (pk PacketBufferPtr) DecRef() {
func (pk *PacketBufferPtr) DecRef() {
pk.packetBufferRefs.DecRef(func() {
if pk.onRelease != nil {
pk.onRelease()
}
pk.buf.Release()
pkPool.Put(pk)
pkPool.Put(pk.packetBuffer)
})
pk.packetBuffer = nil
}
func (pk PacketBufferPtr) reset() {
*pk = PacketBuffer{}
func (pk *packetBuffer) reset() {
*pk = packetBuffer{}
}
// ReservedHeaderBytes returns the number of bytes initially reserved for
@@ -364,7 +374,7 @@ func (pk PacketBufferPtr) headerView(typ headerType) bufferv2.View {
// Clone makes a semi-deep copy of pk. The underlying packet payload is
// shared. Hence, no modifications is done to underlying packet payload.
func (pk PacketBufferPtr) Clone() PacketBufferPtr {
newPk := pkPool.Get().(PacketBufferPtr)
newPk := pkPool.Get().(*packetBuffer)
newPk.reset()
newPk.buf = pk.buf.Clone()
newPk.reserved = pk.reserved
@@ -384,7 +394,9 @@ func (pk PacketBufferPtr) Clone() PacketBufferPtr {
newPk.NetworkPacketInfo = pk.NetworkPacketInfo
newPk.tuple = pk.tuple
newPk.InitRefs()
return newPk
return PacketBufferPtr{
packetBuffer: newPk,
}
}
// ReserveHeaderBytes prepends reserved space for headers at the front
@@ -417,14 +429,16 @@ func (pk PacketBufferPtr) Network() header.Network {
// See PacketBuffer.Data for details about how a packet buffer holds an inbound
// packet.
func (pk PacketBufferPtr) CloneToInbound() PacketBufferPtr {
newPk := pkPool.Get().(PacketBufferPtr)
newPk := pkPool.Get().(*packetBuffer)
newPk.reset()
newPk.buf = pk.buf.Clone()
newPk.InitRefs()
// Treat unfilled header portion as reserved.
newPk.reserved = pk.AvailableHeaderBytes()
newPk.tuple = pk.tuple
return newPk
return PacketBufferPtr{
packetBuffer: newPk,
}
}
// DeepCopyForForwarding creates a deep copy of the packet buffer for
@@ -462,7 +476,7 @@ func (pk PacketBufferPtr) DeepCopyForForwarding(reservedHeaderBytes int) PacketB
// IsNil returns whether the pointer is logically nil.
func (pk PacketBufferPtr) IsNil() bool {
return pk == nil
return pk.packetBuffer == nil
}
// headerInfo stores metadata about a header in a packet.
+1 -1
View File
@@ -37,7 +37,7 @@ func (pl *PacketBufferList) AsSlice() []PacketBufferPtr {
func (pl *PacketBufferList) Reset() {
for i, pb := range pl.pbs {
pb.DecRef()
pl.pbs[i] = nil
pl.pbs[i] = PacketBufferPtr{}
}
pl.pbs = pl.pbs[:0]
}
+1 -1
View File
@@ -17,4 +17,4 @@ package stack
import "unsafe"
// PacketBufferStructSize is the minimal size of the packet buffer overhead.
const PacketBufferStructSize = int(unsafe.Sizeof(PacketBuffer{}))
const PacketBufferStructSize = int(unsafe.Sizeof(packetBuffer{}))
@@ -276,7 +276,7 @@ func (c *WriteContext) TryNewPacketBuffer(reserveHdrBytes int, data bufferv2.Buf
defer e.sendBufferSizeInUseMu.Unlock()
if !e.hasSendSpaceRLocked() {
return nil
return stack.PacketBufferPtr{}
}
// Note that we allow oversubscription - if there is any space at all in the
+1 -3
View File
@@ -116,8 +116,7 @@ func newIncomingSegment(id stack.TransportEndpointID, clock tcpip.Clock, pkt sta
s.window = seqnum.Size(hdr.WindowSize())
s.rcvdTime = clock.NowMonotonic()
s.dataMemSize = pkt.MemSize()
s.pkt = pkt
pkt.IncRef()
s.pkt = pkt.IncRef()
s.csumValid = csumValid
if !s.pkt.RXTransportChecksumValidated {
@@ -195,7 +194,6 @@ func (s *segment) DecRef() {
}
}
s.pkt.DecRef()
s.pkt = nil
segmentPool.Put(s)
})
}