From fd89c0892ba7ff2c1e1d0bc91e181488f2f4ca8f Mon Sep 17 00:00:00 2001 From: Lucas Manning Date: Mon, 20 Dec 2021 10:19:55 -0800 Subject: [PATCH] Remove WritePackets() from LinkEndpoint and NetworkEndpoint. WritePackets ownership is hard to get right, especially when packets can belong to multiple PacketBufferLists, and these lists are modified concurrently. For example, in SendTCPBatch when the qdisc link layer was enabled, packets could belong to multiple lists, like the batch list and the original packet list. If gvisor is concurrently processing a packet in a qdisc batch and that same packet is in a list that is being DecRef'd in the original SendTCPBatch() call, then that packet's entry can point to packets in the batch rather than the ones in the original list. Since this is the only place WritePackets() is used, it is reasonable to just use WritePacket instead. PiperOrigin-RevId: 417432389 --- pkg/tcpip/link/loopback/loopback.go | 11 +- pkg/tcpip/link/qdisc/fifo/BUILD | 15 +- pkg/tcpip/link/qdisc/fifo/fifo.go | 25 --- pkg/tcpip/link/qdisc/fifo/qdisc_test.go | 68 +++++++ pkg/tcpip/network/arp/arp.go | 5 - pkg/tcpip/network/arp/arp_test.go | 8 + pkg/tcpip/network/ip_test.go | 5 - pkg/tcpip/network/ipv4/ipv4.go | 86 --------- pkg/tcpip/network/ipv4/ipv4_test.go | 168 +++--------------- pkg/tcpip/network/ipv6/icmp_test.go | 8 +- pkg/tcpip/network/ipv6/ipv6.go | 85 --------- pkg/tcpip/network/ipv6/ipv6_test.go | 164 +++-------------- pkg/tcpip/stack/nic.go | 58 ++---- pkg/tcpip/stack/nic_test.go | 6 - pkg/tcpip/stack/pending_packets.go | 62 ++----- pkg/tcpip/stack/registration.go | 58 +++--- pkg/tcpip/stack/route.go | 10 -- pkg/tcpip/tests/integration/iptables_test.go | 30 ++-- .../tests/integration/link_resolution_test.go | 23 +-- pkg/tcpip/transport/tcp/connect.go | 24 ++- 20 files changed, 243 insertions(+), 676 deletions(-) create mode 100644 pkg/tcpip/link/qdisc/fifo/qdisc_test.go diff --git a/pkg/tcpip/link/loopback/loopback.go b/pkg/tcpip/link/loopback/loopback.go index 49b0a29a9..68dc30a79 100644 --- a/pkg/tcpip/link/loopback/loopback.go +++ b/pkg/tcpip/link/loopback/loopback.go @@ -81,8 +81,15 @@ func (e *endpoint) WritePacket(_ stack.RouteInfo, _ tcpip.NetworkProtocolNumber, } // WritePackets implements stack.LinkEndpoint.WritePackets. -func (e *endpoint) WritePackets(stack.RouteInfo, stack.PacketBufferList, tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - panic("not implemented") +func (e *endpoint) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { + n := 0 + for p := pkts.Front(); p != nil; p = p.Next() { + if err := e.WriteRawPacket(p); err != nil { + return n, err + } + n++ + } + return n, nil } // ARPHardwareType implements stack.LinkEndpoint.ARPHardwareType. diff --git a/pkg/tcpip/link/qdisc/fifo/BUILD b/pkg/tcpip/link/qdisc/fifo/BUILD index 6ec805dc4..fee5ab450 100644 --- a/pkg/tcpip/link/qdisc/fifo/BUILD +++ b/pkg/tcpip/link/qdisc/fifo/BUILD @@ -1,4 +1,4 @@ -load("//tools:defs.bzl", "go_library") +load("//tools:defs.bzl", "go_library", "go_test") package(licenses = ["notice"]) @@ -16,3 +16,16 @@ go_library( "//pkg/tcpip/stack", ], ) + +go_test( + name = "qdisc_test", + size = "small", + srcs = ["qdisc_test.go"], + deps = [ + ":fifo", + "//pkg/sync", + "//pkg/tcpip", + "//pkg/tcpip/buffer", + "//pkg/tcpip/stack", + ], +) diff --git a/pkg/tcpip/link/qdisc/fifo/fifo.go b/pkg/tcpip/link/qdisc/fifo/fifo.go index e5461db4e..9fc65c811 100644 --- a/pkg/tcpip/link/qdisc/fifo/fifo.go +++ b/pkg/tcpip/link/qdisc/fifo/fifo.go @@ -112,31 +112,6 @@ func (d *discipline) WritePacket(_ stack.RouteInfo, _ tcpip.NetworkProtocolNumbe return nil } -// WritePackets implements stack.QueueingDiscipline.WritePackets. -// -// Each packet in the packet buffer list must have the following fields -// populated: -// - pkt.EgressRoute -// - pkt.GSOOptions -// - pkt.NetworkProtocolNumber -func (d *discipline) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - enqueued := 0 - for pkt := pkts.Front(); pkt != nil; { - qd := &d.dispatchers[int(pkt.Hash)%len(d.dispatchers)] - nxt := pkt.Next() - if !qd.queue.enqueue(pkt) { - if enqueued > 0 { - qd.newPacketWaker.Assert() - } - return enqueued, &tcpip.ErrNoBufferSpace{} - } - pkt = nxt - enqueued++ - qd.newPacketWaker.Assert() - } - return enqueued, nil -} - func (d *discipline) Close() { for i := range d.dispatchers { d.dispatchers[i].closeWaker.Assert() diff --git a/pkg/tcpip/link/qdisc/fifo/qdisc_test.go b/pkg/tcpip/link/qdisc/fifo/qdisc_test.go new file mode 100644 index 000000000..1e17d26c7 --- /dev/null +++ b/pkg/tcpip/link/qdisc/fifo/qdisc_test.go @@ -0,0 +1,68 @@ +// Copyright 2021 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 qdisc_test + +import ( + "math/rand" + "testing" + + "gvisor.dev/gvisor/pkg/sync" + "gvisor.dev/gvisor/pkg/tcpip" + "gvisor.dev/gvisor/pkg/tcpip/buffer" + "gvisor.dev/gvisor/pkg/tcpip/link/qdisc/fifo" + "gvisor.dev/gvisor/pkg/tcpip/stack" +) + +var _ stack.LinkWriter = (*discardWriter)(nil) + +// discardWriter implements LinkWriter. +type discardWriter struct { +} + +func (*discardWriter) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { + return pkts.Len(), nil +} + +// In b/209690936, fast simultaneous writes on qdisc will cause panics. This test +// reproduces the behavior shown in that bug. +func TestFastSimultaneousWrites(t *testing.T) { + lower := &discardWriter{} + linkEP := fifo.New(lower, 16, 1000) + + v := make(buffer.View, 1) + + prot := tcpip.NetworkProtocolNumber(0) + r := stack.RouteInfo{} + + // Simulate many simultaneous writes from various goroutines, similar to TCP's sendTCPBatch(). + nWriters := 100 + nWrites := 100 + var wg sync.WaitGroup + defer wg.Done() + for i := 0; i < nWriters; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < nWrites; j++ { + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ + Data: v.ToVectorisedView(), + }) + pkt.Hash = rand.Uint32() + linkEP.WritePacket(r, prot, pkt) + pkt.DecRef() + } + }() + } +} diff --git a/pkg/tcpip/network/arp/arp.go b/pkg/tcpip/network/arp/arp.go index 8e6be7e26..9ba69ba41 100644 --- a/pkg/tcpip/network/arp/arp.go +++ b/pkg/tcpip/network/arp/arp.go @@ -145,11 +145,6 @@ func (*endpoint) NetworkProtocolNumber() tcpip.NetworkProtocolNumber { return ProtocolNumber } -// WritePackets implements stack.NetworkEndpoint.WritePackets. -func (*endpoint) WritePackets(*stack.Route, stack.PacketBufferList, stack.NetworkHeaderParams) (int, tcpip.Error) { - return 0, &tcpip.ErrNotSupported{} -} - func (*endpoint) WriteHeaderIncludedPacket(*stack.Route, *stack.PacketBuffer) tcpip.Error { return &tcpip.ErrNotSupported{} } diff --git a/pkg/tcpip/network/arp/arp_test.go b/pkg/tcpip/network/arp/arp_test.go index 061cc35ae..ed5b5e5e2 100644 --- a/pkg/tcpip/network/arp/arp_test.go +++ b/pkg/tcpip/network/arp/arp_test.go @@ -420,6 +420,14 @@ type testLinkEndpoint struct { writeErr tcpip.Error } +func (t *testLinkEndpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { + if t.writeErr != nil { + return 0, t.writeErr + } + + return t.LinkEndpoint.WritePackets(r, pkts, protocol) +} + func (t *testLinkEndpoint) WritePacket(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { if t.writeErr != nil { return t.writeErr diff --git a/pkg/tcpip/network/ip_test.go b/pkg/tcpip/network/ip_test.go index 7e8cb262d..f60968575 100644 --- a/pkg/tcpip/network/ip_test.go +++ b/pkg/tcpip/network/ip_test.go @@ -212,11 +212,6 @@ func (t *testObject) WritePacket(_ *stack.Route, protocol tcpip.NetworkProtocolN return nil } -// WritePackets implements stack.LinkEndpoint.WritePackets. -func (*testObject) WritePackets(_ *stack.Route, pkt stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - panic("not implemented") -} - // ARPHardwareType implements stack.LinkEndpoint.ARPHardwareType. func (*testObject) ARPHardwareType() header.ARPHardwareType { panic("not implemented") diff --git a/pkg/tcpip/network/ipv4/ipv4.go b/pkg/tcpip/network/ipv4/ipv4.go index 5aa18f23e..97b58ceab 100644 --- a/pkg/tcpip/network/ipv4/ipv4.go +++ b/pkg/tcpip/network/ipv4/ipv4.go @@ -508,92 +508,6 @@ func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer, headerIn return nil } -// WritePackets implements stack.NetworkEndpoint. -func (e *endpoint) WritePackets(r *stack.Route, pkts stack.PacketBufferList, params stack.NetworkHeaderParams) (int, tcpip.Error) { - if r.Loop()&stack.PacketLoop != 0 { - panic("multiple packets in local loop") - } - if r.Loop()&stack.PacketOut == 0 { - return pkts.Len(), nil - } - - stats := e.stats.ip - - for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { - if err := e.addIPHeader(r.LocalAddress(), r.RemoteAddress(), pkt, params, nil /* options */); err != nil { - return 0, err - } - - networkMTU, err := calculateNetworkMTU(e.nic.MTU(), uint32(pkt.NetworkHeader().View().Size())) - if err != nil { - stats.OutgoingPacketErrors.IncrementBy(uint64(pkts.Len())) - return 0, err - } - - if packetMustBeFragmented(pkt, networkMTU) { - // Keep track of the packet that is about to be fragmented so it can be - // removed once the fragmentation is done. - originalPkt := pkt - if _, _, err := e.handleFragments(r, networkMTU, pkt, func(fragPkt *stack.PacketBuffer) tcpip.Error { - fragPkt.IncRef() - // Modify the packet list in place with the new fragments. - pkts.InsertAfter(pkt, fragPkt) - pkt = fragPkt - return nil - }); err != nil { - panic(fmt.Sprintf("e.handleFragments(_, _, %d, _, _) = %s", networkMTU, err)) - } - // Remove the packet that was just fragmented and process the rest. - pkts.Remove(originalPkt) - } - } - - outNicName := e.protocol.stack.FindNICNameFromID(e.nic.ID()) - // iptables filtering. All packets that reach here are locally - // generated. - outputDropped, natPkts := e.protocol.stack.IPTables().CheckOutputPackets(pkts, r, outNicName) - stats.IPTablesOutputDropped.IncrementBy(uint64(len(outputDropped))) - for pkt := range outputDropped { - pkts.Remove(pkt) - } - - // The NAT-ed packets may now be destined for us. - locallyDelivered := 0 - for pkt := range natPkts { - ep := e.protocol.findEndpointWithAddress(header.IPv4(pkt.NetworkHeader().View()).DestinationAddress()) - if ep == nil { - // The NAT-ed packet is still destined for some remote node. - continue - } - - // Do not send the locally destined packet out the NIC. - pkts.Remove(pkt) - - // Deliver the packet locally. - ep.handleLocalPacket(pkt, true /* canSkipRXChecksum */) - locallyDelivered++ - - } - - // We ignore the list of NAT-ed packets here because Postrouting NAT can only - // change the source address, and does not alter the route or outgoing - // interface of the packet. - postroutingDropped, _ := e.protocol.stack.IPTables().CheckPostroutingPackets(pkts, r, e, outNicName) - stats.IPTablesPostroutingDropped.IncrementBy(uint64(len(postroutingDropped))) - for pkt := range postroutingDropped { - pkts.Remove(pkt) - } - - // The rest of the packets can be delivered to the NIC as a batch. - pktsLen := pkts.Len() - written, err := e.nic.WritePackets(r, pkts, ProtocolNumber) - stats.PacketsSent.IncrementBy(uint64(written)) - stats.OutgoingPacketErrors.IncrementBy(uint64(pktsLen - written)) - - // Dropped packets aren't errors, so include them in the return value. - return locallyDelivered + written + len(outputDropped) + len(postroutingDropped), err -} - // WriteHeaderIncludedPacket implements stack.NetworkEndpoint. func (e *endpoint) WriteHeaderIncludedPacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { // The packet already has an IP header, but there are a few required diff --git a/pkg/tcpip/network/ipv4/ipv4_test.go b/pkg/tcpip/network/ipv4/ipv4_test.go index 50ef0d100..c07ef9a53 100644 --- a/pkg/tcpip/network/ipv4/ipv4_test.go +++ b/pkg/tcpip/network/ipv4/ipv4_test.go @@ -1572,89 +1572,6 @@ func TestFragmentationWritePacket(t *testing.T) { } } -func TestFragmentationWritePackets(t *testing.T) { - const ttl = 42 - writePacketsTests := []struct { - description string - insertBefore int - insertAfter int - }{ - { - description: "Single packet", - insertBefore: 0, - insertAfter: 0, - }, - { - description: "With packet before", - insertBefore: 1, - insertAfter: 0, - }, - { - description: "With packet after", - insertBefore: 0, - insertAfter: 1, - }, - { - description: "With packet before and after", - insertBefore: 1, - insertAfter: 1, - }, - } - tinyPacket := iptestutil.MakeRandPkt(header.TCPMinimumSize, extraHeaderReserve+header.IPv4MinimumSize, []int{1}, header.IPv4ProtocolNumber) - - for _, test := range writePacketsTests { - t.Run(test.description, func(t *testing.T) { - for _, ft := range fragmentationTests { - t.Run(ft.description, func(t *testing.T) { - var pkts stack.PacketBufferList - for i := 0; i < test.insertBefore; i++ { - pkts.PushBack(tinyPacket.Clone()) - } - pkt := iptestutil.MakeRandPkt(ft.transportHeaderLength, extraHeaderReserve+header.IPv4MinimumSize, []int{ft.payloadSize}, header.IPv4ProtocolNumber) - pkts.PushBack(pkt.Clone()) - for i := 0; i < test.insertAfter; i++ { - pkts.PushBack(tinyPacket.Clone()) - } - - ep := iptestutil.NewMockLinkEndpoint(ft.mtu, nil, math.MaxInt32) - r := buildRoute(t, ep) - - wantTotalPackets := len(ft.wantFragments) + test.insertBefore + test.insertAfter - n, err := r.WritePackets(pkts, stack.NetworkHeaderParams{ - Protocol: tcp.ProtocolNumber, - TTL: ttl, - TOS: stack.DefaultTOS, - }) - if err != nil { - t.Errorf("got WritePackets(_, _, _) = (_, %s), want = (_, nil)", err) - } - if n != wantTotalPackets { - t.Errorf("got WritePackets(_, _, _) = (%d, _), want = (%d, _)", n, wantTotalPackets) - } - if got := len(ep.WrittenPackets); got != wantTotalPackets { - t.Errorf("got len(ep.WrittenPackets) = %d, want = %d", got, wantTotalPackets) - } - if got := int(r.Stats().IP.PacketsSent.Value()); got != wantTotalPackets { - t.Errorf("got c.Route.Stats().IP.PacketsSent.Value() = %d, want = %d", got, wantTotalPackets) - } - if got := int(r.Stats().IP.OutgoingPacketErrors.Value()); got != 0 { - t.Errorf("got r.Stats().IP.OutgoingPacketErrors.Value() = %d, want = 0", got) - } - - if wantTotalPackets == 0 { - return - } - - fragments := ep.WrittenPackets[test.insertBefore : len(ft.wantFragments)+test.insertBefore] - if err := compareFragments(fragments, pkt, ft.mtu, ft.wantFragments, tcp.ProtocolNumber, false /* withIPHeader */, extraHeaderReserve); err != nil { - t.Error(err) - } - }) - } - }) - } -} - // TestFragmentationErrors checks that errors are returned from WritePacket // correctly. func TestFragmentationErrors(t *testing.T) { @@ -2929,65 +2846,36 @@ func TestWriteStats(t *testing.T) { }, } - // Parameterize the tests to run with both WritePacket and WritePackets. - writers := []struct { - name string - writePackets func(*stack.Route, stack.PacketBufferList) (int, tcpip.Error) - }{ - { - name: "WritePacket", - writePackets: func(rt *stack.Route, pkts stack.PacketBufferList) (int, tcpip.Error) { - nWritten := 0 - for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { - if err := rt.WritePacket(stack.NetworkHeaderParams{}, pkt); err != nil { - return nWritten, err - } - nWritten++ - } - return nWritten, nil - }, - }, { - name: "WritePackets", - writePackets: func(rt *stack.Route, pkts stack.PacketBufferList) (int, tcpip.Error) { - return rt.WritePackets(pkts, stack.NetworkHeaderParams{}) - }, - }, - } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + ep := iptestutil.NewMockLinkEndpoint(header.IPv4MinimumMTU, &tcpip.ErrInvalidEndpointState{}, test.allowPackets) + rt := buildRoute(t, ep) - for _, writer := range writers { - t.Run(writer.name, func(t *testing.T) { - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - ep := iptestutil.NewMockLinkEndpoint(header.IPv4MinimumMTU, &tcpip.ErrInvalidEndpointState{}, test.allowPackets) - rt := buildRoute(t, ep) - - var pkts stack.PacketBufferList - for i := 0; i < nPackets; i++ { - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ - ReserveHeaderBytes: header.UDPMinimumSize + int(rt.MaxHeaderLength()), - Data: buffer.NewView(0).ToVectorisedView(), - }) - pkt.TransportHeader().Push(header.UDPMinimumSize) - pkts.PushBack(pkt) - } - - test.setup(t, rt.Stack()) - - nWritten, _ := writer.writePackets(rt, pkts) - - if got := int(rt.Stats().IP.PacketsSent.Value()); got != test.expectSent { - t.Errorf("got rt.Stats().IP.PacketsSent.Value() = %d, want = %d", got, test.expectSent) - } - if got := int(rt.Stats().IP.IPTablesOutputDropped.Value()); got != test.expectOutputDropped { - t.Errorf("got rt.Stats().IP.IPTablesOutputDropped.Value() = %d, want = %d", got, test.expectOutputDropped) - } - if got := int(rt.Stats().IP.IPTablesPostroutingDropped.Value()); got != test.expectPostroutingDropped { - t.Errorf("got rt.Stats().IP.IPTablesPostroutingDropped.Value() = %d, want = %d", got, test.expectPostroutingDropped) - } - if nWritten != test.expectWritten { - t.Errorf("got nWritten = %d, want = %d", nWritten, test.expectWritten) - } + test.setup(t, rt.Stack()) + nWritten := 0 + for i := 0; i < nPackets; i++ { + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ + ReserveHeaderBytes: header.UDPMinimumSize + int(rt.MaxHeaderLength()), + Data: buffer.NewView(0).ToVectorisedView(), }) + pkt.TransportHeader().Push(header.UDPMinimumSize) + if err := rt.WritePacket(stack.NetworkHeaderParams{}, pkt); err != nil { + break + } + nWritten++ + } + + if got := int(rt.Stats().IP.PacketsSent.Value()); got != test.expectSent { + t.Errorf("got rt.Stats().IP.PacketsSent.Value() = %d, want = %d", got, test.expectSent) + } + if got := int(rt.Stats().IP.IPTablesOutputDropped.Value()); got != test.expectOutputDropped { + t.Errorf("got rt.Stats().IP.IPTablesOutputDropped.Value() = %d, want = %d", got, test.expectOutputDropped) + } + if got := int(rt.Stats().IP.IPTablesPostroutingDropped.Value()); got != test.expectPostroutingDropped { + t.Errorf("got rt.Stats().IP.IPTablesPostroutingDropped.Value() = %d, want = %d", got, test.expectPostroutingDropped) + } + if nWritten != test.expectWritten { + t.Errorf("got nWritten = %d, want = %d", nWritten, test.expectWritten) } }) } diff --git a/pkg/tcpip/network/ipv6/icmp_test.go b/pkg/tcpip/network/ipv6/icmp_test.go index 03d9f425c..568fba108 100644 --- a/pkg/tcpip/network/ipv6/icmp_test.go +++ b/pkg/tcpip/network/ipv6/icmp_test.go @@ -81,6 +81,10 @@ func (*stubLinkEndpoint) WritePacket(stack.RouteInfo, tcpip.NetworkProtocolNumbe return nil } +func (*stubLinkEndpoint) WritePackets(stack.RouteInfo, stack.PacketBufferList, tcpip.NetworkProtocolNumber) (int, tcpip.Error) { + return 0, nil +} + func (*stubLinkEndpoint) Attach(stack.NetworkDispatcher) {} type stubDispatcher struct { @@ -134,10 +138,6 @@ func (t *testInterface) WritePacket(r *stack.Route, protocol tcpip.NetworkProtoc return t.LinkEndpoint.WritePacket(r.Fields(), protocol, pkt) } -func (t *testInterface) WritePackets(r *stack.Route, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - return t.LinkEndpoint.WritePackets(r.Fields(), pkts, protocol) -} - func (t *testInterface) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { var r stack.RouteInfo r.NetProto = protocol diff --git a/pkg/tcpip/network/ipv6/ipv6.go b/pkg/tcpip/network/ipv6/ipv6.go index a3ea339e5..94fa59f82 100644 --- a/pkg/tcpip/network/ipv6/ipv6.go +++ b/pkg/tcpip/network/ipv6/ipv6.go @@ -830,91 +830,6 @@ func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer, protocol return nil } -// WritePackets implements stack.NetworkEndpoint. -func (e *endpoint) WritePackets(r *stack.Route, pkts stack.PacketBufferList, params stack.NetworkHeaderParams) (int, tcpip.Error) { - if r.Loop()&stack.PacketLoop != 0 { - panic("not implemented") - } - if r.Loop()&stack.PacketOut == 0 { - return pkts.Len(), nil - } - - stats := e.stats.ip - linkMTU := e.nic.MTU() - for pb := pkts.Front(); pb != nil; pb = pb.Next() { - if err := addIPHeader(r.LocalAddress(), r.RemoteAddress(), pb, params, nil /* extensionHeaders */); err != nil { - return 0, err - } - - networkMTU, err := calculateNetworkMTU(linkMTU, uint32(pb.NetworkHeader().View().Size())) - if err != nil { - stats.OutgoingPacketErrors.IncrementBy(uint64(pkts.Len())) - return 0, err - } - if packetMustBeFragmented(pb, networkMTU) { - // Keep track of the packet that is about to be fragmented so it can be - // removed once the fragmentation is done. - originalPkt := pb - if _, _, err := e.handleFragments(r, networkMTU, pb, params.Protocol, func(fragPkt *stack.PacketBuffer) tcpip.Error { - fragPkt.IncRef() - // Modify the packet list in place with the new fragments. - pkts.InsertAfter(pb, fragPkt) - pb = fragPkt - return nil - }); err != nil { - stats.OutgoingPacketErrors.IncrementBy(uint64(pkts.Len())) - return 0, err - } - // Remove the packet that was just fragmented and process the rest. - pkts.Remove(originalPkt) - } - } - - // iptables filtering. All packets that reach here are locally - // generated. - outNicName := e.protocol.stack.FindNICNameFromID(e.nic.ID()) - outputDropped, natPkts := e.protocol.stack.IPTables().CheckOutputPackets(pkts, r, outNicName) - stats.IPTablesOutputDropped.IncrementBy(uint64(len(outputDropped))) - for pkt := range outputDropped { - pkts.Remove(pkt) - } - - // The NAT-ed packets may now be destined for us. - locallyDelivered := 0 - for pkt := range natPkts { - ep := e.protocol.findEndpointWithAddress(header.IPv6(pkt.NetworkHeader().View()).DestinationAddress()) - if ep == nil { - // The NAT-ed packet is still destined for some remote node. - continue - } - - // Do not send the locally destined packet out the NIC. - pkts.Remove(pkt) - - // Deliver the packet locally. - ep.handleLocalPacket(pkt, true /* canSkipRXChecksum */) - locallyDelivered++ - } - - // We ignore the list of NAT-ed packets here because Postrouting NAT can only - // change the source address, and does not alter the route or outgoing - // interface of the packet. - postroutingDropped, _ := e.protocol.stack.IPTables().CheckPostroutingPackets(pkts, r, e, outNicName) - stats.IPTablesPostroutingDropped.IncrementBy(uint64(len(postroutingDropped))) - for pkt := range postroutingDropped { - pkts.Remove(pkt) - } - - // The rest of the packets can be delivered to the NIC as a batch. - pktsLen := pkts.Len() - written, err := e.nic.WritePackets(r, pkts, ProtocolNumber) - stats.PacketsSent.IncrementBy(uint64(written)) - stats.OutgoingPacketErrors.IncrementBy(uint64(pktsLen - written)) - - // Dropped packets aren't errors, so include them in the return value. - return locallyDelivered + written + len(outputDropped) + len(postroutingDropped), err -} - // WriteHeaderIncludedPacket implements stack.NetworkEndpoint. func (e *endpoint) WriteHeaderIncludedPacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { // The packet already has an IP header, but there are a few required checks. diff --git a/pkg/tcpip/network/ipv6/ipv6_test.go b/pkg/tcpip/network/ipv6/ipv6_test.go index e8f157dab..9a434351e 100644 --- a/pkg/tcpip/network/ipv6/ipv6_test.go +++ b/pkg/tcpip/network/ipv6/ipv6_test.go @@ -2603,63 +2603,36 @@ func TestWriteStats(t *testing.T) { }, } - writers := []struct { - name string - writePackets func(*stack.Route, stack.PacketBufferList) (int, tcpip.Error) - }{ - { - name: "WritePacket", - writePackets: func(rt *stack.Route, pkts stack.PacketBufferList) (int, tcpip.Error) { - nWritten := 0 - for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { - if err := rt.WritePacket(stack.NetworkHeaderParams{}, pkt); err != nil { - return nWritten, err - } - nWritten++ - } - return nWritten, nil - }, - }, { - name: "WritePackets", - writePackets: func(rt *stack.Route, pkts stack.PacketBufferList) (int, tcpip.Error) { - return rt.WritePackets(pkts, stack.NetworkHeaderParams{}) - }, - }, - } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + ep := iptestutil.NewMockLinkEndpoint(header.IPv6MinimumMTU, &tcpip.ErrInvalidEndpointState{}, test.allowPackets) + rt := buildRoute(t, ep) + test.setup(t, rt.Stack()) - for _, writer := range writers { - t.Run(writer.name, func(t *testing.T) { - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - ep := iptestutil.NewMockLinkEndpoint(header.IPv6MinimumMTU, &tcpip.ErrInvalidEndpointState{}, test.allowPackets) - rt := buildRoute(t, ep) - var pkts stack.PacketBufferList - for i := 0; i < nPackets; i++ { - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ - ReserveHeaderBytes: header.UDPMinimumSize + int(rt.MaxHeaderLength()), - Data: buffer.NewView(0).ToVectorisedView(), - }) - pkt.TransportHeader().Push(header.UDPMinimumSize) - pkts.PushBack(pkt) - } - - test.setup(t, rt.Stack()) - - nWritten, _ := writer.writePackets(rt, pkts) - - if got := int(rt.Stats().IP.PacketsSent.Value()); got != test.expectSent { - t.Errorf("got rt.Stats().IP.PacketsSent.Value() = %d, want = %d", got, test.expectSent) - } - if got := int(rt.Stats().IP.IPTablesOutputDropped.Value()); got != test.expectOutputDropped { - t.Errorf("got rt.Stats().IP.IPTablesOutputDropped.Value() = %d, want = %d", got, test.expectOutputDropped) - } - if got := int(rt.Stats().IP.IPTablesPostroutingDropped.Value()); got != test.expectPostroutingDropped { - t.Errorf("got r.Stats().IP.IPTablesPostroutingDropped.Value() = %d, want = %d", got, test.expectPostroutingDropped) - } - if nWritten != test.expectWritten { - t.Errorf("got nWritten = %d, want = %d", nWritten, test.expectWritten) - } + nWritten := 0 + for i := 0; i < nPackets; i++ { + pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ + ReserveHeaderBytes: header.UDPMinimumSize + int(rt.MaxHeaderLength()), + Data: buffer.NewView(0).ToVectorisedView(), }) + pkt.TransportHeader().Push(header.UDPMinimumSize) + if err := rt.WritePacket(stack.NetworkHeaderParams{}, pkt); err != nil { + break + } + nWritten++ + } + + if got := int(rt.Stats().IP.PacketsSent.Value()); got != test.expectSent { + t.Errorf("got rt.Stats().IP.PacketsSent.Value() = %d, want = %d", got, test.expectSent) + } + if got := int(rt.Stats().IP.IPTablesOutputDropped.Value()); got != test.expectOutputDropped { + t.Errorf("got rt.Stats().IP.IPTablesOutputDropped.Value() = %d, want = %d", got, test.expectOutputDropped) + } + if got := int(rt.Stats().IP.IPTablesPostroutingDropped.Value()); got != test.expectPostroutingDropped { + t.Errorf("got r.Stats().IP.IPTablesPostroutingDropped.Value() = %d, want = %d", got, test.expectPostroutingDropped) + } + if nWritten != test.expectWritten { + t.Errorf("got nWritten = %d, want = %d", nWritten, test.expectWritten) } }) } @@ -2858,87 +2831,6 @@ func TestFragmentationWritePacket(t *testing.T) { } } -func TestFragmentationWritePackets(t *testing.T) { - const ttl = 42 - tests := []struct { - description string - insertBefore int - insertAfter int - }{ - { - description: "Single packet", - insertBefore: 0, - insertAfter: 0, - }, - { - description: "With packet before", - insertBefore: 1, - insertAfter: 0, - }, - { - description: "With packet after", - insertBefore: 0, - insertAfter: 1, - }, - { - description: "With packet before and after", - insertBefore: 1, - insertAfter: 1, - }, - } - tinyPacket := iptestutil.MakeRandPkt(header.TCPMinimumSize, extraHeaderReserve+header.IPv6MinimumSize, []int{1}, header.IPv6ProtocolNumber) - - for _, test := range tests { - t.Run(test.description, func(t *testing.T) { - for _, ft := range fragmentationTests { - t.Run(ft.description, func(t *testing.T) { - var pkts stack.PacketBufferList - for i := 0; i < test.insertBefore; i++ { - pkts.PushBack(tinyPacket.Clone()) - } - pkt := iptestutil.MakeRandPkt(ft.transHdrLen, extraHeaderReserve+header.IPv6MinimumSize, []int{ft.payloadSize}, header.IPv6ProtocolNumber) - source := pkt - pkts.PushBack(pkt.Clone()) - for i := 0; i < test.insertAfter; i++ { - pkts.PushBack(tinyPacket.Clone()) - } - - ep := iptestutil.NewMockLinkEndpoint(ft.mtu, nil, math.MaxInt32) - r := buildRoute(t, ep) - - wantTotalPackets := len(ft.wantFragments) + test.insertBefore + test.insertAfter - n, err := r.WritePackets(pkts, stack.NetworkHeaderParams{ - Protocol: tcp.ProtocolNumber, - TTL: ttl, - TOS: stack.DefaultTOS, - }) - if n != wantTotalPackets || err != nil { - t.Errorf("got WritePackets(_, _, _) = (%d, %s), want = (%d, nil)", n, err, wantTotalPackets) - } - if got := len(ep.WrittenPackets); got != wantTotalPackets { - t.Errorf("got len(ep.WrittenPackets) = %d, want = %d", got, wantTotalPackets) - } - if got := int(r.Stats().IP.PacketsSent.Value()); got != wantTotalPackets { - t.Errorf("got c.Route.Stats().IP.PacketsSent.Value() = %d, want = %d", got, wantTotalPackets) - } - if got := r.Stats().IP.OutgoingPacketErrors.Value(); got != 0 { - t.Errorf("got r.Stats().IP.OutgoingPacketErrors.Value() = %d, want = 0", got) - } - - if wantTotalPackets == 0 { - return - } - - fragments := ep.WrittenPackets[test.insertBefore : len(ft.wantFragments)+test.insertBefore] - if err := compareFragments(fragments, source, ft.mtu, ft.wantFragments, tcp.ProtocolNumber); err != nil { - t.Error(err) - } - }) - } - }) - } -} - // TestFragmentationErrors checks that errors are returned from WritePacket // correctly. func TestFragmentationErrors(t *testing.T) { diff --git a/pkg/tcpip/stack/nic.go b/pkg/tcpip/stack/nic.go index f7ced0458..f6b9b6179 100644 --- a/pkg/tcpip/stack/nic.go +++ b/pkg/tcpip/stack/nic.go @@ -145,6 +145,14 @@ type delegatingQueueingDiscipline struct { func (*delegatingQueueingDiscipline) Close() {} +// WritePacket passes the packet through to the underlying LinkWriter's WritePackets. +func (qDisc *delegatingQueueingDiscipline) WritePacket(r RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) tcpip.Error { + var pkts PacketBufferList + pkts.PushBack(pkt) + _, err := qDisc.LinkWriter.WritePackets(r, pkts, protocol) + return err +} + // newNIC returns a new NIC using the default NDP configurations from stack. func newNIC(stack *Stack, id tcpip.NICID, ep LinkEndpoint, opts NICOptions) *nic { // TODO(b/141011931): Validate a LinkEndpoint (ep) is valid. For @@ -333,36 +341,17 @@ func (n *nic) IsLoopback() bool { return n.NetworkLinkEndpoint.Capabilities()&CapabilityLoopback != 0 } -// WritePacket implements LinkWriter. -func (n *nic) WritePacket(r *Route, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) tcpip.Error { - _, err := n.enqueuePacketBuffer(r, protocol, pkt) - return err -} - // WriteRawPacket implements LinkRawWriter. func (n *nic) WriteRawPacket(pkt *PacketBuffer) tcpip.Error { return n.rawLinkEP.WriteRawPacket(pkt) } -func (n *nic) writePacketBuffer(r RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt pendingPacketBuffer) (int, tcpip.Error) { - switch pkt := pkt.(type) { - case *PacketBuffer: - if err := n.writePacket(r, protocol, pkt); err != nil { - return 0, err - } - return 1, nil - case *PacketBufferList: - return n.writePackets(r, protocol, *pkt) - default: - panic(fmt.Sprintf("unrecognized pending packet buffer type = %T", pkt)) - } -} - -func (n *nic) enqueuePacketBuffer(r *Route, protocol tcpip.NetworkProtocolNumber, pkt pendingPacketBuffer) (int, tcpip.Error) { +// WritePacket implements LinkWriter. +func (n *nic) WritePacket(r *Route, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) tcpip.Error { routeInfo, _, err := r.resolvedFields(nil) switch err.(type) { case nil: - return n.writePacketBuffer(routeInfo, protocol, pkt) + return n.writePacket(routeInfo, protocol, pkt) case *tcpip.ErrWouldBlock: // As per relevant RFCs, we should queue packets while we wait for link // resolution to complete. @@ -383,7 +372,7 @@ func (n *nic) enqueuePacketBuffer(r *Route, protocol tcpip.NetworkProtocolNumber // completes, the node transmits any queued packets. return n.linkResQueue.enqueue(r, protocol, pkt) default: - return 0, err + return err } } @@ -412,29 +401,6 @@ func (n *nic) writePacket(r RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt return nil } -// WritePackets implements LinkWriter.. -func (n *nic) WritePackets(r *Route, pkts PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - return n.enqueuePacketBuffer(r, protocol, &pkts) -} - -func (n *nic) writePackets(r RouteInfo, protocol tcpip.NetworkProtocolNumber, pkts PacketBufferList) (int, tcpip.Error) { - for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { - pkt.EgressRoute = r - pkt.NetworkProtocolNumber = protocol - n.deliverOutboundPacket(r.RemoteLinkAddress, pkt) - } - - writtenPackets, err := n.qDisc.WritePackets(r, pkts, protocol) - n.stats.tx.packets.IncrementBy(uint64(writtenPackets)) - writtenBytes := 0 - for i, pb := 0, pkts.Front(); i < writtenPackets && pb != nil; i, pb = i+1, pb.Next() { - writtenBytes += pb.Size() - } - - n.stats.tx.bytes.IncrementBy(uint64(writtenBytes)) - return writtenPackets, err -} - // setSpoofing enables or disables address spoofing. func (n *nic) setSpoofing(enable bool) { n.mu.Lock() diff --git a/pkg/tcpip/stack/nic_test.go b/pkg/tcpip/stack/nic_test.go index 88ca9b076..4ac5ff228 100644 --- a/pkg/tcpip/stack/nic_test.go +++ b/pkg/tcpip/stack/nic_test.go @@ -71,12 +71,6 @@ func (*testIPv6Endpoint) WritePacket(*Route, NetworkHeaderParams, *PacketBuffer) return nil } -// WritePackets implements NetworkEndpoint.WritePackets. -func (*testIPv6Endpoint) WritePackets(*Route, PacketBufferList, NetworkHeaderParams) (int, tcpip.Error) { - // Our tests don't use this so we don't support it. - return 0, &tcpip.ErrNotSupported{} -} - // WriteHeaderIncludedPacket implements // NetworkEndpoint.WriteHeaderIncludedPacket. func (*testIPv6Endpoint) WriteHeaderIncludedPacket(*Route, *PacketBuffer) tcpip.Error { diff --git a/pkg/tcpip/stack/pending_packets.go b/pkg/tcpip/stack/pending_packets.go index 7e18d4bc4..358e24de4 100644 --- a/pkg/tcpip/stack/pending_packets.go +++ b/pkg/tcpip/stack/pending_packets.go @@ -28,26 +28,10 @@ const ( maxPendingPacketsPerResolution = 256 ) -// pendingPacketBuffer is a pending packet buffer. -// -// TODO(gvisor.dev/issue/5331): Drop this when we drop WritePacket and only use -// WritePackets so we can use a PacketBufferList everywhere. -type pendingPacketBuffer interface { - len() int -} - -func (*PacketBuffer) len() int { - return 1 -} - -func (p *PacketBufferList) len() int { - return p.Len() -} - type pendingPacket struct { routeInfo RouteInfo proto tcpip.NetworkProtocolNumber - pkt pendingPacketBuffer + pkt *PacketBuffer } // packetsPendingLinkResolution is a queue of packets pending link resolution. @@ -72,12 +56,11 @@ type packetsPendingLinkResolution struct { } } -func (f *packetsPendingLinkResolution) incrementOutgoingPacketErrors(proto tcpip.NetworkProtocolNumber, pkt pendingPacketBuffer) { - n := uint64(pkt.len()) - f.nic.stack.stats.IP.OutgoingPacketErrors.IncrementBy(n) +func (f *packetsPendingLinkResolution) incrementOutgoingPacketErrors(proto tcpip.NetworkProtocolNumber, pkt *PacketBuffer) { + f.nic.stack.stats.IP.OutgoingPacketErrors.Increment() if ipEndpointStats, ok := f.nic.getNetworkEndpoint(proto).Stats().(IPNetworkEndpointStats); ok { - ipEndpointStats.IPStats().OutgoingPacketErrors.IncrementBy(n) + ipEndpointStats.IPStats().OutgoingPacketErrors.Increment() } } @@ -118,7 +101,7 @@ func (f *packetsPendingLinkResolution) dequeue(ch <-chan struct{}, linkAddr tcpi // If the maximum number of pending resolutions is reached, the packets // associated with the oldest link resolution will be dequeued as if they failed // link resolution. -func (f *packetsPendingLinkResolution) enqueue(r *Route, proto tcpip.NetworkProtocolNumber, pkt pendingPacketBuffer) (int, tcpip.Error) { +func (f *packetsPendingLinkResolution) enqueue(r *Route, proto tcpip.NetworkProtocolNumber, pkt *PacketBuffer) tcpip.Error { f.mu.Lock() // Make sure we attempt resolution while holding f's lock so that we avoid // a race where link resolution completes before we enqueue the packets. @@ -136,12 +119,12 @@ func (f *packetsPendingLinkResolution) enqueue(r *Route, proto tcpip.NetworkProt // The route resolved immediately, so we don't need to wait for link // resolution to send the packet. f.mu.Unlock() - return f.nic.writePacketBuffer(routeInfo, proto, pkt) + return f.nic.writePacket(routeInfo, proto, pkt) case *tcpip.ErrWouldBlock: // We need to wait for link resolution to complete. default: f.mu.Unlock() - return 0, err + return err } defer f.mu.Unlock() @@ -152,12 +135,7 @@ func (f *packetsPendingLinkResolution) enqueue(r *Route, proto tcpip.NetworkProt proto: proto, pkt: pkt, }) - switch pkt := pkt.(type) { - case *PacketBuffer: - pkt.IncRef() - case *PacketBufferList: - pkt.IncRef() - } + pkt.IncRef() if len(packets) > maxPendingPacketsPerResolution { f.incrementOutgoingPacketErrors(packets[0].proto, packets[0].pkt) @@ -172,7 +150,7 @@ func (f *packetsPendingLinkResolution) enqueue(r *Route, proto tcpip.NetworkProt f.mu.packets[ch] = packets if ok { - return pkt.len(), nil + return nil } cancelledPackets := f.newCancelChannelLocked(ch) @@ -183,7 +161,7 @@ func (f *packetsPendingLinkResolution) enqueue(r *Route, proto tcpip.NetworkProt go f.dequeuePackets(cancelledPackets, "" /* linkAddr */, &tcpip.ErrAborted{}) } - return pkt.len(), nil + return nil } // newCancelChannelLocked appends the link resolution channel to a FIFO. If the @@ -215,28 +193,14 @@ func (f *packetsPendingLinkResolution) dequeuePackets(packets []pendingPacket, l for _, p := range packets { if err == nil { p.routeInfo.RemoteLinkAddress = linkAddr - _, _ = f.nic.writePacketBuffer(p.routeInfo, p.proto, p.pkt) + _ = f.nic.writePacket(p.routeInfo, p.proto, p.pkt) } else { f.incrementOutgoingPacketErrors(p.proto, p.pkt) if linkResolvableEP, ok := f.nic.getNetworkEndpoint(p.proto).(LinkResolvableNetworkEndpoint); ok { - switch pkt := p.pkt.(type) { - case *PacketBuffer: - linkResolvableEP.HandleLinkResolutionFailure(pkt) - case *PacketBufferList: - for pb := pkt.Front(); pb != nil; pb = pb.Next() { - linkResolvableEP.HandleLinkResolutionFailure(pb) - } - default: - panic(fmt.Sprintf("unrecognized pending packet buffer type = %T", p.pkt)) - } + linkResolvableEP.HandleLinkResolutionFailure(p.pkt) } } - switch pkt := p.pkt.(type) { - case *PacketBuffer: - pkt.DecRef() - case *PacketBufferList: - pkt.DecRef() - } + p.pkt.DecRef() } } diff --git a/pkg/tcpip/stack/registration.go b/pkg/tcpip/stack/registration.go index bf9dcec18..da3cd0749 100644 --- a/pkg/tcpip/stack/registration.go +++ b/pkg/tcpip/stack/registration.go @@ -576,16 +576,6 @@ type NetworkInterface interface { // network and transport header must be set. WritePacket(*Route, tcpip.NetworkProtocolNumber, *PacketBuffer) tcpip.Error - // WritePackets writes packets with the given protocol through the given - // route. Must not be called with an empty list of packet buffers. - // - // WritePackets may modify the packet buffers. - // - // Right now, WritePackets is used only when the software segmentation - // offload is enabled. If it will be used for something else, syscall filters - // may need to be updated. - WritePackets(*Route, PacketBufferList, tcpip.NetworkProtocolNumber) (int, tcpip.Error) - // HandleNeighborProbe processes an incoming neighbor probe (e.g. ARP // request or NDP Neighbor Solicitation). // @@ -642,11 +632,6 @@ type NetworkEndpoint interface { // already been set. WritePacket(r *Route, params NetworkHeaderParams, pkt *PacketBuffer) tcpip.Error - // WritePackets writes packets to the given destination address and - // protocol. pkts must not be zero length. It may modify pkts and - // underlying packets. - WritePackets(r *Route, pkts PacketBufferList, params NetworkHeaderParams) (int, tcpip.Error) - // WriteHeaderIncludedPacket writes a packet that includes a network // header to the given destination address. It may modify pkt. WriteHeaderIncludedPacket(r *Route, pkt *PacketBuffer) tcpip.Error @@ -776,26 +761,14 @@ const ( ) // LinkWriter is an interface that supports sending packets via a data-link -// layer endpoint. +// layer endpoint. It is used with QueueingDiscipline to batch writes from +// upper layer endpoints. type LinkWriter interface { - // WritePacket writes a packet with the given protocol and route. - // - // WritePacket may modify the packet buffer. The packet buffer's - // network and transport header must be set. - // - // To participate in transparent bridging, a LinkEndpoint implementation - // should call eth.Encode with header.EthernetFields.SrcAddr set to - // r.LocalLinkAddress if it is provided. - WritePacket(RouteInfo, tcpip.NetworkProtocolNumber, *PacketBuffer) tcpip.Error - // WritePackets writes packets with the given protocol and route. Must not be // called with an empty list of packet buffers. // - // WritePackets may modify the packet buffers. - // - // Right now, WritePackets is used only when the software segmentation - // offload is enabled. If it will be used for something else, syscall filters - // may need to be updated. + // WritePackets may modify the packet buffers, and takes ownership of the PacketBufferList. + // it is not safe to use the PacketBufferList after a call to WritePackets. WritePackets(RouteInfo, PacketBufferList, tcpip.NetworkProtocolNumber) (int, tcpip.Error) } @@ -867,7 +840,16 @@ type NetworkLinkEndpoint interface { // QueueingDiscipline provides a queueing strategy for outgoing packets (e.g // FIFO, LIFO, Random Early Drop etc). type QueueingDiscipline interface { - LinkWriter + // WritePacket writes a packet with the given protocol and route. + // + // WritePacket may modify the packet buffer. The packet buffer's + // network and transport header must be set. + // + // To participate in transparent bridging, a LinkEndpoint implementation + // should call eth.Encode with header.EthernetFields.SrcAddr set to + // r.LocalLinkAddress if it is provided. + WritePacket(RouteInfo, tcpip.NetworkProtocolNumber, *PacketBuffer) tcpip.Error + Close() } @@ -880,6 +862,18 @@ type LinkEndpoint interface { NetworkLinkEndpoint LinkWriter LinkRawWriter + + // TODO(b/211019749): Remove WritePacket, it's no longer used outside the context of + // tests and LinkEndpoint wrappers. + // WritePacket writes a packet with the given protocol and route. + // + // WritePacket may modify the packet buffer. The packet buffer's + // network and transport header must be set. + // + // To participate in transparent bridging, a LinkEndpoint implementation + // should call eth.Encode with header.EthernetFields.SrcAddr set to + // r.LocalLinkAddress if it is provided. + WritePacket(RouteInfo, tcpip.NetworkProtocolNumber, *PacketBuffer) tcpip.Error } // InjectableLinkEndpoint is a LinkEndpoint where inbound packets are diff --git a/pkg/tcpip/stack/route.go b/pkg/tcpip/stack/route.go index 388fc5b9a..2834bd29b 100644 --- a/pkg/tcpip/stack/route.go +++ b/pkg/tcpip/stack/route.go @@ -465,16 +465,6 @@ func (r *Route) WritePacket(params NetworkHeaderParams, pkt *PacketBuffer) tcpip return r.outgoingNIC.getNetworkEndpoint(r.NetProto()).WritePacket(r, params, pkt) } -// WritePackets writes a list of n packets through the given route and returns -// the number of packets written. -func (r *Route) WritePackets(pkts PacketBufferList, params NetworkHeaderParams) (int, tcpip.Error) { - if !r.isValidForOutgoing() { - return 0, &tcpip.ErrInvalidEndpointState{} - } - - return r.outgoingNIC.getNetworkEndpoint(r.NetProto()).WritePackets(r, pkts, params) -} - // WriteHeaderIncludedPacket writes a packet already containing a network // header through the given route. func (r *Route) WriteHeaderIncludedPacket(pkt *PacketBuffer) tcpip.Error { diff --git a/pkg/tcpip/tests/integration/iptables_test.go b/pkg/tcpip/tests/integration/iptables_test.go index 3ae95aa95..e9f5393e9 100644 --- a/pkg/tcpip/tests/integration/iptables_test.go +++ b/pkg/tcpip/tests/integration/iptables_test.go @@ -355,21 +355,14 @@ func TestIPTablesStatsForInput(t *testing.T) { } } -var _ stack.LinkEndpoint = (*channelEndpointWithoutWritePacket)(nil) +var _ stack.LinkEndpoint = (*channelEndpoint)(nil) -// channelEndpointWithoutWritePacket is a channel endpoint that does not support -// stack.LinkEndpoint.WritePacket. -type channelEndpointWithoutWritePacket struct { +type channelEndpoint struct { *channel.Endpoint t *testing.T } -func (c *channelEndpointWithoutWritePacket) WritePacket(stack.RouteInfo, tcpip.NetworkProtocolNumber, *stack.PacketBuffer) tcpip.Error { - c.t.Error("unexpectedly called WritePacket; all writes should go through WritePackets") - return &tcpip.ErrNotSupported{} -} - var _ stack.Matcher = (*udpSourcePortMatcher)(nil) type udpSourcePortMatcher struct { @@ -610,7 +603,7 @@ func TestIPTableWritePackets(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, }) - e := channelEndpointWithoutWritePacket{ + e := channelEndpoint{ Endpoint: channel.New(4, header.IPv6MinimumMTU, linkAddr), t: t, } @@ -653,13 +646,16 @@ func TestIPTableWritePackets(t *testing.T) { pkts := test.genPacket(r) pktsLen := pkts.Len() - if n, err := r.WritePackets(pkts, stack.NetworkHeaderParams{ - Protocol: header.UDPProtocolNumber, - TTL: 64, - }); err != nil { - t.Fatalf("WritePackets(...): %s", err) - } else if n != pktsLen { - t.Fatalf("got WritePackets(...) = %d, want = %d", n, pktsLen) + for i := 0; i < pktsLen; i++ { + pkt := pkts.Front() + pkts.Remove(pkt) + if err := r.WritePacket(stack.NetworkHeaderParams{ + Protocol: header.UDPProtocolNumber, + TTL: 64, + }, pkt); err != nil { + t.Fatalf("WritePacket(...): %s", err) + } + pkt.DecRef() } if got := s.Stats().IP.PacketsSent.Value(); got != test.expectSent { diff --git a/pkg/tcpip/tests/integration/link_resolution_test.go b/pkg/tcpip/tests/integration/link_resolution_test.go index 93aa5c87d..b13a5d208 100644 --- a/pkg/tcpip/tests/integration/link_resolution_test.go +++ b/pkg/tcpip/tests/integration/link_resolution_test.go @@ -941,8 +941,12 @@ func TestWritePacketsLinkResolution(t *testing.T) { } defer r.Release() + params := stack.NetworkHeaderParams{ + Protocol: udp.ProtocolNumber, + TTL: 64, + TOS: stack.DefaultTOS, + } data := []byte{1, 2} - var pkts stack.PacketBufferList for _, d := range data { pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ ReserveHeaderBytes: header.UDPMinimumSize + int(r.MaxHeaderLength()), @@ -960,19 +964,10 @@ func TestWritePacketsLinkResolution(t *testing.T) { xsum = header.ChecksumCombine(xsum, pkt.Data().AsRange().Checksum()) udpHdr.SetChecksum(^udpHdr.CalculateChecksum(xsum)) - pkts.PushBack(pkt) - } - - params := stack.NetworkHeaderParams{ - Protocol: udp.ProtocolNumber, - TTL: 64, - TOS: stack.DefaultTOS, - } - - if n, err := r.WritePackets(pkts, params); err != nil { - t.Fatalf("r.WritePackets(_, %#v): %s", params, err) - } else if want := pkts.Len(); want != n { - t.Fatalf("got r.WritePackets(_, %#v) = %d, want = %d", params, n, want) + if err := r.WritePacket(params, pkt); err != nil { + t.Fatalf("WritePacket(...): %s", err) + } + pkt.DecRef() } var writer bytes.Buffer diff --git a/pkg/tcpip/transport/tcp/connect.go b/pkg/tcpip/transport/tcp/connect.go index b1ab46c97..ae454dbfb 100644 --- a/pkg/tcpip/transport/tcp/connect.go +++ b/pkg/tcpip/transport/tcp/connect.go @@ -832,7 +832,6 @@ func sendTCPBatch(r *stack.Route, tf tcpFields, data buffer.VectorisedView, gso size := data.Size() hdrSize := header.TCPMinimumSize + int(r.MaxHeaderLength()) + optLen - var pkts stack.PacketBufferList for i := 0; i < n; i++ { packetSize := mss if packetSize > size { @@ -848,19 +847,18 @@ func sendTCPBatch(r *stack.Route, tf tcpFields, data buffer.VectorisedView, gso buildTCPHdr(r, tf, pkt, gso) tf.seq = tf.seq.Add(seqnum.Size(packetSize)) pkt.GSOOptions = gso - pkts.PushBack(pkt) + if tf.ttl == 0 { + tf.ttl = r.DefaultTTL() + } + if err := r.WritePacket(stack.NetworkHeaderParams{Protocol: ProtocolNumber, TTL: tf.ttl, TOS: tf.tos}, pkt); err != nil { + r.Stats().TCP.SegmentSendErrors.Increment() + pkt.DecRef() + return err + } + r.Stats().TCP.SegmentsSent.Increment() + pkt.DecRef() } - defer pkts.DecRef() - - if tf.ttl == 0 { - tf.ttl = r.DefaultTTL() - } - sent, err := r.WritePackets(pkts, stack.NetworkHeaderParams{Protocol: ProtocolNumber, TTL: tf.ttl, TOS: tf.tos}) - if err != nil { - r.Stats().TCP.SegmentSendErrors.IncrementBy(uint64(n - sent)) - } - r.Stats().TCP.SegmentsSent.IncrementBy(uint64(sent)) - return err + return nil } // sendTCP sends a TCP segment with the provided options via the provided