From 6c8187194adffbfe412ab7d0049030440d7f7146 Mon Sep 17 00:00:00 2001 From: Lucas Manning Date: Thu, 15 Jun 2023 13:18:19 -0700 Subject: [PATCH] Automated rollback of changelist 538230394 PiperOrigin-RevId: 540671483 --- pkg/buffer/buffer.go | 6 +-- pkg/buffer/view.go | 15 +++---- .../internal/fragmentation/fragmentation.go | 2 +- pkg/tcpip/network/ipv4/icmp.go | 4 +- pkg/tcpip/network/ipv4/ipv4.go | 38 ++++++++++-------- pkg/tcpip/network/ipv4/ipv4_test.go | 2 +- pkg/tcpip/network/ipv6/icmp.go | 2 +- pkg/tcpip/network/ipv6/ipv6.go | 40 ++++++++++--------- pkg/tcpip/network/ipv6/ipv6_test.go | 2 +- pkg/tcpip/stack/packet_buffer.go | 15 ++++--- pkg/tcpip/transport/raw/endpoint.go | 2 +- 11 files changed, 67 insertions(+), 61 deletions(-) diff --git a/pkg/buffer/buffer.go b/pkg/buffer/buffer.go index 2b1e9269e..cc663ae78 100644 --- a/pkg/buffer/buffer.go +++ b/pkg/buffer/buffer.go @@ -321,10 +321,8 @@ func (b *Buffer) PullUp(offset, length int) (View, bool) { if x := curr.Intersect(tgt); x.Len() == tgt.Len() { // buf covers the whole requested target range. sub := x.Offset(-curr.begin) - // Ensure that v has exclusive ownership over its chunk before returning. - // NAT rules sometimes write directly to the slices backing these views, - // which would break the ownership model if the chunks were shared. - v.unshare() + // Don't increment the reference count of the underlying chunk. Views + // returned by PullUp are explicitly unowned and read only new := View{ read: v.read + sub.begin, write: v.read + sub.end, diff --git a/pkg/buffer/view.go b/pkg/buffer/view.go index 8dc8b9c7d..d7eb2f118 100644 --- a/pkg/buffer/view.go +++ b/pkg/buffer/view.go @@ -272,7 +272,7 @@ func (v *View) ReadFrom(r io.Reader) (n int64, err error) { v.chunk = v.chunk.Clone() } for { - // Check for EOF to avoid an unnecessary allocation. + // Check for EOF to avoid an unnnecesary allocation. if _, e := r.Read(nil); e == io.EOF { return n, nil } @@ -304,7 +304,10 @@ func (v *View) WriteAt(p []byte, off int) (int, error) { if off < 0 || off > v.Size() { return 0, fmt.Errorf("write offset out of bounds: want 0 < off < %d, got off=%d", v.Size(), off) } - v.unshare() + if v.sharesChunk() { + defer v.chunk.DecRef() + v.chunk = v.chunk.Clone() + } n := copy(v.AsSlice()[off:], p) if n < len(p) { return n, io.ErrShortWrite @@ -354,16 +357,10 @@ func (v *View) CapLength(n int) { } func (v *View) availableSlice() []byte { - v.unshare() - return v.chunk.data[v.write:] -} - -// Unshare ensures the backing chunk is exclusively owned by this view. -// This incurs a copy so only use when necessary. -func (v *View) unshare() { if v.sharesChunk() { defer v.chunk.DecRef() c := v.chunk.Clone() v.chunk = c } + return v.chunk.data[v.write:] } diff --git a/pkg/tcpip/network/internal/fragmentation/fragmentation.go b/pkg/tcpip/network/internal/fragmentation/fragmentation.go index bf1d050f0..39dc5ad02 100644 --- a/pkg/tcpip/network/internal/fragmentation/fragmentation.go +++ b/pkg/tcpip/network/internal/fragmentation/fragmentation.go @@ -317,7 +317,7 @@ func MakePacketFragmenter(pkt stack.PacketBufferPtr, fragmentPayloadLen uint32, // supported for outbound packets, the fragmentable data should not include // these headers. var fragmentableData buffer.Buffer - fragmentableData.Append(pkt.TransportHeader().ToView()) + fragmentableData.Append(pkt.TransportHeader().View()) pktBuf := pkt.Data().ToBuffer() fragmentableData.Merge(&pktBuf) fragmentCount := (uint32(fragmentableData.Size()) + fragmentPayloadLen - 1) / fragmentPayloadLen diff --git a/pkg/tcpip/network/ipv4/icmp.go b/pkg/tcpip/network/ipv4/icmp.go index b4092010e..875eca473 100644 --- a/pkg/tcpip/network/ipv4/icmp.go +++ b/pkg/tcpip/network/ipv4/icmp.go @@ -767,8 +767,8 @@ func (p *protocol) returnError(reason icmpReason, pkt stack.PacketBufferPtr, del // required. This is now the payload of the new ICMP packet and no longer // considered a packet in its own right. - payload := buffer.MakeWithView(pkt.NetworkHeader().ToView()) - payload.Append(pkt.TransportHeader().ToView()) + payload := buffer.MakeWithView(pkt.NetworkHeader().View()) + payload.Append(pkt.TransportHeader().View()) if dataCap := payloadLen - int(payload.Size()); dataCap > 0 { buf := pkt.Data().ToBuffer() buf.Truncate(int64(dataCap)) diff --git a/pkg/tcpip/network/ipv4/ipv4.go b/pkg/tcpip/network/ipv4/ipv4.go index 652cd4990..2e5ab0264 100644 --- a/pkg/tcpip/network/ipv4/ipv4.go +++ b/pkg/tcpip/network/ipv4/ipv4.go @@ -22,6 +22,7 @@ import ( "time" "gvisor.dev/gvisor/pkg/atomicbitops" + "gvisor.dev/gvisor/pkg/buffer" "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/header" @@ -726,7 +727,7 @@ func (e *endpoint) forwardPacketWithRoute(route *stack.Route, pkt stack.PacketBu // forwardUnicastPacket attempts to forward a packet to its final destination. func (e *endpoint) forwardUnicastPacket(pkt stack.PacketBufferPtr) ip.ForwardingError { - hView := pkt.NetworkHeader().ToView() + hView := pkt.NetworkHeader().View() defer hView.Release() h := header.IPv4(hView.AsSlice()) @@ -813,11 +814,13 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { return } - if ok := e.protocol.parseAndValidate(pkt); !ok { + hView, ok := e.protocol.parseAndValidate(pkt) + if !ok { stats.MalformedPacketsReceived.Increment() return } - h := header.IPv4(pkt.NetworkHeader().Slice()) + h := header.IPv4(hView.AsSlice()) + defer hView.Release() if !e.nic.IsLoopback() { if !e.protocol.options.AllowExternalLoopbackTraffic { @@ -833,7 +836,7 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } if e.protocol.stack.HandleLocal() { - addressEndpoint := e.AcquireAssignedAddress(h.SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) + addressEndpoint := e.AcquireAssignedAddress(header.IPv4(pkt.NetworkHeader().Slice()).SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) if addressEndpoint != nil { addressEndpoint.DecRef() @@ -854,9 +857,7 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } } - hv := pkt.NetworkHeader().ToView() - defer hv.Release() - e.handleValidatedPacket(hv.AsSlice(), pkt, e.nic.Name() /* inNICName */) + e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) } // handleLocalPacket is like HandlePacket except it does not perform the @@ -870,14 +871,15 @@ func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksu defer pkt.DecRef() pkt.RXChecksumValidated = canSkipRXChecksum - if ok := e.protocol.parseAndValidate(pkt); !ok { + hView, ok := e.protocol.parseAndValidate(pkt) + if !ok { stats.MalformedPacketsReceived.Increment() return } + h := header.IPv4(hView.AsSlice()) + defer hView.Release() - h := pkt.NetworkHeader().ToView() - defer h.Release() - e.handleValidatedPacket(h.AsSlice(), pkt, e.nic.Name() /* inNICName */) + e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) } func validateAddressesForForwarding(h header.IPv4) ip.ForwardingError { @@ -1756,29 +1758,31 @@ func (p *protocol) isSubnetLocalBroadcastAddress(addr tcpip.Address) bool { } // parseAndValidate parses the packet (including its transport layer header) and -// returns true if the IP header was successfully parsed. -func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) bool { +// returns the parsed IP header. +// +// Returns true if the IP header was successfully parsed. +func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (*buffer.View, bool) { transProtoNum, hasTransportHdr, ok := p.Parse(pkt) if !ok { - return false + return nil, false } h := header.IPv4(pkt.NetworkHeader().Slice()) // Do not include the link header's size when calculating the size of the IP // packet. if !h.IsValid(pkt.Size() - len(pkt.LinkHeader().Slice())) { - return false + return nil, false } if !pkt.RXChecksumValidated && !h.IsChecksumValid() { - return false + return nil, false } if hasTransportHdr { p.parseTransport(pkt, transProtoNum) } - return true + return pkt.NetworkHeader().View(), true } func (p *protocol) parseTransport(pkt stack.PacketBufferPtr, transProtoNum tcpip.TransportProtocolNumber) { diff --git a/pkg/tcpip/network/ipv4/ipv4_test.go b/pkg/tcpip/network/ipv4/ipv4_test.go index cab3a1fa0..4b3704eb6 100644 --- a/pkg/tcpip/network/ipv4/ipv4_test.go +++ b/pkg/tcpip/network/ipv4/ipv4_test.go @@ -2010,7 +2010,7 @@ func compareFragments(packets []stack.PacketBufferPtr, sourcePacket stack.Packet } else { sourceCopy.SetFlagsFragmentOffset(sourceCopy.Flags()&^header.IPv4FlagMoreFragments, wantFragments[i].offset) } - reassembledPayload.Append(packet.TransportHeader().ToView()) + reassembledPayload.Append(packet.TransportHeader().View()) reassembledPayload.Append(packet.Data().AsRange().ToView()) // Clear out the checksum and length from the ip because we can't compare // it. diff --git a/pkg/tcpip/network/ipv6/icmp.go b/pkg/tcpip/network/ipv6/icmp.go index 386bb58b6..a98332dc3 100644 --- a/pkg/tcpip/network/ipv6/icmp.go +++ b/pkg/tcpip/network/ipv6/icmp.go @@ -1153,7 +1153,7 @@ func (p *protocol) returnError(reason icmpReason, pkt stack.PacketBufferPtr, del return nil } - network, transport := pkt.NetworkHeader().ToView(), pkt.TransportHeader().ToView() + network, transport := pkt.NetworkHeader().View(), pkt.TransportHeader().View() // As per RFC 4443 section 2.4 // diff --git a/pkg/tcpip/network/ipv6/ipv6.go b/pkg/tcpip/network/ipv6/ipv6.go index c2206ec92..ff5b44653 100644 --- a/pkg/tcpip/network/ipv6/ipv6.go +++ b/pkg/tcpip/network/ipv6/ipv6.go @@ -1082,11 +1082,13 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { return } - if ok := e.protocol.parseAndValidate(pkt); !ok { + hView, ok := e.protocol.parseAndValidate(pkt) + if !ok { stats.MalformedPacketsReceived.Increment() return } - h := header.IPv6(pkt.NetworkHeader().Slice()) + defer hView.Release() + h := header.IPv6(hView.AsSlice()) if !checkV4Mapped(h, stats) { return @@ -1106,7 +1108,7 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } if e.protocol.stack.HandleLocal() { - addressEndpoint := e.AcquireAssignedAddress(h.SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) + addressEndpoint := e.AcquireAssignedAddress(header.IPv6(pkt.NetworkHeader().Slice()).SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) if addressEndpoint != nil { addressEndpoint.DecRef() @@ -1127,9 +1129,7 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } } - hv := pkt.NetworkHeader().ToView() - defer hv.Release() - e.handleValidatedPacket(hv.AsSlice(), pkt, e.nic.Name() /* inNICName */) + e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) } // handleLocalPacket is like HandlePacket except it does not perform the @@ -1143,18 +1143,19 @@ func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksu defer pkt.DecRef() pkt.RXChecksumValidated = canSkipRXChecksum - if ok := e.protocol.parseAndValidate(pkt); !ok { + hView, ok := e.protocol.parseAndValidate(pkt) + if !ok { stats.MalformedPacketsReceived.Increment() return } - h := pkt.NetworkHeader().ToView() - defer h.Release() + defer hView.Release() + h := header.IPv6(hView.AsSlice()) - if !checkV4Mapped(h.AsSlice(), stats) { + if !checkV4Mapped(h, stats) { return } - e.handleValidatedPacket(h.AsSlice(), pkt, e.nic.Name() /* inNICName */) + e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) } // forwardMulticastPacket validates a multicast pkt and attempts to forward it. @@ -1460,12 +1461,12 @@ func (e *endpoint) processExtensionHeaders(h header.IPv6, pkt stack.PacketBuffer // - Any IPv6 header bytes after the first 40 (i.e. extensions). // - The transport header, if present. // - Any other payload data. - v := pkt.NetworkHeader().ToView() + v := pkt.NetworkHeader().View() if v != nil { v.TrimFront(header.IPv6MinimumSize) } buf := buffer.MakeWithView(v) - buf.Append(pkt.TransportHeader().ToView()) + buf.Append(pkt.TransportHeader().View()) dataBuf := pkt.Data().ToBuffer() buf.Merge(&dataBuf) it := header.MakeIPv6PayloadIterator(header.IPv6ExtensionHeaderIdentifier(h.NextHeader()), buf) @@ -2587,25 +2588,28 @@ func (p *protocol) forwardPendingMulticastPacket(pkt stack.PacketBufferPtr, inst func (*protocol) Wait() {} // parseAndValidate parses the packet (including its transport layer header) and -// returns true if the IP header was successfully parsed. -func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) bool { +// returns a view containing the parsed IP header. The caller is responsible +// for releasing the returned View. +// +// Returns true if the IP header was successfully parsed. +func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (*buffer.View, bool) { transProtoNum, hasTransportHdr, ok := p.Parse(pkt) if !ok { - return false + return nil, false } h := header.IPv6(pkt.NetworkHeader().Slice()) // Do not include the link header's size when calculating the size of the IP // packet. if !h.IsValid(pkt.Size() - len(pkt.LinkHeader().Slice())) { - return false + return nil, false } if hasTransportHdr { p.parseTransport(pkt, transProtoNum) } - return true + return pkt.NetworkHeader().View(), true } func (p *protocol) parseTransport(pkt stack.PacketBufferPtr, transProtoNum tcpip.TransportProtocolNumber) { diff --git a/pkg/tcpip/network/ipv6/ipv6_test.go b/pkg/tcpip/network/ipv6/ipv6_test.go index ee18867ee..6d2aa6f93 100644 --- a/pkg/tcpip/network/ipv6/ipv6_test.go +++ b/pkg/tcpip/network/ipv6/ipv6_test.go @@ -239,7 +239,7 @@ func compareFragments(packets []stack.PacketBufferPtr, sourcePacket stack.Packet // Store the reassembled payload as we parse each fragment. The payload // includes the Transport header and everything after. - reassembledPayload.Append(fragment.TransportHeader().ToView()) + reassembledPayload.Append(fragment.TransportHeader().View()) reassembledPayload.Append(fragment.Data().AsRange().ToView()) } diff --git a/pkg/tcpip/stack/packet_buffer.go b/pkg/tcpip/stack/packet_buffer.go index 66863b660..86b756950 100644 --- a/pkg/tcpip/stack/packet_buffer.go +++ b/pkg/tcpip/stack/packet_buffer.go @@ -433,9 +433,11 @@ func (pk PacketBufferPtr) CloneToInbound() PacketBufferPtr { // The returned packet buffer will have the network and transport headers // set if the original packet buffer did. func (pk PacketBufferPtr) DeepCopyForForwarding(reservedHeaderBytes int) PacketBufferPtr { + payload := BufferSince(pk.NetworkHeader()) + defer payload.Release() newPk := NewPacketBuffer(PacketBufferOptions{ ReserveHeaderBytes: reservedHeaderBytes, - Payload: BufferSince(pk.NetworkHeader()), + Payload: payload.DeepClone(), IsForwardedPacket: true, }) @@ -483,8 +485,9 @@ type PacketHeader struct { typ headerType } -// ToView returns an caller-owned copy of the underlying storage of h. -func (h PacketHeader) ToView() *buffer.View { +// View returns an caller-owned copy of the underlying storage of h as a +// *buffer.View. +func (h PacketHeader) View() *buffer.View { view := h.pk.headerView(h.typ) if view.Size() == 0 { return nil @@ -492,9 +495,9 @@ func (h PacketHeader) ToView() *buffer.View { return view.Clone() } -// Slice returns the PacketHeader-owned storage of h as a []byte. The slice is -// guaranteed to be owned exclusively by h, so it's safe to modify the contents -// directly. +// Slice returns the underlying storage of h as a []byte. The returned slice +// should not be modified if the underlying packet could be shared, cloned, or +// borrowed. func (h PacketHeader) Slice() []byte { view := h.pk.headerView(h.typ) return view.AsSlice() diff --git a/pkg/tcpip/transport/raw/endpoint.go b/pkg/tcpip/transport/raw/endpoint.go index d397668f4..476932d2b 100644 --- a/pkg/tcpip/transport/raw/endpoint.go +++ b/pkg/tcpip/transport/raw/endpoint.go @@ -702,7 +702,7 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } } - combinedBuf = buffer.MakeWithView(pkt.TransportHeader().ToView()) + combinedBuf = buffer.MakeWithView(pkt.TransportHeader().View()) pktBuf := pkt.Data().ToBuffer() combinedBuf.Merge(&pktBuf)