From 3be95d62aedc9185c798691220604e3a812eb5c2 Mon Sep 17 00:00:00 2001 From: Lucas Manning Date: Tue, 26 Apr 2022 18:26:22 -0700 Subject: [PATCH] Refactor TCP segments to hold a PacketBuffer as data. HEAD: BenchmarkIperf/operation.Upload-16 38085657 1925 ns/op 34041.85 MB/s 546297856 bandwidth.bytes_per_second BenchmarkIperf/operation.Download-16 58649857 1506 ns/op 43520.02 MB/s 695019520 bandwidth.bytes_per_second With change: BenchmarkIperf/operation.Upload-16 40097910 (-8%) 1771 ns/op 36996.64 MB/s 598457344 bandwidth.bytes_per_second BenchmarkIperf/operation.Download-16 58393820 (-21%) 1196 ns/op 54808.86 MB/s 880225280 bandwidth.bytes_per_second PiperOrigin-RevId: 444723128 --- pkg/tcpip/stack/BUILD | 1 + pkg/tcpip/stack/packet_buffer.go | 59 +++++- pkg/tcpip/stack/packet_buffer_state.go | 28 +++ pkg/tcpip/stack/packet_buffer_test.go | 180 +++++++++++++----- pkg/tcpip/stack/packet_buffer_unsafe.go | 3 +- pkg/tcpip/transport/tcp/accept.go | 18 +- pkg/tcpip/transport/tcp/connect.go | 81 ++++---- pkg/tcpip/transport/tcp/dispatcher.go | 6 +- pkg/tcpip/transport/tcp/endpoint.go | 4 +- pkg/tcpip/transport/tcp/forwarder.go | 7 +- pkg/tcpip/transport/tcp/protocol.go | 15 +- pkg/tcpip/transport/tcp/rack.go | 6 +- pkg/tcpip/transport/tcp/rcv.go | 14 +- pkg/tcpip/transport/tcp/segment.go | 154 +++++++-------- pkg/tcpip/transport/tcp/segment_state.go | 24 --- pkg/tcpip/transport/tcp/segment_test.go | 10 +- pkg/tcpip/transport/tcp/segment_unsafe.go | 3 +- pkg/tcpip/transport/tcp/snd.go | 62 +++--- pkg/tcpip/transport/tcp/test/e2e/tcp_test.go | 8 +- test/packetimpact/tests/BUILD | 1 + .../tests/tcp_zero_receive_window_test.go | 12 +- 21 files changed, 429 insertions(+), 267 deletions(-) create mode 100644 pkg/tcpip/stack/packet_buffer_state.go diff --git a/pkg/tcpip/stack/BUILD b/pkg/tcpip/stack/BUILD index 1ecfd46a7..9751f3147 100644 --- a/pkg/tcpip/stack/BUILD +++ b/pkg/tcpip/stack/BUILD @@ -71,6 +71,7 @@ go_library( "packet_buffer.go", "packet_buffer_list.go", "packet_buffer_refs.go", + "packet_buffer_state.go", "packet_buffer_unsafe.go", "pending_packets.go", "rand.go", diff --git a/pkg/tcpip/stack/packet_buffer.go b/pkg/tcpip/stack/packet_buffer.go index 3fe7c7031..f0bf383ee 100644 --- a/pkg/tcpip/stack/packet_buffer.go +++ b/pkg/tcpip/stack/packet_buffer.go @@ -108,7 +108,7 @@ type PacketBuffer struct { // buf is the underlying buffer for the packet. See struct level docs for // details. - buf buffer.Buffer + buf buffer.Buffer `state:".([]byte)"` reserved int pushed int consumed int @@ -251,7 +251,7 @@ func (pk *PacketBuffer) Size() int { // MemSize returns the estimation size of the pk in memory, including backing // buffer data. func (pk *PacketBuffer) MemSize() int { - return int(pk.buf.Size()) + packetBufferStructSize + return int(pk.buf.Size()) + PacketBufferStructSize } // Data returns the handle to data portion of pk. @@ -348,6 +348,17 @@ func (pk *PacketBuffer) Clone() *PacketBuffer { return newPk } +// ResetHeaders clears headers in the underlying buffer, resets all header +// fields, and prepends reserved space for new headers in pk. +func (pk *PacketBuffer) ResetHeaders(reserved int) { + pk.buf.Remove(0, pk.dataOffset()) + pk.headers = [numHeaderType]headerInfo{} + pk.consumed = 0 + pk.pushed = 0 + pk.reserved = reserved + pk.buf.Prepend(make([]byte, reserved)) +} + // Network returns the network header as a header.Network. // // Network should only be called when NetworkHeader has been set. @@ -572,6 +583,50 @@ func (d PacketData) ReadFromVV(srcVV *tcpipbuffer.VectorisedView, count int) int return done } +// AppendRange appends and takes ownership of the data in r. +func (d PacketData) AppendRange(r Range) { + r.iterate(func(b []byte) { + d.pk.buf.AppendOwned(b) + }) +} + +// Merge clears headers in oth and merges its data with d. +func (d PacketData) Merge(oth PacketData) { + oth.pk.buf.TrimFront(int64(oth.pk.dataOffset())) + d.pk.buf.Merge(&oth.pk.buf) +} + +// ReadFrom moves at most count bytes from the beginning of src to the end +// of d. +func (d PacketData) ReadFrom(src PacketData, count int) { + done := 0 + for _, v := range src.Views() { + if len(v) < count { + count -= len(v) + done += len(v) + // Use AppendOwned to avoid the cost of copying data between buffers. + // This is safe because the buffers are trimmed out of src at the end + // of the function anyways. + d.pk.buf.AppendOwned(v) + } else { + v = v[:count] + count -= len(v) + done += len(v) + d.pk.buf.Append(v) + break + } + } + src.TrimFront(done) +} + +// TrimFront removes up to count bytes from the front of d's payload. +func (d PacketData) TrimFront(count int) { + if count > d.Size() { + count = d.Size() + } + d.pk.buf.Remove(d.pk.dataOffset(), count) +} + // Size returns the number of bytes in the data payload of the packet. func (d PacketData) Size() int { return int(d.pk.buf.Size()) - d.pk.dataOffset() diff --git a/pkg/tcpip/stack/packet_buffer_state.go b/pkg/tcpip/stack/packet_buffer_state.go new file mode 100644 index 000000000..ad7b45cf0 --- /dev/null +++ b/pkg/tcpip/stack/packet_buffer_state.go @@ -0,0 +1,28 @@ +// Copyright 2022 The gVisor Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at // +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package stack + +// saveBuf is invoked by stateify. +func (pk *PacketBuffer) saveBuf() []byte { + var bytes []byte + pk.buf.Apply(func(v []byte) { + bytes = append(bytes, v...) + }) + return bytes +} + +// loadBuf is invoked by stateify. +func (pk *PacketBuffer) loadBuf(data []byte) { + pk.buf.Append(data) +} diff --git a/pkg/tcpip/stack/packet_buffer_test.go b/pkg/tcpip/stack/packet_buffer_test.go index c376ed1a1..48cd11ccb 100644 --- a/pkg/tcpip/stack/packet_buffer_test.go +++ b/pkg/tcpip/stack/packet_buffer_test.go @@ -336,6 +336,30 @@ func TestPacketHeaderConsumeCalledAtMostOnce(t *testing.T) { } } +func TestResetHeadersAllowsDoublePush(t *testing.T) { + link1 := makeView(10) + link2 := makeView(20) + data := makeView(30) + pk := NewPacketBuffer(PacketBufferOptions{ + ReserveHeaderBytes: len(link1), + Data: data.ToVectorisedView(), + }) + + copy(pk.LinkHeader().Push(len(link1)), link1) + checkPacketContents(t, "" /* prefix */, pk, packetContents{ + link: link1, + data: data, + }) + + pk.ResetHeaders(len(link2)) + + copy(pk.LinkHeader().Push(len(link2)), link2) + checkPacketContents(t, "" /* prefix */, pk, packetContents{ + link: link2, + data: data, + }) +} + func TestPacketHeaderPushThenConsumePanics(t *testing.T) { const headerSize = 10 @@ -416,16 +440,19 @@ func TestPacketBufferData(t *testing.T) { } { t.Run(tc.name, func(t *testing.T) { // PullUp - for _, n := range []int{1, len(tc.data)} { - t.Run(fmt.Sprintf("PullUp%d", n), func(t *testing.T) { - pkt := tc.makePkt(t) - v, ok := pkt.Data().PullUp(n) - wantV := []byte(tc.data)[:n] - if !ok || !bytes.Equal(v, wantV) { - t.Errorf("pkt.Data().PullUp(%d) = %q, %t; want %q, true", n, v, ok, wantV) - } - }) - } + t.Run("PullUp", func(t *testing.T) { + for _, n := range []int{1, len(tc.data)} { + t.Run(fmt.Sprintf("%dbytes", n), func(t *testing.T) { + pkt := tc.makePkt(t) + v, ok := pkt.Data().PullUp(n) + wantV := []byte(tc.data)[:n] + if !ok || !bytes.Equal(v, wantV) { + t.Errorf("pkt.Data().PullUp(%d) = %q, %t; want %q, true", n, v, ok, wantV) + } + }) + } + }) + t.Run("PullUpOutOfBounds", func(t *testing.T) { n := len(tc.data) + 1 pkt := tc.makePkt(t) @@ -436,34 +463,38 @@ func TestPacketBufferData(t *testing.T) { }) // Consume. - for _, n := range []int{1, len(tc.data)} { - t.Run(fmt.Sprintf("Consume%d", n), func(t *testing.T) { - pkt := tc.makePkt(t) - v, ok := pkt.Data().Consume(n) - if !ok { - t.Fatalf("Consume failed") - } - if want := []byte(tc.data)[:n]; !bytes.Equal(v, want) { - t.Fatalf("pkt.Data().Consume(n) = 0x%x, want 0x%x", v, want) - } + t.Run("Consume", func(t *testing.T) { + for _, n := range []int{1, len(tc.data)} { + t.Run(fmt.Sprintf("%dbytes", n), func(t *testing.T) { + pkt := tc.makePkt(t) + v, ok := pkt.Data().Consume(n) + if !ok { + t.Fatalf("Consume failed") + } + if want := []byte(tc.data)[:n]; !bytes.Equal(v, want) { + t.Fatalf("pkt.Data().Consume(n) = 0x%x, want 0x%x", v, want) + } - checkData(t, pkt, []byte(tc.data)[n:]) - }) - } + checkData(t, pkt, []byte(tc.data)[n:]) + }) + } + }) // CapLength - for _, n := range []int{0, 1, len(tc.data)} { - t.Run(fmt.Sprintf("CapLength%d", n), func(t *testing.T) { - pkt := tc.makePkt(t) - pkt.Data().CapLength(n) + t.Run("CapLength", func(t *testing.T) { + for _, n := range []int{0, 1, len(tc.data)} { + t.Run("%dbytes", func(t *testing.T) { + pkt := tc.makePkt(t) + pkt.Data().CapLength(n) - want := []byte(tc.data) - if n < len(want) { - want = want[:n] - } - checkData(t, pkt, want) - }) - } + want := []byte(tc.data) + if n < len(want) { + want = want[:n] + } + checkData(t, pkt, want) + }) + } + }) // Views t.Run("Views", func(t *testing.T) { @@ -482,21 +513,23 @@ func TestPacketBufferData(t *testing.T) { }) // ReadFromVV - for _, n := range []int{0, 1, 2, 7, 10, 14, 20} { - t.Run(fmt.Sprintf("ReadFromVV%d", n), func(t *testing.T) { - s := "TO READ" - srcVV := vv(s, s) - s += s + t.Run("ReadFromVV", func(t *testing.T) { + for _, n := range []int{0, 1, 2, 7, 10, 14, 20} { + t.Run(fmt.Sprintf("%dbytes", n), func(t *testing.T) { + s := "TO READ" + srcVV := vv(s, s) + s += s - pkt := tc.makePkt(t) - pkt.Data().ReadFromVV(&srcVV, n) + pkt := tc.makePkt(t) + pkt.Data().ReadFromVV(&srcVV, n) - if n < len(s) { - s = s[:n] - } - checkData(t, pkt, []byte(tc.data+s)) - }) - } + if n < len(s) { + s = s[:n] + } + checkData(t, pkt, []byte(tc.data+s)) + }) + } + }) // ExtractVV t.Run("ExtractVV", func(t *testing.T) { @@ -509,6 +542,59 @@ func TestPacketBufferData(t *testing.T) { t.Errorf("pkt.Data().ExtractVV().ToOwnedView() = %q, want %q", got, want) } }) + + t.Run("AppendRange", func(t *testing.T) { + pkt1 := tc.makePkt(t) + pkt2 := tc.makePkt(t) + subRangeStart := 2 + pkt1.Data().AppendRange(pkt2.Data().AsRange().SubRange(subRangeStart)) + checkData(t, pkt1, []byte(tc.data+tc.data[subRangeStart:])) + }) + + t.Run("Merge", func(t *testing.T) { + pkt1 := tc.makePkt(t) + pkt2 := tc.makePkt(t) + pkt1.Data().Merge(pkt2.Data()) + + checkData(t, pkt1, []byte(tc.data+tc.data)) + if pkt2.buf.Size() != 0 { + t.Errorf("pkt.buf.Size() = %v, want %v", 0, pkt2.buf.Size()) + } + }) + + t.Run("ReadFrom", func(t *testing.T) { + for _, n := range []int{0, 1, 2, 7, 10, 14, 20} { + t.Run(fmt.Sprintf("%dbytes", n), func(t *testing.T) { + pkt1 := tc.makePkt(t) + pkt2 := tc.makePkt(t) + pkt1.Data().ReadFrom(pkt2.Data(), n) + + want1 := tc.data + want2 := "" + if n < len(tc.data) { + want1 = tc.data[:n] + want2 = tc.data[n:] + } + checkData(t, pkt1, []byte(tc.data+want1)) + checkData(t, pkt2, []byte(want2)) + }) + } + }) + + t.Run("TrimFront", func(t *testing.T) { + for _, n := range []int{0, 1, 2, 7, 10, 14, 20} { + t.Run(fmt.Sprintf("%dbytes", n), func(t *testing.T) { + pkt := tc.makePkt(t) + pkt.Data().TrimFront(n) + + want := "" + if n < len(tc.data) { + want = tc.data[n:] + } + checkData(t, pkt, []byte(want)) + }) + } + }) }) } } diff --git a/pkg/tcpip/stack/packet_buffer_unsafe.go b/pkg/tcpip/stack/packet_buffer_unsafe.go index ee3d47270..cd151ff0b 100644 --- a/pkg/tcpip/stack/packet_buffer_unsafe.go +++ b/pkg/tcpip/stack/packet_buffer_unsafe.go @@ -16,4 +16,5 @@ package stack import "unsafe" -const packetBufferStructSize = int(unsafe.Sizeof(PacketBuffer{})) +// PacketBufferStructSize is the minimal size of the packet buffer overhead. +const PacketBufferStructSize = int(unsafe.Sizeof(PacketBuffer{})) diff --git a/pkg/tcpip/transport/tcp/accept.go b/pkg/tcpip/transport/tcp/accept.go index 27994925a..e7cf947f2 100644 --- a/pkg/tcpip/transport/tcp/accept.go +++ b/pkg/tcpip/transport/tcp/accept.go @@ -187,10 +187,10 @@ func (l *listenContext) createConnectingEndpoint(s *segment, rcvdSynOpts header. // Create a new endpoint. netProto := l.netProto if netProto == 0 { - netProto = s.netProto + netProto = s.pkt.NetworkProtocolNumber } - route, err := l.stack.FindRoute(s.nicID, s.dstAddr, s.srcAddr, s.netProto, false /* multicastLoop */) + route, err := l.stack.FindRoute(s.pkt.NICID, s.pkt.Network().DestinationAddress(), s.pkt.Network().SourceAddress(), s.pkt.NetworkProtocolNumber, false /* multicastLoop */) if err != nil { return nil, err // +checklocksignore } @@ -199,9 +199,9 @@ func (l *listenContext) createConnectingEndpoint(s *segment, rcvdSynOpts header. n.mu.Lock() n.ops.SetV6Only(l.v6Only) n.TransportEndpointInfo.ID = s.id - n.boundNICID = s.nicID + n.boundNICID = s.pkt.NICID n.route = route - n.effectiveNetProtos = []tcpip.NetworkProtocolNumber{s.netProto} + n.effectiveNetProtos = []tcpip.NetworkProtocolNumber{s.pkt.NetworkProtocolNumber} n.ops.SetReceiveBufferSize(int64(l.rcvWnd), false /* notify */) n.amss = calculateAdvertisedMSS(n.userMSS, n.route) n.setEndpointState(StateConnecting) @@ -495,7 +495,8 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err return nil } - route, err := e.stack.FindRoute(s.nicID, s.dstAddr, s.srcAddr, s.netProto, false /* multicastLoop */) + net := s.pkt.Network() + route, err := e.stack.FindRoute(s.pkt.NICID, net.DestinationAddress(), net.SourceAddress(), s.pkt.NetworkProtocolNumber, false /* multicastLoop */) if err != nil { return err } @@ -516,7 +517,7 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err MSS: calculateAdvertisedMSS(e.userMSS, route), } if opts.TS { - offset := e.protocol.tsOffset(s.dstAddr, s.srcAddr) + offset := e.protocol.tsOffset(net.DestinationAddress(), net.SourceAddress()) now := e.stack.Clock().NowMonotonic() synOpts.TSVal = offset.TSVal(now) } @@ -648,7 +649,8 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err } n.isRegistered = true - n.TSOffset = n.protocol.tsOffset(s.dstAddr, s.srcAddr) + net := s.pkt.Network() + n.TSOffset = n.protocol.tsOffset(net.DestinationAddress(), net.SourceAddress()) // Switch state to connected. n.isConnectNotified = true @@ -669,7 +671,7 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err // Requeue the segment if the ACK completing the handshake has more info // to be procesed by the newly established endpoint. - if (s.flags.Contains(header.TCPFlagFin) || s.data.Size() > 0) && n.enqueueSegment(s) { + if (s.flags.Contains(header.TCPFlagFin) || s.payloadSize() > 0) && n.enqueueSegment(s) { n.notifyProcessor() } diff --git a/pkg/tcpip/transport/tcp/connect.go b/pkg/tcpip/transport/tcp/connect.go index c13517437..ff20a1478 100644 --- a/pkg/tcpip/transport/tcp/connect.go +++ b/pkg/tcpip/transport/tcp/connect.go @@ -22,7 +22,6 @@ import ( "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" - "gvisor.dev/gvisor/pkg/tcpip/buffer" "gvisor.dev/gvisor/pkg/tcpip/hash/jenkins" "gvisor.dev/gvisor/pkg/tcpip/header" "gvisor.dev/gvisor/pkg/tcpip/seqnum" @@ -453,7 +452,7 @@ func (h *handshake) synRcvdState(s *segment) tcpip.Error { if s.flags.Contains(header.TCPFlagAck) { // If deferAccept is not zero and this is a bare ACK and the // timeout is not hit then drop the ACK. - if h.deferAccept != 0 && s.data.Size() == 0 && h.ep.stack.Clock().NowMonotonic().Sub(h.startTime) < h.deferAccept { + if h.deferAccept != 0 && s.payloadSize() == 0 && h.ep.stack.Clock().NowMonotonic().Sub(h.startTime) < h.deferAccept { h.acked = true h.ep.stack.Stats().DroppedPackets.Increment() return nil @@ -486,7 +485,7 @@ func (h *handshake) synRcvdState(s *segment) tcpip.Error { // Requeue the segment if the ACK completing the handshake has more info // to be procesed by the newly established endpoint. - if (s.flags.Contains(header.TCPFlagFin) || s.data.Size() > 0) && h.ep.enqueueSegment(s) { + if (s.flags.Contains(header.TCPFlagFin) || s.payloadSize() > 0) && h.ep.enqueueSegment(s) { h.ep.protocol.dispatcher.selectProcessor(h.ep.ID).queueEndpoint(h.ep) } @@ -796,16 +795,18 @@ type tcpFields struct { func (e *endpoint) sendSynTCP(r *stack.Route, tf tcpFields, opts header.TCPSynOptions) tcpip.Error { tf.opts = makeSynOptions(opts) // We ignore SYN send errors and let the callers re-attempt send. - if err := e.sendTCP(r, tf, buffer.VectorisedView{}, stack.GSO{}); err != nil { + p := stack.NewPacketBuffer(stack.PacketBufferOptions{}) + defer p.DecRef() + if err := e.sendTCP(r, tf, p, stack.GSO{}); err != nil { e.stats.SendErrors.SynSendToNetworkFailed.Increment() } putOptions(tf.opts) return nil } -func (e *endpoint) sendTCP(r *stack.Route, tf tcpFields, data buffer.VectorisedView, gso stack.GSO) tcpip.Error { +func (e *endpoint) sendTCP(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stack.GSO) tcpip.Error { tf.txHash = e.txHash - if err := sendTCP(r, tf, data, gso, e.owner); err != nil { + if err := sendTCP(r, tf, pkt, gso, e.owner); err != nil { e.stats.SendErrors.SegmentSendToNetworkFailed.Increment() return err } @@ -842,22 +843,16 @@ func buildTCPHdr(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stac } } -func sendTCPBatch(r *stack.Route, tf tcpFields, data buffer.VectorisedView, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { - // We need to shallow clone the VectorisedView here as ReadToView will - // split the VectorisedView and Trim underlying views as it splits. Not - // doing the clone here will cause the underlying views of data itself - // to be altered. - data = data.Clone(nil) - +func sendTCPBatch(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { optLen := len(tf.opts) if tf.rcvWnd > math.MaxUint16 { tf.rcvWnd = math.MaxUint16 } mss := int(gso.MSS) - n := (data.Size() + mss - 1) / mss + n := (pkt.Data().Size() + mss - 1) / mss - size := data.Size() + size := pkt.Data().Size() hdrSize := header.TCPMinimumSize + int(r.MaxHeaderLength()) + optLen for i := 0; i < n; i++ { packetSize := mss @@ -865,43 +860,57 @@ func sendTCPBatch(r *stack.Route, tf tcpFields, data buffer.VectorisedView, gso packetSize = size } size -= packetSize - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ - ReserveHeaderBytes: hdrSize, - }) + + pkt := pkt + // No need to split the packet in the final iteration. The original + // packet already has the truncated data. + shouldSplitPacket := i != n-1 + if shouldSplitPacket { + splitPkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ReserveHeaderBytes: hdrSize}) + splitPkt.Data().ReadFrom(pkt.Data(), packetSize) + pkt = splitPkt + } + pkt.Hash = tf.txHash pkt.Owner = owner - pkt.Data().ReadFromVV(&data, packetSize) buildTCPHdr(r, tf, pkt, gso) tf.seq = tf.seq.Add(seqnum.Size(packetSize)) pkt.GSOOptions = gso if err := r.WritePacket(stack.NetworkHeaderParams{Protocol: ProtocolNumber, TTL: tf.ttl, TOS: tf.tos}, pkt); err != nil { r.Stats().TCP.SegmentSendErrors.Increment() - pkt.DecRef() + if shouldSplitPacket { + pkt.DecRef() + } return err } r.Stats().TCP.SegmentsSent.Increment() - pkt.DecRef() + if shouldSplitPacket { + pkt.DecRef() + } } return nil } // sendTCP sends a TCP segment with the provided options via the provided // network endpoint and under the provided identity. -func sendTCP(r *stack.Route, tf tcpFields, data buffer.VectorisedView, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { +func sendTCP(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { optLen := len(tf.opts) if tf.rcvWnd > math.MaxUint16 { tf.rcvWnd = math.MaxUint16 } + // We need to create a new packet because WritePacket can modify pkt's data, + // and pkt is held by a segment that could be reprocessed later on. + { + sendPkt := pkt.Clone() + defer sendPkt.DecRef() + sendPkt.ResetHeaders(header.TCPMinimumSize + int(r.MaxHeaderLength()) + optLen) + pkt = sendPkt + } - if r.Loop()&stack.PacketLoop == 0 && gso.Type == stack.GSOSW && int(gso.MSS) < data.Size() { - return sendTCPBatch(r, tf, data, gso, owner) + if r.Loop()&stack.PacketLoop == 0 && gso.Type == stack.GSOSW && int(gso.MSS) < pkt.Data().Size() { + return sendTCPBatch(r, tf, pkt, gso, owner) } - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ - ReserveHeaderBytes: header.TCPMinimumSize + int(r.MaxHeaderLength()) + optLen, - Data: data, - }) - defer pkt.DecRef() pkt.GSOOptions = gso pkt.Hash = tf.txHash pkt.Owner = owner @@ -959,11 +968,13 @@ func (e *endpoint) makeOptions(sackBlocks []header.SACKBlock) []byte { // sendEmptyRaw sends a TCP segment to the endpoint's peer. func (e *endpoint) sendEmptyRaw(flags header.TCPFlags, seq, ack seqnum.Value, rcvWnd seqnum.Size) tcpip.Error { - return e.sendRaw(buffer.VectorisedView{}, flags, seq, ack, rcvWnd) + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{}) + defer pkt.DecRef() + return e.sendRaw(pkt, flags, seq, ack, rcvWnd) } // sendRaw sends a TCP segment to the endpoint's peer. -func (e *endpoint) sendRaw(data buffer.VectorisedView, flags header.TCPFlags, seq, ack seqnum.Value, rcvWnd seqnum.Size) tcpip.Error { +func (e *endpoint) sendRaw(pkt *stack.PacketBuffer, 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] @@ -978,7 +989,7 @@ func (e *endpoint) sendRaw(data buffer.VectorisedView, flags header.TCPFlags, se ack: ack, rcvWnd: rcvWnd, opts: options, - }, data, e.gso) + }, pkt, e.gso) putOptions(options) return err } @@ -1059,14 +1070,14 @@ func (e *endpoint) transitionToStateCloseLocked() { // only when the endpoint is in StateClose and we want to deliver the segment // to any other listening endpoint. We reply with RST if we cannot find one. func (e *endpoint) tryDeliverSegmentFromClosedEndpoint(s *segment) { - ep := e.stack.FindTransportEndpoint(e.NetProto, e.TransProto, e.TransportEndpointInfo.ID, s.nicID) + ep := e.stack.FindTransportEndpoint(e.NetProto, e.TransProto, e.TransportEndpointInfo.ID, s.pkt.NICID) if ep == nil && e.NetProto == header.IPv6ProtocolNumber && e.TransportEndpointInfo.ID.LocalAddress.To4() != "" { // Dual-stack socket, try IPv4. ep = e.stack.FindTransportEndpoint( header.IPv4ProtocolNumber, e.TransProto, e.TransportEndpointInfo.ID, - s.nicID, + s.pkt.NICID, ) } if ep == nil { @@ -1393,7 +1404,7 @@ func (e *endpoint) handleTimeWaitSegments() (extendTimeWait bool, reuseTW func() netProtos = []tcpip.NetworkProtocolNumber{header.IPv4ProtocolNumber, header.IPv6ProtocolNumber} } for _, netProto := range netProtos { - if listenEP := e.stack.FindTransportEndpoint(netProto, info.TransProto, newID, s.nicID); listenEP != nil { + if listenEP := e.stack.FindTransportEndpoint(netProto, info.TransProto, newID, s.pkt.NICID); listenEP != nil { tcpEP := listenEP.(*endpoint) if EndpointState(tcpEP.State()) == StateListen { reuseTW = func() { diff --git a/pkg/tcpip/transport/tcp/dispatcher.go b/pkg/tcpip/transport/tcp/dispatcher.go index c066b29a6..39665cbc4 100644 --- a/pkg/tcpip/transport/tcp/dispatcher.go +++ b/pkg/tcpip/transport/tcp/dispatcher.go @@ -415,13 +415,13 @@ func (d *dispatcher) queuePacket(stackEP stack.TransportEndpoint, id stack.Trans ep := stackEP.(*endpoint) - s := newIncomingSegment(id, clock, pkt) - defer s.DecRef() - if !s.parse(pkt.RXTransportChecksumValidated) { + s, err := newIncomingSegment(id, clock, pkt) + if err != nil { ep.stack.Stats().TCP.InvalidSegmentsReceived.Increment() ep.stats.ReceiveErrors.MalformedPacketsReceived.Increment() return } + defer s.DecRef() if !s.csumValid { ep.stack.Stats().TCP.ChecksumErrors.Increment() diff --git a/pkg/tcpip/transport/tcp/endpoint.go b/pkg/tcpip/transport/tcp/endpoint.go index 78a93fce2..e9509f896 100644 --- a/pkg/tcpip/transport/tcp/endpoint.go +++ b/pkg/tcpip/transport/tcp/endpoint.go @@ -1372,7 +1372,7 @@ func (e *endpoint) Read(dst io.Writer, opts tcpip.ReadOptions) (tcpip.ReadResult s := first for s != nil { var n int - n, err = s.data.ReadTo(dst, opts.Peek) + n, err = s.ReadTo(dst, opts.Peek) // Book keeping first then error handling. done += n @@ -1472,7 +1472,7 @@ func (e *endpoint) commitRead(done int) *segment { e.rcvQueueInfo.rcvQueueMu.Lock() memDelta := 0 s := e.rcvQueueInfo.rcvQueue.Front() - for s != nil && s.data.Size() == 0 { + for s != nil && s.payloadSize() == 0 { e.rcvQueueInfo.rcvQueue.Remove(s) // Memory is only considered released when the whole segment has been // read. diff --git a/pkg/tcpip/transport/tcp/forwarder.go b/pkg/tcpip/transport/tcp/forwarder.go index fcc1f9dc7..b3888ff20 100644 --- a/pkg/tcpip/transport/tcp/forwarder.go +++ b/pkg/tcpip/transport/tcp/forwarder.go @@ -65,11 +65,14 @@ 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 { - s := newIncomingSegment(id, f.stack.Clock(), pkt) + s, err := newIncomingSegment(id, f.stack.Clock(), pkt) + if err != nil { + return false + } defer s.DecRef() // We only care about well-formed SYN packets (not SYN-ACK) packets. - if !s.parse(pkt.RXTransportChecksumValidated) || !s.csumValid || !s.flags.Contains(header.TCPFlagSyn) || s.flags.Contains(header.TCPFlagAck) { + if !s.csumValid || !s.flags.Contains(header.TCPFlagSyn) || s.flags.Contains(header.TCPFlagAck) { return false } diff --git a/pkg/tcpip/transport/tcp/protocol.go b/pkg/tcpip/transport/tcp/protocol.go index 34ab1d9b2..26958568d 100644 --- a/pkg/tcpip/transport/tcp/protocol.go +++ b/pkg/tcpip/transport/tcp/protocol.go @@ -157,10 +157,12 @@ func (p *protocol) QueuePacket(ep stack.TransportEndpoint, id stack.TransportEnd // 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 { - s := newIncomingSegment(id, p.stack.Clock(), pkt) + s, err := newIncomingSegment(id, p.stack.Clock(), pkt) + if err != nil { + return stack.UnknownDestinationPacketMalformed + } defer s.DecRef() - - if !s.parse(pkt.RXTransportChecksumValidated) || !s.csumValid { + if !s.csumValid { return stack.UnknownDestinationPacketMalformed } @@ -194,7 +196,8 @@ func (p *protocol) tsOffset(src, dst tcpip.Address) tcp.TSOffset { // If the relevant TTL has its reset value (0 for ipv4TTL, -1 for ipv6HopLimit), // then the route's default TTL will be used. func replyWithReset(st *stack.Stack, s *segment, tos, ipv4TTL uint8, ipv6HopLimit int16) tcpip.Error { - route, err := st.FindRoute(s.nicID, s.dstAddr, s.srcAddr, s.netProto, false /* multicastLoop */) + net := s.pkt.Network() + route, err := st.FindRoute(s.pkt.NICID, net.DestinationAddress(), net.SourceAddress(), s.pkt.NetworkProtocolNumber, false /* multicastLoop */) if err != nil { return err } @@ -224,6 +227,8 @@ func replyWithReset(st *stack.Stack, s *segment, tos, ipv4TTL uint8, ipv6HopLimi ack = s.sequenceNumber.Add(s.logicalLen()) } + p := stack.NewPacketBuffer(stack.PacketBufferOptions{}) + defer p.DecRef() return sendTCP(route, tcpFields{ id: s.id, ttl: ttl, @@ -232,7 +237,7 @@ func replyWithReset(st *stack.Stack, s *segment, tos, ipv4TTL uint8, ipv6HopLimi seq: seq, ack: ack, rcvWnd: 0, - }, buffer.VectorisedView{}, stack.GSO{}, nil /* PacketOwner */) + }, p, stack.GSO{}, nil /* PacketOwner */) } // SetOption implements stack.TransportProtocol.SetOption. diff --git a/pkg/tcpip/transport/tcp/rack.go b/pkg/tcpip/transport/tcp/rack.go index b8d0bb653..6e5852208 100644 --- a/pkg/tcpip/transport/tcp/rack.go +++ b/pkg/tcpip/transport/tcp/rack.go @@ -113,7 +113,7 @@ func (rc *rackControl) update(seg *segment, ackSeg *segment) { // Update rc.xmitTime and rc.endSequence to the transmit time and // ending sequence number of the packet which has been acknowledged // most recently. - endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) + endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) if rc.XmitTime.Before(seg.xmitTime) || (seg.xmitTime == rc.XmitTime && rc.EndSequence.LessThan(endSeq)) { rc.XmitTime = seg.xmitTime rc.EndSequence = endSeq @@ -132,7 +132,7 @@ func (rc *rackControl) update(seg *segment, ackSeg *segment) { // delivered out of order. The sender sets RACK.reord to TRUE if such segment // is identified. func (rc *rackControl) detectReorder(seg *segment) { - endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) + endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) if rc.FACK.LessThan(endSeq) { rc.FACK = endSeq return @@ -366,7 +366,7 @@ func (rc *rackControl) detectLoss(rcvTime tcpip.MonotonicTime) int { continue } - endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) + endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) if seg.xmitTime.Before(rc.XmitTime) || (seg.xmitTime == rc.XmitTime && rc.EndSequence.LessThan(endSeq)) { timeRemaining := seg.xmitTime.Sub(rcvTime) + rc.RTT + rc.ReoWnd if timeRemaining <= 0 { diff --git a/pkg/tcpip/transport/tcp/rcv.go b/pkg/tcpip/transport/tcp/rcv.go index 88251d576..bd5ece56f 100644 --- a/pkg/tcpip/transport/tcp/rcv.go +++ b/pkg/tcpip/transport/tcp/rcv.go @@ -126,7 +126,7 @@ func (r *receiver) getSendParams() (RcvNxt seqnum.Value, rcvWnd seqnum.Size) { // // Also, if the application is reading the data, we keep growing the right // edge, as we are still advertising a window that we think can be serviced. - toGrow := unackLen >= SegSize || bufUsed <= r.prevBufUsed + toGrow := unackLen >= SegOverheadSize || bufUsed <= r.prevBufUsed // Update RcvAcc only if new window is > previously advertised window. We // should never shrink the acceptable sequence space once it has been @@ -216,7 +216,7 @@ func (r *receiver) consumeSegment(s *segment, segSeq seqnum.Value, segLen seqnum segLen -= diff segSeq.UpdateForward(diff) s.sequenceNumber.UpdateForward(diff) - s.data.TrimFront(int(diff)) + s.TrimFront(diff) } // Move segment to ready-to-deliver list. Wakeup any waiters. @@ -400,7 +400,7 @@ func (r *receiver) handleRcvdSegmentClosing(s *segment, state EndpointState, clo // incoming FIN or the user calling shutdown(.., // SHUT_RD) then any data past the RcvNxt should // trigger a RST. - endDataSeq := s.sequenceNumber.Add(seqnum.Size(s.data.Size())) + endDataSeq := s.sequenceNumber.Add(seqnum.Size(s.payloadSize())) if state != StateCloseWait && rcvClosed && r.RcvNxt.LessThan(endDataSeq) { return true, &tcpip.ErrConnectionAborted{} } @@ -441,7 +441,7 @@ func (r *receiver) handleRcvdSegmentClosing(s *segment, state EndpointState, clo // NOTE: We still want to permit a FIN as it's possible only our // end has closed and the peer is yet to send a FIN. Hence we // compare only the payload. - segEnd := s.sequenceNumber.Add(seqnum.Size(s.data.Size())) + segEnd := s.sequenceNumber.Add(seqnum.Size(s.payloadSize())) if rcvClosed && !segEnd.LessThanEq(r.RcvNxt) { return true, nil } @@ -456,7 +456,7 @@ func (r *receiver) handleRcvdSegment(s *segment) (drop bool, err tcpip.Error) { state := r.ep.EndpointState() closed := r.ep.closed - segLen := seqnum.Size(s.data.Size()) + segLen := seqnum.Size(s.payloadSize()) segSeq := s.sequenceNumber // If the sequence number range is outside the acceptable range, just @@ -513,7 +513,7 @@ func (r *receiver) handleRcvdSegment(s *segment) (drop bool, err tcpip.Error) { // now. So try to do it. for !r.closed && r.pendingRcvdSegments.Len() > 0 { s := r.pendingRcvdSegments[0] - segLen := seqnum.Size(s.data.Size()) + segLen := seqnum.Size(s.payloadSize()) segSeq := s.sequenceNumber // Skip segment altogether if it has already been acknowledged. @@ -537,7 +537,7 @@ func (r *receiver) handleRcvdSegment(s *segment) (drop bool, err tcpip.Error) { // +checklocksalias:r.ep.snd.ep.mu=r.ep.mu func (r *receiver) handleTimeWaitSegment(s *segment) (resetTimeWait bool, newSyn bool) { segSeq := s.sequenceNumber - segLen := seqnum.Size(s.data.Size()) + segLen := seqnum.Size(s.payloadSize()) // Just silently drop any RST packets in TIME_WAIT. We do not support // TIME_WAIT assasination as a result we confirm w/ fix 1 as described diff --git a/pkg/tcpip/transport/tcp/segment.go b/pkg/tcpip/transport/tcp/segment.go index 43968704b..d2c567130 100644 --- a/pkg/tcpip/transport/tcp/segment.go +++ b/pkg/tcpip/transport/tcp/segment.go @@ -16,6 +16,7 @@ package tcp import ( "fmt" + "io" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/buffer" @@ -29,7 +30,11 @@ import ( type queueFlags uint8 const ( - recvQ queueFlags = 1 << iota + // SegOverheadSize is the size of an empty seg in memory including packet + // buffer overhead. It is advised to use SegOverheadSize instead of segSize + // in all cases where accounting for segment memory overhead is important. + SegOverheadSize = segSize + stack.PacketBufferStructSize + header.IPv4MaximumHeaderSize + recvQ queueFlags = 1 << iota sendQ ) @@ -46,19 +51,8 @@ type segment struct { qFlags queueFlags id stack.TransportEndpointID `state:"manual"` - // TODO(gvisor.dev/issue/4417): Hold a stack.PacketBuffer instead of - // individual members for link/network packet info. - srcAddr tcpip.Address - dstAddr tcpip.Address - netProto tcpip.NetworkProtocolNumber - nicID tcpip.NICID + pkt *stack.PacketBuffer - data buffer.VectorisedView `state:".(buffer.VectorisedView)"` - - hdr header.TCP - // views is used as buffer for data when its length is large - // enough to store a VectorisedView. - views [8]buffer.View `state:"nosave"` sequenceNumber seqnum.Value ackNumber seqnum.Value flags header.TCPFlags @@ -80,28 +74,56 @@ type segment struct { // acked indicates if the segment has already been SACKed. acked bool - // dataMemSize is the memory used by data initially. + // dataMemSize is the memory used by pkt initially. The value is used for + // memory accounting in the receive buffer instead of pkt.MemSize() because + // packet contents can be modified, so relying on the computed memory size + // to "free" reserved bytes could leak memory in the receiver. dataMemSize int // lost indicates if the segment is marked as lost by RACK. lost bool } -func newIncomingSegment(id stack.TransportEndpointID, clock tcpip.Clock, pkt *stack.PacketBuffer) *segment { - netHdr := pkt.Network() - s := &segment{ - id: id, - srcAddr: netHdr.SourceAddress(), - dstAddr: netHdr.DestinationAddress(), - netProto: pkt.NetworkProtocolNumber, - nicID: pkt.NICID, +func newIncomingSegment(id stack.TransportEndpointID, clock tcpip.Clock, pkt *stack.PacketBuffer) (*segment, error) { + // We check that the offset to the data respects the following constraints: + // 1. That it's at least the minimum header size; if we don't do this + // then part of the header would be delivered to user. + // 2. That the header fits within the buffer; if we don't do this, we + // would panic when we tried to access data beyond the buffer. + if pkt.TransportHeader().View().Size() < header.TCPMinimumSize { + return nil, fmt.Errorf("packet header smaller than minimum TCP header size: minimum size = %d, got size=%d", header.TCPMinimumSize, pkt.TransportHeader().View().Size()) } + hdr := header.TCP(pkt.TransportHeader().View()) + offset := int(hdr.DataOffset()) + if offset < header.TCPMinimumSize || offset > len(hdr) { + return nil, fmt.Errorf("header data offset does not respect size constraints: %d < offset < %d, got offset=%d", header.TCPMinimumSize, len(hdr), offset) + } + + s := &segment{ + id: id, + options: hdr[header.TCPMinimumSize:], + parsedOptions: header.ParseTCPOptions(hdr[header.TCPMinimumSize:]), + sequenceNumber: seqnum.Value(hdr.SequenceNumber()), + ackNumber: seqnum.Value(hdr.AckNumber()), + flags: hdr.Flags(), + window: seqnum.Size(hdr.WindowSize()), + rcvdTime: clock.NowMonotonic(), + dataMemSize: pkt.MemSize(), + pkt: pkt, + } + pkt.IncRef() s.InitRefs() - s.data = pkt.Data().ExtractVV().Clone(s.views[:]) - s.hdr = header.TCP(pkt.TransportHeader().View()) - s.rcvdTime = clock.NowMonotonic() - s.dataMemSize = s.data.Size() - return s + + if s.pkt.RXTransportChecksumValidated { + s.csumValid = true + } else { + s.csum = hdr.Checksum() + payloadChecksum := s.pkt.Data().AsRange().Checksum() + payloadLength := uint16(s.payloadSize()) + net := s.pkt.Network() + s.csumValid = hdr.IsChecksumValid(net.SourceAddress(), net.DestinationAddress(), payloadChecksum, payloadLength) + } + return s, nil } func newOutgoingSegment(id stack.TransportEndpointID, clock tcpip.Clock, v buffer.View) *segment { @@ -110,14 +132,13 @@ func newOutgoingSegment(id stack.TransportEndpointID, clock tcpip.Clock, v buffe } s.InitRefs() s.rcvdTime = clock.NowMonotonic() - if len(v) != 0 { - s.views[0] = v - s.data = buffer.NewVectorisedView(len(v), s.views[:1]) - } - s.dataMemSize = s.data.Size() + s.pkt = stack.NewPacketBuffer(stack.PacketBufferOptions{}) + s.pkt.Data().AppendView(v) + s.dataMemSize = s.pkt.MemSize() return s } +// clone creates a shallow clone of s not including its pkt. func (s *segment) clone() *segment { t := &segment{ id: s.id, @@ -125,8 +146,6 @@ func (s *segment) clone() *segment { ackNumber: s.ackNumber, flags: s.flags, window: s.window, - netProto: s.netProto, - nicID: s.nicID, rcvdTime: s.rcvdTime, xmitTime: s.xmitTime, xmitCount: s.xmitCount, @@ -135,17 +154,15 @@ func (s *segment) clone() *segment { dataMemSize: s.dataMemSize, } t.InitRefs() - t.data = s.data.Clone(t.views[:]) + t.pkt = stack.NewPacketBuffer(stack.PacketBufferOptions{}) return t } // merge merges data in oth and clears oth. func (s *segment) merge(oth *segment) { - s.data.Append(oth.data) - s.dataMemSize = s.data.Size() - - oth.data = buffer.VectorisedView{} - oth.dataMemSize = oth.data.Size() + s.pkt.Data().Merge(oth.pkt.Data()) + s.dataMemSize = s.pkt.MemSize() + oth.dataMemSize = oth.pkt.MemSize() } // setOwner sets the owning endpoint for this segment. Its required @@ -166,6 +183,7 @@ func (s *segment) setOwner(ep *endpoint, qFlags queueFlags) { func (s *segment) DecRef() { s.segmentRefs.DecRef(func() { + defer s.pkt.DecRef() if s.ep != nil { switch s.qFlags { case recvQ: @@ -182,7 +200,7 @@ func (s *segment) DecRef() { // logicalLen is the segment length in the sequence number space. It's defined // as the data length plus one for each of the SYN and FIN bits set. func (s *segment) logicalLen() seqnum.Size { - l := seqnum.Size(s.data.Size()) + l := seqnum.Size(s.payloadSize()) if s.flags.Contains(header.TCPFlagSyn) { l++ } @@ -194,58 +212,24 @@ func (s *segment) logicalLen() seqnum.Size { // payloadSize is the size of s.data. func (s *segment) payloadSize() int { - return s.data.Size() + return s.pkt.Data().Size() } // segMemSize is the amount of memory used to hold the segment data and // the associated metadata. func (s *segment) segMemSize() int { - return SegSize + s.dataMemSize -} - -// parse populates the sequence & ack numbers, flags, and window fields of the -// segment from the TCP header stored in the data. It then updates the view to -// skip the header. -// -// Returns boolean indicating if the parsing was successful. -// -// If checksum verification may not be skipped, parse also verifies the -// TCP checksum and stores the checksum and result of checksum verification in -// the csum and csumValid fields of the segment. -func (s *segment) parse(skipChecksumValidation bool) bool { - // h is the header followed by the payload. We check that the offset to - // the data respects the following constraints: - // 1. That it's at least the minimum header size; if we don't do this - // then part of the header would be delivered to user. - // 2. That the header fits within the buffer; if we don't do this, we - // would panic when we tried to access data beyond the buffer. - // - // N.B. The segment has already been validated as having at least the - // minimum TCP size before reaching here, so it's safe to read the - // fields. - offset := int(s.hdr.DataOffset()) - if offset < header.TCPMinimumSize || offset > len(s.hdr) { - return false - } - - s.options = s.hdr[header.TCPMinimumSize:] - s.parsedOptions = header.ParseTCPOptions(s.options) - if skipChecksumValidation { - s.csumValid = true - } else { - s.csum = s.hdr.Checksum() - payloadChecksum := header.ChecksumVV(s.data, 0) - payloadLength := uint16(s.data.Size()) - s.csumValid = s.hdr.IsChecksumValid(s.srcAddr, s.dstAddr, payloadChecksum, payloadLength) - } - s.sequenceNumber = seqnum.Value(s.hdr.SequenceNumber()) - s.ackNumber = seqnum.Value(s.hdr.AckNumber()) - s.flags = s.hdr.Flags() - s.window = seqnum.Size(s.hdr.WindowSize()) - return true + return segSize + s.dataMemSize } // sackBlock returns a header.SACKBlock that represents this segment. func (s *segment) sackBlock() header.SACKBlock { return header.SACKBlock{Start: s.sequenceNumber, End: s.sequenceNumber.Add(s.logicalLen())} } + +func (s *segment) TrimFront(ackLeft seqnum.Size) { + s.pkt.Data().TrimFront(int(ackLeft)) +} + +func (s *segment) ReadTo(dst io.Writer, peek bool) (int, error) { + return s.pkt.Data().ReadTo(dst, peek) +} diff --git a/pkg/tcpip/transport/tcp/segment_state.go b/pkg/tcpip/transport/tcp/segment_state.go index dcfa80f95..57bbd69ff 100644 --- a/pkg/tcpip/transport/tcp/segment_state.go +++ b/pkg/tcpip/transport/tcp/segment_state.go @@ -14,30 +14,6 @@ package tcp -import ( - "gvisor.dev/gvisor/pkg/tcpip/buffer" -) - -// saveData is invoked by stateify. -func (s *segment) saveData() buffer.VectorisedView { - // We cannot save s.data directly as s.data.views may alias to s.views, - // which is not allowed by state framework (in-struct pointer). - vs := make([]buffer.View, len(s.data.Views())) - for i, v := range s.data.Views() { - vs[i] = v - } - return buffer.NewVectorisedView(s.data.Size(), vs) -} - -// loadData is invoked by stateify. -func (s *segment) loadData(data buffer.VectorisedView) { - // NOTE: We cannot do the s.data = data.Clone(s.views[:]) optimization - // here because data.views is not guaranteed to be loaded by now. Plus, - // data.views will be allocated anyway so there really is little point - // of utilizing s.views for data.views. - s.data = data -} - // saveOptions is invoked by stateify. func (s *segment) saveOptions() []byte { // We cannot save s.options directly as it may point to s.data's trimmed diff --git a/pkg/tcpip/transport/tcp/segment_test.go b/pkg/tcpip/transport/tcp/segment_test.go index 3af4e8a73..c6e9b9d93 100644 --- a/pkg/tcpip/transport/tcp/segment_test.go +++ b/pkg/tcpip/transport/tcp/segment_test.go @@ -31,7 +31,7 @@ type segmentSizeWants struct { func checkSegmentSize(t *testing.T, name string, seg *segment, want segmentSizeWants) { t.Helper() got := segmentSizeWants{ - DataSize: seg.data.Size(), + DataSize: seg.payloadSize(), SegMemSize: seg.segMemSize(), } if diff := cmp.Diff(want, got); diff != "" { @@ -49,21 +49,21 @@ func TestSegmentMerge(t *testing.T) { checkSegmentSize(t, "seg1", seg1, segmentSizeWants{ DataSize: 10, - SegMemSize: SegSize + 10, + SegMemSize: segSize + stack.PacketBufferStructSize + 10, }) checkSegmentSize(t, "seg2", seg2, segmentSizeWants{ DataSize: 20, - SegMemSize: SegSize + 20, + SegMemSize: segSize + stack.PacketBufferStructSize + 20, }) seg1.merge(seg2) checkSegmentSize(t, "seg1", seg1, segmentSizeWants{ DataSize: 30, - SegMemSize: SegSize + 30, + SegMemSize: segSize + stack.PacketBufferStructSize + 30, }) checkSegmentSize(t, "seg2", seg2, segmentSizeWants{ DataSize: 0, - SegMemSize: SegSize, + SegMemSize: segSize + stack.PacketBufferStructSize, }) } diff --git a/pkg/tcpip/transport/tcp/segment_unsafe.go b/pkg/tcpip/transport/tcp/segment_unsafe.go index 392ff0859..0ab7b8f56 100644 --- a/pkg/tcpip/transport/tcp/segment_unsafe.go +++ b/pkg/tcpip/transport/tcp/segment_unsafe.go @@ -19,6 +19,5 @@ import ( ) const ( - // SegSize is the minimal size of the segment overhead. - SegSize = int(unsafe.Sizeof(segment{})) + segSize = int(unsafe.Sizeof(segment{})) ) diff --git a/pkg/tcpip/transport/tcp/snd.go b/pkg/tcpip/transport/tcp/snd.go index 2c8d51a7c..297401c2b 100644 --- a/pkg/tcpip/transport/tcp/snd.go +++ b/pkg/tcpip/transport/tcp/snd.go @@ -22,7 +22,6 @@ import ( "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" - "gvisor.dev/gvisor/pkg/tcpip/buffer" "gvisor.dev/gvisor/pkg/tcpip/header" "gvisor.dev/gvisor/pkg/tcpip/seqnum" "gvisor.dev/gvisor/pkg/tcpip/stack" @@ -314,7 +313,7 @@ func (s *sender) updateMaxPayloadSize(mtu, count int) { break } - if nextSeg == s.writeNext && seg.data.Size() > m { + if nextSeg == s.writeNext && seg.payloadSize() > m { // We found a segment exceeding the MTU. Rewind // writeNext and try to retransmit it. nextSeg = seg @@ -404,7 +403,7 @@ func (s *sender) resendSegment() { // Resend the segment. if seg := s.writeList.Front(); seg != nil { - if seg.data.Size() > s.MaxPayloadSize { + if seg.payloadSize() > s.MaxPayloadSize { s.splitSeg(seg, s.MaxPayloadSize) } @@ -412,8 +411,8 @@ func (s *sender) resendSegment() { // // To prevent retransmission, set both the HighRXT and RescueRXT // to the highest sequence number in the retransmitted segment. - s.FastRecovery.HighRxt = seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) - 1 - s.FastRecovery.RescueRxt = seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) - 1 + s.FastRecovery.HighRxt = seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) - 1 + s.FastRecovery.RescueRxt = seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) - 1 s.sendSegment(seg) s.ep.stack.Stats().TCP.FastRetransmit.Increment() s.ep.stats.SendErrors.FastRetransmit.Increment() @@ -572,7 +571,7 @@ func (s *sender) retransmitTimerExpired() tcpip.Error { // pCount returns the number of packets in the segment. Due to GSO, a segment // can be composed of multiple packets. func (s *sender) pCount(seg *segment, maxPayloadSize int) int { - size := seg.data.Size() + size := seg.payloadSize() if size == 0 { return 1 } @@ -583,12 +582,12 @@ func (s *sender) pCount(seg *segment, maxPayloadSize int) int { // splitSeg splits a given segment at the size specified and inserts the // remainder as a new segment after the current one in the write list. func (s *sender) splitSeg(seg *segment, size int) { - if seg.data.Size() <= size { + if seg.payloadSize() <= size { return } // Split this segment up. nSeg := seg.clone() - nSeg.data.TrimFront(size) + nSeg.pkt.Data().AppendRange(seg.pkt.Data().AsRange().SubRange(size)) nSeg.sequenceNumber.UpdateForward(seqnum.Size(size)) s.writeList.InsertAfter(seg, nSeg) @@ -601,11 +600,10 @@ func (s *sender) splitSeg(seg *segment, size int) { // window space. // ref: net/ipv4/tcp_output.c::tcp_write_xmit(), tcp_mss_split_point() // ref: net/ipv4/tcp_output.c::tcp_write_wakeup(), tcp_snd_wnd_test() - if seg.data.Size() > s.MaxPayloadSize { + if seg.payloadSize() > s.MaxPayloadSize { seg.flags ^= header.TCPFlagPsh } - - seg.data.CapLength(size) + seg.pkt.Data().CapLength(size) } // NextSeg implements the RFC6675 NextSeg() operation. @@ -632,7 +630,7 @@ func (s *sender) NextSeg(nextSegHint *segment) (nextSeg, hint *segment, rescueRt break } segSeq := seg.sequenceNumber - if smss := s.ep.scoreboard.SMSS(); seg.data.Size() > int(smss) { + if smss := s.ep.scoreboard.SMSS(); seg.payloadSize() > int(smss) { s.splitSeg(seg, int(smss)) } @@ -726,7 +724,7 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se // assigned a sequence number to this segment. if !s.isAssignedSequenceNumber(seg) { // Merge segments if allowed. - if seg.data.Size() != 0 { + if seg.payloadSize() != 0 { available := int(s.SndNxt.Size(end)) if available > limit { available = limit @@ -743,8 +741,8 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se // triggering bugs in poorly written DNS // implementations. var nextTooBig bool - for nSeg := seg.Next(); nSeg != nil && nSeg.data.Size() != 0; nSeg = seg.Next() { - if seg.data.Size()+nSeg.data.Size() > available { + for nSeg := seg.Next(); nSeg != nil && nSeg.payloadSize() != 0; nSeg = seg.Next() { + if seg.payloadSize()+nSeg.payloadSize() > available { nextTooBig = true break } @@ -752,7 +750,7 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se s.writeList.Remove(nSeg) nSeg.DecRef() } - if !nextTooBig && seg.data.Size() < available { + if !nextTooBig && seg.payloadSize() < available { // Segment is not full. if s.Outstanding > 0 && s.ep.ops.GetDelayOption() { // Nagle's algorithm. From Wikipedia: @@ -773,7 +771,7 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se // send space and MSS. // TODO(gvisor.dev/issue/2833): Drain the held segments after a // timeout. - if seg.data.Size() < s.MaxPayloadSize && s.ep.ops.GetCorkOption() { + if seg.payloadSize() < s.MaxPayloadSize && s.ep.ops.GetCorkOption() { return false } } @@ -786,7 +784,7 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se } var segEnd seqnum.Value - if seg.data.Size() == 0 { + if seg.payloadSize() == 0 { if s.writeList.Back() != seg { panic("FIN segments must be the final segment in the write list.") } @@ -833,7 +831,7 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se // the retransmit timer handler. if s.SndUna != s.SndNxt { switch { - case available >= seg.data.Size(): + case available >= seg.payloadSize(): // OK to send, the whole segments fits in the // receiver's advertised window. case available >= s.MaxPayloadSize: @@ -860,11 +858,11 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se available = s.MaxPayloadSize } - if seg.data.Size() > available { + if seg.payloadSize() > available { s.splitSeg(seg, available) } - segEnd = seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) + segEnd = seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) } s.sendSegment(seg) @@ -1055,10 +1053,10 @@ func (s *sender) SetPipe() { } pipe := 0 smss := seqnum.Size(s.ep.scoreboard.SMSS()) - for s1 := s.writeList.Front(); s1 != nil && s1.data.Size() != 0 && s.isAssignedSequenceNumber(s1); s1 = s1.Next() { + for s1 := s.writeList.Front(); s1 != nil && s1.payloadSize() != 0 && s.isAssignedSequenceNumber(s1); s1 = s1.Next() { // With GSO each segment can be much larger than SMSS. So check the segment // in SMSS sized ranges. - segEnd := s1.sequenceNumber.Add(seqnum.Size(s1.data.Size())) + segEnd := s1.sequenceNumber.Add(seqnum.Size(s1.payloadSize())) for startSeq := s1.sequenceNumber; startSeq.LessThan(segEnd); startSeq = startSeq.Add(smss) { endSeq := startSeq.Add(smss) if segEnd.LessThan(endSeq) { @@ -1503,7 +1501,7 @@ func (s *sender) handleRcvdSegment(rcvdSeg *segment) { if datalen > ackLeft { prevCount := s.pCount(seg, s.MaxPayloadSize) - seg.data.TrimFront(int(ackLeft)) + seg.TrimFront(ackLeft) seg.sequenceNumber.UpdateForward(ackLeft) s.Outstanding -= prevCount - s.pCount(seg, s.MaxPayloadSize) break @@ -1636,13 +1634,13 @@ func (s *sender) sendSegment(seg *segment) tcpip.Error { seg.xmitTime = s.ep.stack.Clock().NowMonotonic() seg.xmitCount++ seg.lost = false - err := s.sendSegmentFromView(seg.data, seg.flags, seg.sequenceNumber) + err := s.sendSegmentFromPacketBuffer(seg.pkt, seg.flags, seg.sequenceNumber) // Every time a packet containing data is sent (including a // retransmission), if SACK is enabled and we are retransmitting data // then use the conservative timer described in RFC6675 Section 6.0, // otherwise follow the standard time described in RFC6298 Section 5.1. - if err != nil && seg.data.Size() != 0 { + if err != nil && seg.payloadSize() != 0 { if s.FastRecovery.Active && seg.xmitCount > 1 && s.ep.SACKPermitted { s.resendTimer.enable(s.RTO) } else { @@ -1655,11 +1653,11 @@ func (s *sender) sendSegment(seg *segment) tcpip.Error { return err } -// sendSegmentFromView sends a new segment containing the given payload, flags -// and sequence number. +// sendSegmentFromPacketBuffer sends a new segment containing the given payload, +// flags and sequence number. // +checklocks:s.ep.mu // +checklocksalias:s.ep.rcv.ep.mu=s.ep.mu -func (s *sender) sendSegmentFromView(data buffer.VectorisedView, flags header.TCPFlags, seq seqnum.Value) tcpip.Error { +func (s *sender) sendSegmentFromPacketBuffer(pkt *stack.PacketBuffer, flags header.TCPFlags, seq seqnum.Value) tcpip.Error { s.LastSendTime = s.ep.stack.Clock().NowMonotonic() if seq == s.RTTMeasureSeqNum { s.RTTMeasureTime = s.LastSendTime @@ -1670,14 +1668,16 @@ func (s *sender) sendSegmentFromView(data buffer.VectorisedView, flags header.TC // Remember the max sent ack. s.MaxSentAck = rcvNxt - return s.ep.sendRaw(data, flags, seq, rcvNxt, rcvWnd) + return s.ep.sendRaw(pkt, flags, seq, rcvNxt, rcvWnd) } // sendEmptySegment sends a new segment containing the given flags and sequence // number. // +checklocks:s.ep.mu func (s *sender) sendEmptySegment(flags header.TCPFlags, seq seqnum.Value) tcpip.Error { - return s.sendSegmentFromView(buffer.VectorisedView{}, flags, seq) + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{}) + defer pkt.DecRef() + return s.sendSegmentFromPacketBuffer(pkt, flags, seq) } // maybeSendOutOfWindowAck sends an ACK if we are not being rate limited diff --git a/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go b/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go index 49414f15c..ac0f321a6 100644 --- a/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go +++ b/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go @@ -2240,7 +2240,7 @@ func TestSmallReceiveBufferReadiness(t *testing.T) { } for i := 8; i > 0; i /= 2 { - size := int64(i << 10) + size := int64(i << 12) t.Run(fmt.Sprintf("size=%d", size), func(t *testing.T) { var clientWQ waiter.Queue client, err := s.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &clientWQ) @@ -2410,8 +2410,8 @@ func TestSmallSegReceiveWindowAdvertisement(t *testing.T) { // of the window scaled value. This enables the test to perform equality // checks on the incoming receive window. payloadSize := 1 << c.RcvdWindowScale - if payloadSize >= tcp.SegSize { - t.Fatalf("payload size of %d is not less than the segment overhead of %d", payloadSize, tcp.SegSize) + if payloadSize >= tcp.SegOverheadSize { + t.Fatalf("payload size of %d is not less than the segment overhead of %d", payloadSize, tcp.SegOverheadSize) } payload := generateRandomPayload(t, payloadSize) payloadLen := seqnum.Size(len(payload)) @@ -6648,7 +6648,7 @@ func TestReceiveBufferAutoTuningApplicationLimited(t *testing.T) { time.Sleep(latency) // Send an initial payload with atleast segment overhead size. The receive // window would not grow for smaller segments. - rawEP.SendPacketWithTS(make([]byte, tcp.SegSize), tsVal) + rawEP.SendPacketWithTS(make([]byte, tcp.SegOverheadSize), tsVal) pkt := rawEP.VerifyAndReturnACKWithTS(tsVal) rcvWnd := header.TCP(header.IPv4(pkt).Payload()).WindowSize() diff --git a/test/packetimpact/tests/BUILD b/test/packetimpact/tests/BUILD index 2915862ea..9fe53a582 100644 --- a/test/packetimpact/tests/BUILD +++ b/test/packetimpact/tests/BUILD @@ -340,6 +340,7 @@ packetimpact_testbench( srcs = ["tcp_zero_receive_window_test.go"], deps = [ "//pkg/tcpip/header", + "//pkg/tcpip/transport/tcp", "//test/packetimpact/testbench", "@org_golang_x_sys//unix:go_default_library", ], diff --git a/test/packetimpact/tests/tcp_zero_receive_window_test.go b/test/packetimpact/tests/tcp_zero_receive_window_test.go index bd33a2a03..94b6fc707 100644 --- a/test/packetimpact/tests/tcp_zero_receive_window_test.go +++ b/test/packetimpact/tests/tcp_zero_receive_window_test.go @@ -17,11 +17,13 @@ package tcp_zero_receive_window_test import ( "flag" "fmt" + "math" "testing" "time" "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/tcpip/header" + "gvisor.dev/gvisor/pkg/tcpip/transport/tcp" "gvisor.dev/gvisor/test/packetimpact/testbench" ) @@ -31,7 +33,15 @@ func init() { // TestZeroReceiveWindow tests if the DUT sends a zero receive window eventually. func TestZeroReceiveWindow(t *testing.T) { - for _, payloadLen := range []int{64, 512, 1024} { + // minPayloadLen is the smallest size we can use for a payload in this test. + // Any smaller than this and the receive buffer will fill up before the + // receive window can shrink to zero. + + // To solve for minPayloadLen: minPayloadLen(DefaultReceiveBufferSize) = + // maxWndSize(minPayloadLen + segOverheadSize) + maxWndSize := math.MaxUint16 + minPayloadLen := int(math.Ceil(float64(maxWndSize*tcp.SegOverheadSize) / float64(tcp.DefaultReceiveBufferSize-maxWndSize))) + for _, payloadLen := range []int{minPayloadLen, 512, 1024} { t.Run(fmt.Sprintf("TestZeroReceiveWindow_with_%dbytes_payload", payloadLen), func(t *testing.T) { dut := testbench.NewDUT(t) listenFd, remotePort := dut.CreateListener(t, unix.SOCK_STREAM, unix.IPPROTO_TCP, 1)