diff --git a/pkg/tcpip/stack/BUILD b/pkg/tcpip/stack/BUILD index 9751f3147..1ecfd46a7 100644 --- a/pkg/tcpip/stack/BUILD +++ b/pkg/tcpip/stack/BUILD @@ -71,7 +71,6 @@ 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 f0bf383ee..3fe7c7031 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 `state:".([]byte)"` + buf buffer.Buffer 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,17 +348,6 @@ 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. @@ -583,50 +572,6 @@ 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 deleted file mode 100644 index ad7b45cf0..000000000 --- a/pkg/tcpip/stack/packet_buffer_state.go +++ /dev/null @@ -1,28 +0,0 @@ -// 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 48cd11ccb..c376ed1a1 100644 --- a/pkg/tcpip/stack/packet_buffer_test.go +++ b/pkg/tcpip/stack/packet_buffer_test.go @@ -336,30 +336,6 @@ 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 @@ -440,19 +416,16 @@ func TestPacketBufferData(t *testing.T) { } { t.Run(tc.name, func(t *testing.T) { // PullUp - 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) - } - }) - } - }) - + 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("PullUpOutOfBounds", func(t *testing.T) { n := len(tc.data) + 1 pkt := tc.makePkt(t) @@ -463,38 +436,34 @@ func TestPacketBufferData(t *testing.T) { }) // Consume. - 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) - } + 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) + } - checkData(t, pkt, []byte(tc.data)[n:]) - }) - } - }) + checkData(t, pkt, []byte(tc.data)[n:]) + }) + } // CapLength - 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) + 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) - 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) { @@ -513,23 +482,21 @@ func TestPacketBufferData(t *testing.T) { }) // ReadFromVV - 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 + 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 - 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) { @@ -542,59 +509,6 @@ 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 cd151ff0b..ee3d47270 100644 --- a/pkg/tcpip/stack/packet_buffer_unsafe.go +++ b/pkg/tcpip/stack/packet_buffer_unsafe.go @@ -16,5 +16,4 @@ package stack import "unsafe" -// PacketBufferStructSize is the minimal size of the packet buffer overhead. -const PacketBufferStructSize = int(unsafe.Sizeof(PacketBuffer{})) +const packetBufferStructSize = int(unsafe.Sizeof(PacketBuffer{})) diff --git a/pkg/tcpip/transport/tcp/accept.go b/pkg/tcpip/transport/tcp/accept.go index e7cf947f2..27994925a 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.pkt.NetworkProtocolNumber + netProto = s.netProto } - route, err := l.stack.FindRoute(s.pkt.NICID, s.pkt.Network().DestinationAddress(), s.pkt.Network().SourceAddress(), s.pkt.NetworkProtocolNumber, false /* multicastLoop */) + route, err := l.stack.FindRoute(s.nicID, s.dstAddr, s.srcAddr, s.netProto, 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.pkt.NICID + n.boundNICID = s.nicID n.route = route - n.effectiveNetProtos = []tcpip.NetworkProtocolNumber{s.pkt.NetworkProtocolNumber} + n.effectiveNetProtos = []tcpip.NetworkProtocolNumber{s.netProto} n.ops.SetReceiveBufferSize(int64(l.rcvWnd), false /* notify */) n.amss = calculateAdvertisedMSS(n.userMSS, n.route) n.setEndpointState(StateConnecting) @@ -495,8 +495,7 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err return nil } - net := s.pkt.Network() - route, err := e.stack.FindRoute(s.pkt.NICID, net.DestinationAddress(), net.SourceAddress(), s.pkt.NetworkProtocolNumber, false /* multicastLoop */) + route, err := e.stack.FindRoute(s.nicID, s.dstAddr, s.srcAddr, s.netProto, false /* multicastLoop */) if err != nil { return err } @@ -517,7 +516,7 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err MSS: calculateAdvertisedMSS(e.userMSS, route), } if opts.TS { - offset := e.protocol.tsOffset(net.DestinationAddress(), net.SourceAddress()) + offset := e.protocol.tsOffset(s.dstAddr, s.srcAddr) now := e.stack.Clock().NowMonotonic() synOpts.TSVal = offset.TSVal(now) } @@ -649,8 +648,7 @@ func (e *endpoint) handleListenSegment(ctx *listenContext, s *segment) tcpip.Err } n.isRegistered = true - net := s.pkt.Network() - n.TSOffset = n.protocol.tsOffset(net.DestinationAddress(), net.SourceAddress()) + n.TSOffset = n.protocol.tsOffset(s.dstAddr, s.srcAddr) // Switch state to connected. n.isConnectNotified = true @@ -671,7 +669,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.payloadSize() > 0) && n.enqueueSegment(s) { + if (s.flags.Contains(header.TCPFlagFin) || s.data.Size() > 0) && n.enqueueSegment(s) { n.notifyProcessor() } diff --git a/pkg/tcpip/transport/tcp/connect.go b/pkg/tcpip/transport/tcp/connect.go index a86793175..8f5b000e6 100644 --- a/pkg/tcpip/transport/tcp/connect.go +++ b/pkg/tcpip/transport/tcp/connect.go @@ -22,6 +22,7 @@ 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" @@ -452,7 +453,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.payloadSize() == 0 && h.ep.stack.Clock().NowMonotonic().Sub(h.startTime) < h.deferAccept { + if h.deferAccept != 0 && s.data.Size() == 0 && h.ep.stack.Clock().NowMonotonic().Sub(h.startTime) < h.deferAccept { h.acked = true h.ep.stack.Stats().DroppedPackets.Increment() return nil @@ -485,7 +486,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.payloadSize() > 0) && h.ep.enqueueSegment(s) { + if (s.flags.Contains(header.TCPFlagFin) || s.data.Size() > 0) && h.ep.enqueueSegment(s) { h.ep.protocol.dispatcher.selectProcessor(h.ep.ID).queueEndpoint(h.ep) } @@ -795,18 +796,16 @@ 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. - p := stack.NewPacketBuffer(stack.PacketBufferOptions{}) - defer p.DecRef() - if err := e.sendTCP(r, tf, p, stack.GSO{}); err != nil { + if err := e.sendTCP(r, tf, buffer.VectorisedView{}, stack.GSO{}); err != nil { e.stats.SendErrors.SynSendToNetworkFailed.Increment() } putOptions(tf.opts) return nil } -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, data buffer.VectorisedView, gso stack.GSO) tcpip.Error { tf.txHash = e.txHash - if err := sendTCP(r, tf, pkt, gso, e.owner); err != nil { + if err := sendTCP(r, tf, data, gso, e.owner); err != nil { e.stats.SendErrors.SegmentSendToNetworkFailed.Increment() return err } @@ -843,16 +842,22 @@ 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, 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) + optLen := len(tf.opts) if tf.rcvWnd > math.MaxUint16 { tf.rcvWnd = math.MaxUint16 } mss := int(gso.MSS) - n := (pkt.Data().Size() + mss - 1) / mss + n := (data.Size() + mss - 1) / mss - size := pkt.Data().Size() + size := data.Size() hdrSize := header.TCPMinimumSize + int(r.MaxHeaderLength()) + optLen for i := 0; i < n; i++ { packetSize := mss @@ -860,57 +865,43 @@ func sendTCPBatch(r *stack.Route, tf tcpFields, pkt *stack.PacketBuffer, gso sta packetSize = size } size -= packetSize - - 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 := stack.NewPacketBuffer(stack.PacketBufferOptions{ + ReserveHeaderBytes: hdrSize, + }) 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() - if shouldSplitPacket { - pkt.DecRef() - } + pkt.DecRef() return err } r.Stats().TCP.SegmentsSent.Increment() - if shouldSplitPacket { - pkt.DecRef() - } + 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, pkt *stack.PacketBuffer, gso stack.GSO, owner tcpip.PacketOwner) tcpip.Error { +func sendTCP(r *stack.Route, tf tcpFields, data buffer.VectorisedView, 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) < pkt.Data().Size() { - return sendTCPBatch(r, tf, pkt, gso, owner) + if r.Loop()&stack.PacketLoop == 0 && gso.Type == stack.GSOSW && int(gso.MSS) < data.Size() { + return sendTCPBatch(r, tf, data, 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 @@ -968,13 +959,11 @@ 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 { - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{}) - defer pkt.DecRef() - return e.sendRaw(pkt, flags, seq, ack, rcvWnd) + return e.sendRaw(buffer.VectorisedView{}, flags, seq, ack, rcvWnd) } // sendRaw sends a TCP segment to the endpoint's peer. -func (e *endpoint) sendRaw(pkt *stack.PacketBuffer, flags header.TCPFlags, seq, ack seqnum.Value, rcvWnd seqnum.Size) tcpip.Error { +func (e *endpoint) sendRaw(data buffer.VectorisedView, 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] @@ -989,7 +978,7 @@ func (e *endpoint) sendRaw(pkt *stack.PacketBuffer, flags header.TCPFlags, seq, ack: ack, rcvWnd: rcvWnd, opts: options, - }, pkt, e.gso) + }, data, e.gso) putOptions(options) return err } @@ -1070,14 +1059,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.pkt.NICID) + ep := e.stack.FindTransportEndpoint(e.NetProto, e.TransProto, e.TransportEndpointInfo.ID, s.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.pkt.NICID, + s.nicID, ) } if ep == nil { @@ -1409,7 +1398,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.pkt.NICID); listenEP != nil { + if listenEP := e.stack.FindTransportEndpoint(netProto, info.TransProto, newID, s.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 39665cbc4..c066b29a6 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, err := newIncomingSegment(id, clock, pkt) - if err != nil { + s := newIncomingSegment(id, clock, pkt) + defer s.DecRef() + if !s.parse(pkt.RXTransportChecksumValidated) { 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 42b6c8be6..c3d73be89 100644 --- a/pkg/tcpip/transport/tcp/endpoint.go +++ b/pkg/tcpip/transport/tcp/endpoint.go @@ -1375,7 +1375,7 @@ func (e *endpoint) Read(dst io.Writer, opts tcpip.ReadOptions) (tcpip.ReadResult s := first for s != nil { var n int - n, err = s.ReadTo(dst, opts.Peek) + n, err = s.data.ReadTo(dst, opts.Peek) // Book keeping first then error handling. done += n @@ -1475,7 +1475,7 @@ func (e *endpoint) commitRead(done int) *segment { e.rcvQueueInfo.rcvQueueMu.Lock() memDelta := 0 s := e.rcvQueueInfo.rcvQueue.Front() - for s != nil && s.payloadSize() == 0 { + for s != nil && s.data.Size() == 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 b3888ff20..fcc1f9dc7 100644 --- a/pkg/tcpip/transport/tcp/forwarder.go +++ b/pkg/tcpip/transport/tcp/forwarder.go @@ -65,14 +65,11 @@ 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, err := newIncomingSegment(id, f.stack.Clock(), pkt) - if err != nil { - return false - } + s := newIncomingSegment(id, f.stack.Clock(), pkt) defer s.DecRef() // We only care about well-formed SYN packets (not SYN-ACK) packets. - if !s.csumValid || !s.flags.Contains(header.TCPFlagSyn) || s.flags.Contains(header.TCPFlagAck) { + if !s.parse(pkt.RXTransportChecksumValidated) || !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 26958568d..34ab1d9b2 100644 --- a/pkg/tcpip/transport/tcp/protocol.go +++ b/pkg/tcpip/transport/tcp/protocol.go @@ -157,12 +157,10 @@ 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, err := newIncomingSegment(id, p.stack.Clock(), pkt) - if err != nil { - return stack.UnknownDestinationPacketMalformed - } + s := newIncomingSegment(id, p.stack.Clock(), pkt) defer s.DecRef() - if !s.csumValid { + + if !s.parse(pkt.RXTransportChecksumValidated) || !s.csumValid { return stack.UnknownDestinationPacketMalformed } @@ -196,8 +194,7 @@ 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 { - net := s.pkt.Network() - route, err := st.FindRoute(s.pkt.NICID, net.DestinationAddress(), net.SourceAddress(), s.pkt.NetworkProtocolNumber, false /* multicastLoop */) + route, err := st.FindRoute(s.nicID, s.dstAddr, s.srcAddr, s.netProto, false /* multicastLoop */) if err != nil { return err } @@ -227,8 +224,6 @@ 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, @@ -237,7 +232,7 @@ func replyWithReset(st *stack.Stack, s *segment, tos, ipv4TTL uint8, ipv6HopLimi seq: seq, ack: ack, rcvWnd: 0, - }, p, stack.GSO{}, nil /* PacketOwner */) + }, buffer.VectorisedView{}, 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 6e5852208..b8d0bb653 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.payloadSize())) + endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) 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.payloadSize())) + endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) 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.payloadSize())) + endSeq := seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) 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 bd5ece56f..88251d576 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 >= SegOverheadSize || bufUsed <= r.prevBufUsed + toGrow := unackLen >= SegSize || 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.TrimFront(diff) + s.data.TrimFront(int(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.payloadSize())) + endDataSeq := s.sequenceNumber.Add(seqnum.Size(s.data.Size())) 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.payloadSize())) + segEnd := s.sequenceNumber.Add(seqnum.Size(s.data.Size())) 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.payloadSize()) + segLen := seqnum.Size(s.data.Size()) 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.payloadSize()) + segLen := seqnum.Size(s.data.Size()) 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.payloadSize()) + segLen := seqnum.Size(s.data.Size()) // 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 d2c567130..43968704b 100644 --- a/pkg/tcpip/transport/tcp/segment.go +++ b/pkg/tcpip/transport/tcp/segment.go @@ -16,7 +16,6 @@ package tcp import ( "fmt" - "io" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/buffer" @@ -30,11 +29,7 @@ import ( type queueFlags uint8 const ( - // 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 + recvQ queueFlags = 1 << iota sendQ ) @@ -51,8 +46,19 @@ type segment struct { qFlags queueFlags id stack.TransportEndpointID `state:"manual"` - pkt *stack.PacketBuffer + // 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 + 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 @@ -74,56 +80,28 @@ type segment struct { // acked indicates if the segment has already been SACKed. acked bool - // 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 is the memory used by data initially. 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, 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) - } - +func newIncomingSegment(id stack.TransportEndpointID, clock tcpip.Clock, pkt *stack.PacketBuffer) *segment { + netHdr := pkt.Network() 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, + id: id, + srcAddr: netHdr.SourceAddress(), + dstAddr: netHdr.DestinationAddress(), + netProto: pkt.NetworkProtocolNumber, + nicID: pkt.NICID, } - pkt.IncRef() s.InitRefs() - - 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 + 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 } func newOutgoingSegment(id stack.TransportEndpointID, clock tcpip.Clock, v buffer.View) *segment { @@ -132,13 +110,14 @@ func newOutgoingSegment(id stack.TransportEndpointID, clock tcpip.Clock, v buffe } s.InitRefs() s.rcvdTime = clock.NowMonotonic() - s.pkt = stack.NewPacketBuffer(stack.PacketBufferOptions{}) - s.pkt.Data().AppendView(v) - s.dataMemSize = s.pkt.MemSize() + if len(v) != 0 { + s.views[0] = v + s.data = buffer.NewVectorisedView(len(v), s.views[:1]) + } + s.dataMemSize = s.data.Size() return s } -// clone creates a shallow clone of s not including its pkt. func (s *segment) clone() *segment { t := &segment{ id: s.id, @@ -146,6 +125,8 @@ 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, @@ -154,15 +135,17 @@ func (s *segment) clone() *segment { dataMemSize: s.dataMemSize, } t.InitRefs() - t.pkt = stack.NewPacketBuffer(stack.PacketBufferOptions{}) + t.data = s.data.Clone(t.views[:]) return t } // merge merges data in oth and clears oth. func (s *segment) merge(oth *segment) { - s.pkt.Data().Merge(oth.pkt.Data()) - s.dataMemSize = s.pkt.MemSize() - oth.dataMemSize = oth.pkt.MemSize() + s.data.Append(oth.data) + s.dataMemSize = s.data.Size() + + oth.data = buffer.VectorisedView{} + oth.dataMemSize = oth.data.Size() } // setOwner sets the owning endpoint for this segment. Its required @@ -183,7 +166,6 @@ 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: @@ -200,7 +182,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.payloadSize()) + l := seqnum.Size(s.data.Size()) if s.flags.Contains(header.TCPFlagSyn) { l++ } @@ -212,24 +194,58 @@ func (s *segment) logicalLen() seqnum.Size { // payloadSize is the size of s.data. func (s *segment) payloadSize() int { - return s.pkt.Data().Size() + return s.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 + 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 } // 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 57bbd69ff..dcfa80f95 100644 --- a/pkg/tcpip/transport/tcp/segment_state.go +++ b/pkg/tcpip/transport/tcp/segment_state.go @@ -14,6 +14,30 @@ 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 c6e9b9d93..3af4e8a73 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.payloadSize(), + DataSize: seg.data.Size(), 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 + stack.PacketBufferStructSize + 10, + SegMemSize: SegSize + 10, }) checkSegmentSize(t, "seg2", seg2, segmentSizeWants{ DataSize: 20, - SegMemSize: segSize + stack.PacketBufferStructSize + 20, + SegMemSize: SegSize + 20, }) seg1.merge(seg2) checkSegmentSize(t, "seg1", seg1, segmentSizeWants{ DataSize: 30, - SegMemSize: segSize + stack.PacketBufferStructSize + 30, + SegMemSize: SegSize + 30, }) checkSegmentSize(t, "seg2", seg2, segmentSizeWants{ DataSize: 0, - SegMemSize: segSize + stack.PacketBufferStructSize, + SegMemSize: SegSize, }) } diff --git a/pkg/tcpip/transport/tcp/segment_unsafe.go b/pkg/tcpip/transport/tcp/segment_unsafe.go index 0ab7b8f56..392ff0859 100644 --- a/pkg/tcpip/transport/tcp/segment_unsafe.go +++ b/pkg/tcpip/transport/tcp/segment_unsafe.go @@ -19,5 +19,6 @@ import ( ) const ( - segSize = int(unsafe.Sizeof(segment{})) + // SegSize is the minimal size of the segment overhead. + SegSize = int(unsafe.Sizeof(segment{})) ) diff --git a/pkg/tcpip/transport/tcp/snd.go b/pkg/tcpip/transport/tcp/snd.go index 297401c2b..2c8d51a7c 100644 --- a/pkg/tcpip/transport/tcp/snd.go +++ b/pkg/tcpip/transport/tcp/snd.go @@ -22,6 +22,7 @@ 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" @@ -313,7 +314,7 @@ func (s *sender) updateMaxPayloadSize(mtu, count int) { break } - if nextSeg == s.writeNext && seg.payloadSize() > m { + if nextSeg == s.writeNext && seg.data.Size() > m { // We found a segment exceeding the MTU. Rewind // writeNext and try to retransmit it. nextSeg = seg @@ -403,7 +404,7 @@ func (s *sender) resendSegment() { // Resend the segment. if seg := s.writeList.Front(); seg != nil { - if seg.payloadSize() > s.MaxPayloadSize { + if seg.data.Size() > s.MaxPayloadSize { s.splitSeg(seg, s.MaxPayloadSize) } @@ -411,8 +412,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.payloadSize())) - 1 - s.FastRecovery.RescueRxt = seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) - 1 + 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.sendSegment(seg) s.ep.stack.Stats().TCP.FastRetransmit.Increment() s.ep.stats.SendErrors.FastRetransmit.Increment() @@ -571,7 +572,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.payloadSize() + size := seg.data.Size() if size == 0 { return 1 } @@ -582,12 +583,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.payloadSize() <= size { + if seg.data.Size() <= size { return } // Split this segment up. nSeg := seg.clone() - nSeg.pkt.Data().AppendRange(seg.pkt.Data().AsRange().SubRange(size)) + nSeg.data.TrimFront(size) nSeg.sequenceNumber.UpdateForward(seqnum.Size(size)) s.writeList.InsertAfter(seg, nSeg) @@ -600,10 +601,11 @@ 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.payloadSize() > s.MaxPayloadSize { + if seg.data.Size() > s.MaxPayloadSize { seg.flags ^= header.TCPFlagPsh } - seg.pkt.Data().CapLength(size) + + seg.data.CapLength(size) } // NextSeg implements the RFC6675 NextSeg() operation. @@ -630,7 +632,7 @@ func (s *sender) NextSeg(nextSegHint *segment) (nextSeg, hint *segment, rescueRt break } segSeq := seg.sequenceNumber - if smss := s.ep.scoreboard.SMSS(); seg.payloadSize() > int(smss) { + if smss := s.ep.scoreboard.SMSS(); seg.data.Size() > int(smss) { s.splitSeg(seg, int(smss)) } @@ -724,7 +726,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.payloadSize() != 0 { + if seg.data.Size() != 0 { available := int(s.SndNxt.Size(end)) if available > limit { available = limit @@ -741,8 +743,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.payloadSize() != 0; nSeg = seg.Next() { - if seg.payloadSize()+nSeg.payloadSize() > available { + for nSeg := seg.Next(); nSeg != nil && nSeg.data.Size() != 0; nSeg = seg.Next() { + if seg.data.Size()+nSeg.data.Size() > available { nextTooBig = true break } @@ -750,7 +752,7 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se s.writeList.Remove(nSeg) nSeg.DecRef() } - if !nextTooBig && seg.payloadSize() < available { + if !nextTooBig && seg.data.Size() < available { // Segment is not full. if s.Outstanding > 0 && s.ep.ops.GetDelayOption() { // Nagle's algorithm. From Wikipedia: @@ -771,7 +773,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.payloadSize() < s.MaxPayloadSize && s.ep.ops.GetCorkOption() { + if seg.data.Size() < s.MaxPayloadSize && s.ep.ops.GetCorkOption() { return false } } @@ -784,7 +786,7 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se } var segEnd seqnum.Value - if seg.payloadSize() == 0 { + if seg.data.Size() == 0 { if s.writeList.Back() != seg { panic("FIN segments must be the final segment in the write list.") } @@ -831,7 +833,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.payloadSize(): + case available >= seg.data.Size(): // OK to send, the whole segments fits in the // receiver's advertised window. case available >= s.MaxPayloadSize: @@ -858,11 +860,11 @@ func (s *sender) maybeSendSegment(seg *segment, limit int, end seqnum.Value) (se available = s.MaxPayloadSize } - if seg.payloadSize() > available { + if seg.data.Size() > available { s.splitSeg(seg, available) } - segEnd = seg.sequenceNumber.Add(seqnum.Size(seg.payloadSize())) + segEnd = seg.sequenceNumber.Add(seqnum.Size(seg.data.Size())) } s.sendSegment(seg) @@ -1053,10 +1055,10 @@ func (s *sender) SetPipe() { } pipe := 0 smss := seqnum.Size(s.ep.scoreboard.SMSS()) - for s1 := s.writeList.Front(); s1 != nil && s1.payloadSize() != 0 && s.isAssignedSequenceNumber(s1); s1 = s1.Next() { + for s1 := s.writeList.Front(); s1 != nil && s1.data.Size() != 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.payloadSize())) + segEnd := s1.sequenceNumber.Add(seqnum.Size(s1.data.Size())) for startSeq := s1.sequenceNumber; startSeq.LessThan(segEnd); startSeq = startSeq.Add(smss) { endSeq := startSeq.Add(smss) if segEnd.LessThan(endSeq) { @@ -1501,7 +1503,7 @@ func (s *sender) handleRcvdSegment(rcvdSeg *segment) { if datalen > ackLeft { prevCount := s.pCount(seg, s.MaxPayloadSize) - seg.TrimFront(ackLeft) + seg.data.TrimFront(int(ackLeft)) seg.sequenceNumber.UpdateForward(ackLeft) s.Outstanding -= prevCount - s.pCount(seg, s.MaxPayloadSize) break @@ -1634,13 +1636,13 @@ func (s *sender) sendSegment(seg *segment) tcpip.Error { seg.xmitTime = s.ep.stack.Clock().NowMonotonic() seg.xmitCount++ seg.lost = false - err := s.sendSegmentFromPacketBuffer(seg.pkt, seg.flags, seg.sequenceNumber) + err := s.sendSegmentFromView(seg.data, 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.payloadSize() != 0 { + if err != nil && seg.data.Size() != 0 { if s.FastRecovery.Active && seg.xmitCount > 1 && s.ep.SACKPermitted { s.resendTimer.enable(s.RTO) } else { @@ -1653,11 +1655,11 @@ func (s *sender) sendSegment(seg *segment) tcpip.Error { return err } -// sendSegmentFromPacketBuffer sends a new segment containing the given payload, -// flags and sequence number. +// sendSegmentFromView 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) sendSegmentFromPacketBuffer(pkt *stack.PacketBuffer, flags header.TCPFlags, seq seqnum.Value) tcpip.Error { +func (s *sender) sendSegmentFromView(data buffer.VectorisedView, flags header.TCPFlags, seq seqnum.Value) tcpip.Error { s.LastSendTime = s.ep.stack.Clock().NowMonotonic() if seq == s.RTTMeasureSeqNum { s.RTTMeasureTime = s.LastSendTime @@ -1668,16 +1670,14 @@ func (s *sender) sendSegmentFromPacketBuffer(pkt *stack.PacketBuffer, flags head // Remember the max sent ack. s.MaxSentAck = rcvNxt - return s.ep.sendRaw(pkt, flags, seq, rcvNxt, rcvWnd) + return s.ep.sendRaw(data, 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 { - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{}) - defer pkt.DecRef() - return s.sendSegmentFromPacketBuffer(pkt, flags, seq) + return s.sendSegmentFromView(buffer.VectorisedView{}, 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 ac0f321a6..49414f15c 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 << 12) + size := int64(i << 10) 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.SegOverheadSize { - t.Fatalf("payload size of %d is not less than the segment overhead of %d", payloadSize, tcp.SegOverheadSize) + if payloadSize >= tcp.SegSize { + t.Fatalf("payload size of %d is not less than the segment overhead of %d", payloadSize, tcp.SegSize) } 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.SegOverheadSize), tsVal) + rawEP.SendPacketWithTS(make([]byte, tcp.SegSize), 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 9fe53a582..2915862ea 100644 --- a/test/packetimpact/tests/BUILD +++ b/test/packetimpact/tests/BUILD @@ -340,7 +340,6 @@ 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 94b6fc707..bd33a2a03 100644 --- a/test/packetimpact/tests/tcp_zero_receive_window_test.go +++ b/test/packetimpact/tests/tcp_zero_receive_window_test.go @@ -17,13 +17,11 @@ 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" ) @@ -33,15 +31,7 @@ func init() { // TestZeroReceiveWindow tests if the DUT sends a zero receive window eventually. func TestZeroReceiveWindow(t *testing.T) { - // 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} { + for _, payloadLen := range []int{64, 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)