diff --git a/nogo.yaml b/nogo.yaml index 280e46e4e..7fe8f14e7 100644 --- a/nogo.yaml +++ b/nogo.yaml @@ -218,3 +218,7 @@ 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/sentry/socket/netfilter/owner_matcher.go b/pkg/sentry/socket/netfilter/owner_matcher.go index d83d7e535..3a3df0838 100644 --- a/pkg/sentry/socket/netfilter/owner_matcher.go +++ b/pkg/sentry/socket/netfilter/owner_matcher.go @@ -114,7 +114,7 @@ func (*OwnerMatcher) name() string { } // Match implements Matcher.Match. -func (om *OwnerMatcher) Match(hook stack.Hook, pkt *stack.PacketBuffer, _, _ string) (bool, bool) { +func (om *OwnerMatcher) Match(hook stack.Hook, pkt stack.PacketBufferPtr, _, _ string) (bool, bool) { // Support only for OUTPUT chain. if hook != stack.Output { return false, true diff --git a/pkg/sentry/socket/netfilter/targets.go b/pkg/sentry/socket/netfilter/targets.go index 9791efcc1..c0ad654d5 100644 --- a/pkg/sentry/socket/netfilter/targets.go +++ b/pkg/sentry/socket/netfilter/targets.go @@ -644,7 +644,7 @@ func (jt *JumpTarget) id() targetID { } // Action implements stack.Target.Action. -func (jt *JumpTarget) Action(*stack.PacketBuffer, stack.Hook, *stack.Route, stack.AddressableEndpoint) (stack.RuleVerdict, int) { +func (jt *JumpTarget) Action(stack.PacketBufferPtr, stack.Hook, *stack.Route, stack.AddressableEndpoint) (stack.RuleVerdict, int) { return stack.RuleJump, jt.RuleNum } diff --git a/pkg/sentry/socket/netfilter/tcp_matcher.go b/pkg/sentry/socket/netfilter/tcp_matcher.go index df4ff97c3..7d5b95cc7 100644 --- a/pkg/sentry/socket/netfilter/tcp_matcher.go +++ b/pkg/sentry/socket/netfilter/tcp_matcher.go @@ -95,7 +95,7 @@ func (*TCPMatcher) name() string { } // Match implements Matcher.Match. -func (tm *TCPMatcher) Match(hook stack.Hook, pkt *stack.PacketBuffer, _, _ string) (bool, bool) { +func (tm *TCPMatcher) Match(hook stack.Hook, pkt stack.PacketBufferPtr, _, _ string) (bool, bool) { switch pkt.NetworkProtocolNumber { case header.IPv4ProtocolNumber: netHeader := header.IPv4(pkt.NetworkHeader().Slice()) diff --git a/pkg/sentry/socket/netfilter/udp_matcher.go b/pkg/sentry/socket/netfilter/udp_matcher.go index ae5a34b3a..7d7bc3c95 100644 --- a/pkg/sentry/socket/netfilter/udp_matcher.go +++ b/pkg/sentry/socket/netfilter/udp_matcher.go @@ -92,7 +92,7 @@ func (*UDPMatcher) name() string { } // Match implements Matcher.Match. -func (um *UDPMatcher) Match(hook stack.Hook, pkt *stack.PacketBuffer, _, _ string) (bool, bool) { +func (um *UDPMatcher) Match(hook stack.Hook, pkt stack.PacketBufferPtr, _, _ string) (bool, bool) { switch pkt.NetworkProtocolNumber { case header.IPv4ProtocolNumber: netHeader := header.IPv4(pkt.NetworkHeader().Slice()) diff --git a/pkg/tcpip/header/parse/parse.go b/pkg/tcpip/header/parse/parse.go index dc1300528..0e29377cc 100644 --- a/pkg/tcpip/header/parse/parse.go +++ b/pkg/tcpip/header/parse/parse.go @@ -27,7 +27,7 @@ import ( // pkt.Data. // // Returns true if the header was successfully parsed. -func ARP(pkt *stack.PacketBuffer) bool { +func ARP(pkt stack.PacketBufferPtr) bool { _, ok := pkt.NetworkHeader().Consume(header.ARPSize) if ok { pkt.NetworkProtocolNumber = header.ARPProtocolNumber @@ -39,7 +39,7 @@ func ARP(pkt *stack.PacketBuffer) bool { // header with the IPv4 header. // // Returns true if the header was successfully parsed. -func IPv4(pkt *stack.PacketBuffer) bool { +func IPv4(pkt stack.PacketBufferPtr) bool { hdr, ok := pkt.Data().PullUp(header.IPv4MinimumSize) if !ok { return false @@ -71,7 +71,7 @@ func IPv4(pkt *stack.PacketBuffer) bool { // IPv6 parses an IPv6 packet found in pkt.Data and populates pkt's network // header with the IPv6 header. -func IPv6(pkt *stack.PacketBuffer) (proto tcpip.TransportProtocolNumber, fragID uint32, fragOffset uint16, fragMore bool, ok bool) { +func IPv6(pkt stack.PacketBufferPtr) (proto tcpip.TransportProtocolNumber, fragID uint32, fragOffset uint16, fragMore bool, ok bool) { hdr, ok := pkt.Data().PullUp(header.IPv6MinimumSize) if !ok { return 0, 0, 0, false, false @@ -157,7 +157,7 @@ traverseExtensions: // header with the UDP header. // // Returns true if the header was successfully parsed. -func UDP(pkt *stack.PacketBuffer) bool { +func UDP(pkt stack.PacketBufferPtr) bool { _, ok := pkt.TransportHeader().Consume(header.UDPMinimumSize) pkt.TransportProtocolNumber = header.UDPProtocolNumber return ok @@ -167,7 +167,7 @@ func UDP(pkt *stack.PacketBuffer) bool { // header with the TCP header. // // Returns true if the header was successfully parsed. -func TCP(pkt *stack.PacketBuffer) bool { +func TCP(pkt stack.PacketBufferPtr) bool { // TCP header is variable length, peek at it first. hdrLen := header.TCPMinimumSize hdr, ok := pkt.Data().PullUp(hdrLen) @@ -191,7 +191,7 @@ func TCP(pkt *stack.PacketBuffer) bool { // if present. // // Returns true if an ICMPv4 header was successfully parsed. -func ICMPv4(pkt *stack.PacketBuffer) bool { +func ICMPv4(pkt stack.PacketBufferPtr) bool { if _, ok := pkt.TransportHeader().Consume(header.ICMPv4MinimumSize); ok { pkt.TransportProtocolNumber = header.ICMPv4ProtocolNumber return true @@ -203,7 +203,7 @@ func ICMPv4(pkt *stack.PacketBuffer) bool { // if present. // // Returns true if an ICMPv6 header was successfully parsed. -func ICMPv6(pkt *stack.PacketBuffer) bool { +func ICMPv6(pkt stack.PacketBufferPtr) bool { hdr, ok := pkt.Data().PullUp(header.ICMPv6MinimumSize) if !ok { return false diff --git a/pkg/tcpip/link/channel/channel.go b/pkg/tcpip/link/channel/channel.go index 04ba2e278..256c69871 100644 --- a/pkg/tcpip/link/channel/channel.go +++ b/pkg/tcpip/link/channel/channel.go @@ -43,7 +43,7 @@ type NotificationHandle struct { type queue struct { // c is the outbound packet channel. - c chan *stack.PacketBuffer + c chan stack.PacketBufferPtr mu sync.RWMutex // +checklocks:mu notify []*NotificationHandle @@ -58,7 +58,7 @@ func (q *queue) Close() { q.closed = true } -func (q *queue) Read() *stack.PacketBuffer { +func (q *queue) Read() stack.PacketBufferPtr { select { case p := <-q.c: return p @@ -67,7 +67,7 @@ func (q *queue) Read() *stack.PacketBuffer { } } -func (q *queue) ReadContext(ctx context.Context) *stack.PacketBuffer { +func (q *queue) ReadContext(ctx context.Context) stack.PacketBufferPtr { select { case pkt := <-q.c: return pkt @@ -76,7 +76,7 @@ func (q *queue) ReadContext(ctx context.Context) *stack.PacketBuffer { } } -func (q *queue) Write(pkt *stack.PacketBuffer) tcpip.Error { +func (q *queue) Write(pkt stack.PacketBufferPtr) tcpip.Error { // q holds the PacketBuffer. q.mu.RLock() if q.closed { @@ -149,7 +149,7 @@ type Endpoint struct { func New(size int, mtu uint32, linkAddr tcpip.LinkAddress) *Endpoint { return &Endpoint{ q: &queue{ - c: make(chan *stack.PacketBuffer, size), + c: make(chan stack.PacketBufferPtr, size), }, mtu: mtu, linkAddr: linkAddr, @@ -164,20 +164,20 @@ func (e *Endpoint) Close() { } // Read does non-blocking read one packet from the outbound packet queue. -func (e *Endpoint) Read() *stack.PacketBuffer { +func (e *Endpoint) Read() stack.PacketBufferPtr { return e.q.Read() } // ReadContext does blocking read for one packet from the outbound packet queue. // It can be cancelled by ctx, and in this case, it returns nil. -func (e *Endpoint) ReadContext(ctx context.Context) *stack.PacketBuffer { +func (e *Endpoint) ReadContext(ctx context.Context) stack.PacketBufferPtr { return e.q.ReadContext(ctx) } // Drain removes all outbound packets from the channel and counts them. func (e *Endpoint) Drain() int { c := 0 - for pkt := e.Read(); pkt != nil; pkt = e.Read() { + for pkt := e.Read(); !pkt.IsNil(); pkt = e.Read() { pkt.DecRef() c++ } @@ -190,7 +190,7 @@ func (e *Endpoint) NumQueued() int { } // InjectInbound injects an inbound packet. -func (e *Endpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *Endpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { e.dispatcher.DeliverNetworkPacket(protocol, pkt) } @@ -274,4 +274,4 @@ func (*Endpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (*Endpoint) AddHeader(*stack.PacketBuffer) {} +func (*Endpoint) AddHeader(stack.PacketBufferPtr) {} diff --git a/pkg/tcpip/link/ethernet/ethernet.go b/pkg/tcpip/link/ethernet/ethernet.go index a5c7cdae9..3d93cde49 100644 --- a/pkg/tcpip/link/ethernet/ethernet.go +++ b/pkg/tcpip/link/ethernet/ethernet.go @@ -59,7 +59,7 @@ func (e *Endpoint) MTU() uint32 { } // DeliverNetworkPacket implements stack.NetworkDispatcher. -func (e *Endpoint) DeliverNetworkPacket(_ tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *Endpoint) DeliverNetworkPacket(_ tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { hdr, ok := pkt.LinkHeader().Consume(header.EthernetMinimumSize) if !ok { return @@ -93,7 +93,7 @@ func (e *Endpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint. -func (*Endpoint) AddHeader(pkt *stack.PacketBuffer) { +func (*Endpoint) AddHeader(pkt stack.PacketBufferPtr) { eth := header.Ethernet(pkt.LinkHeader().Push(header.EthernetMinimumSize)) fields := header.EthernetFields{ SrcAddr: pkt.EgressRoute.LocalLinkAddress, diff --git a/pkg/tcpip/link/ethernet/ethernet_test.go b/pkg/tcpip/link/ethernet/ethernet_test.go index 842c823eb..ccc0eea10 100644 --- a/pkg/tcpip/link/ethernet/ethernet_test.go +++ b/pkg/tcpip/link/ethernet/ethernet_test.go @@ -35,11 +35,11 @@ type testNetworkDispatcher struct { networkPackets int } -func (t *testNetworkDispatcher) DeliverNetworkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer) { +func (t *testNetworkDispatcher) DeliverNetworkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr) { t.networkPackets++ } -func (*testNetworkDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer, bool) { +func (*testNetworkDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr, bool) { panic("not implemented") } @@ -146,7 +146,7 @@ func TestWritePacketToRemoteAddHeader(t *testing.T) { { pkt := c.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to read a packet") } diff --git a/pkg/tcpip/link/fdbased/endpoint.go b/pkg/tcpip/link/fdbased/endpoint.go index 05c96ada2..c8c162efb 100644 --- a/pkg/tcpip/link/fdbased/endpoint.go +++ b/pkg/tcpip/link/fdbased/endpoint.go @@ -509,7 +509,7 @@ const ( ) // AddHeader implements stack.LinkEndpoint.AddHeader. -func (e *endpoint) AddHeader(pkt *stack.PacketBuffer) { +func (e *endpoint) AddHeader(pkt stack.PacketBufferPtr) { if e.hdrSize > 0 { // Add ethernet header if needed. eth := header.Ethernet(pkt.LinkHeader().Push(header.EthernetMinimumSize)) @@ -523,7 +523,7 @@ func (e *endpoint) AddHeader(pkt *stack.PacketBuffer) { // writePacket writes outbound packets to the file descriptor. If it is not // currently writable, the packet is dropped. -func (e *endpoint) writePacket(pkt *stack.PacketBuffer) tcpip.Error { +func (e *endpoint) writePacket(pkt stack.PacketBufferPtr) tcpip.Error { fdInfo := e.fds[pkt.Hash%uint32(len(e.fds))] fd := fdInfo.fd var vnetHdrBuf []byte @@ -573,7 +573,7 @@ func (e *endpoint) writePacket(pkt *stack.PacketBuffer) tcpip.Error { return rawfile.NonBlockingWriteIovec(fd, iovecs) } -func (e *endpoint) sendBatch(batchFDInfo fdInfo, pkts []*stack.PacketBuffer) (int, tcpip.Error) { +func (e *endpoint) sendBatch(batchFDInfo fdInfo, pkts []stack.PacketBufferPtr) (int, tcpip.Error) { // Degrade to writePacket if underlying fd is not a socket. if !batchFDInfo.isSocket { var written int @@ -690,7 +690,7 @@ func (e *endpoint) sendBatch(batchFDInfo fdInfo, pkts []*stack.PacketBuffer) (in // - pkt.NetworkProtocolNumber func (e *endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { // Preallocate to avoid repeated reallocation as we append to batch. - batch := make([]*stack.PacketBuffer, 0, BatchSize) + batch := make([]stack.PacketBufferPtr, 0, BatchSize) batchFDInfo := fdInfo{fd: -1, isSocket: false} sentPackets := 0 for _, pkt := range pkts.AsSlice() { @@ -775,7 +775,7 @@ func (e *InjectableEndpoint) Attach(dispatcher stack.NetworkDispatcher) { } // InjectInbound injects an inbound packet. -func (e *InjectableEndpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *InjectableEndpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { e.dispatcher.DeliverNetworkPacket(protocol, pkt) } diff --git a/pkg/tcpip/link/fdbased/endpoint_test.go b/pkg/tcpip/link/fdbased/endpoint_test.go index 751fee319..711c230c6 100644 --- a/pkg/tcpip/link/fdbased/endpoint_test.go +++ b/pkg/tcpip/link/fdbased/endpoint_test.go @@ -48,7 +48,7 @@ const ( type packetInfo struct { Proto tcpip.NetworkProtocolNumber - Contents *stack.PacketBuffer + Contents stack.PacketBufferPtr } type packetContents struct { @@ -62,8 +62,8 @@ func checkPacketInfoEqual(t *testing.T, got, want packetInfo) { t.Helper() if diff := cmp.Diff( want, got, - cmp.Transformer("ExtractPacketBuffer", func(pk *stack.PacketBuffer) *packetContents { - if pk == nil { + cmp.Transformer("ExtractPacketBuffer", func(pk stack.PacketBufferPtr) *packetContents { + if pk.IsNil() { return nil } return &packetContents{ @@ -133,12 +133,12 @@ func (c *context) cleanup() { } } -func (c *context) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (c *context) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { pkt.IncRef() c.ch <- packetInfo{protocol, pkt} } -func (c *context) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer, bool) { +func (c *context) DeliverLinkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr, bool) { c.t.Fatal("DeliverLinkPacket not implemented") } @@ -568,15 +568,15 @@ func TestIovecBufferSkipVnetHdr(t *testing.T) { // fakeNetworkDispatcher delivers packets to pkts. type fakeNetworkDispatcher struct { - pkts []*stack.PacketBuffer + pkts []stack.PacketBufferPtr } -func (d *fakeNetworkDispatcher) DeliverNetworkPacket(_ tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (d *fakeNetworkDispatcher) DeliverNetworkPacket(_ tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { pkt.IncRef() d.pkts = append(d.pkts, pkt) } -func (*fakeNetworkDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer, bool) { +func (*fakeNetworkDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr, bool) { panic("not implemented") } diff --git a/pkg/tcpip/link/loopback/loopback.go b/pkg/tcpip/link/loopback/loopback.go index 03a929e10..35cdbdff4 100644 --- a/pkg/tcpip/link/loopback/loopback.go +++ b/pkg/tcpip/link/loopback/loopback.go @@ -93,4 +93,4 @@ func (*endpoint) ARPHardwareType() header.ARPHardwareType { return header.ARPHardwareLoopback } -func (*endpoint) AddHeader(*stack.PacketBuffer) {} +func (*endpoint) AddHeader(stack.PacketBufferPtr) {} diff --git a/pkg/tcpip/link/muxed/injectable.go b/pkg/tcpip/link/muxed/injectable.go index 60f8b1b2d..0a6981ebb 100644 --- a/pkg/tcpip/link/muxed/injectable.go +++ b/pkg/tcpip/link/muxed/injectable.go @@ -81,7 +81,7 @@ func (m *InjectableEndpoint) IsAttached() bool { } // InjectInbound implements stack.InjectableLinkEndpoint. -func (m *InjectableEndpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (m *InjectableEndpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { m.dispatcher.DeliverNetworkPacket(protocol, pkt) } @@ -133,7 +133,7 @@ func (*InjectableEndpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (*InjectableEndpoint) AddHeader(*stack.PacketBuffer) {} +func (*InjectableEndpoint) AddHeader(stack.PacketBufferPtr) {} // NewInjectableEndpoint creates a new multi-endpoint injectable endpoint. func NewInjectableEndpoint(routes map[tcpip.Address]stack.InjectableLinkEndpoint) *InjectableEndpoint { diff --git a/pkg/tcpip/link/nested/nested.go b/pkg/tcpip/link/nested/nested.go index 79b407d8a..cbf1ffc53 100644 --- a/pkg/tcpip/link/nested/nested.go +++ b/pkg/tcpip/link/nested/nested.go @@ -51,7 +51,7 @@ func (e *Endpoint) Init(child stack.LinkEndpoint, embedder stack.NetworkDispatch } // DeliverNetworkPacket implements stack.NetworkDispatcher. -func (e *Endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *Endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { e.mu.RLock() d := e.dispatcher e.mu.RUnlock() @@ -61,7 +61,7 @@ func (e *Endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pk } // DeliverLinkPacket implements stack.NetworkDispatcher. -func (e *Endpoint) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer, incoming bool) { +func (e *Endpoint) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr, incoming bool) { e.mu.RLock() d := e.dispatcher e.mu.RUnlock() @@ -144,6 +144,6 @@ func (e *Endpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (e *Endpoint) AddHeader(pkt *stack.PacketBuffer) { +func (e *Endpoint) AddHeader(pkt stack.PacketBufferPtr) { e.child.AddHeader(pkt) } diff --git a/pkg/tcpip/link/nested/nested_test.go b/pkg/tcpip/link/nested/nested_test.go index d1b7d38f6..aead15713 100644 --- a/pkg/tcpip/link/nested/nested_test.go +++ b/pkg/tcpip/link/nested/nested_test.go @@ -54,11 +54,11 @@ type counterDispatcher struct { var _ stack.NetworkDispatcher = (*counterDispatcher)(nil) -func (d *counterDispatcher) DeliverNetworkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer) { +func (d *counterDispatcher) DeliverNetworkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr) { d.count++ } -func (*counterDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer, bool) { +func (*counterDispatcher) DeliverLinkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr, bool) { panic("not implemented") } diff --git a/pkg/tcpip/link/packetsocket/packetsocket.go b/pkg/tcpip/link/packetsocket/packetsocket.go index 609a14610..12186b877 100644 --- a/pkg/tcpip/link/packetsocket/packetsocket.go +++ b/pkg/tcpip/link/packetsocket/packetsocket.go @@ -40,7 +40,7 @@ func New(lower stack.LinkEndpoint) stack.LinkEndpoint { } // DeliverNetworkPacket implements stack.NetworkDispatcher. -func (e *endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { e.Endpoint.DeliverLinkPacket(protocol, pkt, true /* incoming */) e.Endpoint.DeliverNetworkPacket(protocol, pkt) diff --git a/pkg/tcpip/link/packetsocket/packetsocket_test.go b/pkg/tcpip/link/packetsocket/packetsocket_test.go index ebe05ac12..fb1d9371f 100644 --- a/pkg/tcpip/link/packetsocket/packetsocket_test.go +++ b/pkg/tcpip/link/packetsocket/packetsocket_test.go @@ -53,18 +53,18 @@ func (e *nullEndpoint) Attach(d stack.NetworkDispatcher) { e.disp = d } func (e *nullEndpoint) IsAttached() bool { return e.disp != nil } func (*nullEndpoint) Wait() {} func (*nullEndpoint) ARPHardwareType() header.ARPHardwareType { return header.ARPHardwareNone } -func (*nullEndpoint) AddHeader(*stack.PacketBuffer) {} +func (*nullEndpoint) AddHeader(stack.PacketBufferPtr) {} var _ stack.NetworkDispatcher = (*testNetworkDispatcher)(nil) type linkPacketInfo struct { - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr protocol tcpip.NetworkProtocolNumber incoming bool } type networkPacketInfo struct { - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr protocol tcpip.NetworkProtocolNumber } @@ -77,17 +77,17 @@ type testNetworkDispatcher struct { } func (t *testNetworkDispatcher) reset() { - if pkt := t.linkPacket.pkt; pkt != nil { + if pkt := t.linkPacket.pkt; !pkt.IsNil() { pkt.DecRef() } - if pkt := t.networkPacket.pkt; pkt != nil { + if pkt := t.networkPacket.pkt; !pkt.IsNil() { pkt.DecRef() } *t = testNetworkDispatcher{} } -func (t *testNetworkDispatcher) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (t *testNetworkDispatcher) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { networkPacket := networkPacketInfo{ pkt: pkt.IncRef(), protocol: protocol, @@ -100,7 +100,7 @@ func (t *testNetworkDispatcher) DeliverNetworkPacket(protocol tcpip.NetworkProto t.networkPacket = networkPacket } -func (t *testNetworkDispatcher) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer, incoming bool) { +func (t *testNetworkDispatcher) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr, incoming bool) { linkPacket := linkPacketInfo{ pkt: pkt.IncRef(), protocol: protocol, diff --git a/pkg/tcpip/link/pipe/pipe.go b/pkg/tcpip/link/pipe/pipe.go index 4504735f2..aad1e2d5d 100644 --- a/pkg/tcpip/link/pipe/pipe.go +++ b/pkg/tcpip/link/pipe/pipe.go @@ -110,4 +110,4 @@ func (*Endpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint. -func (*Endpoint) AddHeader(*stack.PacketBuffer) {} +func (*Endpoint) AddHeader(stack.PacketBufferPtr) {} diff --git a/pkg/tcpip/link/qdisc/fifo/fifo.go b/pkg/tcpip/link/qdisc/fifo/fifo.go index 227f5510a..e732c0dc6 100644 --- a/pkg/tcpip/link/qdisc/fifo/fifo.go +++ b/pkg/tcpip/link/qdisc/fifo/fifo.go @@ -97,7 +97,7 @@ func (qd *queueDispatcher) dispatchLoop() { case &qd.newPacketWaker: case &qd.closeWaker: qd.mu.Lock() - for p := qd.queue.removeFront(); p != nil; p = qd.queue.removeFront() { + for p := qd.queue.removeFront(); !p.IsNil(); p = qd.queue.removeFront() { p.DecRef() } qd.queue.decRef() @@ -107,7 +107,7 @@ func (qd *queueDispatcher) dispatchLoop() { panic("unknown waker") } qd.mu.Lock() - for pkt := qd.queue.removeFront(); pkt != nil; pkt = qd.queue.removeFront() { + for pkt := qd.queue.removeFront(); !pkt.IsNil(); pkt = qd.queue.removeFront() { batch.PushBack(pkt) if batch.Len() < BatchSize && !qd.queue.isEmpty() { continue @@ -127,7 +127,7 @@ func (qd *queueDispatcher) dispatchLoop() { // - pkt.EgressRoute // - pkt.GSOOptions // - pkt.NetworkProtocolNumber -func (d *discipline) WritePacket(pkt *stack.PacketBuffer) tcpip.Error { +func (d *discipline) WritePacket(pkt stack.PacketBufferPtr) tcpip.Error { if d.closed.Load() == qDiscClosed { return &tcpip.ErrClosedForSend{} } 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 5b3030b05..462ec058b 100644 --- a/pkg/tcpip/link/qdisc/fifo/packet_buffer_circular_list.go +++ b/pkg/tcpip/link/qdisc/fifo/packet_buffer_circular_list.go @@ -23,14 +23,14 @@ import "gvisor.dev/gvisor/pkg/tcpip/stack" // // +stateify savable type packetBufferCircularList struct { - pbs []*stack.PacketBuffer + pbs []stack.PacketBufferPtr head int size int } // init initializes the list with the given size. func (pl *packetBufferCircularList) init(size int) { - pl.pbs = make([]*stack.PacketBuffer, size) + pl.pbs = make([]stack.PacketBufferPtr, size) } // length returns the number of elements in the list. @@ -60,7 +60,7 @@ func (pl *packetBufferCircularList) isEmpty() bool { // Failing to do so may clobber existing entries. // //go:nosplit -func (pl *packetBufferCircularList) pushBack(pb *stack.PacketBuffer) { +func (pl *packetBufferCircularList) pushBack(pb stack.PacketBufferPtr) { next := (pl.head + pl.size) % len(pl.pbs) pl.pbs[next] = pb pl.size++ @@ -69,7 +69,7 @@ func (pl *packetBufferCircularList) pushBack(pb *stack.PacketBuffer) { // removeFront returns the first element of the list or nil. // //go:nosplit -func (pl *packetBufferCircularList) removeFront() *stack.PacketBuffer { +func (pl *packetBufferCircularList) removeFront() stack.PacketBufferPtr { if pl.isEmpty() { return nil } diff --git a/pkg/tcpip/link/sharedmem/server_tx.go b/pkg/tcpip/link/sharedmem/server_tx.go index ccea4a8ee..abb8fa1cc 100644 --- a/pkg/tcpip/link/sharedmem/server_tx.go +++ b/pkg/tcpip/link/sharedmem/server_tx.go @@ -159,7 +159,7 @@ func (s *serverTx) fillPacket(pktBuffer bufferv2.Buffer, buffers []queue.RxBuffe return bufs, totalCopied } -func (s *serverTx) transmit(pkt *stack.PacketBuffer) bool { +func (s *serverTx) transmit(pkt stack.PacketBufferPtr) bool { buffers := make([]queue.RxBuffer, 8) buffers, totalCopied := s.fillPacket(pkt.ToBuffer(), buffers) if totalCopied == 0 { diff --git a/pkg/tcpip/link/sharedmem/sharedmem.go b/pkg/tcpip/link/sharedmem/sharedmem.go index 1c62251c0..5c0bfe9bd 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem.go +++ b/pkg/tcpip/link/sharedmem/sharedmem.go @@ -318,7 +318,7 @@ func (e *endpoint) LinkAddress() tcpip.LinkAddress { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (e *endpoint) AddHeader(pkt *stack.PacketBuffer) { +func (e *endpoint) AddHeader(pkt stack.PacketBufferPtr) { // Add ethernet header if needed. if len(e.addr) == 0 { return @@ -332,13 +332,13 @@ func (e *endpoint) AddHeader(pkt *stack.PacketBuffer) { }) } -func (e *endpoint) AddVirtioNetHeader(pkt *stack.PacketBuffer) { +func (e *endpoint) AddVirtioNetHeader(pkt stack.PacketBufferPtr) { virtio := header.VirtioNetHeader(pkt.VirtioNetHeader().Push(header.VirtioNetHeaderSize)) virtio.Encode(&header.VirtioNetHeaderFields{}) } // +checklocks:e.mu -func (e *endpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { +func (e *endpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) tcpip.Error { if e.virtioNetHeaderRequired { e.AddVirtioNetHeader(pkt) } diff --git a/pkg/tcpip/link/sharedmem/sharedmem_server.go b/pkg/tcpip/link/sharedmem/sharedmem_server.go index 4558ac062..45421e237 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_server.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_server.go @@ -203,7 +203,7 @@ func (e *serverEndpoint) LinkAddress() tcpip.LinkAddress { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (e *serverEndpoint) AddHeader(pkt *stack.PacketBuffer) { +func (e *serverEndpoint) AddHeader(pkt stack.PacketBufferPtr) { // Add ethernet header if needed. if len(e.addr) == 0 { return @@ -217,13 +217,13 @@ func (e *serverEndpoint) AddHeader(pkt *stack.PacketBuffer) { }) } -func (e *serverEndpoint) AddVirtioNetHeader(pkt *stack.PacketBuffer) { +func (e *serverEndpoint) AddVirtioNetHeader(pkt stack.PacketBufferPtr) { virtio := header.VirtioNetHeader(pkt.VirtioNetHeader().Push(header.VirtioNetHeaderSize)) virtio.Encode(&header.VirtioNetHeaderFields{}) } // +checklocks:e.mu -func (e *serverEndpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { +func (e *serverEndpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) tcpip.Error { if e.virtioNetHeaderRequired { e.AddVirtioNetHeader(pkt) } @@ -239,7 +239,7 @@ func (e *serverEndpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.Net // WritePacket writes outbound packets to the file descriptor. If it is not // currently writable, the packet is dropped. // WritePacket implements stack.LinkEndpoint.WritePacket. -func (e *serverEndpoint) WritePacket(_ stack.RouteInfo, _ tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { +func (e *serverEndpoint) WritePacket(_ stack.RouteInfo, _ tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) tcpip.Error { // Transmit the packet. e.mu.Lock() defer e.mu.Unlock() diff --git a/pkg/tcpip/link/sharedmem/sharedmem_test.go b/pkg/tcpip/link/sharedmem/sharedmem_test.go index 055990377..158663443 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_test.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_test.go @@ -144,7 +144,7 @@ func newTestContext(t *testing.T, mtu, bufferSize uint32, addr tcpip.LinkAddress return c } -func (c *testContext) DeliverNetworkPacket(proto tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (c *testContext) DeliverNetworkPacket(proto tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { c.mu.Lock() c.packets = append(c.packets, packetInfo{ proto: proto, @@ -155,7 +155,7 @@ func (c *testContext) DeliverNetworkPacket(proto tcpip.NetworkProtocolNumber, pk c.packetCh <- struct{}{} } -func (c *testContext) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer, bool) { +func (c *testContext) DeliverLinkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr, bool) { c.t.Fatal("DeliverLinkPacket not implemented") } diff --git a/pkg/tcpip/link/sniffer/pcap.go b/pkg/tcpip/link/sniffer/pcap.go index 491957ac8..648852b24 100644 --- a/pkg/tcpip/link/sniffer/pcap.go +++ b/pkg/tcpip/link/sniffer/pcap.go @@ -50,7 +50,7 @@ var _ encoding.BinaryMarshaler = (*pcapPacket)(nil) type pcapPacket struct { timestamp time.Time - packet *stack.PacketBuffer + packet stack.PacketBufferPtr maxCaptureLen int } diff --git a/pkg/tcpip/link/sniffer/sniffer.go b/pkg/tcpip/link/sniffer/sniffer.go index 566ef5032..470a57bc2 100644 --- a/pkg/tcpip/link/sniffer/sniffer.go +++ b/pkg/tcpip/link/sniffer/sniffer.go @@ -133,12 +133,12 @@ func NewWithWriter(lower stack.LinkEndpoint, writer io.Writer, snapLen uint32) ( // DeliverNetworkPacket implements the stack.NetworkDispatcher interface. It is // called by the link-layer endpoint being wrapped when a packet arrives, and // logs the packet before forwarding to the actual dispatcher. -func (e *endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { e.dumpPacket(DirectionRecv, protocol, pkt) e.Endpoint.DeliverNetworkPacket(protocol, pkt) } -func (e *endpoint) dumpPacket(dir Direction, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *endpoint) dumpPacket(dir Direction, protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { writer := e.writer if writer == nil && LogPackets.Load() == 1 { LogPacket(e.logPrefix, dir, protocol, pkt) @@ -170,7 +170,7 @@ func (e *endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) } // LogPacket logs a packet to stdout. -func LogPacket(prefix string, dir Direction, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func LogPacket(prefix string, dir Direction, protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { // Figure out the network layer info. var transProto uint8 src := tcpip.Address("unknown") @@ -380,7 +380,7 @@ func LogPacket(prefix string, dir Direction, protocol tcpip.NetworkProtocolNumbe // trimmedClone clones the packet buffer to not modify the original. It trims // anything before the network header. -func trimmedClone(pkt *stack.PacketBuffer) *stack.PacketBuffer { +func trimmedClone(pkt stack.PacketBufferPtr) stack.PacketBufferPtr { // We don't clone the original packet buffer so that the new packet buffer // does not have any of its headers set. // diff --git a/pkg/tcpip/link/tun/device.go b/pkg/tcpip/link/tun/device.go index b4720a7fd..dda1cd161 100644 --- a/pkg/tcpip/link/tun/device.go +++ b/pkg/tcpip/link/tun/device.go @@ -254,7 +254,7 @@ func (d *Device) Read() (*bufferv2.View, error) { } pkt := endpoint.Read() - if pkt == nil { + if pkt.IsNil() { return nil, linuxerr.ErrWouldBlock } v := d.encodePkt(pkt) @@ -263,7 +263,7 @@ func (d *Device) Read() (*bufferv2.View, error) { } // encodePkt encodes packet for fd side. -func (d *Device) encodePkt(pkt *stack.PacketBuffer) *bufferv2.View { +func (d *Device) encodePkt(pkt stack.PacketBufferPtr) *bufferv2.View { var view *bufferv2.View // Packet information. @@ -351,7 +351,7 @@ func (e *tunEndpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (e *tunEndpoint) AddHeader(pkt *stack.PacketBuffer) { +func (e *tunEndpoint) AddHeader(pkt stack.PacketBufferPtr) { if !e.isTap { return } diff --git a/pkg/tcpip/link/waitable/waitable.go b/pkg/tcpip/link/waitable/waitable.go index ef5623a2e..12e5daa31 100644 --- a/pkg/tcpip/link/waitable/waitable.go +++ b/pkg/tcpip/link/waitable/waitable.go @@ -53,7 +53,7 @@ func New(lower stack.LinkEndpoint) *Endpoint { // It is called by the link-layer endpoint being wrapped when a packet arrives, // and only forwards to the actual dispatcher if Wait or WaitDispatch haven't // been called. -func (e *Endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *Endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { if !e.dispatchGate.Enter() { return } @@ -63,7 +63,7 @@ func (e *Endpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pk } // DeliverLinkPacket implements stack.NetworkDispatcher. -func (e *Endpoint) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer, incoming bool) { +func (e *Endpoint) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr, incoming bool) { if !e.dispatchGate.Enter() { return } @@ -143,6 +143,6 @@ func (e *Endpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (e *Endpoint) AddHeader(pkt *stack.PacketBuffer) { +func (e *Endpoint) AddHeader(pkt stack.PacketBufferPtr) { e.lower.AddHeader(pkt) } diff --git a/pkg/tcpip/link/waitable/waitable_test.go b/pkg/tcpip/link/waitable/waitable_test.go index f0192bb29..a34bc6e82 100644 --- a/pkg/tcpip/link/waitable/waitable_test.go +++ b/pkg/tcpip/link/waitable/waitable_test.go @@ -40,11 +40,11 @@ type countedEndpoint struct { dispatcher stack.NetworkDispatcher } -func (e *countedEndpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *countedEndpoint) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { e.dispatchCount++ } -func (*countedEndpoint) DeliverLinkPacket(tcpip.NetworkProtocolNumber, *stack.PacketBuffer, bool) { +func (*countedEndpoint) DeliverLinkPacket(tcpip.NetworkProtocolNumber, stack.PacketBufferPtr, bool) { panic("not implemented") } @@ -89,7 +89,7 @@ func (*countedEndpoint) ARPHardwareType() header.ARPHardwareType { func (*countedEndpoint) Wait() {} // AddHeader implements stack.LinkEndpoint.AddHeader. -func (*countedEndpoint) AddHeader(*stack.PacketBuffer) { +func (*countedEndpoint) AddHeader(stack.PacketBufferPtr) { panic("unimplemented") } diff --git a/pkg/tcpip/link/xdp/endpoint.go b/pkg/tcpip/link/xdp/endpoint.go index b532a052f..fad8c6b34 100644 --- a/pkg/tcpip/link/xdp/endpoint.go +++ b/pkg/tcpip/link/xdp/endpoint.go @@ -193,7 +193,7 @@ func (ep *endpoint) Wait() { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (ep *endpoint) AddHeader(pkt *stack.PacketBuffer) { +func (ep *endpoint) AddHeader(pkt stack.PacketBufferPtr) { // Add ethernet header if needed. eth := header.Ethernet(pkt.LinkHeader().Push(header.EthernetMinimumSize)) eth.Encode(&header.EthernetFields{ diff --git a/pkg/tcpip/network/arp/arp.go b/pkg/tcpip/network/arp/arp.go index 674ef5184..1f1489993 100644 --- a/pkg/tcpip/network/arp/arp.go +++ b/pkg/tcpip/network/arp/arp.go @@ -133,7 +133,7 @@ func (e *endpoint) MaxHeaderLength() uint16 { func (*endpoint) Close() {} -func (*endpoint) WritePacket(*stack.Route, stack.NetworkHeaderParams, *stack.PacketBuffer) tcpip.Error { +func (*endpoint) WritePacket(*stack.Route, stack.NetworkHeaderParams, stack.PacketBufferPtr) tcpip.Error { return &tcpip.ErrNotSupported{} } @@ -142,11 +142,11 @@ func (*endpoint) NetworkProtocolNumber() tcpip.NetworkProtocolNumber { return ProtocolNumber } -func (*endpoint) WriteHeaderIncludedPacket(*stack.Route, *stack.PacketBuffer) tcpip.Error { +func (*endpoint) WriteHeaderIncludedPacket(*stack.Route, stack.PacketBufferPtr) tcpip.Error { return &tcpip.ErrNotSupported{} } -func (e *endpoint) HandlePacket(pkt *stack.PacketBuffer) { +func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { stats := e.stats.arp stats.packetsReceived.Increment() @@ -384,7 +384,7 @@ func (*protocol) Close() {} func (*protocol) Wait() {} // Parse implements stack.NetworkProtocol.Parse. -func (*protocol) Parse(pkt *stack.PacketBuffer) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) { +func (*protocol) Parse(pkt stack.PacketBufferPtr) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) { return 0, false, parse.ARP(pkt) } diff --git a/pkg/tcpip/network/arp/arp_test.go b/pkg/tcpip/network/arp/arp_test.go index 6aa7f75ef..6064daa7a 100644 --- a/pkg/tcpip/network/arp/arp_test.go +++ b/pkg/tcpip/network/arp/arp_test.go @@ -320,7 +320,7 @@ func TestDirectRequest(t *testing.T) { // No packets should be sent after receiving an invalid ARP request. // There is no need to perform a blocking read here, since packets are // sent in the same function that handles ARP requests. - if pkt := c.linkEP.Read(); pkt != nil { + if pkt := c.linkEP.Read(); !pkt.IsNil() { t.Errorf("unexpected packet sent: %+v", pkt) } if got, want := c.s.Stats().ARP.RequestsReceivedUnknownTargetAddress.Value(), requestsRecvUnknownAddr+1; got != want { @@ -339,7 +339,7 @@ func TestDirectRequest(t *testing.T) { // Verify an ARP response was sent. pi := c.linkEP.Read() - if pi == nil { + if pi.IsNil() { t.Fatal("expected ARP response to be sent, got none") } @@ -621,7 +621,7 @@ func TestLinkAddressRequest(t *testing.T) { } pkt := linkEP.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to send a link address request") } @@ -680,7 +680,7 @@ func TestDADARPRequestPacket(t *testing.T) { clock.RunImmediatelyScheduledJobs() pkt := e.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to send an ARP request") } diff --git a/pkg/tcpip/network/internal/fragmentation/fragmentation.go b/pkg/tcpip/network/internal/fragmentation/fragmentation.go index fecda178b..db24ac22f 100644 --- a/pkg/tcpip/network/internal/fragmentation/fragmentation.go +++ b/pkg/tcpip/network/internal/fragmentation/fragmentation.go @@ -97,7 +97,7 @@ type TimeoutHandler interface { // OnReassemblyTimeout will be called with the first fragment (or nil, if the // first fragment has not been received) of a packet whose reassembly has // timed out. - OnReassemblyTimeout(pkt *stack.PacketBuffer) + OnReassemblyTimeout(pkt stack.PacketBufferPtr) } // NewFragmentation creates a new Fragmentation. @@ -155,8 +155,8 @@ func NewFragmentation(blockSize uint16, highMemoryLimit, lowMemoryLimit int, rea // to be given here outside of the FragmentID struct because IPv6 should not use // the protocol to identify a fragment. func (f *Fragmentation) Process( - id FragmentID, first, last uint16, more bool, proto uint8, pkt *stack.PacketBuffer) ( - *stack.PacketBuffer, uint8, bool, error) { + 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) } @@ -251,12 +251,12 @@ func (f *Fragmentation) release(r *reassembler, timedOut bool) { if h := f.timeoutHandler; timedOut && h != nil { h.OnReassemblyTimeout(r.pkt) } - if r.pkt != nil { + if !r.pkt.IsNil() { r.pkt.DecRef() r.pkt = nil } for _, h := range r.holes { - if h.pkt != nil { + if !h.pkt.IsNil() { h.pkt.DecRef() h.pkt = nil } @@ -308,7 +308,7 @@ type PacketFragmenter struct { // // reserve is the number of bytes that should be reserved for the headers in // each generated fragment. -func MakePacketFragmenter(pkt *stack.PacketBuffer, fragmentPayloadLen uint32, reserve int) PacketFragmenter { +func MakePacketFragmenter(pkt stack.PacketBufferPtr, fragmentPayloadLen uint32, reserve int) PacketFragmenter { // As per RFC 8200 Section 4.5, some IPv6 extension headers should not be // repeated in each fragment. However we do not currently support any header // of that kind yet, so the following computation is valid for both IPv4 and @@ -339,7 +339,7 @@ func MakePacketFragmenter(pkt *stack.PacketBuffer, fragmentPayloadLen uint32, re // Note that the returned packet will not have its network and link headers // populated, but space for them will be reserved. The transport header will be // stored in the packet's data. -func (pf *PacketFragmenter) BuildNextFragment() (*stack.PacketBuffer, int, int, bool) { +func (pf *PacketFragmenter) BuildNextFragment() (stack.PacketBufferPtr, int, int, bool) { if pf.currentFragment >= pf.fragmentCount { panic("BuildNextFragment should not be called again after the last fragment was returned") } diff --git a/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go b/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go index 74f89fd7a..edbf71757 100644 --- a/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go +++ b/pkg/tcpip/network/internal/fragmentation/fragmentation_test.go @@ -43,7 +43,7 @@ func buf(size int, pieces ...string) bufferv2.Buffer { return buf } -func pkt(size int, pieces ...string) *stack.PacketBuffer { +func pkt(size int, pieces ...string) stack.PacketBufferPtr { return stack.NewPacketBuffer(stack.PacketBufferOptions{ Payload: buf(size, pieces...), }) @@ -55,7 +55,7 @@ type processInput struct { last uint16 more bool proto uint8 - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr } type processOutput struct { @@ -116,7 +116,7 @@ func TestFragmentationProcess(t *testing.T) { 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) - if resPkt != nil { + if !resPkt.IsNil() { defer resPkt.DecRef() } if err != nil { @@ -266,7 +266,7 @@ func TestReassemblingTimeout(t *testing.T) { p := pkt(len(frag.data), frag.data) defer p.DecRef() pkt, _, done, err := f.Process(FragmentID{}, frag.first, frag.last, frag.more, protocol, p) - if pkt != nil { + if !pkt.IsNil() { pkt.DecRef() } if err != nil { @@ -449,7 +449,7 @@ func TestErrors(t *testing.T) { f := NewFragmentation(test.blockSize, HighFragThreshold, LowFragThreshold, reassembleTimeout, c, nil) resPkt, _, done, err := f.Process(FragmentID{}, test.first, test.last, test.more, 0, p0) - if resPkt != nil { + if !resPkt.IsNil() { resPkt.DecRef() } if !errors.Is(err, test.err) { @@ -577,10 +577,10 @@ func TestPacketFragmenter(t *testing.T) { } type testTimeoutHandler struct { - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr } -func (h *testTimeoutHandler) OnReassemblyTimeout(pkt *stack.PacketBuffer) { +func (h *testTimeoutHandler) OnReassemblyTimeout(pkt stack.PacketBufferPtr) { h.pkt = pkt } @@ -598,14 +598,14 @@ func TestTimeoutHandler(t *testing.T) { first uint16 last uint16 more bool - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr } tests := []struct { name string params []processParam wantError bool - wantPkt *stack.PacketBuffer + wantPkt stack.PacketBufferPtr }{ { name: "onTimeout runs", @@ -688,11 +688,11 @@ func TestTimeoutHandler(t *testing.T) { f.release(r, true) } switch { - case handler.pkt != nil && test.wantPkt == nil: + case !handler.pkt.IsNil() && test.wantPkt.IsNil(): t.Errorf("got handler.pkt = not nil (pkt.Data = %x), want = nil", handler.pkt.Data().AsRange().ToSlice()) - case handler.pkt == nil && test.wantPkt != nil: + case handler.pkt.IsNil() && !test.wantPkt.IsNil(): t.Errorf("got handler.pkt = nil, want = not nil (pkt.Data = %x)", test.wantPkt.Data().AsRange().ToSlice()) - case handler.pkt != nil && test.wantPkt != nil: + case !handler.pkt.IsNil() && !test.wantPkt.IsNil(): if diff := cmp.Diff(test.wantPkt.Data().AsRange().ToSlice(), handler.pkt.Data().AsRange().ToSlice()); diff != "" { t.Errorf("pkt.Data mismatch (-want, +got):\n%s", diff) } diff --git a/pkg/tcpip/network/internal/fragmentation/reassembler.go b/pkg/tcpip/network/internal/fragmentation/reassembler.go index b7dace247..d93e6f9d0 100644 --- a/pkg/tcpip/network/internal/fragmentation/reassembler.go +++ b/pkg/tcpip/network/internal/fragmentation/reassembler.go @@ -30,7 +30,7 @@ type hole struct { final bool // pkt is the fragment packet if hole is filled. We keep the whole pkt rather // than the fragmented payload to prevent binding to specific buffer types. - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr } type reassembler struct { @@ -43,7 +43,7 @@ type reassembler struct { filled int done bool createdAt tcpip.MonotonicTime - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr } func newReassembler(id FragmentID, clock tcpip.Clock) *reassembler { @@ -60,7 +60,7 @@ func newReassembler(id FragmentID, clock tcpip.Clock) *reassembler { return r } -func (r *reassembler) process(first, last uint16, more bool, proto uint8, pkt *stack.PacketBuffer) (*stack.PacketBuffer, uint8, bool, int, error) { +func (r *reassembler) process(first, last uint16, more bool, proto uint8, pkt stack.PacketBufferPtr) (stack.PacketBufferPtr, uint8, bool, int, error) { r.mu.Lock() defer r.mu.Unlock() if r.done { @@ -145,7 +145,7 @@ func (r *reassembler) process(first, last uint16, more bool, proto uint8, pkt *s // options received in the first fragment should be used - and they should // override options from following fragments. if first == 0 { - if r.pkt != nil { + if !r.pkt.IsNil() { r.pkt.DecRef() } r.pkt = pkt.IncRef() diff --git a/pkg/tcpip/network/internal/fragmentation/reassembler_test.go b/pkg/tcpip/network/internal/fragmentation/reassembler_test.go index d1e758768..b30845904 100644 --- a/pkg/tcpip/network/internal/fragmentation/reassembler_test.go +++ b/pkg/tcpip/network/internal/fragmentation/reassembler_test.go @@ -29,7 +29,7 @@ type processParams struct { first uint16 last uint16 more bool - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr wantDone bool wantError error } @@ -45,7 +45,7 @@ func TestReassemblerProcess(t *testing.T) { return payload } - pkt := func(sizes ...int) *stack.PacketBuffer { + pkt := func(sizes ...int) stack.PacketBufferPtr { var buf bufferv2.Buffer for _, size := range sizes { buf.Append(v(size)) @@ -59,7 +59,7 @@ func TestReassemblerProcess(t *testing.T) { name string params []processParams want []hole - wantPkt *stack.PacketBuffer + wantPkt stack.PacketBufferPtr }{ { name: "No fragments", @@ -191,19 +191,19 @@ func TestReassemblerProcess(t *testing.T) { // reassembler will leak PacketBuffers. defer func() { for _, h := range r.holes { - if h.pkt != nil { + if !h.pkt.IsNil() { h.pkt.DecRef() } } - if r.pkt != nil { + if !r.pkt.IsNil() { r.pkt.DecRef() } }() - var resPkt *stack.PacketBuffer + var resPkt stack.PacketBufferPtr var isDone bool for _, param := range test.params { pkt, _, done, _, err := r.process(param.first, param.last, param.more, proto, param.pkt) - if pkt != nil { + if !pkt.IsNil() { defer pkt.DecRef() } if done != param.wantDone || err != param.wantError { @@ -215,9 +215,9 @@ func TestReassemblerProcess(t *testing.T) { } } - ignorePkt := func(a, b *stack.PacketBuffer) bool { return true } - cmpPktData := func(a, b *stack.PacketBuffer) bool { - if a == nil || b == nil { + ignorePkt := func(a, b stack.PacketBufferPtr) bool { return true } + cmpPktData := func(a, b stack.PacketBufferPtr) bool { + if a.IsNil() || b.IsNil() { return a == b } return bytes.Equal(a.Data().AsRange().ToSlice(), b.Data().AsRange().ToSlice()) @@ -246,16 +246,16 @@ func TestReassemblerProcess(t *testing.T) { } }) for _, p := range test.params { - if p.pkt != nil { + if !p.pkt.IsNil() { p.pkt.DecRef() } } for _, w := range test.want { - if w.pkt != nil { + if !w.pkt.IsNil() { w.pkt.DecRef() } } - if test.wantPkt != nil { + if !test.wantPkt.IsNil() { test.wantPkt.DecRef() } } diff --git a/pkg/tcpip/network/internal/multicast/example_test.go b/pkg/tcpip/network/internal/multicast/example_test.go index 2dbb891f4..1ad3f6f7b 100644 --- a/pkg/tcpip/network/internal/multicast/example_test.go +++ b/pkg/tcpip/network/internal/multicast/example_test.go @@ -115,7 +115,7 @@ func Example() { // Last used timestamp: 10000000000 } -func forwardPkt(*stack.PacketBuffer, *multicast.InstalledRoute) { +func forwardPkt(stack.PacketBufferPtr, *multicast.InstalledRoute) { fmt.Println("forwardPkt") } @@ -123,11 +123,11 @@ func emitMissingRouteEvent(stack.UnicastSourceAndMulticastDestination) { fmt.Println("emitMissingRouteEvent") } -func deliverPktLocally(*stack.PacketBuffer) { +func deliverPktLocally(stack.PacketBufferPtr) { fmt.Println("deliverPktLocally") } -func newPacketBuffer(body string) *stack.PacketBuffer { +func newPacketBuffer(body string) stack.PacketBufferPtr { return stack.NewPacketBuffer(stack.PacketBufferOptions{ Payload: bufferv2.MakeWithData([]byte(body)), }) diff --git a/pkg/tcpip/network/internal/multicast/route_table.go b/pkg/tcpip/network/internal/multicast/route_table.go index 41227e6ed..5bade5ae7 100644 --- a/pkg/tcpip/network/internal/multicast/route_table.go +++ b/pkg/tcpip/network/internal/multicast/route_table.go @@ -116,7 +116,7 @@ func (r *InstalledRoute) SetLastUsedTimestamp(monotonicTime tcpip.MonotonicTime) // for the entry. For such routes, packets are added to an expiring queue until // a route is installed. type PendingRoute struct { - packets []*stack.PacketBuffer + packets []stack.PacketBufferPtr // expiration is the timestamp at which the pending route should be expired. // @@ -265,7 +265,7 @@ func (r *RouteTable) cleanupPendingRoutes() { func (r *RouteTable) newPendingRoute() PendingRoute { return PendingRoute{ - packets: make([]*stack.PacketBuffer, 0, r.config.MaxPendingQueueSize), + packets: make([]stack.PacketBufferPtr, 0, r.config.MaxPendingQueueSize), expiration: r.config.Clock.NowMonotonic().Add(DefaultPendingRouteExpiration), } } @@ -326,7 +326,7 @@ func (e GetRouteResultState) String() string { // // If the relevant pending route queue is at max capacity, then returns false. // Otherwise, returns true. -func (r *RouteTable) GetRouteOrInsertPending(key stack.UnicastSourceAndMulticastDestination, pkt *stack.PacketBuffer) (GetRouteResult, bool) { +func (r *RouteTable) GetRouteOrInsertPending(key stack.UnicastSourceAndMulticastDestination, pkt stack.PacketBufferPtr) (GetRouteResult, bool) { r.installedMu.RLock() defer r.installedMu.RUnlock() @@ -374,7 +374,7 @@ func (r *RouteTable) getOrCreatePendingRouteRLocked(key stack.UnicastSourceAndMu // returned. The caller assumes ownership of these packets and is responsible // for forwarding and releasing them. If an installed route already exists for // the provided key, then it is overwritten. -func (r *RouteTable) AddInstalledRoute(key stack.UnicastSourceAndMulticastDestination, route *InstalledRoute) []*stack.PacketBuffer { +func (r *RouteTable) AddInstalledRoute(key stack.UnicastSourceAndMulticastDestination, route *InstalledRoute) []stack.PacketBufferPtr { r.installedMu.Lock() defer r.installedMu.Unlock() r.installedRoutes[key] = route diff --git a/pkg/tcpip/network/internal/multicast/route_table_test.go b/pkg/tcpip/network/internal/multicast/route_table_test.go index 31b69e30a..be32da260 100644 --- a/pkg/tcpip/network/internal/multicast/route_table_test.go +++ b/pkg/tcpip/network/internal/multicast/route_table_test.go @@ -45,7 +45,7 @@ var ( defaultRoute = stack.MulticastRoute{inputNICID, defaultOutgoingInterfaces} ) -func newPacketBuffer(body string) *stack.PacketBuffer { +func newPacketBuffer(body string) stack.PacketBufferPtr { return stack.NewPacketBuffer(stack.PacketBufferOptions{ Payload: bufferv2.MakeWithData([]byte(body)), }) @@ -258,7 +258,7 @@ func TestAddInstalledRouteWithPending(t *testing.T) { defer pkt.DecRef() cmpOpts := []cmp.Option{ - cmp.Transformer("AsSlices", func(pkt *stack.PacketBuffer) [][]byte { + cmp.Transformer("AsSlices", func(pkt stack.PacketBufferPtr) [][]byte { return pkt.AsSlices() }), cmp.Comparer(func(a [][]byte, b [][]byte) bool { @@ -269,12 +269,12 @@ func TestAddInstalledRouteWithPending(t *testing.T) { testCases := []struct { name string advance time.Duration - want []*stack.PacketBuffer + want []stack.PacketBufferPtr }{ { name: "not expired", advance: DefaultPendingRouteExpiration, - want: []*stack.PacketBuffer{pkt}, + want: []stack.PacketBufferPtr{pkt}, }, { name: "expired", diff --git a/pkg/tcpip/network/internal/testutil/testutil.go b/pkg/tcpip/network/internal/testutil/testutil.go index 0ea0d0ded..d99901b5f 100644 --- a/pkg/tcpip/network/internal/testutil/testutil.go +++ b/pkg/tcpip/network/internal/testutil/testutil.go @@ -30,7 +30,7 @@ import ( // to it and can mock errors. type MockLinkEndpoint struct { // WrittenPackets is where packets written to the endpoint are stored. - WrittenPackets []*stack.PacketBuffer + WrittenPackets []stack.PacketBufferPtr mtu uint32 err tcpip.Error @@ -88,7 +88,7 @@ func (*MockLinkEndpoint) Wait() {} func (*MockLinkEndpoint) ARPHardwareType() header.ARPHardwareType { return header.ARPHardwareNone } // AddHeader implements LinkEndpoint.AddHeader. -func (*MockLinkEndpoint) AddHeader(*stack.PacketBuffer) {} +func (*MockLinkEndpoint) AddHeader(stack.PacketBufferPtr) {} // Close releases all resources. func (ep *MockLinkEndpoint) Close() { @@ -103,7 +103,7 @@ func (ep *MockLinkEndpoint) Close() { // extraHeaderReserveLength indicates how much extra space will be reserved for // the other headers. The payload is made from Views of the sizes listed in // viewSizes. -func MakeRandPkt(transportHeaderLength int, extraHeaderReserveLength int, viewSizes []int, proto tcpip.NetworkProtocolNumber) *stack.PacketBuffer { +func MakeRandPkt(transportHeaderLength int, extraHeaderReserveLength int, viewSizes []int, proto tcpip.NetworkProtocolNumber) stack.PacketBufferPtr { var buf bufferv2.Buffer for _, s := range viewSizes { diff --git a/pkg/tcpip/network/ip_test.go b/pkg/tcpip/network/ip_test.go index 915b32d4a..81e891fac 100644 --- a/pkg/tcpip/network/ip_test.go +++ b/pkg/tcpip/network/ip_test.go @@ -127,7 +127,7 @@ func (t *testObject) checkValues(protocol tcpip.TransportProtocolNumber, v []byt // DeliverTransportPacket is called by network endpoints after parsing incoming // packets. This is used by the test object to verify that the results of the // parsing are expected. -func (t *testObject) DeliverTransportPacket(protocol tcpip.TransportProtocolNumber, pkt *stack.PacketBuffer) stack.TransportPacketDisposition { +func (t *testObject) DeliverTransportPacket(protocol tcpip.TransportProtocolNumber, pkt stack.PacketBufferPtr) stack.TransportPacketDisposition { netHdr := pkt.Network() v := pkt.Data().AsRange().ToView() defer v.Release() @@ -139,7 +139,7 @@ func (t *testObject) DeliverTransportPacket(protocol tcpip.TransportProtocolNumb // DeliverTransportError is called by network endpoints after parsing // incoming control (ICMP) packets. This is used by the test object to verify // that the results of the parsing are expected. -func (t *testObject) DeliverTransportError(local, remote tcpip.Address, net tcpip.NetworkProtocolNumber, trans tcpip.TransportProtocolNumber, transErr stack.TransportError, pkt *stack.PacketBuffer) { +func (t *testObject) DeliverTransportError(local, remote tcpip.Address, net tcpip.NetworkProtocolNumber, trans tcpip.TransportProtocolNumber, transErr stack.TransportError, pkt stack.PacketBufferPtr) { v := pkt.Data().AsRange().ToView() defer v.Release() t.checkValues(trans, v.AsSlice(), remote, local) @@ -159,7 +159,7 @@ func (t *testObject) DeliverTransportError(local, remote tcpip.Address, net tcpi t.controlCalls++ } -func (t *testObject) DeliverRawPacket(tcpip.TransportProtocolNumber, *stack.PacketBuffer) { +func (t *testObject) DeliverRawPacket(tcpip.TransportProtocolNumber, stack.PacketBufferPtr) { t.rawCalls++ } @@ -198,7 +198,7 @@ func (*testObject) Wait() {} // WritePacket is called by network endpoints after producing a packet and // writing it to the link endpoint. This is used by the test object to verify // that the produced packet is as expected. -func (t *testObject) WritePacket(_ *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { +func (t *testObject) WritePacket(_ *stack.Route, pkt stack.PacketBufferPtr) tcpip.Error { var prot tcpip.TransportProtocolNumber var srcAddr tcpip.Address var dstAddr tcpip.Address @@ -225,7 +225,7 @@ func (*testObject) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (*testObject) AddHeader(*stack.PacketBuffer) { +func (*testObject) AddHeader(stack.PacketBufferPtr) { panic("not implemented") } @@ -354,7 +354,7 @@ func (t *testInterface) setEnabled(v bool) { t.mu.disabled = !v } -func (*testInterface) WritePacketToRemote(tcpip.LinkAddress, *stack.PacketBuffer) tcpip.Error { +func (*testInterface) WritePacketToRemote(tcpip.LinkAddress, stack.PacketBufferPtr) tcpip.Error { return &tcpip.ErrNotSupported{} } @@ -1315,7 +1315,7 @@ func TestIPv6ReceiveControl(t *testing.T) { // after truncation, is large enough to hold a network header, it makes part of // view the packet's NetworkHeader and the rest its Data. Otherwise all of view // becomes Data. -func truncatedPacket(view []byte, trunc, netHdrLen int) *stack.PacketBuffer { +func truncatedPacket(view []byte, trunc, netHdrLen int) stack.PacketBufferPtr { v := view[:len(view)-trunc] pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ Payload: bufferv2.MakeWithData(v), @@ -1367,7 +1367,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { nicAddr tcpip.AddressWithPrefix remoteAddr tcpip.Address pktGen func(*testing.T, tcpip.Address) bufferv2.Buffer - checker func(*testing.T, *stack.PacketBuffer, tcpip.Address) + checker func(*testing.T, stack.PacketBufferPtr, tcpip.Address) expectedErr tcpip.Error }{ { @@ -1391,7 +1391,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { }) return bufferv2.MakeWithData(hdr.View()) }, - checker: func(t *testing.T, pkt *stack.PacketBuffer, src tcpip.Address) { + checker: func(t *testing.T, pkt stack.PacketBufferPtr, src tcpip.Address) { if src == header.IPv4Any { src = localIPv4Addr } @@ -1471,7 +1471,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { }) return bufferv2.MakeWithData(ip) }, - checker: func(t *testing.T, pkt *stack.PacketBuffer, src tcpip.Address) { + checker: func(t *testing.T, pkt stack.PacketBufferPtr, src tcpip.Address) { if src == header.IPv4Any { src = localIPv4Addr } @@ -1516,7 +1516,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { }) return bufferv2.MakeWithData(hdr.View()) }, - checker: func(t *testing.T, pkt *stack.PacketBuffer, src tcpip.Address) { + checker: func(t *testing.T, pkt stack.PacketBufferPtr, src tcpip.Address) { if src == header.IPv4Any { src = localIPv4Addr } @@ -1559,7 +1559,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { buf.Append(bufferv2.NewViewWithData(data)) return buf }, - checker: func(t *testing.T, pkt *stack.PacketBuffer, src tcpip.Address) { + checker: func(t *testing.T, pkt stack.PacketBufferPtr, src tcpip.Address) { if src == header.IPv4Any { src = localIPv4Addr } @@ -1604,7 +1604,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { }) return bufferv2.MakeWithData(hdr.View()) }, - checker: func(t *testing.T, pkt *stack.PacketBuffer, src tcpip.Address) { + checker: func(t *testing.T, pkt stack.PacketBufferPtr, src tcpip.Address) { if src == header.IPv6Any { src = localIPv6Addr } @@ -1651,7 +1651,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { }) return bufferv2.MakeWithData(hdr.View()) }, - checker: func(t *testing.T, pkt *stack.PacketBuffer, src tcpip.Address) { + checker: func(t *testing.T, pkt stack.PacketBufferPtr, src tcpip.Address) { if src == header.IPv6Any { src = localIPv6Addr } @@ -1688,7 +1688,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { }) return bufferv2.MakeWithData(ip) }, - checker: func(t *testing.T, pkt *stack.PacketBuffer, src tcpip.Address) { + checker: func(t *testing.T, pkt stack.PacketBufferPtr, src tcpip.Address) { if src == header.IPv6Any { src = localIPv6Addr } @@ -1792,7 +1792,7 @@ func TestWriteHeaderIncludedPacket(t *testing.T) { } pkt := e.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected a packet to be written") } test.checker(t, pkt, subTest.srcAddr) @@ -1872,7 +1872,7 @@ func TestICMPInclusionSize(t *testing.T) { return v } - v4Checker := func(t *testing.T, pkt *stack.PacketBuffer, payload []byte) { + v4Checker := func(t *testing.T, pkt stack.PacketBufferPtr, payload []byte) { // We already know the entire packet is the right size so we can use its // length to calculate the right payload size to check. expectedPayloadLength := pkt.Size() - header.IPv4MinimumSize - header.ICMPv4MinimumSize @@ -1892,7 +1892,7 @@ func TestICMPInclusionSize(t *testing.T) { ) } - v6Checker := func(t *testing.T, pkt *stack.PacketBuffer, payload []byte) { + v6Checker := func(t *testing.T, pkt stack.PacketBufferPtr, payload []byte) { // We already know the entire packet is the right size so we can use its // length to calculate the right payload size to check. expectedPayloadLength := pkt.Size() - header.IPv6MinimumSize - header.ICMPv6MinimumSize @@ -1913,7 +1913,7 @@ func TestICMPInclusionSize(t *testing.T) { name string srcAddress tcpip.Address injector func(*channel.Endpoint, tcpip.Address, []byte) []byte - checker func(*testing.T, *stack.PacketBuffer, []byte) + checker func(*testing.T, stack.PacketBufferPtr, []byte) payloadLength int // Not including IP header. linkMTU uint32 // Largest IP packet that the link can send as payload. replyLength int // Total size of IP/ICMP packet expected back. @@ -2048,7 +2048,7 @@ func TestICMPInclusionSize(t *testing.T) { }) v := test.injector(e, test.srcAddress, payload) pkt := e.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected a packet to be written") } if got, want := pkt.Size(), test.replyLength; got != want { diff --git a/pkg/tcpip/network/ipv4/icmp.go b/pkg/tcpip/network/ipv4/icmp.go index f71c00d76..afef4ec36 100644 --- a/pkg/tcpip/network/ipv4/icmp.go +++ b/pkg/tcpip/network/ipv4/icmp.go @@ -138,7 +138,7 @@ func (e *endpoint) checkLocalAddress(addr tcpip.Address) bool { // of the original packet that caused the ICMP one to be sent. This information // is used to find out which transport endpoint must be notified about the ICMP // packet. We only expect the payload, not the enclosing ICMP packet. -func (e *endpoint) handleControl(errInfo stack.TransportError, pkt *stack.PacketBuffer) { +func (e *endpoint) handleControl(errInfo stack.TransportError, pkt stack.PacketBufferPtr) { h, ok := pkt.Data().PullUp(header.IPv4MinimumSize) if !ok { return @@ -175,7 +175,7 @@ func (e *endpoint) handleControl(errInfo stack.TransportError, pkt *stack.Packet e.dispatcher.DeliverTransportError(srcAddr, dstAddr, ProtocolNumber, p, errInfo, pkt) } -func (e *endpoint) handleICMP(pkt *stack.PacketBuffer) { +func (e *endpoint) handleICMP(pkt stack.PacketBufferPtr) { received := e.stats.icmp.packetsReceived h := header.ICMPv4(pkt.TransportHeader().Slice()) if len(h) < header.ICMPv4MinimumSize { @@ -484,7 +484,7 @@ func (*icmpReasonHostUnreachable) isICMPReason() {} // the problematic packet. It incorporates as much of that packet as // possible as well as any error metadata as is available. returnError // expects pkt to hold a valid IPv4 packet as per the wire format. -func (p *protocol) returnError(reason icmpReason, pkt *stack.PacketBuffer, deliveredLocally bool) tcpip.Error { +func (p *protocol) returnError(reason icmpReason, pkt stack.PacketBufferPtr, deliveredLocally bool) tcpip.Error { origIPHdr := header.IPv4(pkt.NetworkHeader().Slice()) origIPHdrSrc := origIPHdr.SourceAddress() origIPHdrDst := origIPHdr.DestinationAddress() @@ -684,7 +684,7 @@ func (p *protocol) returnError(reason icmpReason, pkt *stack.PacketBuffer, deliv } // OnReassemblyTimeout implements fragmentation.TimeoutHandler. -func (p *protocol) OnReassemblyTimeout(pkt *stack.PacketBuffer) { +func (p *protocol) OnReassemblyTimeout(pkt stack.PacketBufferPtr) { // OnReassemblyTimeout sends a Time Exceeded Message, as per RFC 792: // // If a host reassembling a fragmented datagram cannot complete the @@ -693,7 +693,7 @@ func (p *protocol) OnReassemblyTimeout(pkt *stack.PacketBuffer) { // // If fragment zero is not available then no time exceeded need be sent at // all. - if pkt != nil { + if !pkt.IsNil() { p.returnError(&icmpReasonReassemblyTimeout{}, pkt, true /* deliveredLocally */) } } diff --git a/pkg/tcpip/network/ipv4/igmp.go b/pkg/tcpip/network/ipv4/igmp.go index 3671c5763..bf338ea5e 100644 --- a/pkg/tcpip/network/ipv4/igmp.go +++ b/pkg/tcpip/network/ipv4/igmp.go @@ -186,7 +186,7 @@ func (igmp *igmpState) isSourceIPValidLocked(src tcpip.Address, messageType head } // +checklocks:igmp.ep.mu -func (igmp *igmpState) isPacketValidLocked(pkt *stack.PacketBuffer, messageType header.IGMPType, hasRouterAlertOption bool) bool { +func (igmp *igmpState) isPacketValidLocked(pkt stack.PacketBufferPtr, messageType header.IGMPType, hasRouterAlertOption bool) bool { // We can safely assume that the IP header is valid if we got this far. iph := header.IPv4(pkt.NetworkHeader().Slice()) @@ -204,7 +204,7 @@ func (igmp *igmpState) isPacketValidLocked(pkt *stack.PacketBuffer, messageType // handleIGMP handles an IGMP packet. // // +checklocks:igmp.ep.mu -func (igmp *igmpState) handleIGMP(pkt *stack.PacketBuffer, hasRouterAlertOption bool) { +func (igmp *igmpState) handleIGMP(pkt stack.PacketBufferPtr, hasRouterAlertOption bool) { received := igmp.ep.stats.igmp.packetsReceived hdr, ok := pkt.Data().PullUp(header.IGMPMinimumSize) if !ok { diff --git a/pkg/tcpip/network/ipv4/igmp_test.go b/pkg/tcpip/network/ipv4/igmp_test.go index c589e0487..c907fde3a 100644 --- a/pkg/tcpip/network/ipv4/igmp_test.go +++ b/pkg/tcpip/network/ipv4/igmp_test.go @@ -46,7 +46,7 @@ var ( // validateIgmpPacket checks that a passed packet is an IPv4 IGMP packet sent // to the provided address with the passed fields set. Raises a t.Error if any // field does not match. -func validateIgmpPacket(t *testing.T, pkt *stack.PacketBuffer, igmpType header.IGMPType, maxRespTime byte, srcAddr, dstAddr, groupAddress tcpip.Address) { +func validateIgmpPacket(t *testing.T, pkt stack.PacketBufferPtr, igmpType header.IGMPType, maxRespTime byte, srcAddr, dstAddr, groupAddress tcpip.Address) { t.Helper() payload := stack.PayloadSince(pkt.NetworkHeader()) @@ -161,7 +161,7 @@ func TestIGMPV1Present(t *testing.T) { // the IGMPv1 General Membership Query in. { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("unable to Read IGMP packet, expected V2MembershipReport") } if got := s.Stats().IGMP.PacketsSent.V2MembershipReport.Value(); got != 1 { @@ -191,13 +191,13 @@ func TestIGMPV1Present(t *testing.T) { // Verify the solicited Membership Report is sent. Now that this NIC has seen // an IGMPv1 query, it should send an IGMPv1 Membership Report. - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet, expected V1MembershipReport only after advancing the clock = %+v", p) } ctx.clock.Advance(ipv4.UnsolicitedReportIntervalMax) { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("unable to Read IGMP packet, expected V1MembershipReport") } if got := s.Stats().IGMP.PacketsSent.V1MembershipReport.Value(); got != 1 { @@ -216,7 +216,7 @@ func TestIGMPV1Present(t *testing.T) { } { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("unable to Read IGMP packet, expected V2MembershipReport") } if got := s.Stats().IGMP.PacketsSent.V2MembershipReport.Value(); got != 2 { @@ -244,7 +244,7 @@ func TestSendQueuedIGMPReports(t *testing.T) { t.Errorf("got reportStat.Value() = %d, want = 0", got) } clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("got unexpected packet = %#v", p) } @@ -263,7 +263,7 @@ func TestSendQueuedIGMPReports(t *testing.T) { if got := reportStat.Value(); got != 1 { t.Errorf("got reportStat.Value() = %d, want = 1", got) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Error("expected to send an IGMP membership report") } else { validateIgmpPacket(t, p, header.IGMPv2MembershipReport, 0, stackAddr, multicastAddr, multicastAddr) @@ -276,7 +276,7 @@ func TestSendQueuedIGMPReports(t *testing.T) { if got := reportStat.Value(); got != 2 { t.Errorf("got reportStat.Value() = %d, want = 2", got) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Error("expected to send an IGMP membership report") } else { validateIgmpPacket(t, p, header.IGMPv2MembershipReport, 0, stackAddr, multicastAddr, multicastAddr) @@ -289,7 +289,7 @@ func TestSendQueuedIGMPReports(t *testing.T) { // Should have no more packets to send after the initial set of unsolicited // reports. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("got unexpected packet = %#v", p) } } diff --git a/pkg/tcpip/network/ipv4/ipv4.go b/pkg/tcpip/network/ipv4/ipv4.go index 4f245f7fe..c9c77fe26 100644 --- a/pkg/tcpip/network/ipv4/ipv4.go +++ b/pkg/tcpip/network/ipv4/ipv4.go @@ -109,7 +109,7 @@ type endpoint struct { } // HandleLinkResolutionFailure implements stack.LinkResolvableNetworkEndpoint. -func (e *endpoint) HandleLinkResolutionFailure(pkt *stack.PacketBuffer) { +func (e *endpoint) HandleLinkResolutionFailure(pkt stack.PacketBufferPtr) { // If we are operating as a router, return an ICMP error to the original // packet's sender. if pkt.NetworkPacketInfo.IsForwardedPacket { @@ -409,7 +409,7 @@ func (e *endpoint) NetworkProtocolNumber() tcpip.NetworkProtocolNumber { return e.protocol.Number() } -func (e *endpoint) addIPHeader(srcAddr, dstAddr tcpip.Address, pkt *stack.PacketBuffer, params stack.NetworkHeaderParams, options header.IPv4OptionsSerializer) tcpip.Error { +func (e *endpoint) addIPHeader(srcAddr, dstAddr tcpip.Address, pkt stack.PacketBufferPtr, params stack.NetworkHeaderParams, options header.IPv4OptionsSerializer) tcpip.Error { hdrLen := header.IPv4MinimumSize var optLen int if options != nil { @@ -447,7 +447,7 @@ func (e *endpoint) addIPHeader(srcAddr, dstAddr tcpip.Address, pkt *stack.Packet // fragment. It returns the number of fragments handled and the number of // fragments left to be processed. The IP header must already be present in the // original packet. -func (e *endpoint) handleFragments(_ *stack.Route, networkMTU uint32, pkt *stack.PacketBuffer, handler func(*stack.PacketBuffer) tcpip.Error) (int, int, tcpip.Error) { +func (e *endpoint) handleFragments(_ *stack.Route, networkMTU uint32, pkt stack.PacketBufferPtr, handler func(stack.PacketBufferPtr) tcpip.Error) (int, int, tcpip.Error) { // Round the MTU down to align to 8 bytes. fragmentPayloadSize := networkMTU &^ 7 networkHeader := header.IPv4(pkt.NetworkHeader().Slice()) @@ -470,7 +470,7 @@ func (e *endpoint) handleFragments(_ *stack.Route, networkMTU uint32, pkt *stack } // WritePacket writes a packet to the given destination address and protocol. -func (e *endpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, pkt *stack.PacketBuffer) tcpip.Error { +func (e *endpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, pkt stack.PacketBufferPtr) tcpip.Error { if err := e.addIPHeader(r.LocalAddress(), r.RemoteAddress(), pkt, params, nil /* options */); err != nil { return err } @@ -478,7 +478,7 @@ func (e *endpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, return e.writePacket(r, pkt) } -func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { +func (e *endpoint) writePacket(r *stack.Route, pkt stack.PacketBufferPtr) tcpip.Error { netHeader := header.IPv4(pkt.NetworkHeader().Slice()) dstAddr := netHeader.DestinationAddress() @@ -510,7 +510,7 @@ func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Er return e.writePacketPostRouting(r, pkt, false /* headerIncluded */) } -func (e *endpoint) writePacketPostRouting(r *stack.Route, pkt *stack.PacketBuffer, headerIncluded bool) tcpip.Error { +func (e *endpoint) writePacketPostRouting(r *stack.Route, pkt stack.PacketBufferPtr, headerIncluded bool) tcpip.Error { if r.Loop()&stack.PacketLoop != 0 { // If the packet was generated by the stack (not a raw/packet endpoint // where a packet may be written with the header included), then we can @@ -545,7 +545,7 @@ func (e *endpoint) writePacketPostRouting(r *stack.Route, pkt *stack.PacketBuffe // is set but the packet must be fragmented for the non-forwarding case. return &tcpip.ErrMessageTooLong{} } - sent, remain, err := e.handleFragments(r, networkMTU, pkt, func(fragPkt *stack.PacketBuffer) tcpip.Error { + sent, remain, err := e.handleFragments(r, networkMTU, pkt, func(fragPkt stack.PacketBufferPtr) tcpip.Error { // TODO(gvisor.dev/issue/3884): Evaluate whether we want to send each // fragment one by one using WritePacket() (current strategy) or if we // want to create a PacketBufferList from the fragments and feed it to @@ -566,7 +566,7 @@ func (e *endpoint) writePacketPostRouting(r *stack.Route, pkt *stack.PacketBuffe } // WriteHeaderIncludedPacket implements stack.NetworkEndpoint. -func (e *endpoint) WriteHeaderIncludedPacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { +func (e *endpoint) WriteHeaderIncludedPacket(r *stack.Route, pkt stack.PacketBufferPtr) tcpip.Error { // The packet already has an IP header, but there are a few required // checks. h, ok := pkt.Data().PullUp(header.IPv4MinimumSize) @@ -628,7 +628,7 @@ func (e *endpoint) WriteHeaderIncludedPacket(r *stack.Route, pkt *stack.PacketBu // updating the options. // // This method should be invoked by the endpoint that received the pkt. -func (e *endpoint) forwardPacketWithRoute(route *stack.Route, pkt *stack.PacketBuffer, updateOptions bool) ip.ForwardingError { +func (e *endpoint) forwardPacketWithRoute(route *stack.Route, pkt stack.PacketBufferPtr, updateOptions bool) ip.ForwardingError { h := header.IPv4(pkt.NetworkHeader().Slice()) stk := e.protocol.stack @@ -696,7 +696,7 @@ func (e *endpoint) forwardPacketWithRoute(route *stack.Route, pkt *stack.PacketB } // forwardUnicastPacket attempts to forward a packet to its final destination. -func (e *endpoint) forwardUnicastPacket(pkt *stack.PacketBuffer) ip.ForwardingError { +func (e *endpoint) forwardUnicastPacket(pkt stack.PacketBufferPtr) ip.ForwardingError { h := header.IPv4(pkt.NetworkHeader().Slice()) dstAddr := h.DestinationAddress() @@ -770,7 +770,7 @@ func (e *endpoint) forwardUnicastPacket(pkt *stack.PacketBuffer) ip.ForwardingEr // HandlePacket is called by the link layer when new ipv4 packets arrive for // this endpoint. -func (e *endpoint) HandlePacket(pkt *stack.PacketBuffer) { +func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { stats := e.stats.ip stats.PacketsReceived.Increment() @@ -827,7 +827,7 @@ func (e *endpoint) HandlePacket(pkt *stack.PacketBuffer) { // handleLocalPacket is like HandlePacket except it does not perform the // prerouting iptables hook or check for loopback traffic that originated from // outside of the netstack (i.e. martian loopback packets). -func (e *endpoint) handleLocalPacket(pkt *stack.PacketBuffer, canSkipRXChecksum bool) { +func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksum bool) { stats := e.stats.ip stats.PacketsReceived.Increment() @@ -868,7 +868,7 @@ func validateAddressesForForwarding(h header.IPv4) ip.ForwardingError { // // This method should be invoked for incoming multicast packets using the // endpoint that received the packet. -func (e *endpoint) forwardMulticastPacket(h header.IPv4, pkt *stack.PacketBuffer) ip.ForwardingError { +func (e *endpoint) forwardMulticastPacket(h header.IPv4, pkt stack.PacketBufferPtr) ip.ForwardingError { if err := validateAddressesForForwarding(h); err != nil { return err } @@ -921,7 +921,7 @@ func (e *endpoint) forwardMulticastPacket(h header.IPv4, pkt *stack.PacketBuffer return &ip.ErrNoRoute{} } -func (e *endpoint) updateOptionsForForwarding(pkt *stack.PacketBuffer) ip.ForwardingError { +func (e *endpoint) updateOptionsForForwarding(pkt stack.PacketBufferPtr) ip.ForwardingError { h := header.IPv4(pkt.NetworkHeader().Slice()) if opts := h.Options(); len(opts) != 0 { newOpts, _, optProblem := e.processIPOptions(pkt, opts, &optionUsageForward{}) @@ -956,7 +956,7 @@ func (e *endpoint) updateOptionsForForwarding(pkt *stack.PacketBuffer) ip.Forwar // provided installedRoute. // // This method should be invoked by the endpoint that received the pkt. -func (e *endpoint) forwardValidatedMulticastPacket(pkt *stack.PacketBuffer, installedRoute *multicast.InstalledRoute) ip.ForwardingError { +func (e *endpoint) forwardValidatedMulticastPacket(pkt stack.PacketBufferPtr, installedRoute *multicast.InstalledRoute) ip.ForwardingError { // Per RFC 1812 section 5.2.1.3, // // Based on the IP source and destination addresses found in the datagram @@ -989,7 +989,7 @@ func (e *endpoint) forwardValidatedMulticastPacket(pkt *stack.PacketBuffer, inst // of the provided outgoingInterface. // // This method should be invoked by the endpoint that received the pkt. -func (e *endpoint) forwardMulticastPacketForOutgoingInterface(pkt *stack.PacketBuffer, outgoingInterface stack.MulticastRouteOutgoingInterface) ip.ForwardingError { +func (e *endpoint) forwardMulticastPacketForOutgoingInterface(pkt stack.PacketBufferPtr, outgoingInterface stack.MulticastRouteOutgoingInterface) ip.ForwardingError { h := header.IPv4(pkt.NetworkHeader().Slice()) // Per RFC 1812 section 5.2.1.3, @@ -1016,7 +1016,7 @@ func (e *endpoint) forwardMulticastPacketForOutgoingInterface(pkt *stack.PacketB return e.forwardPacketWithRoute(route, pkt, true /* updateOptions */) } -func (e *endpoint) handleValidatedPacket(h header.IPv4, pkt *stack.PacketBuffer, inNICName string) { +func (e *endpoint) handleValidatedPacket(h header.IPv4, pkt stack.PacketBufferPtr, inNICName string) { pkt.NICID = e.nic.ID() // Raw socket packets are delivered based solely on the transport protocol @@ -1123,7 +1123,7 @@ func (e *endpoint) handleForwardingError(err ip.ForwardingError) { stats.Forwarding.Errors.Increment() } -func (e *endpoint) deliverPacketLocally(h header.IPv4, pkt *stack.PacketBuffer, inNICName string) { +func (e *endpoint) deliverPacketLocally(h header.IPv4, pkt stack.PacketBufferPtr, inNICName string) { stats := e.stats // iptables filtering. All packets that reach here are intended for // this machine and will not be forwarded. @@ -1633,7 +1633,7 @@ func (p *protocol) MulticastRouteLastUsedTime(addresses stack.UnicastSourceAndMu return timestamp, nil } -func (p *protocol) forwardPendingMulticastPacket(pkt *stack.PacketBuffer, installedRoute *multicast.InstalledRoute) { +func (p *protocol) forwardPendingMulticastPacket(pkt stack.PacketBufferPtr, installedRoute *multicast.InstalledRoute) { defer pkt.DecRef() // Attempt to forward the packet using the endpoint that it originally @@ -1690,7 +1690,7 @@ func (p *protocol) isSubnetLocalBroadcastAddress(addr tcpip.Address) bool { // returns the parsed IP header. // // Returns true if the IP header was successfully parsed. -func (p *protocol) parseAndValidate(pkt *stack.PacketBuffer) (header.IPv4, bool) { +func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (header.IPv4, bool) { transProtoNum, hasTransportHdr, ok := p.Parse(pkt) if !ok { return nil, false @@ -1714,7 +1714,7 @@ func (p *protocol) parseAndValidate(pkt *stack.PacketBuffer) (header.IPv4, bool) return h, true } -func (p *protocol) parseTransport(pkt *stack.PacketBuffer, transProtoNum tcpip.TransportProtocolNumber) { +func (p *protocol) parseTransport(pkt stack.PacketBufferPtr, transProtoNum tcpip.TransportProtocolNumber) { if transProtoNum == header.ICMPv4ProtocolNumber { // The transport layer will handle transport layer parsing errors. _ = parse.ICMPv4(pkt) @@ -1732,7 +1732,7 @@ func (p *protocol) parseTransport(pkt *stack.PacketBuffer, transProtoNum tcpip.T } // Parse implements stack.NetworkProtocol. -func (*protocol) Parse(pkt *stack.PacketBuffer) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) { +func (*protocol) Parse(pkt stack.PacketBufferPtr) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) { if ok := parse.IPv4(pkt); !ok { return 0, false, false } @@ -1759,7 +1759,7 @@ func (p *protocol) allowICMPReply(icmpType header.ICMPv4Type, code header.ICMPv4 } // SendRejectionError implements stack.RejectIPv4WithHandler. -func (p *protocol) SendRejectionError(pkt *stack.PacketBuffer, rejectWith stack.RejectIPv4WithICMPType, inputHook bool) tcpip.Error { +func (p *protocol) SendRejectionError(pkt stack.PacketBufferPtr, rejectWith stack.RejectIPv4WithICMPType, inputHook bool) tcpip.Error { switch rejectWith { case stack.RejectIPv4WithICMPNetUnreachable: return p.returnError(&icmpReasonNetworkUnreachable{}, pkt, inputHook) @@ -1801,7 +1801,7 @@ func calculateNetworkMTU(linkMTU, networkHeaderSize uint32) (uint32, tcpip.Error return networkMTU - networkHeaderSize, nil } -func packetMustBeFragmented(pkt *stack.PacketBuffer, networkMTU uint32) bool { +func packetMustBeFragmented(pkt stack.PacketBufferPtr, networkMTU uint32) bool { payload := len(pkt.TransportHeader().Slice()) + pkt.Data().Size() return pkt.GSOOptions.Type == stack.GSONone && uint32(payload) > networkMTU } @@ -1877,7 +1877,7 @@ func NewProtocol(s *stack.Stack) stack.NetworkProtocol { return NewProtocolWithOptions(Options{})(s) } -func buildNextFragment(pf *fragmentation.PacketFragmenter, originalIPHeader header.IPv4) (*stack.PacketBuffer, bool) { +func buildNextFragment(pf *fragmentation.PacketFragmenter, originalIPHeader header.IPv4) (stack.PacketBufferPtr, bool) { fragPkt, offset, copied, more := pf.BuildNextFragment() fragPkt.NetworkProtocolNumber = ProtocolNumber @@ -2218,7 +2218,7 @@ type optionTracker struct { // // If there were no errors during parsing, the new set of options is returned as // a new buffer. -func (e *endpoint) processIPOptions(pkt *stack.PacketBuffer, opts header.IPv4Options, usage optionsUsage) (header.IPv4Options, optionTracker, *header.IPv4OptParameterProblem) { +func (e *endpoint) processIPOptions(pkt stack.PacketBufferPtr, opts header.IPv4Options, usage optionsUsage) (header.IPv4Options, optionTracker, *header.IPv4OptParameterProblem) { stats := e.stats.ip optIter := opts.MakeIterator() diff --git a/pkg/tcpip/network/ipv4/ipv4_test.go b/pkg/tcpip/network/ipv4/ipv4_test.go index 585322832..5fb4f0677 100644 --- a/pkg/tcpip/network/ipv4/ipv4_test.go +++ b/pkg/tcpip/network/ipv4/ipv4_test.go @@ -249,7 +249,7 @@ type packetOptions struct { options header.IPv4Options } -func newICMPEchoPacket(t *testing.T, srcAddr, dstAddr tcpip.Address, ttl uint8, options packetOptions) (*stack.PacketBuffer, []byte) { +func newICMPEchoPacket(t *testing.T, srcAddr, dstAddr tcpip.Address, ttl uint8, options packetOptions) (stack.PacketBufferPtr, []byte) { const ( arbitraryICMPHeaderSequence = 123 randomIdent = 42 @@ -314,12 +314,12 @@ func max(a, b int) int { return b } -func checkFragements(t *testing.T, ep *channel.Endpoint, expectedFragments []fragmentInfo, requestPkt *stack.PacketBuffer) { +func checkFragements(t *testing.T, ep *channel.Endpoint, expectedFragments []fragmentInfo, requestPkt stack.PacketBufferPtr) { t.Helper() - var fragmentedPackets []*stack.PacketBuffer + var fragmentedPackets []stack.PacketBufferPtr for i := 0; i < len(expectedFragments); i++ { reply := ep.Read() - if reply == nil { + if reply.IsNil() { t.Fatal("Expected ICMP Echo fragment through outgoing NIC") } fragmentedPackets = append(fragmentedPackets, reply) @@ -546,7 +546,7 @@ func TestForwarding(t *testing.T) { reply := incomingEndpoint.Read() if test.icmpError != nil { - if reply == nil { + if reply.IsNil() { t.Fatalf("Expected ICMP packet type %d through incoming NIC", test.icmpError.icmpType) } @@ -564,7 +564,7 @@ func TestForwarding(t *testing.T) { ), ) reply.DecRef() - } else if reply != nil { + } else if !reply.IsNil() { t.Fatalf("Expected no ICMP packet through incoming NIC, instead found: %#v", reply) } @@ -575,7 +575,7 @@ func TestForwarding(t *testing.T) { if test.expectPacketForwarded { reply := outgoingEndpoint.Read() - if reply == nil { + if reply.IsNil() { t.Fatal("Expected ICMP Echo packet through outgoing NIC") } @@ -595,7 +595,7 @@ func TestForwarding(t *testing.T) { ) reply.DecRef() } else { - if reply := outgoingEndpoint.Read(); reply != nil { + if reply := outgoingEndpoint.Read(); !reply.IsNil() { t.Fatalf("Expected no ICMP Echo packet through outgoing NIC, instead found: %#v", reply) } } @@ -732,7 +732,7 @@ func TestFragmentForwarding(t *testing.T) { reply := incomingEndpoint.Read() if test.icmpError != nil { - if reply == nil { + if reply.IsNil() { t.Fatalf("Expected ICMP packet type %d through incoming NIC", test.icmpError.icmpType) } @@ -750,7 +750,7 @@ func TestFragmentForwarding(t *testing.T) { ), ) reply.DecRef() - } else if reply != nil { + } else if !reply.IsNil() { t.Fatalf("Expected no ICMP packet through incoming NIC, instead found: %#v", reply) } @@ -762,7 +762,7 @@ func TestFragmentForwarding(t *testing.T) { if len(test.expectedFragmentsForwarded) > 0 { checkFragements(t, outgoingEndpoint, test.expectedFragmentsForwarded, requestPkt) } else { - if reply := outgoingEndpoint.Read(); reply != nil { + if reply := outgoingEndpoint.Read(); !reply.IsNil() { t.Errorf("Expected no ICMP Echo packet through outgoing NIC, instead found: %#v", reply) } } @@ -898,7 +898,7 @@ func TestMulticastFragmentForwarding(t *testing.T) { incomingEndpoint.InjectInbound(header.IPv4ProtocolNumber, requestPkt) reply := incomingEndpoint.Read() - if reply != nil { + if !reply.IsNil() { // An ICMP error should never be sent in response to a multicast packet. t.Errorf("Expected no ICMP packet through incoming NIC, instead found: %#v", reply) } @@ -911,7 +911,7 @@ func TestMulticastFragmentForwarding(t *testing.T) { if len(test.expectedFragmentsForwarded) > 0 { checkFragements(t, outgoingEndpoint, test.expectedFragmentsForwarded, requestPkt) } else { - if reply := outgoingEndpoint.Read(); reply != nil { + if reply := outgoingEndpoint.Read(); !reply.IsNil() { t.Errorf("Expected no ICMP Echo packet through outgoing NIC, instead found: %#v", reply) } } @@ -1063,7 +1063,7 @@ func TestMulticastForwardingOptions(t *testing.T) { incomingEndpoint.InjectInbound(header.IPv4ProtocolNumber, requestPkt) reply := incomingEndpoint.Read() - if reply != nil { + if !reply.IsNil() { // An ICMP error should never be sent in response to a multicast packet. t.Errorf("Expected no ICMP packet through incoming NIC, instead found: %#v", reply) } @@ -1075,7 +1075,7 @@ func TestMulticastForwardingOptions(t *testing.T) { if test.expectPacketForwarded { reply := outgoingEndpoint.Read() - if reply == nil { + if reply.IsNil() { t.Fatal("Expected ICMP Echo packet through outgoing NIC") } @@ -1095,7 +1095,7 @@ func TestMulticastForwardingOptions(t *testing.T) { ) reply.DecRef() } else { - if reply := outgoingEndpoint.Read(); reply != nil { + if reply := outgoingEndpoint.Read(); !reply.IsNil() { t.Fatalf("Expected no ICMP Echo packet through outgoing NIC, instead found: %#v", reply) } } @@ -1816,7 +1816,7 @@ func TestIPv4Sanity(t *testing.T) { defer requestPkt.DecRef() e.InjectInbound(header.IPv4ProtocolNumber, requestPkt) reply := e.Read() - if reply == nil { + if reply.IsNil() { if test.shouldFail { if test.expectErrorICMP { t.Fatalf("ICMP error response (type %d, code %d) missing", test.ICMPType, test.ICMPCode) @@ -1937,7 +1937,7 @@ func TestIPv4Sanity(t *testing.T) { // If withIPHeader is set to true, we will validate the fragmented packets' IP // headers against the source packet's IP header. If set to false, we validate // the fragmented packets' IP headers against each other. -func compareFragments(packets []*stack.PacketBuffer, sourcePacket *stack.PacketBuffer, mtu uint32, wantFragments []fragmentInfo, proto tcpip.TransportProtocolNumber, withIPHeader bool, expectedAvailableHeaderBytes int) error { +func compareFragments(packets []stack.PacketBufferPtr, sourcePacket stack.PacketBufferPtr, mtu uint32, wantFragments []fragmentInfo, proto tcpip.TransportProtocolNumber, withIPHeader bool, expectedAvailableHeaderBytes int) error { // Make a complete array of the sourcePacket packet. var source header.IPv4 buf := sourcePacket.ToBuffer() @@ -2783,12 +2783,12 @@ func TestFragmentReassemblyTimeout(t *testing.T) { reply := e.Read() if !test.expectICMP { - if reply != nil { + if !reply.IsNil() { t.Fatalf("unexpected ICMP error message received: %#v", reply) } return } - if reply == nil { + if reply.IsNil() { t.Fatal("expected ICMP error message missing") } if firstFragmentSent.Size() == 0 { @@ -3507,7 +3507,7 @@ func (*limitedMatcher) Name() string { } // Match implements Matcher.Match. -func (lm *limitedMatcher) Match(stack.Hook, *stack.PacketBuffer, string, string) (bool, bool) { +func (lm *limitedMatcher) Match(stack.Hook, stack.PacketBufferPtr, string, string) (bool, bool) { if lm.limit == 0 { return true, false } @@ -3573,7 +3573,7 @@ func TestPacketQueuing(t *testing.T) { }, checkResp: func(t *testing.T, e *channel.Endpoint) { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("timed out waiting for packet") } defer p.DecRef() @@ -3621,7 +3621,7 @@ func TestPacketQueuing(t *testing.T) { }, checkResp: func(t *testing.T, e *channel.Endpoint) { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("timed out waiting for packet") } defer p.DecRef() @@ -3675,7 +3675,7 @@ func TestPacketQueuing(t *testing.T) { { clock.RunImmediatelyScheduledJobs() p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("timed out waiting for packet") } if p.NetworkProtocolNumber != arp.ProtocolNumber { @@ -3920,7 +3920,7 @@ func TestIcmpRateLimit(t *testing.T) { }, check: func(t *testing.T, e *channel.Endpoint, round int) { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("expected echo response, no packet read in endpoint in round %d", round) } defer p.DecRef() @@ -3962,13 +3962,13 @@ func TestIcmpRateLimit(t *testing.T) { check: func(t *testing.T, e *channel.Endpoint, round int) { p := e.Read() if round >= icmpBurst { - if p != nil { + if !p.IsNil() { t.Errorf("got packet %x in round %d, expected ICMP rate limit to stop it", p.Data().AsRange().ToSlice(), round) p.DecRef() } return } - if p == nil { + if p.IsNil() { t.Fatalf("expected unreachable in round %d, no packet read in endpoint", round) } defer p.DecRef() diff --git a/pkg/tcpip/network/ipv6/icmp.go b/pkg/tcpip/network/ipv6/icmp.go index c19d6bf81..d629d7696 100644 --- a/pkg/tcpip/network/ipv6/icmp.go +++ b/pkg/tcpip/network/ipv6/icmp.go @@ -164,7 +164,7 @@ func (e *endpoint) checkLocalAddress(addr tcpip.Address) bool { // the original packet that caused the ICMP one to be sent. This information is // used to find out which transport endpoint must be notified about the ICMP // packet. -func (e *endpoint) handleControl(transErr stack.TransportError, pkt *stack.PacketBuffer) { +func (e *endpoint) handleControl(transErr stack.TransportError, pkt stack.PacketBufferPtr) { h, ok := pkt.Data().PullUp(header.IPv6MinimumSize) if !ok { return @@ -267,7 +267,7 @@ func getTargetLinkAddr(it header.NDPOptionIterator) (tcpip.LinkAddress, bool) { }) } -func isMLDValid(pkt *stack.PacketBuffer, iph header.IPv6, routerAlert *header.IPv6RouterAlertOption) bool { +func isMLDValid(pkt stack.PacketBufferPtr, iph header.IPv6, routerAlert *header.IPv6RouterAlertOption) bool { // As per RFC 2710 section 3: // All MLD messages described in this document are sent with a link-local // IPv6 Source Address, an IPv6 Hop Limit of 1, and an IPv6 Router Alert @@ -287,7 +287,7 @@ func isMLDValid(pkt *stack.PacketBuffer, iph header.IPv6, routerAlert *header.IP return true } -func (e *endpoint) handleICMP(pkt *stack.PacketBuffer, hasFragmentHeader bool, routerAlert *header.IPv6RouterAlertOption) { +func (e *endpoint) handleICMP(pkt stack.PacketBufferPtr, hasFragmentHeader bool, routerAlert *header.IPv6RouterAlertOption) { sent := e.stats.icmp.packetsSent received := e.stats.icmp.packetsReceived h := header.ICMPv6(pkt.TransportHeader().Slice()) @@ -1042,7 +1042,7 @@ func (*icmpReasonReassemblyTimeout) respondsToMulticast() bool { // returnError takes an error descriptor and generates the appropriate ICMP // error packet for IPv6 and sends it. -func (p *protocol) returnError(reason icmpReason, pkt *stack.PacketBuffer, deliveredLocally bool) tcpip.Error { +func (p *protocol) returnError(reason icmpReason, pkt stack.PacketBufferPtr, deliveredLocally bool) tcpip.Error { origIPHdr := header.IPv6(pkt.NetworkHeader().Slice()) origIPHdrSrc := origIPHdr.SourceAddress() origIPHdrDst := origIPHdr.DestinationAddress() @@ -1207,14 +1207,14 @@ func (p *protocol) returnError(reason icmpReason, pkt *stack.PacketBuffer, deliv } // OnReassemblyTimeout implements fragmentation.TimeoutHandler. -func (p *protocol) OnReassemblyTimeout(pkt *stack.PacketBuffer) { +func (p *protocol) OnReassemblyTimeout(pkt stack.PacketBufferPtr) { // OnReassemblyTimeout sends a Time Exceeded Message as per RFC 2460 Section // 4.5: // // If the first fragment (i.e., the one with a Fragment Offset of zero) has // been received, an ICMP Time Exceeded -- Fragment Reassembly Time Exceeded // message should be sent to the source of that fragment. - if pkt != nil { + if !pkt.IsNil() { p.returnError(&icmpReasonReassemblyTimeout{}, pkt, true /* deliveredLocally */) } } diff --git a/pkg/tcpip/network/ipv6/icmp_test.go b/pkg/tcpip/network/ipv6/icmp_test.go index a5a0ed20a..59d790f65 100644 --- a/pkg/tcpip/network/ipv6/icmp_test.go +++ b/pkg/tcpip/network/ipv6/icmp_test.go @@ -86,7 +86,7 @@ func (*stubLinkEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.E func (*stubLinkEndpoint) Attach(stack.NetworkDispatcher) {} -func (*stubLinkEndpoint) AddHeader(*stack.PacketBuffer) {} +func (*stubLinkEndpoint) AddHeader(stack.PacketBufferPtr) {} func (*stubLinkEndpoint) Wait() {} @@ -94,11 +94,11 @@ type stubDispatcher struct { stack.TransportDispatcher } -func (*stubDispatcher) DeliverTransportPacket(tcpip.TransportProtocolNumber, *stack.PacketBuffer) stack.TransportPacketDisposition { +func (*stubDispatcher) DeliverTransportPacket(tcpip.TransportProtocolNumber, stack.PacketBufferPtr) stack.TransportPacketDisposition { return stack.TransportPacketHandled } -func (*stubDispatcher) DeliverRawPacket(tcpip.TransportProtocolNumber, *stack.PacketBuffer) { +func (*stubDispatcher) DeliverRawPacket(tcpip.TransportProtocolNumber, stack.PacketBufferPtr) { // No-op. } @@ -137,7 +137,7 @@ func (*testInterface) Spoofing() bool { return false } -func (t *testInterface) WritePacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { +func (t *testInterface) WritePacket(r *stack.Route, pkt stack.PacketBufferPtr) tcpip.Error { pkt.EgressRoute = r.Fields() var pkts stack.PacketBufferList pkts.PushBack(pkt) @@ -145,7 +145,7 @@ func (t *testInterface) WritePacket(r *stack.Route, pkt *stack.PacketBuffer) tcp return err } -func (t *testInterface) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt *stack.PacketBuffer) tcpip.Error { +func (t *testInterface) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt stack.PacketBufferPtr) tcpip.Error { pkt.EgressRoute.NetProto = pkt.NetworkProtocolNumber pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr var pkts stack.PacketBufferList @@ -521,7 +521,7 @@ func routeICMPv6Packet(t *testing.T, clock *faketime.ManualClock, args routeArgs clock.RunImmediatelyScheduledJobs() pi := args.src.Read() - if pi == nil { + if pi.IsNil() { t.Fatal("packet didn't arrive") } defer pi.DecRef() @@ -1341,7 +1341,7 @@ func TestLinkAddressRequest(t *testing.T) { } pkt := linkEP.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to send a link address request") } defer pkt.DecRef() @@ -1425,7 +1425,7 @@ func TestPacketQueing(t *testing.T) { }, checkResp: func(t *testing.T, e *channel.Endpoint) { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("timed out waiting for packet") } defer p.DecRef() @@ -1476,7 +1476,7 @@ func TestPacketQueing(t *testing.T) { }, checkResp: func(t *testing.T, e *channel.Endpoint) { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("timed out waiting for packet") } defer p.DecRef() @@ -1531,7 +1531,7 @@ func TestPacketQueing(t *testing.T) { { c.clock.RunImmediatelyScheduledJobs() p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("timed out waiting for packet") } if p.NetworkProtocolNumber != ProtocolNumber { diff --git a/pkg/tcpip/network/ipv6/ipv6.go b/pkg/tcpip/network/ipv6/ipv6.go index d5b8f9ded..190e28835 100644 --- a/pkg/tcpip/network/ipv6/ipv6.go +++ b/pkg/tcpip/network/ipv6/ipv6.go @@ -285,7 +285,7 @@ func (*endpoint) DuplicateAddressProtocol() tcpip.NetworkProtocolNumber { } // HandleLinkResolutionFailure implements stack.LinkResolvableNetworkEndpoint. -func (e *endpoint) HandleLinkResolutionFailure(pkt *stack.PacketBuffer) { +func (e *endpoint) HandleLinkResolutionFailure(pkt stack.PacketBufferPtr) { // If we are operating as a router, we should return an ICMP error to the // original packet's sender. if pkt.NetworkPacketInfo.IsForwardedPacket { @@ -709,7 +709,7 @@ func (e *endpoint) MaxHeaderLength() uint16 { return e.nic.MaxHeaderLength() + header.IPv6MinimumSize } -func addIPHeader(srcAddr, dstAddr tcpip.Address, pkt *stack.PacketBuffer, params stack.NetworkHeaderParams, extensionHeaders header.IPv6ExtHdrSerializer) tcpip.Error { +func addIPHeader(srcAddr, dstAddr tcpip.Address, pkt stack.PacketBufferPtr, params stack.NetworkHeaderParams, extensionHeaders header.IPv6ExtHdrSerializer) tcpip.Error { extHdrsLen := extensionHeaders.Length() length := pkt.Size() + extensionHeaders.Length() if length > math.MaxUint16 { @@ -728,7 +728,7 @@ func addIPHeader(srcAddr, dstAddr tcpip.Address, pkt *stack.PacketBuffer, params return nil } -func packetMustBeFragmented(pkt *stack.PacketBuffer, networkMTU uint32) bool { +func packetMustBeFragmented(pkt stack.PacketBufferPtr, networkMTU uint32) bool { payload := len(pkt.TransportHeader().Slice()) + pkt.Data().Size() return pkt.GSOOptions.Type == stack.GSONone && uint32(payload) > networkMTU } @@ -738,7 +738,7 @@ func packetMustBeFragmented(pkt *stack.PacketBuffer, networkMTU uint32) bool { // fragments left to be processed. The IP header must already be present in the // original packet. The transport header protocol number is required to avoid // parsing the IPv6 extension headers. -func (e *endpoint) handleFragments(r *stack.Route, networkMTU uint32, pkt *stack.PacketBuffer, transProto tcpip.TransportProtocolNumber, handler func(*stack.PacketBuffer) tcpip.Error) (int, int, tcpip.Error) { +func (e *endpoint) handleFragments(r *stack.Route, networkMTU uint32, pkt stack.PacketBufferPtr, transProto tcpip.TransportProtocolNumber, handler func(stack.PacketBufferPtr) tcpip.Error) (int, int, tcpip.Error) { networkHeader := header.IPv6(pkt.NetworkHeader().Slice()) // TODO(gvisor.dev/issue/3912): Once the Authentication or ESP Headers are @@ -780,7 +780,7 @@ func (e *endpoint) handleFragments(r *stack.Route, networkMTU uint32, pkt *stack } // WritePacket writes a packet to the given destination address and protocol. -func (e *endpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, pkt *stack.PacketBuffer) tcpip.Error { +func (e *endpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, pkt stack.PacketBufferPtr) tcpip.Error { dstAddr := r.RemoteAddress() if err := addIPHeader(r.LocalAddress(), dstAddr, pkt, params, nil /* extensionHeaders */); err != nil { return err @@ -814,7 +814,7 @@ func (e *endpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, return e.writePacket(r, pkt, params.Protocol, false /* headerIncluded */) } -func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer, protocol tcpip.TransportProtocolNumber, headerIncluded bool) tcpip.Error { +func (e *endpoint) writePacket(r *stack.Route, pkt stack.PacketBufferPtr, protocol tcpip.TransportProtocolNumber, headerIncluded bool) tcpip.Error { if r.Loop()&stack.PacketLoop != 0 { // If the packet was generated by the stack (not a raw/packet endpoint // where a packet may be written with the header included), then we can @@ -848,7 +848,7 @@ func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer, protocol // not by routers along a packet's delivery path. return &tcpip.ErrMessageTooLong{} } - sent, remain, err := e.handleFragments(r, networkMTU, pkt, protocol, func(fragPkt *stack.PacketBuffer) tcpip.Error { + sent, remain, err := e.handleFragments(r, networkMTU, pkt, protocol, func(fragPkt stack.PacketBufferPtr) tcpip.Error { // TODO(gvisor.dev/issue/3884): Evaluate whether we want to send each // fragment one by one using WritePacket() (current strategy) or if we // want to create a PacketBufferList from the fragments and feed it to @@ -870,7 +870,7 @@ func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer, protocol } // WriteHeaderIncludedPacket implements stack.NetworkEndpoint. -func (e *endpoint) WriteHeaderIncludedPacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { +func (e *endpoint) WriteHeaderIncludedPacket(r *stack.Route, pkt stack.PacketBufferPtr) tcpip.Error { // The packet already has an IP header, but there are a few required checks. h, ok := pkt.Data().PullUp(header.IPv6MinimumSize) if !ok { @@ -917,7 +917,7 @@ func validateAddressesForForwarding(h header.IPv6) ip.ForwardingError { // forwardUnicastPacket attempts to forward a unicast packet to its final // destination. -func (e *endpoint) forwardUnicastPacket(pkt *stack.PacketBuffer) ip.ForwardingError { +func (e *endpoint) forwardUnicastPacket(pkt stack.PacketBufferPtr) ip.ForwardingError { h := header.IPv6(pkt.NetworkHeader().Slice()) if err := validateAddressesForForwarding(h); err != nil { @@ -984,7 +984,7 @@ func (e *endpoint) forwardUnicastPacket(pkt *stack.PacketBuffer) ip.ForwardingEr // forwardPacketWithRoute emits the pkt using the provided route. // // This method should be invoked by the endpoint that received the pkt. -func (e *endpoint) forwardPacketWithRoute(route *stack.Route, pkt *stack.PacketBuffer) ip.ForwardingError { +func (e *endpoint) forwardPacketWithRoute(route *stack.Route, pkt stack.PacketBufferPtr) ip.ForwardingError { h := header.IPv6(pkt.NetworkHeader().Slice()) stk := e.protocol.stack @@ -1034,7 +1034,7 @@ func (e *endpoint) forwardPacketWithRoute(route *stack.Route, pkt *stack.PacketB // HandlePacket is called by the link layer when new ipv6 packets arrive for // this endpoint. -func (e *endpoint) HandlePacket(pkt *stack.PacketBuffer) { +func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { stats := e.stats.ip stats.PacketsReceived.Increment() @@ -1091,7 +1091,7 @@ func (e *endpoint) HandlePacket(pkt *stack.PacketBuffer) { // handleLocalPacket is like HandlePacket except it does not perform the // prerouting iptables hook or check for loopback traffic that originated from // outside of the netstack (i.e. martian loopback packets). -func (e *endpoint) handleLocalPacket(pkt *stack.PacketBuffer, canSkipRXChecksum bool) { +func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksum bool) { stats := e.stats.ip stats.PacketsReceived.Increment() @@ -1112,7 +1112,7 @@ func (e *endpoint) handleLocalPacket(pkt *stack.PacketBuffer, canSkipRXChecksum // // This method should be invoked for incoming multicast packets using the // endpoint that received the packet. -func (e *endpoint) forwardMulticastPacket(h header.IPv6, pkt *stack.PacketBuffer) ip.ForwardingError { +func (e *endpoint) forwardMulticastPacket(h header.IPv6, pkt stack.PacketBufferPtr) ip.ForwardingError { if err := validateAddressesForForwarding(h); err != nil { return err } @@ -1158,7 +1158,7 @@ func (e *endpoint) forwardMulticastPacket(h header.IPv6, pkt *stack.PacketBuffer // provided installedRoute. // // This method should be invoked by the endpoint that received the pkt. -func (e *endpoint) forwardValidatedMulticastPacket(pkt *stack.PacketBuffer, installedRoute *multicast.InstalledRoute) ip.ForwardingError { +func (e *endpoint) forwardValidatedMulticastPacket(pkt stack.PacketBufferPtr, installedRoute *multicast.InstalledRoute) ip.ForwardingError { // Per RFC 1812 section 5.2.1.3, // // Based on the IP source and destination addresses found in the datagram @@ -1191,7 +1191,7 @@ func (e *endpoint) forwardValidatedMulticastPacket(pkt *stack.PacketBuffer, inst // of the provided outgoing interface. // // This method should be invoked by the endpoint that received the pkt. -func (e *endpoint) forwardMulticastPacketForOutgoingInterface(pkt *stack.PacketBuffer, outgoingInterface stack.MulticastRouteOutgoingInterface) ip.ForwardingError { +func (e *endpoint) forwardMulticastPacketForOutgoingInterface(pkt stack.PacketBufferPtr, outgoingInterface stack.MulticastRouteOutgoingInterface) ip.ForwardingError { h := header.IPv6(pkt.NetworkHeader().Slice()) // Per RFC 1812 section 5.2.1.3, @@ -1248,7 +1248,7 @@ func (e *endpoint) handleForwardingError(err ip.ForwardingError) { stats.Forwarding.Errors.Increment() } -func (e *endpoint) handleValidatedPacket(h header.IPv6, pkt *stack.PacketBuffer, inNICName string) { +func (e *endpoint) handleValidatedPacket(h header.IPv6, pkt stack.PacketBufferPtr, inNICName string) { pkt.NICID = e.nic.ID() // Raw socket packets are delivered based solely on the transport protocol @@ -1307,7 +1307,7 @@ func (e *endpoint) handleValidatedPacket(h header.IPv6, pkt *stack.PacketBuffer, } } -func (e *endpoint) deliverPacketLocally(h header.IPv6, pkt *stack.PacketBuffer, inNICName string) { +func (e *endpoint) deliverPacketLocally(h header.IPv6, pkt stack.PacketBufferPtr, inNICName string) { stats := e.stats.ip // iptables filtering. All packets that reach here are intended for @@ -1323,7 +1323,7 @@ func (e *endpoint) deliverPacketLocally(h header.IPv6, pkt *stack.PacketBuffer, _ = e.processExtensionHeaders(h, pkt, false /* forwarding */) } -func (e *endpoint) processExtensionHeader(it *header.IPv6PayloadIterator, pkt **stack.PacketBuffer, h header.IPv6, routerAlert **header.IPv6RouterAlertOption, hasFragmentHeader *bool, forwarding bool) (bool, error) { +func (e *endpoint) processExtensionHeader(it *header.IPv6PayloadIterator, pkt *stack.PacketBufferPtr, h header.IPv6, routerAlert **header.IPv6RouterAlertOption, hasFragmentHeader *bool, forwarding bool) (bool, error) { stats := e.stats.ip dstAddr := h.DestinationAddress() // Keep track of the start of the previous header so we can report the @@ -1401,7 +1401,7 @@ func (e *endpoint) processExtensionHeader(it *header.IPv6PayloadIterator, pkt ** // processExtensionHeaders processes the extension headers in the given packet. // Returns an error if the processing of a header failed or if the packet should // be discarded. -func (e *endpoint) processExtensionHeaders(h header.IPv6, pkt *stack.PacketBuffer, forwarding bool) error { +func (e *endpoint) processExtensionHeaders(h header.IPv6, pkt stack.PacketBufferPtr, forwarding bool) error { // Create a VV to parse the packet. We don't plan to modify anything here. // vv consists of: // - Any IPv6 header bytes after the first 40 (i.e. extensions). @@ -1437,7 +1437,7 @@ func (e *endpoint) processExtensionHeaders(h header.IPv6, pkt *stack.PacketBuffe } } -func (e *endpoint) processIPv6RawPayloadHeader(extHdr *header.IPv6RawPayloadHeader, it *header.IPv6PayloadIterator, pkt *stack.PacketBuffer, routerAlert *header.IPv6RouterAlertOption, previousHeaderStart uint32, hasFragmentHeader bool) error { +func (e *endpoint) processIPv6RawPayloadHeader(extHdr *header.IPv6RawPayloadHeader, it *header.IPv6PayloadIterator, pkt stack.PacketBufferPtr, routerAlert *header.IPv6RouterAlertOption, previousHeaderStart uint32, hasFragmentHeader bool) error { stats := e.stats.ip // If the last header in the payload isn't a known IPv6 extension header, // handle it as if it is transport layer data.Ã¥ @@ -1520,7 +1520,7 @@ func (e *endpoint) processIPv6RawPayloadHeader(extHdr *header.IPv6RawPayloadHead } } -func (e *endpoint) processIPv6RoutingExtHeader(extHdr *header.IPv6RoutingExtHdr, it *header.IPv6PayloadIterator, pkt *stack.PacketBuffer) error { +func (e *endpoint) processIPv6RoutingExtHeader(extHdr *header.IPv6RoutingExtHdr, it *header.IPv6PayloadIterator, pkt stack.PacketBufferPtr) error { // As per RFC 8200 section 4.4, if a node encounters a routing header with // an unrecognized routing type value, with a non-zero Segments Left // value, the node must discard the packet and send an ICMP Parameter @@ -1543,7 +1543,7 @@ func (e *endpoint) processIPv6RoutingExtHeader(extHdr *header.IPv6RoutingExtHdr, return fmt.Errorf("found unrecognized routing type with non-zero segments left in header = %#v", extHdr) } -func (e *endpoint) processIPv6DestinationOptionsExtHdr(extHdr *header.IPv6DestinationOptionsExtHdr, it *header.IPv6PayloadIterator, pkt *stack.PacketBuffer, dstAddr tcpip.Address) error { +func (e *endpoint) processIPv6DestinationOptionsExtHdr(extHdr *header.IPv6DestinationOptionsExtHdr, it *header.IPv6PayloadIterator, pkt stack.PacketBufferPtr, dstAddr tcpip.Address) error { stats := e.stats.ip optsIt := extHdr.Iter() var uopt *header.IPv6UnknownExtHdrOption @@ -1606,7 +1606,7 @@ func (e *endpoint) processIPv6DestinationOptionsExtHdr(extHdr *header.IPv6Destin return nil } -func (e *endpoint) processIPv6HopByHopOptionsExtHdr(extHdr *header.IPv6HopByHopOptionsExtHdr, it *header.IPv6PayloadIterator, pkt *stack.PacketBuffer, dstAddr tcpip.Address, routerAlert **header.IPv6RouterAlertOption, previousHeaderStart uint32, forwarding bool) error { +func (e *endpoint) processIPv6HopByHopOptionsExtHdr(extHdr *header.IPv6HopByHopOptionsExtHdr, it *header.IPv6PayloadIterator, pkt stack.PacketBufferPtr, dstAddr tcpip.Address, routerAlert **header.IPv6RouterAlertOption, previousHeaderStart uint32, forwarding bool) error { stats := e.stats.ip // As per RFC 8200 section 4.1, the Hop By Hop extension header is // restricted to appear immediately after an IPv6 fixed header. @@ -1688,7 +1688,7 @@ func (e *endpoint) processIPv6HopByHopOptionsExtHdr(extHdr *header.IPv6HopByHopO return nil } -func (e *endpoint) processFragmentExtHdr(extHdr *header.IPv6FragmentExtHdr, it *header.IPv6PayloadIterator, pkt **stack.PacketBuffer, h header.IPv6) error { +func (e *endpoint) processFragmentExtHdr(extHdr *header.IPv6FragmentExtHdr, it *header.IPv6PayloadIterator, pkt *stack.PacketBufferPtr, h header.IPv6) error { stats := e.stats.ip fragmentFieldOffset := it.ParseOffset() @@ -2509,7 +2509,7 @@ func (p *protocol) DisableMulticastForwarding() { p.multicastRouteTable.RemoveAllInstalledRoutes() } -func (p *protocol) forwardPendingMulticastPacket(pkt *stack.PacketBuffer, installedRoute *multicast.InstalledRoute) { +func (p *protocol) forwardPendingMulticastPacket(pkt stack.PacketBufferPtr, installedRoute *multicast.InstalledRoute) { defer pkt.DecRef() // Attempt to forward the packet using the endpoint that it originally @@ -2538,7 +2538,7 @@ func (*protocol) Wait() {} // returns the parsed IP header. // // Returns true if the IP header was successfully parsed. -func (p *protocol) parseAndValidate(pkt *stack.PacketBuffer) (header.IPv6, bool) { +func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (header.IPv6, bool) { transProtoNum, hasTransportHdr, ok := p.Parse(pkt) if !ok { return nil, false @@ -2558,7 +2558,7 @@ func (p *protocol) parseAndValidate(pkt *stack.PacketBuffer) (header.IPv6, bool) return h, true } -func (p *protocol) parseTransport(pkt *stack.PacketBuffer, transProtoNum tcpip.TransportProtocolNumber) { +func (p *protocol) parseTransport(pkt stack.PacketBufferPtr, transProtoNum tcpip.TransportProtocolNumber) { if transProtoNum == header.ICMPv6ProtocolNumber { // The transport layer will handle transport layer parsing errors. _ = parse.ICMPv6(pkt) @@ -2576,7 +2576,7 @@ func (p *protocol) parseTransport(pkt *stack.PacketBuffer, transProtoNum tcpip.T } // Parse implements stack.NetworkProtocol. -func (*protocol) Parse(pkt *stack.PacketBuffer) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) { +func (*protocol) Parse(pkt stack.PacketBufferPtr) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) { proto, _, fragOffset, fragMore, ok := parse.IPv6(pkt) if !ok { return 0, false, false @@ -2598,7 +2598,7 @@ func (p *protocol) allowICMPReply(icmpType header.ICMPv6Type) bool { } // SendRejectionError implements stack.RejectIPv6WithHandler. -func (p *protocol) SendRejectionError(pkt *stack.PacketBuffer, rejectWith stack.RejectIPv6WithICMPType, inputHook bool) tcpip.Error { +func (p *protocol) SendRejectionError(pkt stack.PacketBufferPtr, rejectWith stack.RejectIPv6WithICMPType, inputHook bool) tcpip.Error { switch rejectWith { case stack.RejectIPv6WithICMPNoRoute: return p.returnError(&icmpReasonNetUnreachable{}, pkt, inputHook) @@ -2742,7 +2742,7 @@ func NewProtocol(s *stack.Stack) stack.NetworkProtocol { return NewProtocolWithOptions(Options{})(s) } -func calculateFragmentReserve(pkt *stack.PacketBuffer) int { +func calculateFragmentReserve(pkt stack.PacketBufferPtr) int { return pkt.AvailableHeaderBytes() + len(pkt.NetworkHeader().Slice()) + header.IPv6FragmentHeaderSize } @@ -2769,7 +2769,7 @@ func hashRoute(r *stack.Route, hashIV uint32) uint32 { return h.Sum32() } -func buildNextFragment(pf *fragmentation.PacketFragmenter, originalIPHeaders header.IPv6, transportProto tcpip.TransportProtocolNumber, id uint32) (*stack.PacketBuffer, bool) { +func buildNextFragment(pf *fragmentation.PacketFragmenter, originalIPHeaders header.IPv6, transportProto tcpip.TransportProtocolNumber, id uint32) (stack.PacketBufferPtr, bool) { fragPkt, offset, copied, more := pf.BuildNextFragment() fragPkt.NetworkProtocolNumber = ProtocolNumber diff --git a/pkg/tcpip/network/ipv6/ipv6_test.go b/pkg/tcpip/network/ipv6/ipv6_test.go index 8df942d89..38de76caa 100644 --- a/pkg/tcpip/network/ipv6/ipv6_test.go +++ b/pkg/tcpip/network/ipv6/ipv6_test.go @@ -166,7 +166,7 @@ func testReceiveUDP(t *testing.T, s *stack.Stack, e *channel.Endpoint, src, dst } } -func compareFragments(packets []*stack.PacketBuffer, sourcePacket *stack.PacketBuffer, mtu uint32, wantFragments []fragmentInfo, proto tcpip.TransportProtocolNumber) error { +func compareFragments(packets []stack.PacketBufferPtr, sourcePacket stack.PacketBufferPtr, mtu uint32, wantFragments []fragmentInfo, proto tcpip.TransportProtocolNumber) error { // sourcePacket does not have its IP Header populated. Let's copy the one // from the first fragment. source := header.IPv6(packets[0].NetworkHeader().Slice()) @@ -1032,7 +1032,7 @@ func TestReceiveIPv6ExtHdrs(t *testing.T) { } if !test.expectICMP { - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("unexpected packet received: %#v", p) } return @@ -1040,7 +1040,7 @@ func TestReceiveIPv6ExtHdrs(t *testing.T) { // ICMP required. p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("expected packet wasn't written out") } defer p.DecRef() @@ -2073,12 +2073,12 @@ func TestInvalidIPv6Fragments(t *testing.T) { reply := e.Read() if !test.expectICMP { - if reply != nil { + if !reply.IsNil() { t.Fatalf("unexpected ICMP error message received: %#v", reply) } return } - if reply == nil { + if reply.IsNil() { t.Fatal("expected ICMP error message missing") } @@ -2327,12 +2327,12 @@ func TestFragmentReassemblyTimeout(t *testing.T) { reply := e.Read() if !test.expectICMP { - if reply != nil { + if !reply.IsNil() { t.Fatalf("unexpected ICMP error message received: %#v", reply) } return } - if reply == nil { + if reply.IsNil() { t.Fatal("expected ICMP error message missing") } if firstFragmentSent == nil { @@ -2545,7 +2545,7 @@ func (*limitedMatcher) Name() string { } // Match implements Matcher.Match. -func (lm *limitedMatcher) Match(stack.Hook, *stack.PacketBuffer, string, string) (bool, bool) { +func (lm *limitedMatcher) Match(stack.Hook, stack.PacketBufferPtr, string, string) (bool, bool) { if lm.limit == 0 { return true, false } @@ -3152,7 +3152,7 @@ func TestForwarding(t *testing.T) { } if test.expectedICMPError != nil { - if reply == nil { + if reply.IsNil() { t.Fatalf("Expected ICMP packet type %d through incoming NIC", test.expectedICMPError.icmpType) } @@ -3186,13 +3186,13 @@ func TestForwarding(t *testing.T) { if n := outgoingEndpoint.Drain(); n != 0 { t.Fatalf("e2.Drain() = %d, want = 0", n) } - } else if reply != nil { + } else if !reply.IsNil() { t.Fatalf("Expected no ICMP packet through incoming NIC, instead found: %#v", reply) } reply = outgoingEndpoint.Read() if test.expectPacketForwarded { - if reply == nil { + if reply.IsNil() { t.Fatal("Expected ICMP Echo Request packet through outgoing NIC") } @@ -3214,7 +3214,7 @@ func TestForwarding(t *testing.T) { if n := incomingEndpoint.Drain(); n != 0 { t.Fatalf("e1.Drain() = %d, want = 0", n) } - } else if reply != nil { + } else if !reply.IsNil() { t.Fatalf("Expected no ICMP Echo packet through outgoing NIC, instead found: %#v", reply) } @@ -3483,7 +3483,7 @@ func TestMulticastForwarding(t *testing.T) { } if test.expectedICMPError != nil { - if reply == nil { + if reply.IsNil() { t.Fatalf("Expected ICMP packet type %d through incoming NIC", test.expectedICMPError.icmpType) } @@ -3517,13 +3517,13 @@ func TestMulticastForwarding(t *testing.T) { if n := outgoingEndpoint.Drain(); n != 0 { t.Fatalf("e2.Drain() = %d, want = 0", n) } - } else if reply != nil { + } else if !reply.IsNil() { t.Fatalf("Expected no ICMP packet through incoming NIC, instead found: %#v", reply) } reply = outgoingEndpoint.Read() if test.expectPacketForwarded { - if reply == nil { + if reply.IsNil() { t.Fatal("Expected ICMP Echo Request packet through outgoing NIC") } @@ -3545,7 +3545,7 @@ func TestMulticastForwarding(t *testing.T) { if n := incomingEndpoint.Drain(); n != 0 { t.Fatalf("e1.Drain() = %d, want = 0", n) } - } else if reply != nil { + } else if !reply.IsNil() { t.Fatalf("Expected no ICMP Echo packet through outgoing NIC, instead found: %#v", reply) } @@ -3664,7 +3664,7 @@ func TestIcmpRateLimit(t *testing.T) { }, check: func(t *testing.T, e *channel.Endpoint, round int) { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("expected echo response, no packet read in endpoint in round %d", round) } defer p.DecRef() @@ -3712,13 +3712,13 @@ func TestIcmpRateLimit(t *testing.T) { check: func(t *testing.T, e *channel.Endpoint, round int) { p := e.Read() if round >= icmpBurst { - if p != nil { + if !p.IsNil() { t.Errorf("got packet %x in round %d, expected ICMP rate limit to stop it", p.Data().AsRange().ToSlice(), round) p.DecRef() } return } - if p == nil { + if p.IsNil() { t.Fatalf("expected unreachable in round %d, no packet read in endpoint", round) } payload := stack.PayloadSince(p.NetworkHeader()) diff --git a/pkg/tcpip/network/ipv6/mld_test.go b/pkg/tcpip/network/ipv6/mld_test.go index c4c8e279d..c57f6cd03 100644 --- a/pkg/tcpip/network/ipv6/mld_test.go +++ b/pkg/tcpip/network/ipv6/mld_test.go @@ -104,7 +104,7 @@ func TestIPv6JoinLeaveSolicitedNodeAddressPerformsMLD(t *testing.T) { if err := s.AddProtocolAddress(nicID, protocolAddr, stack.AddressProperties{}); err != nil { t.Fatalf("AddProtocolAddress(%d, %+v, {}): %s", nicID, protocolAddr, err) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { validateMLDPacket(t, stack.PayloadSince(p.NetworkHeader()), linkLocalAddr, linkLocalAddrSNMC, header.ICMPv6MulticastListenerReport, linkLocalAddrSNMC) @@ -117,7 +117,7 @@ func TestIPv6JoinLeaveSolicitedNodeAddressPerformsMLD(t *testing.T) { if err := s.RemoveAddress(nicID, linkLocalAddr); err != nil { t.Fatalf("RemoveAddress(%d, %s) = %s", nicID, linkLocalAddr, err) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a done message to be sent") } else { validateMLDPacket(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, header.IPv6AllRoutersLinkLocalMulticastAddress, header.ICMPv6MulticastListenerDone, linkLocalAddrSNMC) @@ -196,7 +196,7 @@ func TestSendQueuedMLDReports(t *testing.T) { resolveDAD := func(addr, snmc tcpip.Address) { clock.Advance(dadResolutionTime) - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected DAD packet") } else { payload := stack.PayloadSince(p.NetworkHeader()) @@ -233,14 +233,14 @@ func TestSendQueuedMLDReports(t *testing.T) { if got := reportStat.Value(); got != reportCounter { t.Errorf("got reportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Errorf("expected MLD report for %s", globalMulticastAddr) } else { validateMLDPacket(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, globalMulticastAddr, header.ICMPv6MulticastListenerReport, globalMulticastAddr) p.DecRef() } clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Errorf("got unexpected packet = %#v", p) p.DecRef() } @@ -267,7 +267,7 @@ func TestSendQueuedMLDReports(t *testing.T) { if got := reportStat.Value(); got != reportCounter { t.Errorf("got reportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Errorf("expected MLD report for %s", globalAddrSNMC) } else { validateMLDPacket(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, globalAddrSNMC, header.ICMPv6MulticastListenerReport, globalAddrSNMC) @@ -288,7 +288,7 @@ func TestSendQueuedMLDReports(t *testing.T) { if got := doneStat.Value(); got != doneCounter { t.Errorf("got doneStat.Value() = %d, want = %d", got, doneCounter) } - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Errorf("got unexpected packet = %#v", p) p.DecRef() } @@ -310,7 +310,7 @@ func TestSendQueuedMLDReports(t *testing.T) { if got := reportStat.Value(); got != reportCounter { t.Errorf("got reportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Errorf("expected MLD report for %s", linkLocalAddrSNMC) } else { validateMLDPacket(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, linkLocalAddrSNMC, header.ICMPv6MulticastListenerReport, linkLocalAddrSNMC) @@ -336,7 +336,7 @@ func TestSendQueuedMLDReports(t *testing.T) { } for range addrs { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("expected MLD report for %s and %s; addrs = %#v", globalMulticastAddr, linkLocalAddrSNMC, addrs) } @@ -359,7 +359,7 @@ func TestSendQueuedMLDReports(t *testing.T) { // Should not send any more reports. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Errorf("got unexpected packet = %#v", p) p.DecRef() } @@ -629,7 +629,7 @@ func TestMLDSkipProtocol(t *testing.T) { if err := s.AddProtocolAddress(nicID, protocolAddr, stack.AddressProperties{}); err != nil { t.Fatalf("AddProtocolAddress(%d, %+v, {}): %s", nicID, protocolAddr, err) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { validateMLDPacket(t, stack.PayloadSince(p.NetworkHeader()), linkLocalAddr, linkLocalAddrSNMC, header.ICMPv6MulticastListenerReport, linkLocalAddrSNMC) @@ -646,14 +646,14 @@ func TestMLDSkipProtocol(t *testing.T) { } if !test.expectReport { - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("got e.Read() = (%#v, true), want = (_, false)", p) } return } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { validateMLDPacket(t, stack.PayloadSince(p.NetworkHeader()), linkLocalAddr, test.group, header.ICMPv6MulticastListenerReport, test.group) diff --git a/pkg/tcpip/network/ipv6/ndp_test.go b/pkg/tcpip/network/ipv6/ndp_test.go index fa29d1b2e..b47182528 100644 --- a/pkg/tcpip/network/ipv6/ndp_test.go +++ b/pkg/tcpip/network/ipv6/ndp_test.go @@ -461,7 +461,7 @@ func TestNeighborSolicitationResponse(t *testing.T) { t.Fatalf("got invalid = %d, want = 1", got) } - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("unexpected response to an invalid NS = %+v", p) } @@ -476,7 +476,7 @@ func TestNeighborSolicitationResponse(t *testing.T) { if test.performsLinkResolution { c.clock.RunImmediatelyScheduledJobs() p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("expected an NDP NS response") } @@ -538,7 +538,7 @@ func TestNeighborSolicitationResponse(t *testing.T) { c.clock.RunImmediatelyScheduledJobs() p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("expected an NDP NA response") } defer p.DecRef() @@ -1305,7 +1305,7 @@ func TestCheckDuplicateAddress(t *testing.T) { checkDADMsg := func() { clock.RunImmediatelyScheduledJobs() p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("expected %d-th DAD message", dadPacketsSent) } defer p.DecRef() @@ -1385,7 +1385,7 @@ func TestCheckDuplicateAddress(t *testing.T) { } // Should have no more packets. - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Errorf("got unexpected packet = %#v", p) } } diff --git a/pkg/tcpip/network/multicast_group_test.go b/pkg/tcpip/network/multicast_group_test.go index 28e539c1c..0b9a2bbfb 100644 --- a/pkg/tcpip/network/multicast_group_test.go +++ b/pkg/tcpip/network/multicast_group_test.go @@ -80,7 +80,7 @@ var ( // validateMLDPacket checks that a passed PacketInfo is an IPv6 MLD packet // sent to the provided address with the passed fields set. -func validateMLDPacket(t *testing.T, p *stack.PacketBuffer, remoteAddress tcpip.Address, mldType uint8, maxRespTime byte, groupAddress tcpip.Address) { +func validateMLDPacket(t *testing.T, p stack.PacketBufferPtr, remoteAddress tcpip.Address, mldType uint8, maxRespTime byte, groupAddress tcpip.Address) { t.Helper() payload := stack.PayloadSince(p.NetworkHeader()) @@ -102,7 +102,7 @@ func validateMLDPacket(t *testing.T, p *stack.PacketBuffer, remoteAddress tcpip. // validateIGMPPacket checks that a passed PacketInfo is an IPv4 IGMP packet // sent to the provided address with the passed fields set. -func validateIGMPPacket(t *testing.T, p *stack.PacketBuffer, remoteAddress tcpip.Address, igmpType uint8, maxRespTime byte, groupAddress tcpip.Address) { +func validateIGMPPacket(t *testing.T, p stack.PacketBufferPtr, remoteAddress tcpip.Address, igmpType uint8, maxRespTime byte, groupAddress tcpip.Address) { t.Helper() payload := stack.PayloadSince(p.NetworkHeader()) @@ -207,7 +207,7 @@ func checkInitialIPv6Groups(t *testing.T, e *channel.Endpoint, s *stack.Stack, c if got := stats.MulticastListenerReport.Value(); got != reportCounter { t.Errorf("got stats.MulticastListenerReport.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { validateMLDPacket(t, p, ipv6AddrSNMC, mldReport, 0, ipv6AddrSNMC) @@ -223,7 +223,7 @@ func checkInitialIPv6Groups(t *testing.T, e *channel.Endpoint, s *stack.Stack, c if got := stats.MulticastListenerDone.Value(); got != leaveCounter { t.Errorf("got stats.MulticastListenerDone.Value() = %d, want = %d", got, leaveCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { validateMLDPacket(t, p, header.IPv6AllRoutersLinkLocalMulticastAddress, mldDone, 0, ipv6AddrSNMC) @@ -232,7 +232,7 @@ func checkInitialIPv6Groups(t *testing.T, e *channel.Endpoint, s *stack.Stack, c // Should not send any more packets. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } @@ -367,7 +367,7 @@ func TestMGPDisabled(t *testing.T) { t.Fatalf("got sentReportStat.Value() = %d, want = 0", got) } clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet, stack with disabled MGP sent packet = %#v", p) } @@ -380,7 +380,7 @@ func TestMGPDisabled(t *testing.T) { t.Fatalf("got sentReportStat.Value() = %d, want = 0", got) } clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet, stack with disabled IGMP sent packet = %#v", p) } @@ -391,7 +391,7 @@ func TestMGPDisabled(t *testing.T) { t.Fatalf("got receivedQueryStat(_).Value() = %d, want = 1", got) } clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet, stack with disabled IGMP sent packet = %+v", p) } }) @@ -502,7 +502,7 @@ func TestMGPJoinGroup(t *testing.T) { maxUnsolicitedResponseDelay time.Duration sentReportStat func(*stack.Stack) *tcpip.StatCounter receivedQueryStat func(*stack.Stack) *tcpip.StatCounter - validateReport func(*testing.T, *stack.PacketBuffer) + validateReport func(*testing.T, stack.PacketBufferPtr) checkInitialGroups func(*testing.T, *channel.Endpoint, *stack.Stack, *faketime.ManualClock) (uint64, uint64) }{ { @@ -516,7 +516,7 @@ func TestMGPJoinGroup(t *testing.T) { receivedQueryStat: func(s *stack.Stack) *tcpip.StatCounter { return s.Stats().IGMP.PacketsReceived.MembershipQuery }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateIGMPPacket(t, p, ipv4MulticastAddr1, igmpv2MembershipReport, 0, ipv4MulticastAddr1) @@ -533,7 +533,7 @@ func TestMGPJoinGroup(t *testing.T) { receivedQueryStat: func(s *stack.Stack) *tcpip.StatCounter { return s.Stats().ICMP.V6.PacketsReceived.MulticastListenerQuery }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateMLDPacket(t, p, ipv6MulticastAddr1, mldReport, 0, ipv6MulticastAddr1) @@ -563,7 +563,7 @@ func TestMGPJoinGroup(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { test.validateReport(t, p) @@ -576,7 +576,7 @@ func TestMGPJoinGroup(t *testing.T) { // Verify the second report is sent by the maximum unsolicited response // interval. p := e.Read() - if p != nil { + if !p.IsNil() { t.Fatalf("sent unexpected packet, expected report only after advancing the clock = %#v", p) } clock.Advance(test.maxUnsolicitedResponseDelay) @@ -584,7 +584,7 @@ func TestMGPJoinGroup(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { test.validateReport(t, p) @@ -593,7 +593,7 @@ func TestMGPJoinGroup(t *testing.T) { // Should not send any more packets. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } }) @@ -609,8 +609,8 @@ func TestMGPLeaveGroup(t *testing.T) { multicastAddr tcpip.Address sentReportStat func(*stack.Stack) *tcpip.StatCounter sentLeaveStat func(*stack.Stack) *tcpip.StatCounter - validateReport func(*testing.T, *stack.PacketBuffer) - validateLeave func(*testing.T, *stack.PacketBuffer) + validateReport func(*testing.T, stack.PacketBufferPtr) + validateLeave func(*testing.T, stack.PacketBufferPtr) checkInitialGroups func(*testing.T, *channel.Endpoint, *stack.Stack, *faketime.ManualClock) (uint64, uint64) }{ { @@ -623,12 +623,12 @@ func TestMGPLeaveGroup(t *testing.T) { sentLeaveStat: func(s *stack.Stack) *tcpip.StatCounter { return s.Stats().IGMP.PacketsSent.LeaveGroup }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateIGMPPacket(t, p, ipv4MulticastAddr1, igmpv2MembershipReport, 0, ipv4MulticastAddr1) }, - validateLeave: func(t *testing.T, p *stack.PacketBuffer) { + validateLeave: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateIGMPPacket(t, p, header.IPv4AllRoutersGroup, igmpLeaveGroup, 0, ipv4MulticastAddr1) @@ -644,12 +644,12 @@ func TestMGPLeaveGroup(t *testing.T) { sentLeaveStat: func(s *stack.Stack) *tcpip.StatCounter { return s.Stats().ICMP.V6.PacketsSent.MulticastListenerDone }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateMLDPacket(t, p, ipv6MulticastAddr1, mldReport, 0, ipv6MulticastAddr1) }, - validateLeave: func(t *testing.T, p *stack.PacketBuffer) { + validateLeave: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateMLDPacket(t, p, header.IPv6AllRoutersLinkLocalMulticastAddress, mldDone, 0, ipv6MulticastAddr1) @@ -677,7 +677,7 @@ func TestMGPLeaveGroup(t *testing.T) { if got := test.sentReportStat(s).Value(); got != reportCounter { t.Errorf("got sentReportStat(_).Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { test.validateReport(t, p) @@ -695,7 +695,7 @@ func TestMGPLeaveGroup(t *testing.T) { if got := test.sentLeaveStat(s).Value(); got != leaveCounter { t.Fatalf("got sentLeaveStat(_).Value() = %d, want = %d", got, leaveCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a leave message to be sent") } else { test.validateLeave(t, p) @@ -704,7 +704,7 @@ func TestMGPLeaveGroup(t *testing.T) { // Should not send any more packets. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } }) @@ -722,7 +722,7 @@ func TestMGPQueryMessages(t *testing.T) { sentReportStat func(*stack.Stack) *tcpip.StatCounter receivedQueryStat func(*stack.Stack) *tcpip.StatCounter rxQuery func(*channel.Endpoint, uint8, tcpip.Address) - validateReport func(*testing.T, *stack.PacketBuffer) + validateReport func(*testing.T, stack.PacketBufferPtr) maxRespTimeToDuration func(uint8) time.Duration checkInitialGroups func(*testing.T, *channel.Endpoint, *stack.Stack, *faketime.ManualClock) (uint64, uint64) }{ @@ -740,7 +740,7 @@ func TestMGPQueryMessages(t *testing.T) { rxQuery: func(e *channel.Endpoint, maxRespTime uint8, groupAddress tcpip.Address) { createAndInjectIGMPPacket(e, igmpMembershipQuery, maxRespTime, groupAddress) }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateIGMPPacket(t, p, ipv4MulticastAddr1, igmpv2MembershipReport, 0, ipv4MulticastAddr1) @@ -761,7 +761,7 @@ func TestMGPQueryMessages(t *testing.T) { rxQuery: func(e *channel.Endpoint, maxRespTime uint8, groupAddress tcpip.Address) { createAndInjectMLDPacket(e, mldQuery, maxRespTime, groupAddress) }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateMLDPacket(t, p, ipv6MulticastAddr1, mldReport, 0, ipv6MulticastAddr1) @@ -822,7 +822,7 @@ func TestMGPQueryMessages(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("(i=%d) got sentReportStat.Value() = %d, want = %d", i, got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatalf("expected %d-th report message to be sent", i) } else { test.validateReport(t, p) @@ -836,7 +836,7 @@ func TestMGPQueryMessages(t *testing.T) { // Should not send any more packets until a query. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } @@ -845,7 +845,7 @@ func TestMGPQueryMessages(t *testing.T) { // targeted at the host. const maxRespTime = 100 test.rxQuery(e, maxRespTime, subTest.multicastAddr) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } @@ -855,7 +855,7 @@ func TestMGPQueryMessages(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { test.validateReport(t, p) @@ -865,7 +865,7 @@ func TestMGPQueryMessages(t *testing.T) { // Should not send any more packets. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } }) @@ -884,7 +884,7 @@ func TestMGPReportMessages(t *testing.T) { sentReportStat func(*stack.Stack) *tcpip.StatCounter sentLeaveStat func(*stack.Stack) *tcpip.StatCounter rxReport func(*channel.Endpoint) - validateReport func(*testing.T, *stack.PacketBuffer) + validateReport func(*testing.T, stack.PacketBufferPtr) maxRespTimeToDuration func(uint8) time.Duration checkInitialGroups func(*testing.T, *channel.Endpoint, *stack.Stack, *faketime.ManualClock) (uint64, uint64) }{ @@ -901,7 +901,7 @@ func TestMGPReportMessages(t *testing.T) { rxReport: func(e *channel.Endpoint) { createAndInjectIGMPPacket(e, igmpv2MembershipReport, 0, ipv4MulticastAddr1) }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateIGMPPacket(t, p, ipv4MulticastAddr1, igmpv2MembershipReport, 0, ipv4MulticastAddr1) @@ -921,7 +921,7 @@ func TestMGPReportMessages(t *testing.T) { rxReport: func(e *channel.Endpoint) { createAndInjectMLDPacket(e, mldReport, 0, ipv6MulticastAddr1) }, - validateReport: func(t *testing.T, p *stack.PacketBuffer) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr) { t.Helper() validateMLDPacket(t, p, ipv6MulticastAddr1, mldReport, 0, ipv6MulticastAddr1) @@ -953,7 +953,7 @@ func TestMGPReportMessages(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { test.validateReport(t, p) @@ -970,7 +970,7 @@ func TestMGPReportMessages(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Errorf("sent unexpected packet = %#v", p) } if t.Failed() { @@ -989,7 +989,7 @@ func TestMGPReportMessages(t *testing.T) { // Should not send any more packets. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } }) @@ -1005,9 +1005,9 @@ func TestMGPWithNICLifecycle(t *testing.T) { maxUnsolicitedResponseDelay time.Duration sentReportStat func(*stack.Stack) *tcpip.StatCounter sentLeaveStat func(*stack.Stack) *tcpip.StatCounter - validateReport func(*testing.T, *stack.PacketBuffer, tcpip.Address) - validateLeave func(*testing.T, *stack.PacketBuffer, tcpip.Address) - getAndCheckGroupAddress func(*testing.T, map[tcpip.Address]bool, *stack.PacketBuffer) tcpip.Address + validateReport func(*testing.T, stack.PacketBufferPtr, tcpip.Address) + validateLeave func(*testing.T, stack.PacketBufferPtr, tcpip.Address) + getAndCheckGroupAddress func(*testing.T, map[tcpip.Address]bool, stack.PacketBufferPtr) tcpip.Address checkInitialGroups func(*testing.T, *channel.Endpoint, *stack.Stack, *faketime.ManualClock) (uint64, uint64) }{ { @@ -1022,17 +1022,17 @@ func TestMGPWithNICLifecycle(t *testing.T) { sentLeaveStat: func(s *stack.Stack) *tcpip.StatCounter { return s.Stats().IGMP.PacketsSent.LeaveGroup }, - validateReport: func(t *testing.T, p *stack.PacketBuffer, addr tcpip.Address) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr, addr tcpip.Address) { t.Helper() validateIGMPPacket(t, p, addr, igmpv2MembershipReport, 0, addr) }, - validateLeave: func(t *testing.T, p *stack.PacketBuffer, addr tcpip.Address) { + validateLeave: func(t *testing.T, p stack.PacketBufferPtr, addr tcpip.Address) { t.Helper() validateIGMPPacket(t, p, header.IPv4AllRoutersGroup, igmpLeaveGroup, 0, addr) }, - getAndCheckGroupAddress: func(t *testing.T, seen map[tcpip.Address]bool, p *stack.PacketBuffer) tcpip.Address { + getAndCheckGroupAddress: func(t *testing.T, seen map[tcpip.Address]bool, p stack.PacketBufferPtr) tcpip.Address { t.Helper() payload := stack.PayloadSince(p.NetworkHeader()) @@ -1065,17 +1065,17 @@ func TestMGPWithNICLifecycle(t *testing.T) { sentLeaveStat: func(s *stack.Stack) *tcpip.StatCounter { return s.Stats().ICMP.V6.PacketsSent.MulticastListenerDone }, - validateReport: func(t *testing.T, p *stack.PacketBuffer, addr tcpip.Address) { + validateReport: func(t *testing.T, p stack.PacketBufferPtr, addr tcpip.Address) { t.Helper() validateMLDPacket(t, p, addr, mldReport, 0, addr) }, - validateLeave: func(t *testing.T, p *stack.PacketBuffer, addr tcpip.Address) { + validateLeave: func(t *testing.T, p stack.PacketBufferPtr, addr tcpip.Address) { t.Helper() validateMLDPacket(t, p, header.IPv6AllRoutersLinkLocalMulticastAddress, mldDone, 0, addr) }, - getAndCheckGroupAddress: func(t *testing.T, seen map[tcpip.Address]bool, p *stack.PacketBuffer) tcpip.Address { + getAndCheckGroupAddress: func(t *testing.T, seen map[tcpip.Address]bool, p stack.PacketBufferPtr) tcpip.Address { t.Helper() payload := stack.PayloadSince(p.NetworkHeader()) defer payload.Release() @@ -1145,7 +1145,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatalf("expected a report message to be sent for %s", a) } else { test.validateReport(t, p, a) @@ -1174,7 +1174,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { for i := range test.multicastAddrs { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("expected (%d-th) leave message to be sent", i) } @@ -1202,7 +1202,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { for i := range test.multicastAddrs { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatalf("expected (%d-th) report message to be sent", i) } @@ -1223,7 +1223,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { t.Errorf("got sentLeaveStat.Value() = %d, want = %d", got, leaveCounter) } for i := range test.multicastAddrs { - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatalf("expected (%d-th) leave message to be sent", i) } else { p.DecRef() @@ -1236,7 +1236,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { if got := sentLeaveStat.Value(); got != leaveCounter { t.Errorf("got sentLeaveStat.Value() = %d, want = %d", got, leaveCounter) } - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("leaving group %s on disabled NIC sent unexpected packet = %#v", a, p) } } @@ -1246,7 +1246,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("joining group %s on disabled NIC sent unexpected packet = %#v", test.finalMulticastAddr, p) } @@ -1259,7 +1259,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { test.validateReport(t, p, test.finalMulticastAddr) @@ -1271,7 +1271,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { if got := sentReportStat.Value(); got != reportCounter { t.Errorf("got sentReportStat.Value() = %d, want = %d", got, reportCounter) } - if p := e.Read(); p == nil { + if p := e.Read(); p.IsNil() { t.Fatal("expected a report message to be sent") } else { test.validateReport(t, p, test.finalMulticastAddr) @@ -1280,7 +1280,7 @@ func TestMGPWithNICLifecycle(t *testing.T) { // Should not send any more packets. clock.Advance(time.Hour) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("sent unexpected packet = %#v", p) } }) diff --git a/pkg/tcpip/stack/conntrack.go b/pkg/tcpip/stack/conntrack.go index cbd7956e8..b70035784 100644 --- a/pkg/tcpip/stack/conntrack.go +++ b/pkg/tcpip/stack/conntrack.go @@ -177,7 +177,7 @@ func (cn *conn) timedOut(now tcpip.MonotonicTime) bool { } // update the connection tracking state. -func (cn *conn) update(pkt *PacketBuffer, reply bool) { +func (cn *conn) update(pkt PacketBufferPtr, reply bool) { cn.stateMu.Lock() defer cn.stateMu.Unlock() @@ -269,7 +269,7 @@ func v6NetAndTransHdr(icmpPayload []byte, minTransHdrLen int) (header.Network, [ return netHdr, transHdr[:minTransHdrLen] } -func getEmbeddedNetAndTransHeaders(pkt *PacketBuffer, netHdrLength int, getNetAndTransHdr netAndTransHeadersFunc, transProto tcpip.TransportProtocolNumber) (header.Network, header.ChecksummableTransport, bool) { +func getEmbeddedNetAndTransHeaders(pkt PacketBufferPtr, netHdrLength int, getNetAndTransHdr netAndTransHeadersFunc, transProto tcpip.TransportProtocolNumber) (header.Network, header.ChecksummableTransport, bool) { switch transProto { case header.TCPProtocolNumber: if netAndTransHeader, ok := pkt.Data().PullUp(netHdrLength + header.TCPMinimumSize); ok { @@ -285,7 +285,7 @@ func getEmbeddedNetAndTransHeaders(pkt *PacketBuffer, netHdrLength int, getNetAn return nil, nil, false } -func getHeaders(pkt *PacketBuffer) (netHdr header.Network, transHdr header.Transport, isICMPError bool, ok bool) { +func getHeaders(pkt PacketBufferPtr) (netHdr header.Network, transHdr header.Transport, isICMPError bool, ok bool) { switch pkt.TransportProtocolNumber { case header.TCPProtocolNumber: if tcpHeader := header.TCP(pkt.TransportHeader().Slice()); len(tcpHeader) >= header.TCPMinimumSize { @@ -373,7 +373,7 @@ func getTupleIDForRegularPacket(netHdr header.Network, netProto tcpip.NetworkPro } } -func getTupleIDForPacketInICMPError(pkt *PacketBuffer, getNetAndTransHdr netAndTransHeadersFunc, netProto tcpip.NetworkProtocolNumber, netLen int, transProto tcpip.TransportProtocolNumber) (tupleID, bool) { +func getTupleIDForPacketInICMPError(pkt PacketBufferPtr, getNetAndTransHdr netAndTransHeadersFunc, netProto tcpip.NetworkProtocolNumber, netLen int, transProto tcpip.TransportProtocolNumber) (tupleID, bool) { if netHdr, transHdr, ok := getEmbeddedNetAndTransHeaders(pkt, netLen, getNetAndTransHdr, transProto); ok { return tupleID{ srcAddr: netHdr.DestinationAddress(), @@ -396,7 +396,7 @@ const ( getTupleIDOKAndDontAllowNewConn ) -func getTupleIDForEchoPacket(pkt *PacketBuffer, ident uint16, request bool) tupleID { +func getTupleIDForEchoPacket(pkt PacketBufferPtr, ident uint16, request bool) tupleID { netHdr := pkt.Network() tid := tupleID{ srcAddr: netHdr.SourceAddress(), @@ -414,7 +414,7 @@ func getTupleIDForEchoPacket(pkt *PacketBuffer, ident uint16, request bool) tupl return tid } -func getTupleID(pkt *PacketBuffer) (tupleID, getTupleIDDisposition) { +func getTupleID(pkt PacketBufferPtr) (tupleID, getTupleIDDisposition) { switch pkt.TransportProtocolNumber { case header.TCPProtocolNumber: if transHeader := header.TCP(pkt.TransportHeader().Slice()); len(transHeader) >= header.TCPMinimumSize { @@ -504,7 +504,7 @@ func (ct *ConnTrack) init() { // // If the packet's protocol is trackable, the connection's state is updated to // match the contents of the packet. -func (ct *ConnTrack) getConnAndUpdate(pkt *PacketBuffer, skipChecksumValidation bool) *tuple { +func (ct *ConnTrack) getConnAndUpdate(pkt PacketBufferPtr, skipChecksumValidation bool) *tuple { // Get or (maybe) create a connection. t := func() *tuple { var allowNewConn bool @@ -729,7 +729,7 @@ type portOrIdentRange struct { // // Generally, only the first packet of a connection reaches this method; other // packets will be manipulated without needing to modify the connection. -func (cn *conn) performNAT(pkt *PacketBuffer, hook Hook, r *Route, portsOrIdents portOrIdentRange, natAddress tcpip.Address, dnat bool) { +func (cn *conn) performNAT(pkt PacketBufferPtr, hook Hook, r *Route, portsOrIdents portOrIdentRange, natAddress tcpip.Address, dnat bool) { lastPortOrIdent := func() uint16 { lastPortOrIdent := uint32(portsOrIdents.start) + portsOrIdents.size - 1 if lastPortOrIdent > math.MaxUint16 { @@ -830,7 +830,7 @@ func (cn *conn) performNAT(pkt *PacketBuffer, hook Hook, r *Route, portsOrIdents // has had NAT performed on it. // // Returns true if the packet can skip the NAT table. -func (cn *conn) handlePacket(pkt *PacketBuffer, hook Hook, rt *Route) bool { +func (cn *conn) handlePacket(pkt PacketBufferPtr, hook Hook, rt *Route) bool { netHdr, transHdr, isICMPError, ok := getHeaders(pkt) if !ok { return false diff --git a/pkg/tcpip/stack/conntrack_test.go b/pkg/tcpip/stack/conntrack_test.go index c7993193e..1da184ca3 100644 --- a/pkg/tcpip/stack/conntrack_test.go +++ b/pkg/tcpip/stack/conntrack_test.go @@ -280,7 +280,7 @@ type genTCPOpts struct { } // genTCPPacket returns an initialized IPv4 TCP packet. -func genTCPPacket(opts genTCPOpts) *PacketBuffer { +func genTCPPacket(opts genTCPOpts) PacketBufferPtr { // Get values from opts. windowSize := uint16(50000) if opts.windowSize != nil { diff --git a/pkg/tcpip/stack/forwarding_test.go b/pkg/tcpip/stack/forwarding_test.go index 73cc800f0..59130ea3e 100644 --- a/pkg/tcpip/stack/forwarding_test.go +++ b/pkg/tcpip/stack/forwarding_test.go @@ -80,7 +80,7 @@ func (*fwdTestNetworkEndpoint) DefaultTTL() uint8 { return 123 } -func (f *fwdTestNetworkEndpoint) HandlePacket(pkt *PacketBuffer) { +func (f *fwdTestNetworkEndpoint) HandlePacket(pkt PacketBufferPtr) { if _, _, ok := f.proto.Parse(pkt); !ok { return } @@ -118,7 +118,7 @@ func (f *fwdTestNetworkEndpoint) NetworkProtocolNumber() tcpip.NetworkProtocolNu return f.proto.Number() } -func (f *fwdTestNetworkEndpoint) WritePacket(r *Route, params NetworkHeaderParams, pkt *PacketBuffer) tcpip.Error { +func (f *fwdTestNetworkEndpoint) WritePacket(r *Route, params NetworkHeaderParams, pkt PacketBufferPtr) tcpip.Error { // Add the protocol's header to the packet and send it to the link // endpoint. b := pkt.NetworkHeader().Push(fwdTestNetHeaderLen) @@ -130,7 +130,7 @@ func (f *fwdTestNetworkEndpoint) WritePacket(r *Route, params NetworkHeaderParam return f.nic.WritePacket(r, pkt) } -func (f *fwdTestNetworkEndpoint) WriteHeaderIncludedPacket(r *Route, pkt *PacketBuffer) tcpip.Error { +func (f *fwdTestNetworkEndpoint) WriteHeaderIncludedPacket(r *Route, pkt PacketBufferPtr) tcpip.Error { // The network header should not already be populated. if _, ok := pkt.NetworkHeader().Consume(fwdTestNetHeaderLen); !ok { return &tcpip.ErrMalformedHeader{} @@ -181,7 +181,7 @@ func (*fwdTestNetworkProtocol) ParseAddresses(v []byte) (src, dst tcpip.Address) return tcpip.Address(v[srcAddrOffset : srcAddrOffset+1]), tcpip.Address(v[dstAddrOffset : dstAddrOffset+1]) } -func (*fwdTestNetworkProtocol) Parse(pkt *PacketBuffer) (tcpip.TransportProtocolNumber, bool, bool) { +func (*fwdTestNetworkProtocol) Parse(pkt PacketBufferPtr) (tcpip.TransportProtocolNumber, bool, bool) { netHeader, ok := pkt.NetworkHeader().Consume(fwdTestNetHeaderLen) if !ok { return 0, false, false @@ -256,16 +256,16 @@ type fwdTestLinkEndpoint struct { linkAddr tcpip.LinkAddress // C is where outbound packets are queued. - C chan *PacketBuffer + C chan PacketBufferPtr } // InjectInbound injects an inbound packet. -func (e *fwdTestLinkEndpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) { +func (e *fwdTestLinkEndpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt PacketBufferPtr) { e.InjectLinkAddr(protocol, "", pkt) } // InjectLinkAddr injects an inbound packet with a remote link address. -func (e *fwdTestLinkEndpoint) InjectLinkAddr(protocol tcpip.NetworkProtocolNumber, remote tcpip.LinkAddress, pkt *PacketBuffer) { +func (e *fwdTestLinkEndpoint) InjectLinkAddr(protocol tcpip.NetworkProtocolNumber, remote tcpip.LinkAddress, pkt PacketBufferPtr) { e.dispatcher.DeliverNetworkPacket(protocol, pkt) } @@ -327,7 +327,7 @@ func (*fwdTestLinkEndpoint) ARPHardwareType() header.ARPHardwareType { } // AddHeader implements stack.LinkEndpoint.AddHeader. -func (*fwdTestLinkEndpoint) AddHeader(*PacketBuffer) {} +func (*fwdTestLinkEndpoint) AddHeader(PacketBufferPtr) {} func fwdTestNetFactory(t *testing.T, proto *fwdTestNetworkProtocol) (*faketime.ManualClock, *fwdTestLinkEndpoint, *fwdTestLinkEndpoint) { clock := faketime.NewManualClock() @@ -347,7 +347,7 @@ func fwdTestNetFactory(t *testing.T, proto *fwdTestNetworkProtocol) (*faketime.M // NIC 1 has the link address "a", and added the network address 1. ep1 := &fwdTestLinkEndpoint{ - C: make(chan *PacketBuffer, 300), + C: make(chan PacketBufferPtr, 300), mtu: fwdTestNetDefaultMTU, linkAddr: "a", } @@ -367,7 +367,7 @@ func fwdTestNetFactory(t *testing.T, proto *fwdTestNetworkProtocol) (*faketime.M // NIC 2 has the link address "b", and added the network address 2. ep2 := &fwdTestLinkEndpoint{ - C: make(chan *PacketBuffer, 300), + C: make(chan PacketBufferPtr, 300), mtu: fwdTestNetDefaultMTU, linkAddr: "b", } @@ -429,7 +429,7 @@ func TestForwardingWithStaticResolver(t *testing.T) { Payload: bufferv2.MakeWithData(buf), })) - var p *PacketBuffer + var p PacketBufferPtr clock.Advance(proto.addrResolveDelay) select { @@ -473,7 +473,7 @@ func TestForwardingWithFakeResolver(t *testing.T) { Payload: bufferv2.MakeWithData(buf), })) - var p *PacketBuffer + var p PacketBufferPtr clock.Advance(proto.addrResolveDelay) select { @@ -584,7 +584,7 @@ func TestForwardingWithFakeResolverPartialTimeout(t *testing.T) { Payload: bufferv2.MakeWithData(buf), })) - var p *PacketBuffer + var p PacketBufferPtr clock.Advance(proto.addrResolveDelay) select { @@ -636,7 +636,7 @@ func TestForwardingWithFakeResolverTwoPackets(t *testing.T) { } for i := 0; i < 2; i++ { - var p *PacketBuffer + var p PacketBufferPtr clock.Advance(proto.addrResolveDelay) select { @@ -691,7 +691,7 @@ func TestForwardingWithFakeResolverManyPackets(t *testing.T) { } for i := 0; i < maxPendingPacketsPerResolution; i++ { - var p *PacketBuffer + var p PacketBufferPtr clock.Advance(proto.addrResolveDelay) select { @@ -757,7 +757,7 @@ func TestForwardingWithFakeResolverManyResolutions(t *testing.T) { } for i := 0; i < maxPendingResolutions; i++ { - var p *PacketBuffer + var p PacketBufferPtr clock.Advance(proto.addrResolveDelay) select { diff --git a/pkg/tcpip/stack/iptables.go b/pkg/tcpip/stack/iptables.go index eb50f1347..0ad5303c8 100644 --- a/pkg/tcpip/stack/iptables.go +++ b/pkg/tcpip/stack/iptables.go @@ -281,7 +281,7 @@ type checkTable struct { // - Calls to dynamic functions, which can allocate. // // +checkescape:hard -func (it *IPTables) shouldSkipOrPopulateTables(tables []checkTable, pkt *PacketBuffer) bool { +func (it *IPTables) shouldSkipOrPopulateTables(tables []checkTable, pkt PacketBufferPtr) bool { switch pkt.NetworkProtocolNumber { case header.IPv4ProtocolNumber, header.IPv6ProtocolNumber: default: @@ -315,8 +315,8 @@ func (it *IPTables) shouldSkipOrPopulateTables(tables []checkTable, pkt *PacketB // This is called in the hot path even when iptables are disabled, so we ensure // that it does not allocate. Note that called functions (e.g. // getConnAndUpdate) can allocate. -// +checkescape -func (it *IPTables) CheckPrerouting(pkt *PacketBuffer, addressEP AddressableEndpoint, inNicName string) bool { +// TODO(b/233951539): checkescape fails on arm sometimes. Fix and re-add. +func (it *IPTables) CheckPrerouting(pkt PacketBufferPtr, addressEP AddressableEndpoint, inNicName string) bool { tables := [...]checkTable{ { fn: check, @@ -353,8 +353,8 @@ func (it *IPTables) CheckPrerouting(pkt *PacketBuffer, addressEP AddressableEndp // This is called in the hot path even when iptables are disabled, so we ensure // that it does not allocate. Note that called functions (e.g. // getConnAndUpdate) can allocate. -// +checkescape -func (it *IPTables) CheckInput(pkt *PacketBuffer, inNicName string) bool { +// TODO(b/233951539): checkescape fails on arm sometimes. Fix and re-add. +func (it *IPTables) CheckInput(pkt PacketBufferPtr, inNicName string) bool { tables := [...]checkTable{ { fn: checkNAT, @@ -393,8 +393,8 @@ func (it *IPTables) CheckInput(pkt *PacketBuffer, inNicName string) bool { // This is called in the hot path even when iptables are disabled, so we ensure // that it does not allocate. Note that called functions (e.g. // getConnAndUpdate) can allocate. -// +checkescape -func (it *IPTables) CheckForward(pkt *PacketBuffer, inNicName, outNicName string) bool { +// TODO(b/233951539): checkescape fails on arm sometimes. Fix and re-add. +func (it *IPTables) CheckForward(pkt PacketBufferPtr, inNicName, outNicName string) bool { tables := [...]checkTable{ { fn: check, @@ -425,8 +425,8 @@ func (it *IPTables) CheckForward(pkt *PacketBuffer, inNicName, outNicName string // This is called in the hot path even when iptables are disabled, so we ensure // that it does not allocate. Note that called functions (e.g. // getConnAndUpdate) can allocate. -// +checkescape -func (it *IPTables) CheckOutput(pkt *PacketBuffer, r *Route, outNicName string) bool { +// TODO(b/233951539): checkescape fails on arm sometimes. Fix and re-add. +func (it *IPTables) CheckOutput(pkt PacketBufferPtr, r *Route, outNicName string) bool { tables := [...]checkTable{ { fn: check, @@ -469,8 +469,8 @@ func (it *IPTables) CheckOutput(pkt *PacketBuffer, r *Route, outNicName string) // This is called in the hot path even when iptables are disabled, so we ensure // that it does not allocate. Note that called functions (e.g. // getConnAndUpdate) can allocate. -// +checkescape -func (it *IPTables) CheckPostrouting(pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, outNicName string) bool { +// TODO(b/233951539): checkescape fails on arm sometimes. Fix and re-add. +func (it *IPTables) CheckPostrouting(pkt PacketBufferPtr, r *Route, addressEP AddressableEndpoint, outNicName string) bool { tables := [...]checkTable{ { fn: check, @@ -501,16 +501,16 @@ func (it *IPTables) CheckPostrouting(pkt *PacketBuffer, r *Route, addressEP Addr // Note: this used to omit the *IPTables parameter, but doing so caused // unnecessary allocations. -type checkTableFn func(it *IPTables, table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool +type checkTableFn func(it *IPTables, table Table, hook Hook, pkt PacketBufferPtr, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool -func checkNAT(it *IPTables, table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { +func checkNAT(it *IPTables, table Table, hook Hook, pkt PacketBufferPtr, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { return it.checkNAT(table, hook, pkt, r, addressEP, inNicName, outNicName) } // checkNAT runs the packet through the NAT table. // // See check. -func (it *IPTables) checkNAT(table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { +func (it *IPTables) checkNAT(table Table, hook Hook, pkt PacketBufferPtr, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { t := pkt.tuple if t != nil && t.conn.handlePacket(pkt, hook, r) { return true @@ -548,7 +548,7 @@ func (it *IPTables) checkNAT(table Table, hook Hook, pkt *PacketBuffer, r *Route return true } -func check(it *IPTables, table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { +func check(it *IPTables, table Table, hook Hook, pkt PacketBufferPtr, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { return it.check(table, hook, pkt, r, addressEP, inNicName, outNicName) } @@ -557,7 +557,7 @@ func check(it *IPTables, table Table, hook Hook, pkt *PacketBuffer, r *Route, ad // network stack or tables, or false when it must be dropped. // // Precondition: The packet's network and transport header must be set. -func (it *IPTables) check(table Table, hook Hook, pkt *PacketBuffer, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { +func (it *IPTables) check(table Table, hook Hook, pkt PacketBufferPtr, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) bool { ruleIdx := table.BuiltinChains[hook] switch verdict := it.checkChain(hook, pkt, table, ruleIdx, r, addressEP, inNicName, outNicName); verdict { // If the table returns Accept, move on to the next table. @@ -610,7 +610,7 @@ func (it *IPTables) startReaper(interval time.Duration) { // Preconditions: // - pkt is a IPv4 packet of at least length header.IPv4MinimumSize. // - pkt.NetworkHeader is not nil. -func (it *IPTables) checkChain(hook Hook, pkt *PacketBuffer, table Table, ruleIdx int, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) chainVerdict { +func (it *IPTables) checkChain(hook Hook, pkt PacketBufferPtr, table Table, ruleIdx int, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) chainVerdict { // Start from ruleIdx and walk the list of rules until a rule gives us // a verdict. for ruleIdx < len(table.Rules) { @@ -657,7 +657,10 @@ func (it *IPTables) checkChain(hook Hook, pkt *PacketBuffer, table Table, ruleId // Preconditions: // - pkt is a IPv4 packet of at least length header.IPv4MinimumSize. // - pkt.NetworkHeader is not nil. -func (it *IPTables) checkRule(hook Hook, pkt *PacketBuffer, table Table, ruleIdx int, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) (RuleVerdict, int) { +// +// * pkt is a IPv4 packet of at least length header.IPv4MinimumSize. +// * pkt.NetworkHeader is not nil. +func (it *IPTables) checkRule(hook Hook, pkt PacketBufferPtr, table Table, ruleIdx int, r *Route, addressEP AddressableEndpoint, inNicName, outNicName string) (RuleVerdict, int) { rule := table.Rules[ruleIdx] // Check whether the packet matches the IP header filter. diff --git a/pkg/tcpip/stack/iptables_targets.go b/pkg/tcpip/stack/iptables_targets.go index d93670d6b..5b4a736fe 100644 --- a/pkg/tcpip/stack/iptables_targets.go +++ b/pkg/tcpip/stack/iptables_targets.go @@ -30,7 +30,7 @@ type AcceptTarget struct { } // Action implements Target.Action. -func (*AcceptTarget) Action(*PacketBuffer, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { +func (*AcceptTarget) Action(PacketBufferPtr, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { return RuleAccept, 0 } @@ -41,14 +41,14 @@ type DropTarget struct { } // Action implements Target.Action. -func (*DropTarget) Action(*PacketBuffer, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { +func (*DropTarget) Action(PacketBufferPtr, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { return RuleDrop, 0 } // RejectIPv4WithHandler handles rejecting a packet. type RejectIPv4WithHandler interface { // SendRejectionError sends an error packet in response to the packet. - SendRejectionError(pkt *PacketBuffer, rejectWith RejectIPv4WithICMPType, inputHook bool) tcpip.Error + SendRejectionError(pkt PacketBufferPtr, rejectWith RejectIPv4WithICMPType, inputHook bool) tcpip.Error } // RejectIPv4WithICMPType indicates the type of ICMP error that should be sent. @@ -73,7 +73,7 @@ type RejectIPv4Target struct { } // Action implements Target.Action. -func (rt *RejectIPv4Target) Action(pkt *PacketBuffer, hook Hook, _ *Route, _ AddressableEndpoint) (RuleVerdict, int) { +func (rt *RejectIPv4Target) Action(pkt PacketBufferPtr, hook Hook, _ *Route, _ AddressableEndpoint) (RuleVerdict, int) { switch hook { case Input, Forward, Output: // There is nothing reasonable for us to do in response to an error here; @@ -90,7 +90,7 @@ func (rt *RejectIPv4Target) Action(pkt *PacketBuffer, hook Hook, _ *Route, _ Add // RejectIPv6WithHandler handles rejecting a packet. type RejectIPv6WithHandler interface { // SendRejectionError sends an error packet in response to the packet. - SendRejectionError(pkt *PacketBuffer, rejectWith RejectIPv6WithICMPType, forwardingHook bool) tcpip.Error + SendRejectionError(pkt PacketBufferPtr, rejectWith RejectIPv6WithICMPType, forwardingHook bool) tcpip.Error } // RejectIPv6WithICMPType indicates the type of ICMP error that should be sent. @@ -113,7 +113,7 @@ type RejectIPv6Target struct { } // Action implements Target.Action. -func (rt *RejectIPv6Target) Action(pkt *PacketBuffer, hook Hook, _ *Route, _ AddressableEndpoint) (RuleVerdict, int) { +func (rt *RejectIPv6Target) Action(pkt PacketBufferPtr, hook Hook, _ *Route, _ AddressableEndpoint) (RuleVerdict, int) { switch hook { case Input, Forward, Output: // There is nothing reasonable for us to do in response to an error here; @@ -135,7 +135,7 @@ type ErrorTarget struct { } // Action implements Target.Action. -func (*ErrorTarget) Action(*PacketBuffer, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { +func (*ErrorTarget) Action(PacketBufferPtr, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { log.Debugf("ErrorTarget triggered.") return RuleDrop, 0 } @@ -150,7 +150,7 @@ type UserChainTarget struct { } // Action implements Target.Action. -func (*UserChainTarget) Action(*PacketBuffer, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { +func (*UserChainTarget) Action(PacketBufferPtr, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { panic("UserChainTarget should never be called.") } @@ -162,7 +162,7 @@ type ReturnTarget struct { } // Action implements Target.Action. -func (*ReturnTarget) Action(*PacketBuffer, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { +func (*ReturnTarget) Action(PacketBufferPtr, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) { return RuleReturn, 0 } @@ -185,7 +185,7 @@ type DNATTarget struct { } // Action implements Target.Action. -func (rt *DNATTarget) Action(pkt *PacketBuffer, hook Hook, r *Route, addressEP AddressableEndpoint) (RuleVerdict, int) { +func (rt *DNATTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, addressEP AddressableEndpoint) (RuleVerdict, int) { // Sanity check. if rt.NetworkProtocol != pkt.NetworkProtocolNumber { panic(fmt.Sprintf( @@ -219,7 +219,7 @@ type RedirectTarget struct { } // Action implements Target.Action. -func (rt *RedirectTarget) Action(pkt *PacketBuffer, hook Hook, r *Route, addressEP AddressableEndpoint) (RuleVerdict, int) { +func (rt *RedirectTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, addressEP AddressableEndpoint) (RuleVerdict, int) { // Sanity check. if rt.NetworkProtocol != pkt.NetworkProtocolNumber { panic(fmt.Sprintf( @@ -257,7 +257,7 @@ type SNATTarget struct { NetworkProtocol tcpip.NetworkProtocolNumber } -func dnatAction(pkt *PacketBuffer, hook Hook, r *Route, port uint16, address tcpip.Address) (RuleVerdict, int) { +func dnatAction(pkt PacketBufferPtr, hook Hook, r *Route, port uint16, address tcpip.Address) (RuleVerdict, int) { return natAction(pkt, hook, r, portOrIdentRange{start: port, size: 1}, address, true /* dnat */) } @@ -278,7 +278,7 @@ func targetPortRangeForTCPAndUDP(originalSrcPort uint16) portOrIdentRange { } } -func snatAction(pkt *PacketBuffer, hook Hook, r *Route, port uint16, address tcpip.Address) (RuleVerdict, int) { +func snatAction(pkt PacketBufferPtr, hook Hook, r *Route, port uint16, address tcpip.Address) (RuleVerdict, int) { portsOrIdents := portOrIdentRange{start: port, size: 1} switch pkt.TransportProtocolNumber { @@ -301,7 +301,7 @@ func snatAction(pkt *PacketBuffer, hook Hook, r *Route, port uint16, address tcp return natAction(pkt, hook, r, portsOrIdents, address, false /* dnat */) } -func natAction(pkt *PacketBuffer, hook Hook, r *Route, portsOrIdents portOrIdentRange, address tcpip.Address, dnat bool) (RuleVerdict, int) { +func natAction(pkt PacketBufferPtr, hook Hook, r *Route, portsOrIdents portOrIdentRange, address tcpip.Address, dnat bool) (RuleVerdict, int) { // Drop the packet if network and transport header are not set. if len(pkt.NetworkHeader().Slice()) == 0 || len(pkt.TransportHeader().Slice()) == 0 { return RuleDrop, 0 @@ -316,7 +316,7 @@ func natAction(pkt *PacketBuffer, hook Hook, r *Route, portsOrIdents portOrIdent } // Action implements Target.Action. -func (st *SNATTarget) Action(pkt *PacketBuffer, hook Hook, r *Route, _ AddressableEndpoint) (RuleVerdict, int) { +func (st *SNATTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, _ AddressableEndpoint) (RuleVerdict, int) { // Sanity check. if st.NetworkProtocol != pkt.NetworkProtocolNumber { panic(fmt.Sprintf( @@ -343,7 +343,7 @@ type MasqueradeTarget struct { } // Action implements Target.Action. -func (mt *MasqueradeTarget) Action(pkt *PacketBuffer, hook Hook, r *Route, addressEP AddressableEndpoint) (RuleVerdict, int) { +func (mt *MasqueradeTarget) Action(pkt PacketBufferPtr, hook Hook, r *Route, addressEP AddressableEndpoint) (RuleVerdict, int) { // Sanity check. if mt.NetworkProtocol != pkt.NetworkProtocolNumber { panic(fmt.Sprintf( diff --git a/pkg/tcpip/stack/iptables_test.go b/pkg/tcpip/stack/iptables_test.go index 0c54c0811..af93f9201 100644 --- a/pkg/tcpip/stack/iptables_test.go +++ b/pkg/tcpip/stack/iptables_test.go @@ -41,7 +41,7 @@ var ( dstAddr = testutil.MustParse6("c::3") ) -func v6PacketBufferWithSrcAddr(srcAddr tcpip.Address) *PacketBuffer { +func v6PacketBufferWithSrcAddr(srcAddr tcpip.Address) PacketBufferPtr { pkt := NewPacketBuffer(PacketBufferOptions{ ReserveHeaderBytes: header.IPv6MinimumSize + header.UDPMinimumSize, }) @@ -69,7 +69,7 @@ func v6PacketBufferWithSrcAddr(srcAddr tcpip.Address) *PacketBuffer { return pkt } -func v6PacketBuffer() *PacketBuffer { +func v6PacketBuffer() PacketBufferPtr { return v6PacketBufferWithSrcAddr(srcAddr) } @@ -239,19 +239,19 @@ func TestNATedConnectionReap(t *testing.T) { func TestNATAlwaysPerformed(t *testing.T) { tests := []struct { name string - dnatHook func(*testing.T, *IPTables, *PacketBuffer) - snatHook func(*testing.T, *IPTables, *PacketBuffer) + dnatHook func(*testing.T, *IPTables, PacketBufferPtr) + snatHook func(*testing.T, *IPTables, PacketBufferPtr) }{ { name: "Prerouting and Input", - dnatHook: func(t *testing.T, iptables *IPTables, pkt *PacketBuffer) { + dnatHook: func(t *testing.T, iptables *IPTables, pkt PacketBufferPtr) { t.Helper() if !iptables.CheckPrerouting(pkt, nil /* addressEP */, "" /* inNicName */) { t.Fatal("got iptables.CheckPrerouting(...) = false, want = true") } }, - snatHook: func(t *testing.T, iptables *IPTables, pkt *PacketBuffer) { + snatHook: func(t *testing.T, iptables *IPTables, pkt PacketBufferPtr) { t.Helper() if !iptables.CheckInput(pkt, "" /* inNicName */) { @@ -261,7 +261,7 @@ func TestNATAlwaysPerformed(t *testing.T) { }, { name: "Output and Postrouting", - dnatHook: func(t *testing.T, iptables *IPTables, pkt *PacketBuffer) { + dnatHook: func(t *testing.T, iptables *IPTables, pkt PacketBufferPtr) { t.Helper() // Output hook depends on a route but if the route is local, we don't @@ -275,7 +275,7 @@ func TestNATAlwaysPerformed(t *testing.T) { t.Fatal("got iptables.CheckOutput(...) = false, want = true") } }, - snatHook: func(t *testing.T, iptables *IPTables, pkt *PacketBuffer) { + snatHook: func(t *testing.T, iptables *IPTables, pkt PacketBufferPtr) { t.Helper() // Postrouting hook depends on a route but if the route is local, we @@ -327,11 +327,11 @@ func TestNATConflict(t *testing.T) { tests := []struct { name string - checkIPTables func(*testing.T, *IPTables, *PacketBuffer, bool) + checkIPTables func(*testing.T, *IPTables, PacketBufferPtr, bool) }{ { name: "Prerouting and Input", - checkIPTables: func(t *testing.T, iptables *IPTables, pkt *PacketBuffer, lastHookOK bool) { + checkIPTables: func(t *testing.T, iptables *IPTables, pkt PacketBufferPtr, lastHookOK bool) { t.Helper() if !iptables.CheckPrerouting(pkt, nil /* addressEP */, "" /* inNicName */) { @@ -344,7 +344,7 @@ func TestNATConflict(t *testing.T) { }, { name: "Output and Postrouting", - checkIPTables: func(t *testing.T, iptables *IPTables, pkt *PacketBuffer, lastHookOK bool) { + checkIPTables: func(t *testing.T, iptables *IPTables, pkt PacketBufferPtr, lastHookOK bool) { t.Helper() // Output and Postrouting hooks depends on a route but if the route is diff --git a/pkg/tcpip/stack/iptables_types.go b/pkg/tcpip/stack/iptables_types.go index 12def5ccb..99b5c9684 100644 --- a/pkg/tcpip/stack/iptables_types.go +++ b/pkg/tcpip/stack/iptables_types.go @@ -240,7 +240,7 @@ type IPHeaderFilter struct { // // Preconditions: pkt.NetworkHeader is set and is at least of the minimal IPv4 // or IPv6 header length. -func (fl IPHeaderFilter) match(pkt *PacketBuffer, hook Hook, inNicName, outNicName string) bool { +func (fl IPHeaderFilter) match(pkt PacketBufferPtr, hook Hook, inNicName, outNicName string) bool { // Extract header fields. var ( transProto tcpip.TransportProtocolNumber @@ -345,7 +345,7 @@ type Matcher interface { // used for suspicious packets. // // Precondition: packet.NetworkHeader is set. - Match(hook Hook, packet *PacketBuffer, inputInterfaceName, outputInterfaceName string) (matches bool, hotdrop bool) + Match(hook Hook, packet PacketBufferPtr, inputInterfaceName, outputInterfaceName string) (matches bool, hotdrop bool) } // A Target is the interface for taking an action for a packet. @@ -353,5 +353,5 @@ type Target interface { // Action takes an action on the packet and returns a verdict on how // traversal should (or should not) continue. If the return value is // Jump, it also returns the index of the rule to jump to. - Action(*PacketBuffer, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) + Action(PacketBufferPtr, Hook, *Route, AddressableEndpoint) (RuleVerdict, int) } diff --git a/pkg/tcpip/stack/ndp_test.go b/pkg/tcpip/stack/ndp_test.go index 7462b680c..31629b451 100644 --- a/pkg/tcpip/stack/ndp_test.go +++ b/pkg/tcpip/stack/ndp_test.go @@ -697,7 +697,7 @@ func TestDADResolve(t *testing.T) { // Validate the sent Neighbor Solicitation messages. for i := uint8(0); i < test.dupAddrDetectTransmits; i++ { p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("packet didn't arrive") } @@ -1230,7 +1230,7 @@ func TestSetNDPConfigurations(t *testing.T) { // raBuf returns a valid NDP Router Advertisement with options, router // preference and DHCPv6 configurations specified. -func raBuf(ip tcpip.Address, rl uint16, managedAddress, otherConfigurations bool, prf header.NDPRoutePreference, optSer header.NDPOptionsSerializer) *stack.PacketBuffer { +func raBuf(ip tcpip.Address, rl uint16, managedAddress, otherConfigurations bool, prf header.NDPRoutePreference, optSer header.NDPOptionsSerializer) stack.PacketBufferPtr { const flagsByte = 1 const routerLifetimeOffset = 2 @@ -1281,7 +1281,7 @@ func raBuf(ip tcpip.Address, rl uint16, managedAddress, otherConfigurations bool // // Note, raBufWithOpts does not populate any of the RA fields other than the // Router Lifetime. -func raBufWithOpts(ip tcpip.Address, rl uint16, optSer header.NDPOptionsSerializer) *stack.PacketBuffer { +func raBufWithOpts(ip tcpip.Address, rl uint16, optSer header.NDPOptionsSerializer) stack.PacketBufferPtr { return raBuf(ip, rl, false /* managedAddress */, false /* otherConfigurations */, 0 /* prf */, optSer) } @@ -1290,7 +1290,7 @@ func raBufWithOpts(ip tcpip.Address, rl uint16, optSer header.NDPOptionsSerializ // // Note, raBufWithDHCPv6 does not populate any of the RA fields other than the // DHCPv6 related ones. -func raBufWithDHCPv6(ip tcpip.Address, managedAddresses, otherConfigurations bool) *stack.PacketBuffer { +func raBufWithDHCPv6(ip tcpip.Address, managedAddresses, otherConfigurations bool) stack.PacketBufferPtr { return raBuf(ip, 0, managedAddresses, otherConfigurations, 0 /* prf */, header.NDPOptionsSerializer{}) } @@ -1298,7 +1298,7 @@ func raBufWithDHCPv6(ip tcpip.Address, managedAddresses, otherConfigurations boo // // Note, raBuf does not populate any of the RA fields other than the // Router Lifetime. -func raBufSimple(ip tcpip.Address, rl uint16) *stack.PacketBuffer { +func raBufSimple(ip tcpip.Address, rl uint16) stack.PacketBufferPtr { return raBufWithOpts(ip, rl, header.NDPOptionsSerializer{}) } @@ -1306,7 +1306,7 @@ func raBufSimple(ip tcpip.Address, rl uint16) *stack.PacketBuffer { // // Note, raBufWithPrf does not populate any of the RA fields other than the // Router Lifetime and Default Router Preference fields. -func raBufWithPrf(ip tcpip.Address, rl uint16, prf header.NDPRoutePreference) *stack.PacketBuffer { +func raBufWithPrf(ip tcpip.Address, rl uint16, prf header.NDPRoutePreference) stack.PacketBufferPtr { return raBuf(ip, rl, false /* managedAddress */, false /* otherConfigurations */, prf, header.NDPOptionsSerializer{}) } @@ -1315,7 +1315,7 @@ func raBufWithPrf(ip tcpip.Address, rl uint16, prf header.NDPRoutePreference) *s // // Note, raBufWithPI does not populate any of the RA fields other than the // Router Lifetime. -func raBufWithPI(ip tcpip.Address, rl uint16, prefix tcpip.AddressWithPrefix, onLink, auto bool, vl, pl uint32) *stack.PacketBuffer { +func raBufWithPI(ip tcpip.Address, rl uint16, prefix tcpip.AddressWithPrefix, onLink, auto bool, vl, pl uint32) stack.PacketBufferPtr { flags := uint8(0) if onLink { // The OnLink flag is the 7th bit in the flags byte. @@ -1352,7 +1352,7 @@ func raBufWithPI(ip tcpip.Address, rl uint16, prefix tcpip.AddressWithPrefix, on // Information option. // // All fields in the RA will be zero except the RIO option. -func raBufWithRIO(t *testing.T, ip tcpip.Address, prefix tcpip.AddressWithPrefix, lifetimeSeconds uint32, prf header.NDPRoutePreference) *stack.PacketBuffer { +func raBufWithRIO(t *testing.T, ip tcpip.Address, prefix tcpip.AddressWithPrefix, lifetimeSeconds uint32, prf header.NDPRoutePreference) stack.PacketBufferPtr { // buf will hold the route information option after the Type and Length // fields. // @@ -1395,7 +1395,7 @@ func TestDynamicConfigurationsDisabled(t *testing.T) { tests := []struct { name string config func(bool) ipv6.NDPConfigurations - ra *stack.PacketBuffer + ra stack.PacketBufferPtr }{ { name: "No Router Discovery", @@ -1481,10 +1481,10 @@ func TestDynamicConfigurationsDisabled(t *testing.T) { t.Errorf("got v6Stats.ICMP.PacketsSent.RouterSolicit.Value() = %d, want = %d", got, want) } if handleRAsDisabled { - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Errorf("unexpectedly got a packet = %#v", p) } - } else if p := e.Read(); p == nil { + } else if p := e.Read(); p.IsNil() { t.Error("expected router solicitation packet") } else if p.NetworkProtocolNumber != header.IPv6ProtocolNumber { t.Errorf("got Proto = %d, want = %d", p.NetworkProtocolNumber, header.IPv6ProtocolNumber) @@ -1584,14 +1584,14 @@ func TestOffLinkRouteDiscovery(t *testing.T) { discoverMoreSpecificRoutes bool dest tcpip.Subnet - ra func(*testing.T, tcpip.Address, uint16, header.NDPRoutePreference) *stack.PacketBuffer + ra func(*testing.T, tcpip.Address, uint16, header.NDPRoutePreference) stack.PacketBufferPtr }{ { name: "Default router discovery", discoverDefaultRouters: true, discoverMoreSpecificRoutes: false, dest: header.IPv6EmptySubnet, - ra: func(_ *testing.T, router tcpip.Address, lifetimeSeconds uint16, prf header.NDPRoutePreference) *stack.PacketBuffer { + ra: func(_ *testing.T, router tcpip.Address, lifetimeSeconds uint16, prf header.NDPRoutePreference) stack.PacketBufferPtr { return raBufWithPrf(router, lifetimeSeconds, prf) }, }, @@ -1600,7 +1600,7 @@ func TestOffLinkRouteDiscovery(t *testing.T) { discoverDefaultRouters: false, discoverMoreSpecificRoutes: true, dest: moreSpecificPrefix.Subnet(), - ra: func(t *testing.T, router tcpip.Address, lifetimeSeconds uint16, prf header.NDPRoutePreference) *stack.PacketBuffer { + ra: func(t *testing.T, router tcpip.Address, lifetimeSeconds uint16, prf header.NDPRoutePreference) stack.PacketBufferPtr { return raBufWithRIO(t, router, moreSpecificPrefix, uint32(lifetimeSeconds), prf) }, }, @@ -5733,7 +5733,7 @@ func TestRouterSolicitation(t *testing.T) { clock.Advance(timeout) p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("expected router solicitation packet") } defer p.DecRef() @@ -5762,7 +5762,7 @@ func TestRouterSolicitation(t *testing.T) { t.Helper() clock.Advance(timeout) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { t.Fatalf("unexpectedly got a packet = %#v", p) } } @@ -5924,7 +5924,7 @@ func TestStopStartSolicitingRouters(t *testing.T) { clock.Advance(timeout) p := e.Read() - if p == nil { + if p.IsNil() { t.Fatal("timed out waiting for packet") } @@ -5957,11 +5957,11 @@ func TestStopStartSolicitingRouters(t *testing.T) { // Stop soliciting routers. test.stopFn(t, s, true /* first */) clock.Advance(delay) - if p := e.Read(); p != nil { + if p := e.Read(); !p.IsNil() { p.DecRef() // A single RS may have been sent before solicitations were stopped. clock.Advance(interval) - if e.Read() != nil { + if pb := e.Read(); !pb.IsNil() { t.Fatal("should not have sent more than one RS message") } } @@ -5970,7 +5970,7 @@ func TestStopStartSolicitingRouters(t *testing.T) { // do nothing. test.stopFn(t, s, false /* first */) clock.Advance(delay) - if e.Read() != nil { + if pb := e.Read(); !pb.IsNil() { t.Fatal("unexpectedly got a packet after router solicitation has been stopepd") } @@ -5985,7 +5985,7 @@ func TestStopStartSolicitingRouters(t *testing.T) { waitForPkt(clock, interval) waitForPkt(clock, interval) clock.Advance(interval) - if e.Read() != nil { + if pb := e.Read(); !pb.IsNil() { t.Fatal("unexpectedly got an extra packet after sending out the expected RSs") } @@ -5993,7 +5993,7 @@ func TestStopStartSolicitingRouters(t *testing.T) { // nothing. test.startFn(t, s) clock.Advance(interval) - if e.Read() != nil { + if pb := e.Read(); !pb.IsNil() { t.Fatal("unexpectedly got a packet after finishing router solicitations") } }) diff --git a/pkg/tcpip/stack/nic.go b/pkg/tcpip/stack/nic.go index c9f5ce5e1..d7258dd95 100644 --- a/pkg/tcpip/stack/nic.go +++ b/pkg/tcpip/stack/nic.go @@ -141,7 +141,7 @@ type delegatingQueueingDiscipline struct { func (*delegatingQueueingDiscipline) Close() {} // WritePacket passes the packet through to the underlying LinkWriter's WritePackets. -func (qDisc *delegatingQueueingDiscipline) WritePacket(pkt *PacketBuffer) tcpip.Error { +func (qDisc *delegatingQueueingDiscipline) WritePacket(pkt PacketBufferPtr) tcpip.Error { var pkts PacketBufferList pkts.PushBack(pkt) _, err := qDisc.LinkWriter.WritePackets(pkts) @@ -340,7 +340,7 @@ func (n *nic) IsLoopback() bool { } // WritePacket implements NetworkEndpoint. -func (n *nic) WritePacket(r *Route, pkt *PacketBuffer) tcpip.Error { +func (n *nic) WritePacket(r *Route, pkt PacketBufferPtr) tcpip.Error { routeInfo, _, err := r.resolvedFields(nil) switch err.(type) { case nil: @@ -371,7 +371,7 @@ func (n *nic) WritePacket(r *Route, pkt *PacketBuffer) tcpip.Error { } // WritePacketToRemote implements NetworkInterface. -func (n *nic) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt *PacketBuffer) tcpip.Error { +func (n *nic) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt PacketBufferPtr) tcpip.Error { pkt.EgressRoute = RouteInfo{ routeInfo: routeInfo{ NetProto: pkt.NetworkProtocolNumber, @@ -382,12 +382,12 @@ func (n *nic) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt *PacketB return n.writePacket(pkt) } -func (n *nic) writePacket(pkt *PacketBuffer) tcpip.Error { +func (n *nic) writePacket(pkt PacketBufferPtr) tcpip.Error { n.NetworkLinkEndpoint.AddHeader(pkt) return n.writeRawPacket(pkt) } -func (n *nic) writeRawPacket(pkt *PacketBuffer) tcpip.Error { +func (n *nic) writeRawPacket(pkt PacketBufferPtr) tcpip.Error { if err := n.qDisc.WritePacket(pkt); err != nil { if _, ok := err.(*tcpip.ErrNoBufferSpace); ok { n.stats.txPacketsDroppedNoBufferSpace.Increment() @@ -717,7 +717,7 @@ func (n *nic) isInGroup(addr tcpip.Address) bool { // DeliverNetworkPacket finds the appropriate network protocol endpoint and // hands the packet over for further processing. This function is called when // the NIC receives a packet from the link endpoint. -func (n *nic) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) { +func (n *nic) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt PacketBufferPtr) { enabled := n.Enabled() // If the NIC is not yet enabled, don't receive any packets. if !enabled { @@ -740,16 +740,16 @@ func (n *nic) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *Pa networkEndpoint.HandlePacket(pkt) } -func (n *nic) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer, incoming bool) { +func (n *nic) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt PacketBufferPtr, incoming bool) { // Deliver to interested packet endpoints without holding NIC lock. - var packetEPPkt *PacketBuffer + var packetEPPkt PacketBufferPtr defer func() { - if packetEPPkt != nil { + if !packetEPPkt.IsNil() { packetEPPkt.DecRef() } }() deliverPacketEPs := func(ep PacketEndpoint) { - if packetEPPkt == nil { + if packetEPPkt.IsNil() { // Packet endpoints hold the full packet. // // We perform a deep copy because higher-level endpoints may point to @@ -797,7 +797,7 @@ func (n *nic) DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *Packe // DeliverTransportPacket delivers the packets to the appropriate transport // protocol endpoint. -func (n *nic) DeliverTransportPacket(protocol tcpip.TransportProtocolNumber, pkt *PacketBuffer) TransportPacketDisposition { +func (n *nic) DeliverTransportPacket(protocol tcpip.TransportProtocolNumber, pkt PacketBufferPtr) TransportPacketDisposition { state, ok := n.stack.transportProtocols[protocol] if !ok { n.stats.unknownL4ProtocolRcvdPacketCounts.Increment(uint64(protocol)) @@ -857,7 +857,7 @@ func (n *nic) DeliverTransportPacket(protocol tcpip.TransportProtocolNumber, pkt } // DeliverTransportError implements TransportDispatcher. -func (n *nic) DeliverTransportError(local, remote tcpip.Address, net tcpip.NetworkProtocolNumber, trans tcpip.TransportProtocolNumber, transErr TransportError, pkt *PacketBuffer) { +func (n *nic) DeliverTransportError(local, remote tcpip.Address, net tcpip.NetworkProtocolNumber, trans tcpip.TransportProtocolNumber, transErr TransportError, pkt PacketBufferPtr) { state, ok := n.stack.transportProtocols[trans] if !ok { return @@ -885,7 +885,7 @@ func (n *nic) DeliverTransportError(local, remote tcpip.Address, net tcpip.Netwo } // DeliverRawPacket implements TransportDispatcher. -func (n *nic) DeliverRawPacket(protocol tcpip.TransportProtocolNumber, pkt *PacketBuffer) { +func (n *nic) DeliverRawPacket(protocol tcpip.TransportProtocolNumber, pkt PacketBufferPtr) { // For ICMPv4 only we validate the header length for compatibility with // raw(7) ICMP_FILTER. The same check is made in Linux here: // https://github.com/torvalds/linux/blob/70585216/net/ipv4/raw.c#L189. diff --git a/pkg/tcpip/stack/nic_test.go b/pkg/tcpip/stack/nic_test.go index 1beb38cf2..9a5674823 100644 --- a/pkg/tcpip/stack/nic_test.go +++ b/pkg/tcpip/stack/nic_test.go @@ -68,19 +68,19 @@ func (e *testIPv6Endpoint) MaxHeaderLength() uint16 { } // WritePacket implements NetworkEndpoint.WritePacket. -func (*testIPv6Endpoint) WritePacket(*Route, NetworkHeaderParams, *PacketBuffer) tcpip.Error { +func (*testIPv6Endpoint) WritePacket(*Route, NetworkHeaderParams, PacketBufferPtr) tcpip.Error { return nil } // WriteHeaderIncludedPacket implements // NetworkEndpoint.WriteHeaderIncludedPacket. -func (*testIPv6Endpoint) WriteHeaderIncludedPacket(*Route, *PacketBuffer) tcpip.Error { +func (*testIPv6Endpoint) WriteHeaderIncludedPacket(*Route, PacketBufferPtr) tcpip.Error { // Our tests don't use this so we don't support it. return &tcpip.ErrNotSupported{} } // HandlePacket implements NetworkEndpoint.HandlePacket. -func (*testIPv6Endpoint) HandlePacket(*PacketBuffer) {} +func (*testIPv6Endpoint) HandlePacket(PacketBufferPtr) {} // Close implements NetworkEndpoint.Close. func (e *testIPv6Endpoint) Close() { @@ -155,7 +155,7 @@ func (*testIPv6Protocol) Close() {} func (*testIPv6Protocol) Wait() {} // Parse implements NetworkProtocol.Parse. -func (*testIPv6Protocol) Parse(*PacketBuffer) (tcpip.TransportProtocolNumber, bool, bool) { +func (*testIPv6Protocol) Parse(PacketBufferPtr) (tcpip.TransportProtocolNumber, bool, bool) { return 0, false, false } diff --git a/pkg/tcpip/stack/packet_buffer.go b/pkg/tcpip/stack/packet_buffer.go index eed38448e..2d721cfa6 100644 --- a/pkg/tcpip/stack/packet_buffer.go +++ b/pkg/tcpip/stack/packet_buffer.go @@ -58,6 +58,9 @@ type PacketBufferOptions struct { OnRelease func() } +// PacketBufferPtr is a pointer to a PacketBuffer. +type PacketBufferPtr = *PacketBuffer + // 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 @@ -169,8 +172,8 @@ type PacketBuffer struct { } // NewPacketBuffer creates a new PacketBuffer with opts. -func NewPacketBuffer(opts PacketBufferOptions) *PacketBuffer { - pk := pkPool.Get().(*PacketBuffer) +func NewPacketBuffer(opts PacketBufferOptions) PacketBufferPtr { + pk := pkPool.Get().(PacketBufferPtr) pk.reset() if opts.ReserveHeaderBytes != 0 { v := bufferv2.NewViewSize(opts.ReserveHeaderBytes) @@ -187,7 +190,7 @@ func NewPacketBuffer(opts PacketBufferOptions) *PacketBuffer { } // IncRef increments the PacketBuffer's refcount. -func (pk *PacketBuffer) IncRef() *PacketBuffer { +func (pk PacketBufferPtr) IncRef() PacketBufferPtr { pk.packetBufferRefs.IncRef() return pk } @@ -195,7 +198,7 @@ func (pk *PacketBuffer) IncRef() *PacketBuffer { // DecRef decrements the PacketBuffer's refcount. If the refcount is // decremented to zero, the PacketBuffer is returned to the PacketBuffer // pool. -func (pk *PacketBuffer) DecRef() { +func (pk PacketBufferPtr) DecRef() { pk.packetBufferRefs.DecRef(func() { if pk.onRelease != nil { pk.onRelease() @@ -206,24 +209,24 @@ func (pk *PacketBuffer) DecRef() { }) } -func (pk *PacketBuffer) reset() { +func (pk PacketBufferPtr) reset() { *pk = PacketBuffer{} } // ReservedHeaderBytes returns the number of bytes initially reserved for // headers. -func (pk *PacketBuffer) ReservedHeaderBytes() int { +func (pk PacketBufferPtr) ReservedHeaderBytes() int { return pk.reserved } // AvailableHeaderBytes returns the number of bytes currently available for // headers. This is relevant to PacketHeader.Push method only. -func (pk *PacketBuffer) AvailableHeaderBytes() int { +func (pk PacketBufferPtr) AvailableHeaderBytes() int { return pk.reserved - pk.pushed } // VirtioNetHeader returns the handle to virtio-layer header. -func (pk *PacketBuffer) VirtioNetHeader() PacketHeader { +func (pk PacketBufferPtr) VirtioNetHeader() PacketHeader { return PacketHeader{ pk: pk, typ: virtioNetHeader, @@ -231,7 +234,7 @@ func (pk *PacketBuffer) VirtioNetHeader() PacketHeader { } // LinkHeader returns the handle to link-layer header. -func (pk *PacketBuffer) LinkHeader() PacketHeader { +func (pk PacketBufferPtr) LinkHeader() PacketHeader { return PacketHeader{ pk: pk, typ: linkHeader, @@ -239,7 +242,7 @@ func (pk *PacketBuffer) LinkHeader() PacketHeader { } // NetworkHeader returns the handle to network-layer header. -func (pk *PacketBuffer) NetworkHeader() PacketHeader { +func (pk PacketBufferPtr) NetworkHeader() PacketHeader { return PacketHeader{ pk: pk, typ: networkHeader, @@ -247,7 +250,7 @@ func (pk *PacketBuffer) NetworkHeader() PacketHeader { } // TransportHeader returns the handle to transport-layer header. -func (pk *PacketBuffer) TransportHeader() PacketHeader { +func (pk PacketBufferPtr) TransportHeader() PacketHeader { return PacketHeader{ pk: pk, typ: transportHeader, @@ -255,28 +258,28 @@ func (pk *PacketBuffer) TransportHeader() PacketHeader { } // HeaderSize returns the total size of all headers in bytes. -func (pk *PacketBuffer) HeaderSize() int { +func (pk PacketBufferPtr) HeaderSize() int { return pk.pushed + pk.consumed } // Size returns the size of packet in bytes. -func (pk *PacketBuffer) Size() int { +func (pk PacketBufferPtr) Size() int { return int(pk.buf.Size()) - pk.headerOffset() } // MemSize returns the estimation size of the pk in memory, including backing // buffer data. -func (pk *PacketBuffer) MemSize() int { +func (pk PacketBufferPtr) MemSize() int { return int(pk.buf.Size()) + PacketBufferStructSize } // Data returns the handle to data portion of pk. -func (pk *PacketBuffer) Data() PacketData { +func (pk PacketBufferPtr) Data() PacketData { return PacketData{pk: pk} } // AsSlices returns the underlying storage of the whole packet. -func (pk *PacketBuffer) AsSlices() [][]byte { +func (pk PacketBufferPtr) AsSlices() [][]byte { var views [][]byte offset := pk.headerOffset() pk.buf.SubApply(offset, int(pk.buf.Size())-offset, func(v *bufferv2.View) { @@ -287,7 +290,7 @@ func (pk *PacketBuffer) AsSlices() [][]byte { // ToBuffer returns a caller-owned copy of the underlying storage of the whole // packet. -func (pk *PacketBuffer) ToBuffer() bufferv2.Buffer { +func (pk PacketBufferPtr) ToBuffer() bufferv2.Buffer { b := pk.buf.Clone() b.TrimFront(int64(pk.headerOffset())) return b @@ -295,7 +298,7 @@ func (pk *PacketBuffer) ToBuffer() bufferv2.Buffer { // ToView returns a caller-owned copy of the underlying storage of the whole // packet as a view. -func (pk *PacketBuffer) ToView() *bufferv2.View { +func (pk PacketBufferPtr) ToView() *bufferv2.View { p := bufferv2.NewView(int(pk.buf.Size())) offset := pk.headerOffset() pk.buf.SubApply(offset, int(pk.buf.Size())-offset, func(v *bufferv2.View) { @@ -304,19 +307,19 @@ func (pk *PacketBuffer) ToView() *bufferv2.View { return p } -func (pk *PacketBuffer) headerOffset() int { +func (pk PacketBufferPtr) headerOffset() int { return pk.reserved - pk.pushed } -func (pk *PacketBuffer) headerOffsetOf(typ headerType) int { +func (pk PacketBufferPtr) headerOffsetOf(typ headerType) int { return pk.reserved + pk.headers[typ].offset } -func (pk *PacketBuffer) dataOffset() int { +func (pk PacketBufferPtr) dataOffset() int { return pk.reserved + pk.consumed } -func (pk *PacketBuffer) push(typ headerType, size int) []byte { +func (pk PacketBufferPtr) push(typ headerType, size int) []byte { h := &pk.headers[typ] if h.length > 0 { panic(fmt.Sprintf("push(%s, %d) called after previous push", typ, size)) @@ -331,7 +334,7 @@ func (pk *PacketBuffer) push(typ headerType, size int) []byte { return view.AsSlice() } -func (pk *PacketBuffer) consume(typ headerType, size int) (v []byte, consumed bool) { +func (pk PacketBufferPtr) consume(typ headerType, size int) (v []byte, consumed bool) { h := &pk.headers[typ] if h.length > 0 { panic(fmt.Sprintf("consume must not be called twice: type %s", typ)) @@ -346,7 +349,7 @@ func (pk *PacketBuffer) consume(typ headerType, size int) (v []byte, consumed bo return view.AsSlice(), true } -func (pk *PacketBuffer) headerView(typ headerType) bufferv2.View { +func (pk PacketBufferPtr) headerView(typ headerType) bufferv2.View { h := &pk.headers[typ] if h.length == 0 { return bufferv2.View{} @@ -360,8 +363,8 @@ func (pk *PacketBuffer) 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 *PacketBuffer) Clone() *PacketBuffer { - newPk := pkPool.Get().(*PacketBuffer) +func (pk PacketBufferPtr) Clone() PacketBufferPtr { + newPk := pkPool.Get().(PacketBufferPtr) newPk.reset() newPk.buf = pk.buf.Clone() newPk.reserved = pk.reserved @@ -386,7 +389,7 @@ func (pk *PacketBuffer) Clone() *PacketBuffer { // ReserveHeaderBytes prepends reserved space for headers at the front // of the underlying buf. Can only be called once per packet. -func (pk *PacketBuffer) ReserveHeaderBytes(reserved int) { +func (pk PacketBufferPtr) ReserveHeaderBytes(reserved int) { if pk.reserved != 0 { panic(fmt.Sprintf("ReserveHeaderBytes(...) called on packet with reserved=%d, want reserved=0", pk.reserved)) } @@ -397,7 +400,7 @@ func (pk *PacketBuffer) ReserveHeaderBytes(reserved int) { // Network returns the network header as a header.Network. // // Network should only be called when NetworkHeader has been set. -func (pk *PacketBuffer) Network() header.Network { +func (pk PacketBufferPtr) Network() header.Network { switch netProto := pk.NetworkProtocolNumber; netProto { case header.IPv4ProtocolNumber: return header.IPv4(pk.NetworkHeader().Slice()) @@ -413,8 +416,8 @@ func (pk *PacketBuffer) Network() header.Network { // // See PacketBuffer.Data for details about how a packet buffer holds an inbound // packet. -func (pk *PacketBuffer) CloneToInbound() *PacketBuffer { - newPk := pkPool.Get().(*PacketBuffer) +func (pk PacketBufferPtr) CloneToInbound() PacketBufferPtr { + newPk := pkPool.Get().(PacketBufferPtr) newPk.reset() newPk.buf = pk.buf.Clone() newPk.InitRefs() @@ -429,7 +432,7 @@ func (pk *PacketBuffer) CloneToInbound() *PacketBuffer { // // The returned packet buffer will have the network and transport headers // set if the original packet buffer did. -func (pk *PacketBuffer) DeepCopyForForwarding(reservedHeaderBytes int) *PacketBuffer { +func (pk PacketBufferPtr) DeepCopyForForwarding(reservedHeaderBytes int) PacketBufferPtr { newPk := NewPacketBuffer(PacketBufferOptions{ ReserveHeaderBytes: reservedHeaderBytes, Payload: BufferSince(pk.NetworkHeader()), @@ -457,6 +460,11 @@ func (pk *PacketBuffer) DeepCopyForForwarding(reservedHeaderBytes int) *PacketBu return newPk } +// IsNil returns whether the pointer is logically nil. +func (pk PacketBufferPtr) IsNil() bool { + return pk == nil +} + // headerInfo stores metadata about a header in a packet. // // +stateify savable @@ -471,7 +479,7 @@ type headerInfo struct { // PacketHeader is a handle object to a header in the underlying packet. type PacketHeader struct { - pk *PacketBuffer + pk PacketBufferPtr typ headerType } @@ -513,7 +521,7 @@ func (h PacketHeader) Consume(size int) (v []byte, consumed bool) { // // +stateify savable type PacketData struct { - pk *PacketBuffer + pk PacketBufferPtr } // PullUp returns a contiguous slice of size bytes from the beginning of d. @@ -591,7 +599,7 @@ func (d PacketData) MergeBuffer(b *bufferv2.Buffer) { // MergeFragment appends the data portion of frag to dst. It modifies // frag and frag should not be used again. -func MergeFragment(dst, frag *PacketBuffer) { +func MergeFragment(dst, frag PacketBufferPtr) { frag.buf.TrimFront(int64(frag.dataOffset())) dst.buf.Merge(&frag.buf) } @@ -658,7 +666,7 @@ func (d PacketData) Checksum() uint16 { // Range represents a contiguous subportion of a PacketBuffer. type Range struct { - pk *PacketBuffer + pk PacketBufferPtr offset int length int } diff --git a/pkg/tcpip/stack/packet_buffer_list.go b/pkg/tcpip/stack/packet_buffer_list.go index fbd8ab92c..31107c3ba 100644 --- a/pkg/tcpip/stack/packet_buffer_list.go +++ b/pkg/tcpip/stack/packet_buffer_list.go @@ -21,13 +21,13 @@ package stack // // +stateify savable type PacketBufferList struct { - pbs []*PacketBuffer + pbs []PacketBufferPtr } // AsSlice returns a slice containing the packets in the list. // //go:nosplit -func (pl *PacketBufferList) AsSlice() []*PacketBuffer { +func (pl *PacketBufferList) AsSlice() []PacketBufferPtr { return pl.pbs } @@ -52,7 +52,7 @@ func (pl *PacketBufferList) Len() int { // PushBack inserts the PacketBuffer at the back of the list. // //go:nosplit -func (pl *PacketBufferList) PushBack(pb *PacketBuffer) { +func (pl *PacketBufferList) PushBack(pb PacketBufferPtr) { pl.pbs = append(pl.pbs, pb) } diff --git a/pkg/tcpip/stack/packet_buffer_test.go b/pkg/tcpip/stack/packet_buffer_test.go index 97b7b9de2..23e7b74f0 100644 --- a/pkg/tcpip/stack/packet_buffer_test.go +++ b/pkg/tcpip/stack/packet_buffer_test.go @@ -393,12 +393,12 @@ func TestPacketHeaderConsumeThenPushPanics(t *testing.T) { func TestPacketBufferData(t *testing.T) { for _, tc := range []struct { name string - makePkt func(*testing.T) *PacketBuffer + makePkt func(*testing.T) PacketBufferPtr data string }{ { name: "inbound packet", - makePkt: func(*testing.T) *PacketBuffer { + makePkt: func(*testing.T) PacketBufferPtr { pkt := NewPacketBuffer(PacketBufferOptions{ Payload: buf("aabbbbccccccDATA"), }) @@ -411,7 +411,7 @@ func TestPacketBufferData(t *testing.T) { }, { name: "outbound packet", - makePkt: func(*testing.T) *PacketBuffer { + makePkt: func(*testing.T) PacketBufferPtr { pkt := NewPacketBuffer(PacketBufferOptions{ ReserveHeaderBytes: 12, Payload: buf("DATA"), @@ -551,7 +551,7 @@ type packetContents struct { data []byte } -func checkPacketContents(t *testing.T, prefix string, pk *PacketBuffer, want packetContents) { +func checkPacketContents(t *testing.T, prefix string, pk PacketBufferPtr, want packetContents) { t.Helper() // Headers. checkPacketHeader(t, prefix+"pk.LinkHeader", pk.LinkHeader(), want.link) @@ -591,7 +591,7 @@ func checkPacketContents(t *testing.T, prefix string, pk *PacketBuffer, want pac concatViews(want.transport, want.data)) } -func checkInitialPacketBuffer(t *testing.T, pk *PacketBuffer, opts PacketBufferOptions) { +func checkInitialPacketBuffer(t *testing.T, pk PacketBufferPtr, opts PacketBufferOptions) { t.Helper() reserved := opts.ReserveHeaderBytes if got, want := pk.ReservedHeaderBytes(), reserved; got != want { @@ -624,7 +624,7 @@ func checkViewEqual(t *testing.T, what string, got, want []byte) { } } -func checkData(t *testing.T, pkt *PacketBuffer, want []byte) { +func checkData(t *testing.T, pkt PacketBufferPtr, want []byte) { t.Helper() if got := pkt.Data().AsRange().ToSlice(); !bytes.Equal(got, want) { t.Errorf("pkt.Data().Slices() = 0x%x, want 0x%x", got, want) diff --git a/pkg/tcpip/stack/pending_packets.go b/pkg/tcpip/stack/pending_packets.go index 4b45f6758..8faa5b60a 100644 --- a/pkg/tcpip/stack/pending_packets.go +++ b/pkg/tcpip/stack/pending_packets.go @@ -30,7 +30,7 @@ const ( type pendingPacket struct { routeInfo RouteInfo - pkt *PacketBuffer + pkt PacketBufferPtr } // packetsPendingLinkResolution is a queue of packets pending link resolution. @@ -55,7 +55,7 @@ type packetsPendingLinkResolution struct { } } -func (f *packetsPendingLinkResolution) incrementOutgoingPacketErrors(pkt *PacketBuffer) { +func (f *packetsPendingLinkResolution) incrementOutgoingPacketErrors(pkt PacketBufferPtr) { f.nic.stack.stats.IP.OutgoingPacketErrors.Increment() if ipEndpointStats, ok := f.nic.getNetworkEndpoint(pkt.NetworkProtocolNumber).Stats().(IPNetworkEndpointStats); ok { @@ -114,7 +114,7 @@ func (f *packetsPendingLinkResolution) dequeue(ch <-chan struct{}, linkAddr tcpi // If the maximum number of pending resolutions is reached, the packets // associated with the oldest link resolution will be dequeued as if they failed // link resolution. -func (f *packetsPendingLinkResolution) enqueue(r *Route, pkt *PacketBuffer) tcpip.Error { +func (f *packetsPendingLinkResolution) enqueue(r *Route, pkt PacketBufferPtr) tcpip.Error { f.mu.Lock() // Make sure we attempt resolution while holding f's lock so that we avoid // a race where link resolution completes before we enqueue the packets. diff --git a/pkg/tcpip/stack/registration.go b/pkg/tcpip/stack/registration.go index e21d94510..cf3bfe3b2 100644 --- a/pkg/tcpip/stack/registration.go +++ b/pkg/tcpip/stack/registration.go @@ -105,12 +105,12 @@ type TransportEndpoint interface { // transport endpoint. It sets the packet buffer's transport header. // // HandlePacket may modify the packet. - HandlePacket(TransportEndpointID, *PacketBuffer) + HandlePacket(TransportEndpointID, PacketBufferPtr) // HandleError is called when the transport endpoint receives an error. // // HandleError takes may modify the packet buffer. - HandleError(TransportError, *PacketBuffer) + HandleError(TransportError, PacketBufferPtr) // Abort initiates an expedited endpoint teardown. It puts the endpoint // in a closed state and frees all resources associated with it. This @@ -138,7 +138,7 @@ type RawTransportEndpoint interface { // layer up. // // HandlePacket may modify the packet. - HandlePacket(*PacketBuffer) + HandlePacket(PacketBufferPtr) } // PacketEndpoint is the interface that needs to be implemented by packet @@ -156,7 +156,7 @@ type PacketEndpoint interface { // should construct its own ethernet header for applications. // // HandlePacket may modify pkt. - HandlePacket(nicID tcpip.NICID, netProto tcpip.NetworkProtocolNumber, pkt *PacketBuffer) + HandlePacket(nicID tcpip.NICID, netProto tcpip.NetworkProtocolNumber, pkt PacketBufferPtr) } // UnknownDestinationPacketDisposition enumerates the possible return values from @@ -206,7 +206,7 @@ type TransportProtocol interface { // // HandleUnknownDestinationPacket may modify the packet if it handles // the issue. - HandleUnknownDestinationPacket(TransportEndpointID, *PacketBuffer) UnknownDestinationPacketDisposition + HandleUnknownDestinationPacket(TransportEndpointID, PacketBufferPtr) UnknownDestinationPacketDisposition // SetOption allows enabling/disabling protocol specific features. // SetOption returns an error if the option is not supported or the @@ -235,7 +235,7 @@ type TransportProtocol interface { // Parse sets pkt.TransportHeader and trims pkt.Data appropriately. It does // neither and returns false if pkt.Data is too small, i.e. pkt.Data.Size() < // MinimumPacketSize() - Parse(pkt *PacketBuffer) (ok bool) + Parse(pkt PacketBufferPtr) (ok bool) } // TransportPacketDisposition is the result from attempting to deliver a packet @@ -267,18 +267,18 @@ type TransportDispatcher interface { // pkt.NetworkHeader must be set before calling DeliverTransportPacket. // // DeliverTransportPacket may modify the packet. - DeliverTransportPacket(tcpip.TransportProtocolNumber, *PacketBuffer) TransportPacketDisposition + DeliverTransportPacket(tcpip.TransportProtocolNumber, PacketBufferPtr) TransportPacketDisposition // DeliverTransportError delivers an error to the appropriate transport // endpoint. // // DeliverTransportError may modify the packet buffer. - DeliverTransportError(local, remote tcpip.Address, _ tcpip.NetworkProtocolNumber, _ tcpip.TransportProtocolNumber, _ TransportError, _ *PacketBuffer) + DeliverTransportError(local, remote tcpip.Address, _ tcpip.NetworkProtocolNumber, _ tcpip.TransportProtocolNumber, _ TransportError, _ PacketBufferPtr) // DeliverRawPacket delivers a packet to any subscribed raw sockets. // // DeliverRawPacket does NOT take ownership of the packet buffer. - DeliverRawPacket(tcpip.TransportProtocolNumber, *PacketBuffer) + DeliverRawPacket(tcpip.TransportProtocolNumber, PacketBufferPtr) } // PacketLooping specifies where an outbound packet should be sent. @@ -725,13 +725,13 @@ type NetworkInterface interface { CheckLocalAddress(tcpip.NetworkProtocolNumber, tcpip.Address) bool // WritePacketToRemote writes the packet to the given remote link address. - WritePacketToRemote(tcpip.LinkAddress, *PacketBuffer) tcpip.Error + WritePacketToRemote(tcpip.LinkAddress, PacketBufferPtr) tcpip.Error // WritePacket writes a packet through the given route. // // WritePacket may modify the packet buffer. The packet buffer's // network and transport header must be set. - WritePacket(*Route, *PacketBuffer) tcpip.Error + WritePacket(*Route, PacketBufferPtr) tcpip.Error // HandleNeighborProbe processes an incoming neighbor probe (e.g. ARP // request or NDP Neighbor Solicitation). @@ -749,7 +749,7 @@ type NetworkInterface interface { type LinkResolvableNetworkEndpoint interface { // HandleLinkResolutionFailure is called when link resolution prevents the // argument from having been sent. - HandleLinkResolutionFailure(*PacketBuffer) + HandleLinkResolutionFailure(PacketBufferPtr) } // NetworkEndpoint is the interface that needs to be implemented by endpoints @@ -787,17 +787,17 @@ type NetworkEndpoint interface { // WritePacket writes a packet to the given destination address and // protocol. It may modify pkt. pkt.TransportHeader must have // already been set. - WritePacket(r *Route, params NetworkHeaderParams, pkt *PacketBuffer) tcpip.Error + WritePacket(r *Route, params NetworkHeaderParams, pkt PacketBufferPtr) tcpip.Error // WriteHeaderIncludedPacket writes a packet that includes a network // header to the given destination address. It may modify pkt. - WriteHeaderIncludedPacket(r *Route, pkt *PacketBuffer) tcpip.Error + WriteHeaderIncludedPacket(r *Route, pkt PacketBufferPtr) tcpip.Error // HandlePacket is called by the link layer when new packets arrive to // this network endpoint. It sets pkt.NetworkHeader. // // HandlePacket may modify pkt. - HandlePacket(pkt *PacketBuffer) + HandlePacket(pkt PacketBufferPtr) // Close is called when the endpoint is removed from a stack. Close() @@ -896,7 +896,7 @@ type NetworkProtocol interface { // - Whether there is an encapsulated transport protocol payload (e.g. ARP // does not encapsulate anything). // - Whether pkt.Data was large enough to parse and set pkt.NetworkHeader. - Parse(pkt *PacketBuffer) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) + Parse(pkt PacketBufferPtr) (proto tcpip.TransportProtocolNumber, hasTransportHdr bool, ok bool) } // UnicastSourceAndMulticastDestination is a tuple that represents a unicast @@ -1012,14 +1012,14 @@ type NetworkDispatcher interface { // If the link-layer has a header, the packet's link header must be populated. // // DeliverNetworkPacket may modify pkt. - DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) + DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt PacketBufferPtr) // DeliverLinkPacket delivers a packet to any interested packet endpoints. // // This method should be called with both incoming and outgoing packets. // // If the link-layer has a header, the packet's link header must be populated. - DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer, incoming bool) + DeliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt PacketBufferPtr, incoming bool) } // LinkEndpointCapabilities is the type associated with the capabilities @@ -1105,7 +1105,7 @@ type NetworkLinkEndpoint interface { ARPHardwareType() header.ARPHardwareType // AddHeader adds a link layer header to the packet if required. - AddHeader(*PacketBuffer) + AddHeader(PacketBufferPtr) } // QueueingDiscipline provides a queueing strategy for outgoing packets (e.g @@ -1119,7 +1119,7 @@ type QueueingDiscipline interface { // To participate in transparent bridging, a LinkEndpoint implementation // should call eth.Encode with header.EthernetFields.SrcAddr set to // pkg.EgressRoute.LocalLinkAddress if it is provided. - WritePacket(*PacketBuffer) tcpip.Error + WritePacket(PacketBufferPtr) tcpip.Error Close() } @@ -1140,7 +1140,7 @@ type InjectableLinkEndpoint interface { LinkEndpoint // InjectInbound injects an inbound packet. - InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) + InjectInbound(protocol tcpip.NetworkProtocolNumber, pkt PacketBufferPtr) // InjectOutbound writes a fully formed outbound packet directly to the // link. diff --git a/pkg/tcpip/stack/route.go b/pkg/tcpip/stack/route.go index c8db589a1..4d4c01898 100644 --- a/pkg/tcpip/stack/route.go +++ b/pkg/tcpip/stack/route.go @@ -490,7 +490,7 @@ func (r *Route) isValidForOutgoingRLocked() bool { } // WritePacket writes the packet through the given route. -func (r *Route) WritePacket(params NetworkHeaderParams, pkt *PacketBuffer) tcpip.Error { +func (r *Route) WritePacket(params NetworkHeaderParams, pkt PacketBufferPtr) tcpip.Error { if !r.isValidForOutgoing() { return &tcpip.ErrInvalidEndpointState{} } @@ -500,7 +500,7 @@ func (r *Route) WritePacket(params NetworkHeaderParams, pkt *PacketBuffer) tcpip // WriteHeaderIncludedPacket writes a packet already containing a network // header through the given route. -func (r *Route) WriteHeaderIncludedPacket(pkt *PacketBuffer) tcpip.Error { +func (r *Route) WriteHeaderIncludedPacket(pkt PacketBufferPtr) tcpip.Error { if !r.isValidForOutgoing() { return &tcpip.ErrInvalidEndpointState{} } diff --git a/pkg/tcpip/stack/stack.go b/pkg/tcpip/stack/stack.go index 0f173a10c..5b444a47e 100644 --- a/pkg/tcpip/stack/stack.go +++ b/pkg/tcpip/stack/stack.go @@ -46,7 +46,7 @@ const ( type transportProtocolState struct { proto TransportProtocol - defaultHandler func(id TransportEndpointID, pkt *PacketBuffer) bool + defaultHandler func(id TransportEndpointID, pkt PacketBufferPtr) bool } // ResumableEndpoint is an endpoint that needs to be resumed after restore. @@ -487,7 +487,7 @@ func (s *Stack) TransportProtocolOption(transport tcpip.TransportProtocolNumber, // // It must be called only during initialization of the stack. Changing it as the // stack is operating is not supported. -func (s *Stack) SetTransportProtocolHandler(p tcpip.TransportProtocolNumber, h func(TransportEndpointID, *PacketBuffer) bool) { +func (s *Stack) SetTransportProtocolHandler(p tcpip.TransportProtocolNumber, h func(TransportEndpointID, PacketBufferPtr) bool) { state := s.transportProtocols[p] if state != nil { state.defaultHandler = h @@ -2062,7 +2062,7 @@ const ( // ParsePacketBufferTransport parses the provided packet buffer's transport // header. -func (s *Stack) ParsePacketBufferTransport(protocol tcpip.TransportProtocolNumber, pkt *PacketBuffer) ParseResult { +func (s *Stack) ParsePacketBufferTransport(protocol tcpip.TransportProtocolNumber, pkt PacketBufferPtr) ParseResult { pkt.TransportProtocolNumber = protocol // Parse the transport header if present. state, ok := s.transportProtocols[protocol] diff --git a/pkg/tcpip/stack/stack_test.go b/pkg/tcpip/stack/stack_test.go index 0bb5250af..42dcd8506 100644 --- a/pkg/tcpip/stack/stack_test.go +++ b/pkg/tcpip/stack/stack_test.go @@ -122,7 +122,7 @@ func (*fakeNetworkEndpoint) DefaultTTL() uint8 { return 123 } -func (f *fakeNetworkEndpoint) HandlePacket(pkt *stack.PacketBuffer) { +func (f *fakeNetworkEndpoint) HandlePacket(pkt stack.PacketBufferPtr) { if _, _, ok := f.proto.Parse(pkt); !ok { return } @@ -179,7 +179,7 @@ func (f *fakeNetworkEndpoint) NetworkProtocolNumber() tcpip.NetworkProtocolNumbe return f.proto.Number() } -func (f *fakeNetworkEndpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, pkt *stack.PacketBuffer) tcpip.Error { +func (f *fakeNetworkEndpoint) WritePacket(r *stack.Route, params stack.NetworkHeaderParams, pkt stack.PacketBufferPtr) tcpip.Error { // Increment the sent packet count in the protocol descriptor. f.proto.sendPacketCount[int(r.RemoteAddress()[0])%len(f.proto.sendPacketCount)]++ @@ -206,7 +206,7 @@ func (*fakeNetworkEndpoint) WritePackets(*stack.Route, stack.PacketBufferList, s panic("not implemented") } -func (*fakeNetworkEndpoint) WriteHeaderIncludedPacket(*stack.Route, *stack.PacketBuffer) tcpip.Error { +func (*fakeNetworkEndpoint) WriteHeaderIncludedPacket(*stack.Route, stack.PacketBufferPtr) tcpip.Error { return &tcpip.ErrNotSupported{} } @@ -307,7 +307,7 @@ func (*fakeNetworkProtocol) Close() {} func (*fakeNetworkProtocol) Wait() {} // Parse implements NetworkProtocol.Parse. -func (*fakeNetworkProtocol) Parse(pkt *stack.PacketBuffer) (tcpip.TransportProtocolNumber, bool, bool) { +func (*fakeNetworkProtocol) Parse(pkt stack.PacketBufferPtr) (tcpip.TransportProtocolNumber, bool, bool) { hdr, ok := pkt.NetworkHeader().Consume(fakeNetHeaderLen) if !ok { return 0, false, false @@ -4796,7 +4796,7 @@ func TestFindRouteWithForwarding(t *testing.T) { t.Errorf("got %d unexpected packets from ep1", n) } pkt := ep2.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("packet not sent through ep2") } defer pkt.DecRef() @@ -5277,7 +5277,7 @@ func TestWritePacketToRemote(t *testing.T) { } pkt := e.Read() - if got, want := pkt != nil, true; got != want { + if got, want := !pkt.IsNil(), true; got != want { t.Fatalf("e.Read() = %t, want %t", got, want) } defer pkt.DecRef() @@ -5299,7 +5299,7 @@ func TestWritePacketToRemote(t *testing.T) { t.Fatalf("s.WritePacketToRemote(_, _, _, _) = %s, want = %s", err, &tcpip.ErrUnknownDevice{}) } pkt := e.Read() - if got, want := pkt != nil, false; got != want { + if got, want := !pkt.IsNil(), false; got != want { t.Fatalf("e.Read() = %t, %v; want %t", got, pkt, want) } }) diff --git a/pkg/tcpip/stack/transport_demuxer.go b/pkg/tcpip/stack/transport_demuxer.go index 8996c25eb..a9602ce1f 100644 --- a/pkg/tcpip/stack/transport_demuxer.go +++ b/pkg/tcpip/stack/transport_demuxer.go @@ -156,7 +156,7 @@ func (epsByNIC *endpointsByNIC) transportEndpoints() []TransportEndpoint { // handlePacket is called by the stack when new packets arrive to this transport // endpoint. It returns false if the packet could not be matched to any // transport endpoint, true otherwise. -func (epsByNIC *endpointsByNIC) handlePacket(id TransportEndpointID, pkt *PacketBuffer) bool { +func (epsByNIC *endpointsByNIC) handlePacket(id TransportEndpointID, pkt PacketBufferPtr) bool { epsByNIC.mu.RLock() mpep, ok := epsByNIC.endpoints[pkt.NICID] @@ -188,7 +188,7 @@ func (epsByNIC *endpointsByNIC) handlePacket(id TransportEndpointID, pkt *Packet } // handleError delivers an error to the transport endpoint identified by id. -func (epsByNIC *endpointsByNIC) handleError(n *nic, id TransportEndpointID, transErr TransportError, pkt *PacketBuffer) { +func (epsByNIC *endpointsByNIC) handleError(n *nic, id TransportEndpointID, transErr TransportError, pkt PacketBufferPtr) { epsByNIC.mu.RLock() mpep, ok := epsByNIC.endpoints[n.ID()] @@ -279,7 +279,7 @@ type transportDemuxer struct { // the dispatcher to delivery packets to the QueuePacket method instead of // calling HandlePacket directly on the endpoint. type queuedTransportProtocol interface { - QueuePacket(ep TransportEndpoint, id TransportEndpointID, pkt *PacketBuffer) + QueuePacket(ep TransportEndpoint, id TransportEndpointID, pkt PacketBufferPtr) } func newTransportDemuxer(stack *Stack) *transportDemuxer { @@ -401,7 +401,7 @@ func (ep *multiPortEndpoint) selectEndpoint(id TransportEndpointID, seed uint32) return ep.endpoints[idx] } -func (ep *multiPortEndpoint) handlePacketAll(id TransportEndpointID, pkt *PacketBuffer) { +func (ep *multiPortEndpoint) handlePacketAll(id TransportEndpointID, pkt PacketBufferPtr) { ep.mu.RLock() queuedProtocol, mustQueue := ep.demux.queuedProtocols[protocolIDs{ep.netProto, ep.transProto}] // HandlePacket may modify pkt, so each endpoint needs @@ -547,7 +547,7 @@ func (d *transportDemuxer) unregisterEndpoint(netProtos []tcpip.NetworkProtocolN // deliverPacket attempts to find one or more matching transport endpoints, and // then, if matches are found, delivers the packet to them. Returns true if // the packet no longer needs to be handled. -func (d *transportDemuxer) deliverPacket(protocol tcpip.TransportProtocolNumber, pkt *PacketBuffer, id TransportEndpointID) bool { +func (d *transportDemuxer) deliverPacket(protocol tcpip.TransportProtocolNumber, pkt PacketBufferPtr, id TransportEndpointID) bool { eps, ok := d.protocol[protocolIDs{pkt.NetworkProtocolNumber, protocol}] if !ok { return false @@ -600,7 +600,7 @@ func (d *transportDemuxer) deliverPacket(protocol tcpip.TransportProtocolNumber, // deliverRawPacket attempts to deliver the given packet and returns whether it // was delivered successfully. -func (d *transportDemuxer) deliverRawPacket(protocol tcpip.TransportProtocolNumber, pkt *PacketBuffer) bool { +func (d *transportDemuxer) deliverRawPacket(protocol tcpip.TransportProtocolNumber, pkt PacketBufferPtr) bool { eps, ok := d.protocol[protocolIDs{pkt.NetworkProtocolNumber, protocol}] if !ok { return false @@ -634,7 +634,7 @@ func (d *transportDemuxer) deliverRawPacket(protocol tcpip.TransportProtocolNumb // endpoint. // // Returns true if the error was delivered. -func (d *transportDemuxer) deliverError(n *nic, net tcpip.NetworkProtocolNumber, trans tcpip.TransportProtocolNumber, transErr TransportError, pkt *PacketBuffer, id TransportEndpointID) bool { +func (d *transportDemuxer) deliverError(n *nic, net tcpip.NetworkProtocolNumber, trans tcpip.TransportProtocolNumber, transErr TransportError, pkt PacketBufferPtr, id TransportEndpointID) bool { eps, ok := d.protocol[protocolIDs{net, trans}] if !ok { return false @@ -719,7 +719,7 @@ func (d *transportDemuxer) unregisterRawEndpoint(netProto tcpip.NetworkProtocolN eps.mu.Unlock() } -func isInboundMulticastOrBroadcast(pkt *PacketBuffer, localAddr tcpip.Address) bool { +func isInboundMulticastOrBroadcast(pkt PacketBufferPtr, localAddr tcpip.Address) bool { return pkt.NetworkPacketInfo.LocalAddressBroadcast || header.IsV4MulticastAddress(localAddr) || header.IsV6MulticastAddress(localAddr) } diff --git a/pkg/tcpip/stack/transport_test.go b/pkg/tcpip/stack/transport_test.go index 5ac859b8d..61796c718 100644 --- a/pkg/tcpip/stack/transport_test.go +++ b/pkg/tcpip/stack/transport_test.go @@ -213,7 +213,7 @@ func (*fakeTransportEndpoint) GetRemoteAddress() (tcpip.FullAddress, tcpip.Error return tcpip.FullAddress{}, nil } -func (f *fakeTransportEndpoint) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) { +func (f *fakeTransportEndpoint) HandlePacket(id stack.TransportEndpointID, pkt stack.PacketBufferPtr) { // Increment the number of received packets. f.proto.packetCount++ if f.acceptQueue == nil { @@ -239,7 +239,7 @@ func (f *fakeTransportEndpoint) HandlePacket(id stack.TransportEndpointID, pkt * f.acceptQueue = append(f.acceptQueue, ep) } -func (f *fakeTransportEndpoint) HandleError(stack.TransportError, *stack.PacketBuffer) { +func (f *fakeTransportEndpoint) HandleError(stack.TransportError, stack.PacketBufferPtr) { // Increment the number of received control packets. f.proto.controlCount++ } @@ -298,7 +298,7 @@ func (*fakeTransportProtocol) ParsePorts([]byte) (src, dst uint16, err tcpip.Err return 0, 0, nil } -func (*fakeTransportProtocol) HandleUnknownDestinationPacket(stack.TransportEndpointID, *stack.PacketBuffer) stack.UnknownDestinationPacketDisposition { +func (*fakeTransportProtocol) HandleUnknownDestinationPacket(stack.TransportEndpointID, stack.PacketBufferPtr) stack.UnknownDestinationPacketDisposition { return stack.UnknownDestinationPacketHandled } @@ -338,7 +338,7 @@ func (*fakeTransportProtocol) Pause() {} func (*fakeTransportProtocol) Resume() {} // Parse implements TransportProtocol.Parse. -func (*fakeTransportProtocol) Parse(pkt *stack.PacketBuffer) bool { +func (*fakeTransportProtocol) Parse(pkt stack.PacketBufferPtr) bool { if _, ok := pkt.TransportHeader().Consume(fakeTransHeaderLen); ok { pkt.TransportProtocolNumber = fakeTransNumber return true diff --git a/pkg/tcpip/tests/integration/forward_test.go b/pkg/tcpip/tests/integration/forward_test.go index edbffcb27..7f5280ba7 100644 --- a/pkg/tcpip/tests/integration/forward_test.go +++ b/pkg/tcpip/tests/integration/forward_test.go @@ -469,7 +469,7 @@ func TestUnicastForwarding(t *testing.T) { test.rx(e1, test.srcAddr, test.dstAddr) p := e2.Read() - if (p != nil) != test.expectForward { + if (!p.IsNil()) != test.expectForward { t.Fatalf("got e2.Read() = %#v, want = (_ == nil) = %t", p, test.expectForward) } @@ -629,15 +629,15 @@ func TestPerInterfaceForwarding(t *testing.T) { }) test.rx(subTest.nicEP, test.srcAddr, test.dstAddr) - if p := subTest.nicEP.Read(); p != nil { + if p := subTest.nicEP.Read(); !p.IsNil() { t.Errorf("unexpectedly got a response from the interface the packet arrived on: %#v", p) p.DecRef() } p := subTest.otherNICEP.Read() - if (p != nil) != subTest.expectForwarding { + if (!p.IsNil()) != subTest.expectForwarding { t.Errorf("got otherNICEP.Read() = (%#v, %t), want = (_, %t)", p, ok, subTest.expectForwarding) } - if p != nil { + if !p.IsNil() { payload := stack.PayloadSince(p.NetworkHeader()) defer payload.Release() test.checker(t, payload) diff --git a/pkg/tcpip/tests/integration/iptables_test.go b/pkg/tcpip/tests/integration/iptables_test.go index e200e3a15..078a246df 100644 --- a/pkg/tcpip/tests/integration/iptables_test.go +++ b/pkg/tcpip/tests/integration/iptables_test.go @@ -51,7 +51,7 @@ func (*inputIfNameMatcher) Name() string { return "inputIfNameMatcher" } -func (im *inputIfNameMatcher) Match(hook stack.Hook, _ *stack.PacketBuffer, inNicName, _ string) (bool, bool) { +func (im *inputIfNameMatcher) Match(hook stack.Hook, _ stack.PacketBufferPtr, inNicName, _ string) (bool, bool) { return (hook == stack.Input && im.name != "" && im.name == inNicName), false } @@ -109,7 +109,7 @@ func genStackV4(t *testing.T) (*stack.Stack, *channel.Endpoint) { return s, e } -func genPacketV6() *stack.PacketBuffer { +func genPacketV6() stack.PacketBufferPtr { pktSize := header.IPv6MinimumSize + payloadSize hdr := prependable.New(pktSize) ip := header.IPv6(hdr.Prepend(pktSize)) @@ -124,7 +124,7 @@ func genPacketV6() *stack.PacketBuffer { return stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buf}) } -func genPacketV4() *stack.PacketBuffer { +func genPacketV4() stack.PacketBufferPtr { pktSize := header.IPv4MinimumSize + payloadSize hdr := prependable.New(pktSize) ip := header.IPv4(hdr.Prepend(pktSize)) @@ -150,7 +150,7 @@ func TestIPTablesStatsForInput(t *testing.T) { name string setupStack func(*testing.T) (*stack.Stack, *channel.Endpoint) setupFilter func(*testing.T, *stack.Stack) - genPacket func() *stack.PacketBuffer + genPacket func() stack.PacketBufferPtr proto tcpip.NetworkProtocolNumber expectReceived int expectInputDropped int @@ -362,7 +362,7 @@ func (*udpSourcePortMatcher) Name() string { return "udpSourcePortMatcher" } -func (m *udpSourcePortMatcher) Match(_ stack.Hook, pkt *stack.PacketBuffer, _, _ string) (matches, hotdrop bool) { +func (m *udpSourcePortMatcher) Match(_ stack.Hook, pkt stack.PacketBufferPtr, _, _ string) (matches, hotdrop bool) { udp := header.UDP(pkt.TransportHeader().Slice()) if len(udp) < header.UDPMinimumSize { // Drop immediately as the packet is invalid. @@ -938,7 +938,7 @@ func TestForwardingHook(t *testing.T) { } p := e2.Read() - if (p != nil) != expectTransmitPacket { + if (!p.IsNil()) != expectTransmitPacket { t.Fatalf("got e2.Read() = %#v, want = (_ == nil) = %t", p, expectTransmitPacket) } if expectTransmitPacket { @@ -1179,16 +1179,16 @@ func TestFilteringEchoPacketsWithLocalForwarding(t *testing.T) { expectPacket := subTest.expectResult == noneDropped p := e1.Read() - if (p != nil) != expectPacket { + if (!p.IsNil()) != expectPacket { t.Errorf("got e1.Read() = %#v, want = (_ == nil) = %t", p, expectPacket) } - if p != nil { + if !p.IsNil() { payload := stack.PayloadSince(p.NetworkHeader()) defer payload.Release() test.checker(t, payload) p.DecRef() } - if p := e2.Read(); p != nil { + if p := e2.Read(); !p.IsNil() { t.Errorf("got e1.Read() = %#v, want = nil)", p) p.DecRef() } @@ -1536,7 +1536,7 @@ func TestNATEcho(t *testing.T) { Payload: bufferv2.MakeWithData(test.echoPkt(natTypeTest.requestSrc, natTypeTest.requestDst, false /* reply */)), })) pkt := ep1.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to read a packet on ep1") } payload := stack.PayloadSince(pkt.NetworkHeader()) @@ -1555,7 +1555,7 @@ func TestNATEcho(t *testing.T) { Payload: bufferv2.MakeWithData(test.echoPkt(natTypeTest.expectedRequestDst, natTypeTest.expectedRequestSrc, true /* reply */)), })) pkt := ep2.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to read a packet on ep2") } payload := stack.PayloadSince(pkt.NetworkHeader()) @@ -2525,7 +2525,7 @@ func TestNATICMPError(t *testing.T) { { pkt := ep1.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to read a packet on ep1") } pktView := stack.PayloadSince(pkt.NetworkHeader()) @@ -2546,7 +2546,7 @@ func TestNATICMPError(t *testing.T) { pkt := ep2.Read() expectResponse := icmpType.expectResponse && trimTest.expectNATedICMP - if (pkt != nil) != expectResponse { + if (!pkt.IsNil()) != expectResponse { t.Fatalf("got ep2.Read() = %#v, want = (_ == nil) = %t", pkt, expectResponse) } if !expectResponse { @@ -2894,7 +2894,7 @@ func TestSNATHandlePortOrIdentConflicts(t *testing.T) { })) pkt := ep1.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to read a packet on ep1") } pktView := stack.PayloadSince(pkt.NetworkHeader()) @@ -3035,7 +3035,7 @@ type icmpv4Matcher struct { icmpType header.ICMPv4Type } -func (m *icmpv4Matcher) Match(_ stack.Hook, pkt *stack.PacketBuffer, _, _ string) (matches bool, hotdrop bool) { +func (m *icmpv4Matcher) Match(_ stack.Hook, pkt stack.PacketBufferPtr, _, _ string) (matches bool, hotdrop bool) { if pkt.NetworkProtocolNumber != header.IPv4ProtocolNumber { return false, false } @@ -3051,7 +3051,7 @@ type icmpv6Matcher struct { icmpType header.ICMPv6Type } -func (m *icmpv6Matcher) Match(_ stack.Hook, pkt *stack.PacketBuffer, _, _ string) (matches bool, hotdrop bool) { +func (m *icmpv6Matcher) Match(_ stack.Hook, pkt stack.PacketBufferPtr, _, _ string) (matches bool, hotdrop bool) { if pkt.NetworkProtocolNumber != header.IPv6ProtocolNumber { return false, false } @@ -3301,7 +3301,7 @@ func TestRejectWith(t *testing.T) { { pkt := ep1.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected to read a packet on ep1") } payload := stack.PayloadSince(pkt.NetworkHeader()) @@ -3330,7 +3330,7 @@ func TestInvalidTransportHeader(t *testing.T) { tests := []struct { name string setupStack func(*testing.T) (*stack.Stack, *channel.Endpoint) - genPacket func(int8) *stack.PacketBuffer + genPacket func(int8) stack.PacketBufferPtr offset int8 }{ { @@ -3398,7 +3398,7 @@ func TestInvalidTransportHeader(t *testing.T) { } } -func genTCP4(offset int8) *stack.PacketBuffer { +func genTCP4(offset int8) stack.PacketBufferPtr { pktSize := header.IPv4MinimumSize + header.TCPMinimumSize hdr := prependable.New(pktSize) @@ -3430,7 +3430,7 @@ func genTCP4(offset int8) *stack.PacketBuffer { return stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buf}) } -func genTCP6(offset int8) *stack.PacketBuffer { +func genTCP6(offset int8) stack.PacketBufferPtr { pktSize := header.IPv6MinimumSize + header.TCPMinimumSize hdr := prependable.New(pktSize) @@ -3456,7 +3456,7 @@ func genTCP6(offset int8) *stack.PacketBuffer { return stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buf}) } -func genUDP4(offset int8) *stack.PacketBuffer { +func genUDP4(offset int8) stack.PacketBufferPtr { pktSize := header.IPv4MinimumSize + header.UDPMinimumSize hdr := prependable.New(pktSize) @@ -3487,7 +3487,7 @@ func genUDP4(offset int8) *stack.PacketBuffer { return stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buf}) } -func genUDP6(offset int8) *stack.PacketBuffer { +func genUDP6(offset int8) stack.PacketBufferPtr { pktSize := header.IPv6MinimumSize + header.UDPMinimumSize hdr := prependable.New(pktSize) diff --git a/pkg/tcpip/tests/integration/link_resolution_test.go b/pkg/tcpip/tests/integration/link_resolution_test.go index d668936b2..8fa0cc3da 100644 --- a/pkg/tcpip/tests/integration/link_resolution_test.go +++ b/pkg/tcpip/tests/integration/link_resolution_test.go @@ -439,7 +439,7 @@ func TestForwardingWithLinkResolutionFailure(t *testing.T) { utils.RxICMPv6EchoRequest(e, src, dst, ttl) } - arpChecker := func(t *testing.T, request *stack.PacketBuffer, src, dst tcpip.Address) { + arpChecker := func(t *testing.T, request stack.PacketBufferPtr, src, dst tcpip.Address) { if request.NetworkProtocolNumber != arp.ProtocolNumber { t.Errorf("got request.NetworkProtocolNumber = %d, want = %d", request.NetworkProtocolNumber, arp.ProtocolNumber) } @@ -461,7 +461,7 @@ func TestForwardingWithLinkResolutionFailure(t *testing.T) { } } - ndpChecker := func(t *testing.T, request *stack.PacketBuffer, src, dst tcpip.Address) { + ndpChecker := func(t *testing.T, request stack.PacketBufferPtr, src, dst tcpip.Address) { if request.NetworkProtocolNumber != header.IPv6ProtocolNumber { t.Fatalf("got Proto = %d, want = %d", request.NetworkProtocolNumber, header.IPv6ProtocolNumber) } @@ -517,7 +517,7 @@ func TestForwardingWithLinkResolutionFailure(t *testing.T) { outgoingAddr tcpip.AddressWithPrefix transportProtocol func(*stack.Stack) stack.TransportProtocol rx func(*channel.Endpoint, tcpip.Address, tcpip.Address) - linkResolutionRequestChecker func(*testing.T, *stack.PacketBuffer, tcpip.Address, tcpip.Address) + linkResolutionRequestChecker func(*testing.T, stack.PacketBufferPtr, tcpip.Address, tcpip.Address) icmpReplyChecker func(*testing.T, *bufferv2.View, tcpip.Address, tcpip.Address) mtu uint32 }{ @@ -625,7 +625,7 @@ func TestForwardingWithLinkResolutionFailure(t *testing.T) { for i := 0; i < int(nudConfigs.MaxMulticastProbes); i++ { request := outgoingEndpoint.Read() - if request == nil { + if request.IsNil() { t.Fatal("expected ARP packet through outgoing NIC") } @@ -641,7 +641,7 @@ func TestForwardingWithLinkResolutionFailure(t *testing.T) { // link resolution fails, and this dequeue is what triggers the ICMP // error. reply := incomingEndpoint.Read() - if reply == nil { + if reply.IsNil() { t.Fatal("expected ICMP packet through incoming NIC") } @@ -653,7 +653,7 @@ func TestForwardingWithLinkResolutionFailure(t *testing.T) { // Since link resolution failed, we don't expect the packet to be // forwarded. forwardedPacket := outgoingEndpoint.Read() - if forwardedPacket != nil { + if !forwardedPacket.IsNil() { t.Fatalf("expected no ICMP Echo packet through outgoing NIC, instead found: %#v", forwardedPacket) } diff --git a/pkg/tcpip/tests/integration/multicast_broadcast_test.go b/pkg/tcpip/tests/integration/multicast_broadcast_test.go index 0fe229fc6..57416b607 100644 --- a/pkg/tcpip/tests/integration/multicast_broadcast_test.go +++ b/pkg/tcpip/tests/integration/multicast_broadcast_test.go @@ -146,7 +146,7 @@ func TestPingMulticastBroadcast(t *testing.T) { test.rxICMP(e, test.srcAddr, test.dstAddr, ttl) pkt := e.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("expected ICMP response") } defer pkt.DecRef() diff --git a/pkg/tcpip/tests/integration/multicast_forward_test.go b/pkg/tcpip/tests/integration/multicast_forward_test.go index 26da29ec2..7662ec181 100644 --- a/pkg/tcpip/tests/integration/multicast_forward_test.go +++ b/pkg/tcpip/tests/integration/multicast_forward_test.go @@ -165,7 +165,7 @@ func getEndpointAddr(protocol tcpip.NetworkProtocolNumber, addrType endpointAddr } } -func checkEchoRequest(t *testing.T, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer, srcAddr, dstAddr tcpip.Address, ttl uint8) { +func checkEchoRequest(t *testing.T, protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr, srcAddr, dstAddr tcpip.Address, ttl uint8) { payload := stack.PayloadSince(pkt.NetworkHeader()) defer payload.Release() switch protocol { @@ -192,7 +192,7 @@ func checkEchoRequest(t *testing.T, protocol tcpip.NetworkProtocolNumber, pkt *s } } -func checkEchoReply(t *testing.T, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer, srcAddr, dstAddr tcpip.Address) { +func checkEchoReply(t *testing.T, protocol tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr, srcAddr, dstAddr tcpip.Address) { payload := stack.PayloadSince(pkt.NetworkHeader()) defer payload.Release() switch protocol { @@ -463,7 +463,7 @@ func TestAddMulticastRoute(t *testing.T) { injectPacket(incomingEp, protocol, srcAddr, dstAddr, packetTTL) p := incomingEp.Read() - if p != nil { + if !p.IsNil() { // An ICMP error should never be sent in response to a multicast packet. t.Fatalf("got incomingEp.Read() = %#v, want = nil", p) } @@ -502,7 +502,7 @@ func TestAddMulticastRoute(t *testing.T) { p := outgoingEp.Read() - if (p != nil) != test.expectForward { + if (!p.IsNil()) != test.expectForward { t.Fatalf("got outgoingEp.Read() = %#v, want = (_ == nil) = %t", p, test.expectForward) } @@ -698,7 +698,7 @@ func TestMulticastRouteLastUsedTime(t *testing.T) { injectPacket(incomingEp, protocol, srcAddr, dstAddr, packetTTL) p := incomingEp.Read() - if p != nil { + if !p.IsNil() { t.Fatalf("Expected no ICMP packet through incoming NIC, instead found: %#v", p) } @@ -862,7 +862,7 @@ func TestRemoveMulticastRoute(t *testing.T) { injectPacket(incomingEp, protocol, srcAddr, dstAddr, packetTTL) p := incomingEp.Read() - if p != nil { + if !p.IsNil() { // An ICMP error should never be sent in response to a multicast // packet. t.Errorf("expected no ICMP packet through incoming NIC, instead found: %#v", p) @@ -878,7 +878,7 @@ func TestRemoveMulticastRoute(t *testing.T) { // If the route was successfully removed, then the packet should not be // forwarded. expectForward := test.wantErr != nil - if (p != nil) != expectForward { + if (!p.IsNil()) != expectForward { t.Fatalf("got outgoingEp.Read() = %#v, want = (_ == nil) = %t", p, expectForward) } @@ -1139,7 +1139,7 @@ func TestMulticastForwarding(t *testing.T) { injectPacket(incomingEp, protocol, srcAddr, dstAddr, test.ttl) p := incomingEp.Read() - if p != nil { + if !p.IsNil() { // An ICMP error should never be sent in response to a multicast packet. t.Fatalf("expected no ICMP packet through incoming NIC, instead found: %#v", p) } @@ -1154,7 +1154,7 @@ func TestMulticastForwarding(t *testing.T) { expectForward := contains(nicID, test.expectedForwardingInterfaces) - if (p != nil) != expectForward { + if (!p.IsNil()) != expectForward { t.Fatalf("got outgoingEp.Read() = %#v, want = (_ == nil) = %t", p, expectForward) } @@ -1171,7 +1171,7 @@ func TestMulticastForwarding(t *testing.T) { p = otherEp.Read() - if (p != nil) != test.joinMulticastGroup { + if (!p.IsNil()) != test.joinMulticastGroup { t.Fatalf("got otherEp.Read() = %#v, want = (_ == nil) = %t", p, test.joinMulticastGroup) } diff --git a/pkg/tcpip/tests/utils/utils.go b/pkg/tcpip/tests/utils/utils.go index d83951361..0d0884804 100644 --- a/pkg/tcpip/tests/utils/utils.go +++ b/pkg/tcpip/tests/utils/utils.go @@ -209,7 +209,7 @@ var _ stack.NetworkDispatcher = (*EndpointWithDestinationCheck)(nil) var _ stack.LinkEndpoint = (*EndpointWithDestinationCheck)(nil) // DeliverNetworkPacket implements stack.NetworkDispatcher. -func (e *EndpointWithDestinationCheck) DeliverNetworkPacket(proto tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (e *EndpointWithDestinationCheck) DeliverNetworkPacket(proto tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { if dst := header.Ethernet(pkt.LinkHeader().Slice()).DestinationAddress(); dst == e.Endpoint.LinkAddress() || dst == header.EthernetBroadcastAddress || header.IsMulticastEthernetAddress(dst) { e.Endpoint.DeliverNetworkPacket(proto, pkt) } diff --git a/pkg/tcpip/transport/datagram_test.go b/pkg/tcpip/transport/datagram_test.go index bbcf8cf74..135867369 100644 --- a/pkg/tcpip/transport/datagram_test.go +++ b/pkg/tcpip/transport/datagram_test.go @@ -168,7 +168,7 @@ func (e *mockEndpoint) Attach(d stack.NetworkDispatcher) { e.disp = d } func (e *mockEndpoint) IsAttached() bool { return e.disp != nil } func (*mockEndpoint) Wait() {} func (*mockEndpoint) ARPHardwareType() header.ARPHardwareType { return header.ARPHardwareNone } -func (*mockEndpoint) AddHeader(*stack.PacketBuffer) {} +func (*mockEndpoint) AddHeader(stack.PacketBufferPtr) {} func (e *mockEndpoint) releasePackets() { e.pkts.DecRef() e.pkts = stack.PacketBufferList{} @@ -1096,7 +1096,7 @@ func TestIPv6PacketInfo(t *testing.T) { { p := e1.Read() - if p == nil { + if p.IsNil() { t.Fatal("packet didn't arrive at ep1") } @@ -1106,7 +1106,7 @@ func TestIPv6PacketInfo(t *testing.T) { ) } - if p := e2.Read(); p != nil { + if p := e2.Read(); !p.IsNil() { t.Errorf("unexpected packet from ep2 = %#v", p) } }) diff --git a/pkg/tcpip/transport/icmp/endpoint.go b/pkg/tcpip/transport/icmp/endpoint.go index b0de7d628..9a348b084 100644 --- a/pkg/tcpip/transport/icmp/endpoint.go +++ b/pkg/tcpip/transport/icmp/endpoint.go @@ -36,7 +36,7 @@ type icmpPacket struct { icmpPacketEntry senderAddress tcpip.FullAddress packetInfo tcpip.IPPacketInfo - data *stack.PacketBuffer + data stack.PacketBufferPtr receivedAt time.Time `state:".(int64)"` // tosOrTClass stores either the Type of Service for IPv4 or the Traffic Class @@ -412,7 +412,7 @@ func send4(s *stack.Stack, ctx *network.WriteContext, ident uint16, data *buffer } pkt := ctx.TryNewPacketBuffer(header.ICMPv4MinimumSize+int(maxHeaderLength), bufferv2.Buffer{}) - if pkt == nil { + if pkt.IsNil() { return &tcpip.ErrWouldBlock{} } defer pkt.DecRef() @@ -454,7 +454,7 @@ func send6(s *stack.Stack, ctx *network.WriteContext, ident uint16, data *buffer } pkt := ctx.TryNewPacketBuffer(header.ICMPv6MinimumSize+int(maxHeaderLength), bufferv2.Buffer{}) - if pkt == nil { + if pkt.IsNil() { return &tcpip.ErrWouldBlock{} } defer pkt.DecRef() @@ -696,7 +696,7 @@ func (e *endpoint) Readiness(mask waiter.EventMask) waiter.EventMask { // HandlePacket is called by the stack when new packets arrive to this transport // endpoint. -func (e *endpoint) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) { +func (e *endpoint) HandlePacket(id stack.TransportEndpointID, pkt stack.PacketBufferPtr) { // Only accept echo replies. switch e.net.NetProto() { case header.IPv4ProtocolNumber: @@ -784,7 +784,7 @@ func (e *endpoint) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketB } // HandleError implements stack.TransportEndpoint. -func (*endpoint) HandleError(stack.TransportError, *stack.PacketBuffer) {} +func (*endpoint) HandleError(stack.TransportError, stack.PacketBufferPtr) {} // State implements tcpip.Endpoint.State. The ICMP endpoint currently doesn't // expose internal socket state. diff --git a/pkg/tcpip/transport/icmp/icmp_test.go b/pkg/tcpip/transport/icmp/icmp_test.go index 7d82a00d6..4084caa8e 100644 --- a/pkg/tcpip/transport/icmp/icmp_test.go +++ b/pkg/tcpip/transport/icmp/icmp_test.go @@ -142,7 +142,7 @@ func TestWriteUnboundWithBindToDevice(t *testing.T) { // Verify the packet was sent out the default NIC. p := defaultEP.Read() - if p == nil { + if p.IsNil() { t.Fatalf("got defaultEP.Read(_) = _, false; want = _, true (packet wasn't written out)") } defer p.DecRef() @@ -159,7 +159,7 @@ func TestWriteUnboundWithBindToDevice(t *testing.T) { }...) // Verify the packet was not sent out the alternate NIC. - if p := alternateEP.Read(); p != nil { + if p := alternateEP.Read(); !p.IsNil() { t.Fatalf("got alternateEP.Read(_) = %+v, true; want = _, false", p) } } @@ -184,13 +184,13 @@ func TestWriteUnboundWithBindToDevice(t *testing.T) { } // Verify the packet was not sent out the default NIC. - if p := defaultEP.Read(); p != nil { + if p := defaultEP.Read(); !p.IsNil() { t.Fatalf("got defaultEP.Read(_) = %+v, true; want = _, false", p) } // Verify the packet was sent out the alternate NIC. p := alternateEP.Read() - if p == nil { + if p.IsNil() { t.Fatalf("got alternateEP.Read(_) = _, false; want = _, true (packet wasn't written out)") } defer p.DecRef() @@ -228,7 +228,7 @@ func TestWriteUnboundWithBindToDevice(t *testing.T) { // Verify the packet was sent out the default NIC. p := defaultEP.Read() - if p == nil { + if p.IsNil() { t.Fatalf("got defaultEP.Read(_) = _, false; want = _, true (packet wasn't written out)") } defer p.DecRef() @@ -245,7 +245,7 @@ func TestWriteUnboundWithBindToDevice(t *testing.T) { }...) // Verify the packet was not sent out the alternate NIC. - if p := alternateEP.Read(); p != nil { + if p := alternateEP.Read(); !p.IsNil() { t.Fatalf("got alternateEP.Read(_) = %+v, true; want = _, false", p) } } diff --git a/pkg/tcpip/transport/icmp/protocol.go b/pkg/tcpip/transport/icmp/protocol.go index 7e6e3db18..d9833e478 100644 --- a/pkg/tcpip/transport/icmp/protocol.go +++ b/pkg/tcpip/transport/icmp/protocol.go @@ -100,7 +100,7 @@ func (p *protocol) ParsePorts(v []byte) (src, dst uint16, err tcpip.Error) { // HandleUnknownDestinationPacket handles packets targeted at this protocol but // that don't match any existing endpoint. -func (*protocol) HandleUnknownDestinationPacket(stack.TransportEndpointID, *stack.PacketBuffer) stack.UnknownDestinationPacketDisposition { +func (*protocol) HandleUnknownDestinationPacket(stack.TransportEndpointID, stack.PacketBufferPtr) stack.UnknownDestinationPacketDisposition { return stack.UnknownDestinationPacketHandled } @@ -127,7 +127,7 @@ func (*protocol) Pause() {} func (*protocol) Resume() {} // Parse implements stack.TransportProtocol.Parse. -func (*protocol) Parse(pkt *stack.PacketBuffer) bool { +func (*protocol) Parse(pkt stack.PacketBufferPtr) bool { // Right now, the Parse() method is tied to enabled protocols passed into // stack.New. This works for UDP and TCP, but we handle ICMP traffic even // when netstack users don't pass ICMP as a supported protocol. diff --git a/pkg/tcpip/transport/internal/network/endpoint.go b/pkg/tcpip/transport/internal/network/endpoint.go index c5a5b8f91..aab27a52a 100644 --- a/pkg/tcpip/transport/internal/network/endpoint.go +++ b/pkg/tcpip/transport/internal/network/endpoint.go @@ -269,7 +269,7 @@ func (c *WriteContext) PacketInfo() WritePacketInfo { // // If this method returns nil, the caller should wait for the endpoint to become // writable. -func (c *WriteContext) TryNewPacketBuffer(reserveHdrBytes int, data bufferv2.Buffer) *stack.PacketBuffer { +func (c *WriteContext) TryNewPacketBuffer(reserveHdrBytes int, data bufferv2.Buffer) stack.PacketBufferPtr { e := c.e e.sendBufferSizeInUseMu.Lock() @@ -312,7 +312,7 @@ func (c *WriteContext) TryNewPacketBuffer(reserveHdrBytes int, data bufferv2.Buf } // WritePacket attempts to write the packet. -func (c *WriteContext) WritePacket(pkt *stack.PacketBuffer, headerIncluded bool) tcpip.Error { +func (c *WriteContext) WritePacket(pkt stack.PacketBufferPtr, headerIncluded bool) tcpip.Error { c.e.mu.RLock() pkt.Owner = c.e.owner c.e.mu.RUnlock() diff --git a/pkg/tcpip/transport/internal/network/endpoint_test.go b/pkg/tcpip/transport/internal/network/endpoint_test.go index 142e5a52b..d757d0f7a 100644 --- a/pkg/tcpip/transport/internal/network/endpoint_test.go +++ b/pkg/tcpip/transport/internal/network/endpoint_test.go @@ -211,7 +211,7 @@ func TestEndpointStateTransitions(t *testing.T) { if err := ctx.WritePacket(injectPkt, false /* headerIncluded */); err != nil { t.Fatalf("ctx.WritePacket(_, false): %s", err) } - if pkt := e.Read(); pkt == nil { + if pkt := e.Read(); pkt.IsNil() { t.Fatalf("expected packet to be read from link endpoint") } else { payload := stack.PayloadSince(pkt.NetworkHeader()) diff --git a/pkg/tcpip/transport/internal/noop/endpoint.go b/pkg/tcpip/transport/internal/noop/endpoint.go index be2adae1c..3e9c4c4bb 100644 --- a/pkg/tcpip/transport/internal/noop/endpoint.go +++ b/pkg/tcpip/transport/internal/noop/endpoint.go @@ -137,7 +137,7 @@ func (*endpoint) GetSockOptInt(tcpip.SockOptInt) (int, tcpip.Error) { } // HandlePacket implements stack.RawTransportEndpoint.HandlePacket. -func (*endpoint) HandlePacket(pkt *stack.PacketBuffer) { +func (*endpoint) HandlePacket(pkt stack.PacketBufferPtr) { panic(fmt.Sprintf("unreachable: noop.endpoint should never be registered, but got packet: %+v", pkt)) } diff --git a/pkg/tcpip/transport/packet/endpoint.go b/pkg/tcpip/transport/packet/endpoint.go index 3000d44b8..a39af2dbb 100644 --- a/pkg/tcpip/transport/packet/endpoint.go +++ b/pkg/tcpip/transport/packet/endpoint.go @@ -40,7 +40,7 @@ import ( type packet struct { packetEntry // data holds the actual packet data, including any headers and payload. - data *stack.PacketBuffer + data stack.PacketBufferPtr receivedAt time.Time `state:".(int64)"` // senderAddr is the network address of the sender. senderAddr tcpip.FullAddress @@ -417,7 +417,7 @@ func (ep *endpoint) GetSockOptInt(opt tcpip.SockOptInt) (int, tcpip.Error) { } // HandlePacket implements stack.PacketEndpoint.HandlePacket. -func (ep *endpoint) HandlePacket(nicID tcpip.NICID, netProto tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { +func (ep *endpoint) HandlePacket(nicID tcpip.NICID, netProto tcpip.NetworkProtocolNumber, pkt stack.PacketBufferPtr) { ep.rcvMu.Lock() // Drop the packet if our buffer is currently full. diff --git a/pkg/tcpip/transport/raw/endpoint.go b/pkg/tcpip/transport/raw/endpoint.go index 1192d7b1b..c9ea3621d 100644 --- a/pkg/tcpip/transport/raw/endpoint.go +++ b/pkg/tcpip/transport/raw/endpoint.go @@ -46,7 +46,7 @@ type rawPacket struct { rawPacketEntry // data holds the actual packet data, including any headers and // payload. - data *stack.PacketBuffer + data stack.PacketBufferPtr receivedAt time.Time `state:".(int64)"` // senderAddr is the network address of the sender. senderAddr tcpip.FullAddress @@ -376,7 +376,7 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, tcp } pkt := ctx.TryNewPacketBuffer(int(ctx.PacketInfo().MaxHeaderLength), payload.Clone()) - if pkt == nil { + if pkt.IsNil() { return 0, &tcpip.ErrWouldBlock{} } defer pkt.DecRef() @@ -585,7 +585,7 @@ func (e *endpoint) GetSockOptInt(opt tcpip.SockOptInt) (int, tcpip.Error) { } // HandlePacket implements stack.RawTransportEndpoint.HandlePacket. -func (e *endpoint) HandlePacket(pkt *stack.PacketBuffer) { +func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { notifyReadableEvents := func() bool { e.mu.RLock() defer e.mu.RUnlock() diff --git a/pkg/tcpip/transport/tcp/connect.go b/pkg/tcpip/transport/tcp/connect.go index 7847484db..b6c3c768f 100644 --- a/pkg/tcpip/transport/tcp/connect.go +++ b/pkg/tcpip/transport/tcp/connect.go @@ -805,7 +805,7 @@ func (e *endpoint) sendSynTCP(r *stack.Route, tf tcpFields, opts header.TCPSynOp } // This method takes ownership of pkt. -func (e *endpoint) sendTCP(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stack.GSO) tcpip.Error { +func (e *endpoint) sendTCP(r *stack.Route, tf tcpFields, pkt stack.PacketBufferPtr, gso stack.GSO) tcpip.Error { tf.txHash = e.txHash if err := sendTCP(r, tf, pkt, gso, e.owner); err != nil { e.stats.SendErrors.SegmentSendToNetworkFailed.Increment() @@ -815,7 +815,7 @@ func (e *endpoint) sendTCP(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer return nil } -func buildTCPHdr(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stack.GSO) { +func buildTCPHdr(r *stack.Route, tf tcpFields, pkt stack.PacketBufferPtr, gso stack.GSO) { optLen := len(tf.opts) tcp := header.TCP(pkt.TransportHeader().Push(header.TCPMinimumSize + optLen)) pkt.TransportProtocolNumber = header.TCPProtocolNumber @@ -844,7 +844,7 @@ func buildTCPHdr(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stac } } -func sendTCPBatch(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { +func sendTCPBatch(r *stack.Route, tf tcpFields, pkt stack.PacketBufferPtr, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { optLen := len(tf.opts) if tf.rcvWnd > math.MaxUint16 { tf.rcvWnd = math.MaxUint16 @@ -895,7 +895,7 @@ func sendTCPBatch(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso sta // sendTCP sends a TCP segment with the provided options via the provided // network endpoint and under the provided identity. This method takes // ownership of pkt. -func sendTCP(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { +func sendTCP(r *stack.Route, tf tcpFields, pkt stack.PacketBufferPtr, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { if tf.rcvWnd > math.MaxUint16 { tf.rcvWnd = math.MaxUint16 } @@ -968,7 +968,7 @@ func (e *endpoint) sendEmptyRaw(flags header.TCPFlags, seq, ack seqnum.Value, rc // sendRaw sends a TCP segment to the endpoint's peer. This method takes // ownership of pkt. pkt must not have any headers set. -func (e *endpoint) sendRaw(pkt *stack.PacketBuffer, flags header.TCPFlags, seq, ack seqnum.Value, rcvWnd seqnum.Size) tcpip.Error { +func (e *endpoint) sendRaw(pkt stack.PacketBufferPtr, flags header.TCPFlags, seq, ack seqnum.Value, rcvWnd seqnum.Size) tcpip.Error { var sackBlocks []header.SACKBlock if e.EndpointState() == StateEstablished && e.rcv.pendingRcvdSegments.Len() > 0 && (flags&header.TCPFlagAck != 0) { sackBlocks = e.sack.Blocks[:e.sack.NumBlocks] diff --git a/pkg/tcpip/transport/tcp/dispatcher.go b/pkg/tcpip/transport/tcp/dispatcher.go index 6e4ab5b86..1c7a3c44e 100644 --- a/pkg/tcpip/transport/tcp/dispatcher.go +++ b/pkg/tcpip/transport/tcp/dispatcher.go @@ -409,7 +409,7 @@ func (d *dispatcher) wait() { // queuePacket queues an incoming packet to the matching tcp endpoint and // also queues the endpoint to a processor queue for processing. -func (d *dispatcher) queuePacket(stackEP stack.TransportEndpoint, id stack.TransportEndpointID, clock tcpip.Clock, pkt *stack.PacketBuffer) { +func (d *dispatcher) queuePacket(stackEP stack.TransportEndpoint, id stack.TransportEndpointID, clock tcpip.Clock, pkt stack.PacketBufferPtr) { d.mu.Lock() closed := d.closed d.mu.Unlock() diff --git a/pkg/tcpip/transport/tcp/endpoint.go b/pkg/tcpip/transport/tcp/endpoint.go index 6d1a530db..a4efb2486 100644 --- a/pkg/tcpip/transport/tcp/endpoint.go +++ b/pkg/tcpip/transport/tcp/endpoint.go @@ -2807,7 +2807,7 @@ func (e *endpoint) getRemoteAddress() tcpip.FullAddress { } } -func (*endpoint) HandlePacket(stack.TransportEndpointID, *stack.PacketBuffer) { +func (*endpoint) HandlePacket(stack.TransportEndpointID, stack.PacketBufferPtr) { // TCP HandlePacket is not required anymore as inbound packets first // land at the Dispatcher which then can either deliver using the // worker go routine or directly do the invoke the tcp processing inline @@ -2825,7 +2825,7 @@ func (e *endpoint) enqueueSegment(s *segment) bool { return true } -func (e *endpoint) onICMPError(err tcpip.Error, transErr stack.TransportError, pkt *stack.PacketBuffer) { +func (e *endpoint) onICMPError(err tcpip.Error, transErr stack.TransportError, pkt stack.PacketBufferPtr) { // Update last error first. e.lastErrorMu.Lock() e.lastError = err @@ -2882,7 +2882,7 @@ func (e *endpoint) onICMPError(err tcpip.Error, transErr stack.TransportError, p } // HandleError implements stack.TransportEndpoint. -func (e *endpoint) HandleError(transErr stack.TransportError, pkt *stack.PacketBuffer) { +func (e *endpoint) HandleError(transErr stack.TransportError, pkt stack.PacketBufferPtr) { handlePacketTooBig := func(mtu uint32) { e.sndQueueInfo.sndQueueMu.Lock() update := false diff --git a/pkg/tcpip/transport/tcp/forwarder.go b/pkg/tcpip/transport/tcp/forwarder.go index b3888ff20..3d632939d 100644 --- a/pkg/tcpip/transport/tcp/forwarder.go +++ b/pkg/tcpip/transport/tcp/forwarder.go @@ -64,7 +64,7 @@ func NewForwarder(s *stack.Stack, rcvWnd, maxInFlight int, handler func(*Forward // // This function is expected to be passed as an argument to the // stack.SetTransportProtocolHandler function. -func (f *Forwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool { +func (f *Forwarder) HandlePacket(id stack.TransportEndpointID, pkt stack.PacketBufferPtr) bool { s, err := newIncomingSegment(id, f.stack.Clock(), pkt) if err != nil { return false diff --git a/pkg/tcpip/transport/tcp/protocol.go b/pkg/tcpip/transport/tcp/protocol.go index 500c8ab95..e38ecddea 100644 --- a/pkg/tcpip/transport/tcp/protocol.go +++ b/pkg/tcpip/transport/tcp/protocol.go @@ -144,7 +144,7 @@ func (*protocol) ParsePorts(v []byte) (src, dst uint16, err tcpip.Error) { // to a specific processing queue. Each queue is serviced by its own processor // goroutine which is responsible for dequeuing and doing full TCP dispatch of // the packet. -func (p *protocol) QueuePacket(ep stack.TransportEndpoint, id stack.TransportEndpointID, pkt *stack.PacketBuffer) { +func (p *protocol) QueuePacket(ep stack.TransportEndpoint, id stack.TransportEndpointID, pkt stack.PacketBufferPtr) { p.dispatcher.queuePacket(ep, id, p.stack.Clock(), pkt) } @@ -155,7 +155,7 @@ func (p *protocol) QueuePacket(ep stack.TransportEndpoint, id stack.TransportEnd // a reset is sent in response to any incoming segment except another reset. In // particular, SYNs addressed to a non-existent connection are rejected by this // means." -func (p *protocol) HandleUnknownDestinationPacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) stack.UnknownDestinationPacketDisposition { +func (p *protocol) HandleUnknownDestinationPacket(id stack.TransportEndpointID, pkt stack.PacketBufferPtr) stack.UnknownDestinationPacketDisposition { s, err := newIncomingSegment(id, p.stack.Clock(), pkt) if err != nil { return stack.UnknownDestinationPacketMalformed @@ -501,7 +501,7 @@ func (p *protocol) Resume() { } // Parse implements stack.TransportProtocol.Parse. -func (*protocol) Parse(pkt *stack.PacketBuffer) bool { +func (*protocol) Parse(pkt stack.PacketBufferPtr) bool { return parse.TCP(pkt) } diff --git a/pkg/tcpip/transport/tcp/segment.go b/pkg/tcpip/transport/tcp/segment.go index 40946b319..024a1b4db 100644 --- a/pkg/tcpip/transport/tcp/segment.go +++ b/pkg/tcpip/transport/tcp/segment.go @@ -59,7 +59,7 @@ type segment struct { qFlags queueFlags id stack.TransportEndpointID `state:"manual"` - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr sequenceNumber seqnum.Value ackNumber seqnum.Value @@ -92,7 +92,7 @@ type segment struct { lost bool } -func newIncomingSegment(id stack.TransportEndpointID, clock tcpip.Clock, pkt *stack.PacketBuffer) (*segment, error) { +func newIncomingSegment(id stack.TransportEndpointID, clock tcpip.Clock, pkt stack.PacketBufferPtr) (*segment, error) { hdr := header.TCP(pkt.TransportHeader().Slice()) netHdr := pkt.Network() csum, csumValid, ok := header.TCPValid( diff --git a/pkg/tcpip/transport/tcp/snd.go b/pkg/tcpip/transport/tcp/snd.go index 3f20cbe21..d762c410e 100644 --- a/pkg/tcpip/transport/tcp/snd.go +++ b/pkg/tcpip/transport/tcp/snd.go @@ -1660,7 +1660,7 @@ func (s *sender) sendSegment(seg *segment) tcpip.Error { // flags and sequence number. // +checklocks:s.ep.mu // +checklocksalias:s.ep.rcv.ep.mu=s.ep.mu -func (s *sender) sendSegmentFromPacketBuffer(pkt *stack.PacketBuffer, flags header.TCPFlags, seq seqnum.Value) tcpip.Error { +func (s *sender) sendSegmentFromPacketBuffer(pkt stack.PacketBufferPtr, flags header.TCPFlags, seq seqnum.Value) tcpip.Error { s.LastSendTime = s.ep.stack.Clock().NowMonotonic() if seq == s.RTTMeasureSeqNum { s.RTTMeasureTime = s.LastSendTime diff --git a/pkg/tcpip/transport/tcp/testing/context/context.go b/pkg/tcpip/transport/tcp/testing/context/context.go index 824924b18..8f718eb02 100644 --- a/pkg/tcpip/transport/tcp/testing/context/context.go +++ b/pkg/tcpip/transport/tcp/testing/context/context.go @@ -307,7 +307,7 @@ func (c *Context) CheckNoPacketTimeout(errMsg string, wait time.Duration) { ctx, cancel := context.WithTimeout(context.Background(), wait) defer cancel() - if c.linkEP.ReadContext(ctx) != nil { + if pkt := c.linkEP.ReadContext(ctx); !pkt.IsNil() { c.t.Fatal(errMsg) } } @@ -327,7 +327,7 @@ func (c *Context) GetPacketWithTimeout(timeout time.Duration) *bufferv2.View { ctx, cancel := context.WithTimeout(context.Background(), timeout) defer cancel() pkt := c.linkEP.ReadContext(ctx) - if pkt == nil { + if pkt.IsNil() { return nil } defer pkt.DecRef() @@ -377,7 +377,7 @@ func (c *Context) GetPacketNonBlocking() *bufferv2.View { c.t.Helper() pkt := c.linkEP.Read() - if pkt == nil { + if pkt.IsNil() { return nil } defer pkt.DecRef() @@ -624,7 +624,7 @@ func (c *Context) GetV6Packet() *bufferv2.View { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() pkt := c.linkEP.ReadContext(ctx) - if pkt == nil { + if pkt.IsNil() { c.t.Fatalf("Packet wasn't written out") return nil } diff --git a/pkg/tcpip/transport/udp/endpoint.go b/pkg/tcpip/transport/udp/endpoint.go index 3bb44ae57..78783f582 100644 --- a/pkg/tcpip/transport/udp/endpoint.go +++ b/pkg/tcpip/transport/udp/endpoint.go @@ -40,7 +40,7 @@ type udpPacket struct { senderAddress tcpip.FullAddress destinationAddress tcpip.FullAddress packetInfo tcpip.IPPacketInfo - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr receivedAt time.Time `state:".(int64)"` // tosOrTClass stores either the Type of Service for IPv4 or the Traffic Class // for IPv6. @@ -473,7 +473,7 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, tcp dataSz := udpInfo.data.Size() pktInfo := udpInfo.ctx.PacketInfo() pkt := udpInfo.ctx.TryNewPacketBuffer(header.UDPMinimumSize+int(pktInfo.MaxHeaderLength), udpInfo.data) - if pkt == nil { + if pkt.IsNil() { return 0, &tcpip.ErrWouldBlock{} } defer pkt.DecRef() @@ -902,7 +902,7 @@ func (e *endpoint) Readiness(mask waiter.EventMask) waiter.EventMask { // HandlePacket is called by the stack when new packets arrive to this transport // endpoint. -func (e *endpoint) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) { +func (e *endpoint) HandlePacket(id stack.TransportEndpointID, pkt stack.PacketBufferPtr) { // Get the header then trim it from the view. hdr := header.UDP(pkt.TransportHeader().Slice()) netHdr := pkt.Network() @@ -994,7 +994,7 @@ func (e *endpoint) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketB } } -func (e *endpoint) onICMPError(err tcpip.Error, transErr stack.TransportError, pkt *stack.PacketBuffer) { +func (e *endpoint) onICMPError(err tcpip.Error, transErr stack.TransportError, pkt stack.PacketBufferPtr) { // Update last error first. e.lastErrorMu.Lock() e.lastError = err @@ -1042,7 +1042,7 @@ func (e *endpoint) onICMPError(err tcpip.Error, transErr stack.TransportError, p } // HandleError implements stack.TransportEndpoint. -func (e *endpoint) HandleError(transErr stack.TransportError, pkt *stack.PacketBuffer) { +func (e *endpoint) HandleError(transErr stack.TransportError, pkt stack.PacketBufferPtr) { // TODO(gvisor.dev/issues/5270): Handle all transport errors. switch transErr.Kind() { case stack.DestinationPortUnreachableTransportError: diff --git a/pkg/tcpip/transport/udp/forwarder.go b/pkg/tcpip/transport/udp/forwarder.go index 7238fc019..4997fa11c 100644 --- a/pkg/tcpip/transport/udp/forwarder.go +++ b/pkg/tcpip/transport/udp/forwarder.go @@ -43,7 +43,7 @@ func NewForwarder(s *stack.Stack, handler func(*ForwarderRequest)) *Forwarder { // // This function is expected to be passed as an argument to the // stack.SetTransportProtocolHandler function. -func (f *Forwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool { +func (f *Forwarder) HandlePacket(id stack.TransportEndpointID, pkt stack.PacketBufferPtr) bool { f.handler(&ForwarderRequest{ stack: f.stack, id: id, @@ -59,7 +59,7 @@ func (f *Forwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Packet type ForwarderRequest struct { stack *stack.Stack id stack.TransportEndpointID - pkt *stack.PacketBuffer + pkt stack.PacketBufferPtr } // ID returns the 4-tuple (src address, src port, dst address, dst port) that diff --git a/pkg/tcpip/transport/udp/protocol.go b/pkg/tcpip/transport/udp/protocol.go index 45b7114e7..2a8518efd 100644 --- a/pkg/tcpip/transport/udp/protocol.go +++ b/pkg/tcpip/transport/udp/protocol.go @@ -77,7 +77,7 @@ func (*protocol) ParsePorts(v []byte) (src, dst uint16, err tcpip.Error) { // HandleUnknownDestinationPacket handles packets that are targeted at this // protocol but don't match any existing endpoint. -func (p *protocol) HandleUnknownDestinationPacket(id stack.TransportEndpointID, pkt *stack.PacketBuffer) stack.UnknownDestinationPacketDisposition { +func (p *protocol) HandleUnknownDestinationPacket(id stack.TransportEndpointID, pkt stack.PacketBufferPtr) stack.UnknownDestinationPacketDisposition { hdr := header.UDP(pkt.TransportHeader().Slice()) netHdr := pkt.Network() lengthValid, csumValid := header.UDPValid( @@ -124,7 +124,7 @@ func (*protocol) Pause() {} func (*protocol) Resume() {} // Parse implements stack.TransportProtocol.Parse. -func (*protocol) Parse(pkt *stack.PacketBuffer) bool { +func (*protocol) Parse(pkt stack.PacketBufferPtr) bool { return parse.UDP(pkt) } diff --git a/pkg/tcpip/transport/udp/udp_test.go b/pkg/tcpip/transport/udp/udp_test.go index 556c5ff66..53f486fcd 100644 --- a/pkg/tcpip/transport/udp/udp_test.go +++ b/pkg/tcpip/transport/udp/udp_test.go @@ -591,7 +591,7 @@ func testWriteAndVerifyInternal(c *context.Context, flow context.TestFlow, setDe // Received the packet and check the payload. p := c.LinkEP.Read() - if p == nil { + if p.IsNil() { c.T.Fatalf("Packet wasn't written out") } defer p.DecRef() @@ -1526,7 +1526,7 @@ func TestV4UnknownDestination(t *testing.T) { } } if !tc.icmpRequired { - if p := c.LinkEP.Read(); p != nil { + if p := c.LinkEP.Read(); !p.IsNil() { t.Fatalf("unexpected packet received: %+v", p) } return @@ -1534,7 +1534,7 @@ func TestV4UnknownDestination(t *testing.T) { // ICMP required. p := c.LinkEP.Read() - if p == nil { + if p.IsNil() { t.Fatalf("packet wasn't written out") } @@ -1623,7 +1623,7 @@ func TestV6UnknownDestination(t *testing.T) { } } if !tc.icmpRequired { - if p := c.LinkEP.Read(); p != nil { + if p := c.LinkEP.Read(); !p.IsNil() { t.Fatalf("unexpected packet received: %+v", p) } return @@ -1631,7 +1631,7 @@ func TestV6UnknownDestination(t *testing.T) { // ICMP required. p := c.LinkEP.Read() - if p == nil { + if p.IsNil() { t.Fatalf("packet wasn't written out") } @@ -2184,7 +2184,7 @@ func TestChecksumWithZeroValueOnesComplementSum(t *testing.T) { } pkt := c.LinkEP.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("Packet wasn't written out") } @@ -2221,7 +2221,7 @@ func TestChecksumWithZeroValueOnesComplementSum(t *testing.T) { { pkt := c.LinkEP.Read() - if pkt == nil { + if pkt.IsNil() { t.Fatal("Packet wasn't written out") } defer pkt.DecRef()