From 639ca440e6917a9c57ea3c6e1434f97afa97288d Mon Sep 17 00:00:00 2001 From: Lucas Manning Date: Tue, 6 Jun 2023 10:48:22 -0700 Subject: [PATCH] Change Buffer.PullUp so that it returns views that are not shared. The underlying chunks returned from PullUp should not be shared, since the underlying slice can sometimes be directly modified. This change also reworks some of the network parsing code so that ownership of views is more explicit to the reader. PiperOrigin-RevId: 538230394 --- 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, 61 insertions(+), 67 deletions(-) diff --git a/pkg/buffer/buffer.go b/pkg/buffer/buffer.go index cc663ae78..2b1e9269e 100644 --- a/pkg/buffer/buffer.go +++ b/pkg/buffer/buffer.go @@ -321,8 +321,10 @@ 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) - // Don't increment the reference count of the underlying chunk. Views - // returned by PullUp are explicitly unowned and read only + // 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() 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 d7eb2f118..8dc8b9c7d 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 unnnecesary allocation. + // Check for EOF to avoid an unnecessary allocation. if _, e := r.Read(nil); e == io.EOF { return n, nil } @@ -304,10 +304,7 @@ 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) } - if v.sharesChunk() { - defer v.chunk.DecRef() - v.chunk = v.chunk.Clone() - } + v.unshare() n := copy(v.AsSlice()[off:], p) if n < len(p) { return n, io.ErrShortWrite @@ -357,10 +354,16 @@ 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 39dc5ad02..bf1d050f0 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().View()) + fragmentableData.Append(pkt.TransportHeader().ToView()) 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 875eca473..b4092010e 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().View()) - payload.Append(pkt.TransportHeader().View()) + payload := buffer.MakeWithView(pkt.NetworkHeader().ToView()) + payload.Append(pkt.TransportHeader().ToView()) 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 2e5ab0264..652cd4990 100644 --- a/pkg/tcpip/network/ipv4/ipv4.go +++ b/pkg/tcpip/network/ipv4/ipv4.go @@ -22,7 +22,6 @@ 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" @@ -727,7 +726,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().View() + hView := pkt.NetworkHeader().ToView() defer hView.Release() h := header.IPv4(hView.AsSlice()) @@ -814,13 +813,11 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { return } - hView, ok := e.protocol.parseAndValidate(pkt) - if !ok { + if ok := e.protocol.parseAndValidate(pkt); !ok { stats.MalformedPacketsReceived.Increment() return } - h := header.IPv4(hView.AsSlice()) - defer hView.Release() + h := header.IPv4(pkt.NetworkHeader().Slice()) if !e.nic.IsLoopback() { if !e.protocol.options.AllowExternalLoopbackTraffic { @@ -836,7 +833,7 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } if e.protocol.stack.HandleLocal() { - addressEndpoint := e.AcquireAssignedAddress(header.IPv4(pkt.NetworkHeader().Slice()).SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) + addressEndpoint := e.AcquireAssignedAddress(h.SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) if addressEndpoint != nil { addressEndpoint.DecRef() @@ -857,7 +854,9 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } } - e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) + hv := pkt.NetworkHeader().ToView() + defer hv.Release() + e.handleValidatedPacket(hv.AsSlice(), pkt, e.nic.Name() /* inNICName */) } // handleLocalPacket is like HandlePacket except it does not perform the @@ -871,15 +870,14 @@ func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksu defer pkt.DecRef() pkt.RXChecksumValidated = canSkipRXChecksum - hView, ok := e.protocol.parseAndValidate(pkt) - if !ok { + if ok := e.protocol.parseAndValidate(pkt); !ok { stats.MalformedPacketsReceived.Increment() return } - h := header.IPv4(hView.AsSlice()) - defer hView.Release() - e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) + h := pkt.NetworkHeader().ToView() + defer h.Release() + e.handleValidatedPacket(h.AsSlice(), pkt, e.nic.Name() /* inNICName */) } func validateAddressesForForwarding(h header.IPv4) ip.ForwardingError { @@ -1758,31 +1756,29 @@ func (p *protocol) isSubnetLocalBroadcastAddress(addr tcpip.Address) bool { } // parseAndValidate parses the packet (including its transport layer header) and -// returns the parsed IP header. -// -// Returns true if the IP header was successfully parsed. -func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) (*buffer.View, bool) { +// returns true if the IP header was successfully parsed. +func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) bool { transProtoNum, hasTransportHdr, ok := p.Parse(pkt) if !ok { - return nil, false + return 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 nil, false + return false } if !pkt.RXChecksumValidated && !h.IsChecksumValid() { - return nil, false + return false } if hasTransportHdr { p.parseTransport(pkt, transProtoNum) } - return pkt.NetworkHeader().View(), true + return 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 4b3704eb6..cab3a1fa0 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().View()) + reassembledPayload.Append(packet.TransportHeader().ToView()) 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 a98332dc3..386bb58b6 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().View(), pkt.TransportHeader().View() + network, transport := pkt.NetworkHeader().ToView(), pkt.TransportHeader().ToView() // As per RFC 4443 section 2.4 // diff --git a/pkg/tcpip/network/ipv6/ipv6.go b/pkg/tcpip/network/ipv6/ipv6.go index ff5b44653..c2206ec92 100644 --- a/pkg/tcpip/network/ipv6/ipv6.go +++ b/pkg/tcpip/network/ipv6/ipv6.go @@ -1082,13 +1082,11 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { return } - hView, ok := e.protocol.parseAndValidate(pkt) - if !ok { + if ok := e.protocol.parseAndValidate(pkt); !ok { stats.MalformedPacketsReceived.Increment() return } - defer hView.Release() - h := header.IPv6(hView.AsSlice()) + h := header.IPv6(pkt.NetworkHeader().Slice()) if !checkV4Mapped(h, stats) { return @@ -1108,7 +1106,7 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } if e.protocol.stack.HandleLocal() { - addressEndpoint := e.AcquireAssignedAddress(header.IPv6(pkt.NetworkHeader().Slice()).SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) + addressEndpoint := e.AcquireAssignedAddress(h.SourceAddress(), e.nic.Promiscuous(), stack.CanBePrimaryEndpoint) if addressEndpoint != nil { addressEndpoint.DecRef() @@ -1129,7 +1127,9 @@ func (e *endpoint) HandlePacket(pkt stack.PacketBufferPtr) { } } - e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) + hv := pkt.NetworkHeader().ToView() + defer hv.Release() + e.handleValidatedPacket(hv.AsSlice(), pkt, e.nic.Name() /* inNICName */) } // handleLocalPacket is like HandlePacket except it does not perform the @@ -1143,19 +1143,18 @@ func (e *endpoint) handleLocalPacket(pkt stack.PacketBufferPtr, canSkipRXChecksu defer pkt.DecRef() pkt.RXChecksumValidated = canSkipRXChecksum - hView, ok := e.protocol.parseAndValidate(pkt) - if !ok { + if ok := e.protocol.parseAndValidate(pkt); !ok { stats.MalformedPacketsReceived.Increment() return } - defer hView.Release() - h := header.IPv6(hView.AsSlice()) + h := pkt.NetworkHeader().ToView() + defer h.Release() - if !checkV4Mapped(h, stats) { + if !checkV4Mapped(h.AsSlice(), stats) { return } - e.handleValidatedPacket(h, pkt, e.nic.Name() /* inNICName */) + e.handleValidatedPacket(h.AsSlice(), pkt, e.nic.Name() /* inNICName */) } // forwardMulticastPacket validates a multicast pkt and attempts to forward it. @@ -1461,12 +1460,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().View() + v := pkt.NetworkHeader().ToView() if v != nil { v.TrimFront(header.IPv6MinimumSize) } buf := buffer.MakeWithView(v) - buf.Append(pkt.TransportHeader().View()) + buf.Append(pkt.TransportHeader().ToView()) dataBuf := pkt.Data().ToBuffer() buf.Merge(&dataBuf) it := header.MakeIPv6PayloadIterator(header.IPv6ExtensionHeaderIdentifier(h.NextHeader()), buf) @@ -2588,28 +2587,25 @@ func (p *protocol) forwardPendingMulticastPacket(pkt stack.PacketBufferPtr, inst func (*protocol) Wait() {} // parseAndValidate parses the packet (including its transport layer header) and -// 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) { +// returns true if the IP header was successfully parsed. +func (p *protocol) parseAndValidate(pkt stack.PacketBufferPtr) bool { transProtoNum, hasTransportHdr, ok := p.Parse(pkt) if !ok { - return nil, false + return 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 nil, false + return false } if hasTransportHdr { p.parseTransport(pkt, transProtoNum) } - return pkt.NetworkHeader().View(), true + return 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 6d2aa6f93..ee18867ee 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().View()) + reassembledPayload.Append(fragment.TransportHeader().ToView()) reassembledPayload.Append(fragment.Data().AsRange().ToView()) } diff --git a/pkg/tcpip/stack/packet_buffer.go b/pkg/tcpip/stack/packet_buffer.go index 86b756950..66863b660 100644 --- a/pkg/tcpip/stack/packet_buffer.go +++ b/pkg/tcpip/stack/packet_buffer.go @@ -433,11 +433,9 @@ 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: payload.DeepClone(), + Payload: BufferSince(pk.NetworkHeader()), IsForwardedPacket: true, }) @@ -485,9 +483,8 @@ type PacketHeader struct { typ headerType } -// View returns an caller-owned copy of the underlying storage of h as a -// *buffer.View. -func (h PacketHeader) View() *buffer.View { +// ToView returns an caller-owned copy of the underlying storage of h. +func (h PacketHeader) ToView() *buffer.View { view := h.pk.headerView(h.typ) if view.Size() == 0 { return nil @@ -495,9 +492,9 @@ func (h PacketHeader) View() *buffer.View { return view.Clone() } -// 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. +// 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. 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 476932d2b..d397668f4 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().View()) + combinedBuf = buffer.MakeWithView(pkt.TransportHeader().ToView()) pktBuf := pkt.Data().ToBuffer() combinedBuf.Merge(&pktBuf)