diff --git a/nogo.yaml b/nogo.yaml index 7fe8f14e7..280e46e4e 100644 --- a/nogo.yaml +++ b/nogo.yaml @@ -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. diff --git a/pkg/tcpip/link/channel/channel.go b/pkg/tcpip/link/channel/channel.go index 256c69871..c1e48db20 100644 --- a/pkg/tcpip/link/channel/channel.go +++ b/pkg/tcpip/link/channel/channel.go @@ -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{} } } diff --git a/pkg/tcpip/link/qdisc/fifo/packet_buffer_circular_list.go b/pkg/tcpip/link/qdisc/fifo/packet_buffer_circular_list.go index 462ec058b..9c9661efb 100644 --- a/pkg/tcpip/link/qdisc/fifo/packet_buffer_circular_list.go +++ b/pkg/tcpip/link/qdisc/fifo/packet_buffer_circular_list.go @@ -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 diff --git a/pkg/tcpip/link/qdisc/fifo/qdisc_test.go b/pkg/tcpip/link/qdisc/fifo/qdisc_test.go index a533cb0dd..660ddbadc 100644 --- a/pkg/tcpip/link/qdisc/fifo/qdisc_test.go +++ b/pkg/tcpip/link/qdisc/fifo/qdisc_test.go @@ -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{}) diff --git a/pkg/tcpip/network/internal/fragmentation/fragmentation.go b/pkg/tcpip/network/internal/fragmentation/fragmentation.go index db24ac22f..86695d73b 100644 --- a/pkg/tcpip/network/internal/fragmentation/fragmentation.go +++ b/pkg/tcpip/network/internal/fragmentation/fragmentation.go @@ -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 diff --git a/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go b/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go index edbf71757..4a3fd6656 100644 --- a/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go +++ b/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go @@ -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") diff --git a/pkg/tcpip/network/internal/fragmentation/reassembler.go b/pkg/tcpip/network/internal/fragmentation/reassembler.go index d93e6f9d0..17d056f1c 100644 --- a/pkg/tcpip/network/internal/fragmentation/reassembler.go +++ b/pkg/tcpip/network/internal/fragmentation/reassembler.go @@ -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 { diff --git a/pkg/tcpip/network/ipv4/icmp.go b/pkg/tcpip/network/ipv4/icmp.go index afef4ec36..7823f991a 100644 --- a/pkg/tcpip/network/ipv4/icmp.go +++ b/pkg/tcpip/network/ipv4/icmp.go @@ -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) { diff --git a/pkg/tcpip/stack/BUILD b/pkg/tcpip/stack/BUILD index 45569b067..a8a9d7b88 100644 --- a/pkg/tcpip/stack/BUILD +++ b/pkg/tcpip/stack/BUILD @@ -34,7 +34,7 @@ go_template_instance( prefix = "packetBuffer", template = "//pkg/refsvfs2:refs_template", types = { - "T": "PacketBuffer", + "T": "packetBuffer", }, ) diff --git a/pkg/tcpip/stack/packet_buffer.go b/pkg/tcpip/stack/packet_buffer.go index 2d721cfa6..a849df8a5 100644 --- a/pkg/tcpip/stack/packet_buffer.go +++ b/pkg/tcpip/stack/packet_buffer.go @@ -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. diff --git a/pkg/tcpip/stack/packet_buffer_list.go b/pkg/tcpip/stack/packet_buffer_list.go index 31107c3ba..6f09b804d 100644 --- a/pkg/tcpip/stack/packet_buffer_list.go +++ b/pkg/tcpip/stack/packet_buffer_list.go @@ -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] } diff --git a/pkg/tcpip/stack/packet_buffer_unsafe.go b/pkg/tcpip/stack/packet_buffer_unsafe.go index cd151ff0b..6aea4702c 100644 --- a/pkg/tcpip/stack/packet_buffer_unsafe.go +++ b/pkg/tcpip/stack/packet_buffer_unsafe.go @@ -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{})) diff --git a/pkg/tcpip/transport/internal/network/endpoint.go b/pkg/tcpip/transport/internal/network/endpoint.go index aab27a52a..1a5fbc6a6 100644 --- a/pkg/tcpip/transport/internal/network/endpoint.go +++ b/pkg/tcpip/transport/internal/network/endpoint.go @@ -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 diff --git a/pkg/tcpip/transport/tcp/segment.go b/pkg/tcpip/transport/tcp/segment.go index 024a1b4db..abdac9077 100644 --- a/pkg/tcpip/transport/tcp/segment.go +++ b/pkg/tcpip/transport/tcp/segment.go @@ -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) }) }