From 196baa62ca9492d93f17b176859f6eaf539155be Mon Sep 17 00:00:00 2001 From: Ghanan Gowripalan Date: Fri, 14 Jan 2022 13:20:15 -0800 Subject: [PATCH] Don't pass route info and net proto to write fns ...as the packet buffer already holds that information. Updates #3810. Fixes #6537. PiperOrigin-RevId: 421898143 --- pkg/tcpip/link/channel/channel.go | 2 +- pkg/tcpip/link/ethernet/ethernet.go | 4 +- pkg/tcpip/link/ethernet/ethernet_test.go | 6 +- pkg/tcpip/link/fdbased/endpoint.go | 2 +- pkg/tcpip/link/fdbased/endpoint_test.go | 8 +- pkg/tcpip/link/loopback/loopback.go | 2 +- pkg/tcpip/link/muxed/injectable.go | 29 ++++- pkg/tcpip/link/muxed/injectable_test.go | 12 +- pkg/tcpip/link/nested/nested.go | 4 +- pkg/tcpip/link/pipe/pipe.go | 11 +- pkg/tcpip/link/qdisc/fifo/fifo.go | 7 +- pkg/tcpip/link/qdisc/fifo/qdisc_test.go | 7 +- pkg/tcpip/link/sharedmem/sharedmem.go | 2 +- pkg/tcpip/link/sharedmem/sharedmem_server.go | 2 +- pkg/tcpip/link/sharedmem/sharedmem_test.go | 45 ++++++-- pkg/tcpip/link/sniffer/sniffer.go | 4 +- pkg/tcpip/link/waitable/waitable.go | 4 +- pkg/tcpip/link/waitable/waitable_test.go | 20 +++- pkg/tcpip/network/arp/arp.go | 4 +- pkg/tcpip/network/arp/arp_test.go | 4 +- .../network/internal/testutil/testutil.go | 2 +- pkg/tcpip/network/ip_test.go | 4 +- pkg/tcpip/network/ipv4/icmp.go | 4 - pkg/tcpip/network/ipv4/igmp.go | 2 +- pkg/tcpip/network/ipv4/ipv4.go | 4 +- pkg/tcpip/network/ipv6/icmp_test.go | 18 +-- pkg/tcpip/network/ipv6/ipv6.go | 4 +- pkg/tcpip/network/ipv6/mld.go | 2 +- pkg/tcpip/network/ipv6/ndp.go | 4 +- pkg/tcpip/stack/forwarding_test.go | 108 ++++++++---------- pkg/tcpip/stack/nic.go | 27 ++--- pkg/tcpip/stack/pending_packets.go | 20 ++-- pkg/tcpip/stack/registration.go | 19 ++- pkg/tcpip/stack/stack.go | 2 +- pkg/tcpip/stack/stack_test.go | 2 +- 35 files changed, 212 insertions(+), 189 deletions(-) diff --git a/pkg/tcpip/link/channel/channel.go b/pkg/tcpip/link/channel/channel.go index 83531f6f8..36ef274f5 100644 --- a/pkg/tcpip/link/channel/channel.go +++ b/pkg/tcpip/link/channel/channel.go @@ -240,7 +240,7 @@ func (e *Endpoint) LinkAddress() tcpip.LinkAddress { } // WritePackets stores outbound packets into the channel. -func (e *Endpoint) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *Endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { n := 0 for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { if !e.q.Write(pkt) { diff --git a/pkg/tcpip/link/ethernet/ethernet.go b/pkg/tcpip/link/ethernet/ethernet.go index 8913d6d49..bf7814245 100644 --- a/pkg/tcpip/link/ethernet/ethernet.go +++ b/pkg/tcpip/link/ethernet/ethernet.go @@ -81,14 +81,14 @@ func (e *Endpoint) Capabilities() stack.LinkEndpointCapabilities { } // WritePackets implements stack.LinkEndpoint. -func (e *Endpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, proto tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *Endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { linkAddr := e.LinkAddress() for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { e.AddHeader(linkAddr, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) } - return e.Endpoint.WritePackets(r, pkts, proto) + return e.Endpoint.WritePackets(pkts) } // MaxHeaderLength implements stack.LinkEndpoint. diff --git a/pkg/tcpip/link/ethernet/ethernet_test.go b/pkg/tcpip/link/ethernet/ethernet_test.go index 4361fe034..ac66bdf47 100644 --- a/pkg/tcpip/link/ethernet/ethernet_test.go +++ b/pkg/tcpip/link/ethernet/ethernet_test.go @@ -140,10 +140,10 @@ func TestWritePacketsAddHeader(t *testing.T) { var pkts stack.PacketBufferList pkts.PushFront(pkt) - if n, err := e.WritePackets(stack.RouteInfo{}, pkts, 0 /* protocol */); err != nil { - t.Fatalf("e.WritePackets({}, _, 0): %s", err) + if n, err := e.WritePackets(pkts); err != nil { + t.Fatalf("e.WritePackets(_): %s", err) } else if n != 1 { - t.Fatalf("got e.WritePackets({}, _, 0) = %d, want = 1", n) + t.Fatalf("got e.WritePackets(_) = %d, want = 1", n) } } diff --git a/pkg/tcpip/link/fdbased/endpoint.go b/pkg/tcpip/link/fdbased/endpoint.go index 83e156885..490cd9dc5 100644 --- a/pkg/tcpip/link/fdbased/endpoint.go +++ b/pkg/tcpip/link/fdbased/endpoint.go @@ -668,7 +668,7 @@ func (e *endpoint) sendBatch(batchFD int, pkts []*stack.PacketBuffer) (int, tcpi // - pkt.EgressRoute // - pkt.GSOOptions // - pkt.NetworkProtocolNumber -func (e *endpoint) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { // Preallocate to avoid repeated reallocation as we append to batch. // batchSz is 47 because when SWGSO is in use then a single 65KB TCP // segment can get split into 46 segments of 1420 bytes and a single 216 diff --git a/pkg/tcpip/link/fdbased/endpoint_test.go b/pkg/tcpip/link/fdbased/endpoint_test.go index aa13ba096..391012356 100644 --- a/pkg/tcpip/link/fdbased/endpoint_test.go +++ b/pkg/tcpip/link/fdbased/endpoint_test.go @@ -224,8 +224,8 @@ func testWritePacket(t *testing.T, plen int, eth bool, gsoMaxSize uint32, hash u } var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, proto); err != nil { - t.Fatalf("WritePacket failed: %v", err) + if _, err := c.ep.WritePackets(pkts); err != nil { + t.Fatalf("WritePackets failed: %s", err) } // Read from the corresponding FD, then compare with what we wrote. @@ -345,8 +345,8 @@ func TestPreserveSrcAddress(t *testing.T) { pkt.EgressRoute = r var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, proto); err != nil { - t.Fatalf("WritePacket failed: %v", err) + if _, err := c.ep.WritePackets(pkts); err != nil { + t.Fatalf("WritePackets failed: %s", err) } // Read from the FD, then compare with what we wrote. diff --git a/pkg/tcpip/link/loopback/loopback.go b/pkg/tcpip/link/loopback/loopback.go index 9b4d0ca1e..7b4635271 100644 --- a/pkg/tcpip/link/loopback/loopback.go +++ b/pkg/tcpip/link/loopback/loopback.go @@ -75,7 +75,7 @@ func (*endpoint) LinkAddress() tcpip.LinkAddress { func (*endpoint) Wait() {} // WritePackets implements stack.LinkEndpoint.WritePackets. -func (e *endpoint) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { n := 0 for p := pkts.Front(); p != nil; p = p.Next() { if err := e.WriteRawPacket(p); err != nil { diff --git a/pkg/tcpip/link/muxed/injectable.go b/pkg/tcpip/link/muxed/injectable.go index 0dae3eca7..863fb097c 100644 --- a/pkg/tcpip/link/muxed/injectable.go +++ b/pkg/tcpip/link/muxed/injectable.go @@ -86,13 +86,30 @@ func (m *InjectableEndpoint) InjectInbound(protocol tcpip.NetworkProtocolNumber, // WritePackets writes outbound packets to the appropriate // LinkInjectableEndpoint based on the RemoteAddress. HandleLocal only works if -// r.RemoteAddress has a route registered in this endpoint. -func (m *InjectableEndpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - endpoint, ok := m.routes[r.RemoteAddress] - if !ok { - return 0, &tcpip.ErrNoRoute{} +// pkt.EgressRoute.RemoteAddress has a route registered in this endpoint. +func (m *InjectableEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { + i := 0 + for pkt := pkts.Front(); pkt != nil; { + nextPkt := pkt.Next() + + endpoint, ok := m.routes[pkt.EgressRoute.RemoteAddress] + if !ok { + return i, &tcpip.ErrNoRoute{} + } + + var tmpPkts stack.PacketBufferList + tmpPkts.PushFront(pkt) + + n, err := endpoint.WritePackets(tmpPkts) + if err != nil { + return i, err + } + + i += n + pkt = nextPkt } - return endpoint.WritePackets(r, pkts, protocol) + + return i, nil } // InjectOutbound writes outbound packets to the appropriate diff --git a/pkg/tcpip/link/muxed/injectable_test.go b/pkg/tcpip/link/muxed/injectable_test.go index f129748d8..a57373622 100644 --- a/pkg/tcpip/link/muxed/injectable_test.go +++ b/pkg/tcpip/link/muxed/injectable_test.go @@ -51,12 +51,12 @@ func TestInjectableEndpointDispatch(t *testing.T) { Data: buffer.NewViewFromBytes([]byte{0xFB}).ToVectorisedView(), }) pkt.TransportHeader().Push(1)[0] = 0xFA - var packetRoute stack.RouteInfo - packetRoute.RemoteAddress = dstIP + pkt.EgressRoute.RemoteAddress = dstIP + pkt.NetworkProtocolNumber = ipv4.ProtocolNumber var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := endpoint.WritePackets(packetRoute, pkts, ipv4.ProtocolNumber); err != nil { + if _, err := endpoint.WritePackets(pkts); err != nil { t.Fatalf("Unable to write packets: %s", err) } @@ -78,12 +78,12 @@ func TestInjectableEndpointDispatchHdrOnly(t *testing.T) { Data: buffer.NewView(0).ToVectorisedView(), }) pkt.TransportHeader().Push(1)[0] = 0xFA - var packetRoute stack.RouteInfo - packetRoute.RemoteAddress = dstIP + pkt.EgressRoute.RemoteAddress = dstIP + pkt.NetworkProtocolNumber = ipv4.ProtocolNumber var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := endpoint.WritePackets(packetRoute, pkts, ipv4.ProtocolNumber); err != nil { + if _, err := endpoint.WritePackets(pkts); err != nil { t.Fatalf("Unable to write packets: %s", err) } buf := make([]byte, 6500) diff --git a/pkg/tcpip/link/nested/nested.go b/pkg/tcpip/link/nested/nested.go index 99c5b9bfc..b5975ec51 100644 --- a/pkg/tcpip/link/nested/nested.go +++ b/pkg/tcpip/link/nested/nested.go @@ -103,8 +103,8 @@ func (e *Endpoint) LinkAddress() tcpip.LinkAddress { } // WritePackets implements stack.LinkEndpoint. -func (e *Endpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - return e.child.WritePackets(r, pkts, protocol) +func (e *Endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { + return e.child.WritePackets(pkts) } // Wait implements stack.LinkEndpoint. diff --git a/pkg/tcpip/link/pipe/pipe.go b/pkg/tcpip/link/pipe/pipe.go index a7620e709..33a0b2eab 100644 --- a/pkg/tcpip/link/pipe/pipe.go +++ b/pkg/tcpip/link/pipe/pipe.go @@ -48,7 +48,7 @@ type Endpoint struct { mtu uint32 } -func (e *Endpoint) deliverPackets(r stack.RouteInfo, proto tcpip.NetworkProtocolNumber, pkts stack.PacketBufferList) { +func (e *Endpoint) deliverPackets(pkts stack.PacketBufferList) { if !e.linked.IsAttached() { return } @@ -65,15 +65,16 @@ func (e *Endpoint) deliverPackets(r stack.RouteInfo, proto tcpip.NetworkProtocol newPkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ Data: buffer.NewVectorisedView(pkt.Size(), pkt.Views()), }) - e.linked.dispatcher.DeliverNetworkPacket(r.LocalLinkAddress /* remote */, r.RemoteLinkAddress /* local */, proto, newPkt) + r := pkt.EgressRoute + e.linked.dispatcher.DeliverNetworkPacket(r.LocalLinkAddress /* remote */, r.RemoteLinkAddress /* local */, pkt.NetworkProtocolNumber, newPkt) newPkt.DecRef() } } // WritePackets implements stack.LinkEndpoint. -func (e *Endpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, proto tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *Endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { n := pkts.Len() - e.deliverPackets(r, proto, pkts) + e.deliverPackets(pkts) return n, nil } @@ -123,6 +124,6 @@ func (*Endpoint) AddHeader(_, _ tcpip.LinkAddress, _ tcpip.NetworkProtocolNumber func (e *Endpoint) WriteRawPacket(pkt *stack.PacketBuffer) tcpip.Error { var pkts stack.PacketBufferList pkts.PushBack(pkt) - _, err := e.WritePackets(stack.RouteInfo{}, pkts, 0) + _, err := e.WritePackets(pkts) return err } diff --git a/pkg/tcpip/link/qdisc/fifo/fifo.go b/pkg/tcpip/link/qdisc/fifo/fifo.go index 73f351ada..b90bb0a99 100644 --- a/pkg/tcpip/link/qdisc/fifo/fifo.go +++ b/pkg/tcpip/link/qdisc/fifo/fifo.go @@ -96,15 +96,14 @@ func (qd *queueDispatcher) dispatchLoop() { if pkt == nil { break } + qd.queue.Remove(pkt) qd.used-- batch.PushBack(pkt) } qd.mu.Unlock() - // We pass a protocol of zero here because each packet carries its - // NetworkProtocol. - _, _ = qd.lower.WritePackets(stack.RouteInfo{}, batch, 0 /* protocol */) + _, _ = qd.lower.WritePackets(batch) batch.DecRef() batch.Reset() } @@ -116,7 +115,7 @@ func (qd *queueDispatcher) dispatchLoop() { // - pkt.EgressRoute // - pkt.GSOOptions // - pkt.NetworkProtocolNumber -func (d *discipline) WritePacket(_ stack.RouteInfo, _ tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { +func (d *discipline) WritePacket(pkt *stack.PacketBuffer) tcpip.Error { qd := &d.dispatchers[int(pkt.Hash)%len(d.dispatchers)] qd.mu.Lock() haveSpace := qd.used < qd.limit diff --git a/pkg/tcpip/link/qdisc/fifo/qdisc_test.go b/pkg/tcpip/link/qdisc/fifo/qdisc_test.go index 1e17d26c7..c8ce95deb 100644 --- a/pkg/tcpip/link/qdisc/fifo/qdisc_test.go +++ b/pkg/tcpip/link/qdisc/fifo/qdisc_test.go @@ -31,7 +31,7 @@ var _ stack.LinkWriter = (*discardWriter)(nil) type discardWriter struct { } -func (*discardWriter) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (*discardWriter) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { return pkts.Len(), nil } @@ -43,9 +43,6 @@ func TestFastSimultaneousWrites(t *testing.T) { 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 @@ -60,7 +57,7 @@ func TestFastSimultaneousWrites(t *testing.T) { Data: v.ToVectorisedView(), }) pkt.Hash = rand.Uint32() - linkEP.WritePacket(r, prot, pkt) + linkEP.WritePacket(pkt) pkt.DecRef() } }() diff --git a/pkg/tcpip/link/sharedmem/sharedmem.go b/pkg/tcpip/link/sharedmem/sharedmem.go index e304149b9..1b4da2521 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem.go +++ b/pkg/tcpip/link/sharedmem/sharedmem.go @@ -364,7 +364,7 @@ func (e *endpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.NetworkPr } // WritePackets implements stack.LinkEndpoint.WritePackets. -func (e *endpoint) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { n := 0 var err tcpip.Error e.mu.Lock() diff --git a/pkg/tcpip/link/sharedmem/sharedmem_server.go b/pkg/tcpip/link/sharedmem/sharedmem_server.go index fb0b7b8d7..aff6aa40d 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_server.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_server.go @@ -274,7 +274,7 @@ func (e *serverEndpoint) WritePacket(_ stack.RouteInfo, _ tcpip.NetworkProtocolN } // WritePackets implements stack.LinkEndpoint.WritePackets. -func (e *serverEndpoint) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *serverEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { n := 0 var err tcpip.Error e.mu.Lock() diff --git a/pkg/tcpip/link/sharedmem/sharedmem_test.go b/pkg/tcpip/link/sharedmem/sharedmem_test.go index 16759c5fa..cf58fe513 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_test.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_test.go @@ -235,7 +235,7 @@ func TestSimpleSend(t *testing.T) { pkt.NetworkProtocolNumber = proto var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, proto); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed: %s", err) } @@ -311,7 +311,7 @@ func TestPreserveSrcAddressInSend(t *testing.T) { var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, proto); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed: %s", err) } @@ -366,10 +366,12 @@ func TestFillTxQueue(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } @@ -387,9 +389,12 @@ func TestFillTxQueue(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + var pkts stack.PacketBufferList pkts.PushBack(pkt) - _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber) + _, err := c.ep.WritePackets(pkts) if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } @@ -420,8 +425,10 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { Data: buf.ToVectorisedView(), }) pkts.PushBack(pkt) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber } - if _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } } @@ -444,9 +451,11 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } @@ -464,9 +473,11 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber var pkts stack.PacketBufferList pkts.PushBack(pkt) - _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber) + _, err := c.ep.WritePackets(pkts) if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } @@ -492,9 +503,11 @@ func TestFillTxMemory(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber var pkts stack.PacketBufferList pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } @@ -513,9 +526,11 @@ func TestFillTxMemory(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + pkt.EgressRoute = r var pkts stack.PacketBufferList pkts.PushBack(pkt) - _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber) + _, err := c.ep.WritePackets(pkts) if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } @@ -543,8 +558,10 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { Data: buf.ToVectorisedView(), }) var pkts stack.PacketBufferList + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } @@ -560,8 +577,10 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buffer.NewView(bufferSize).ToVectorisedView(), }) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber pkts.PushBack(pkt) - _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber) + _, err := c.ep.WritePackets(pkts) if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } @@ -574,8 +593,10 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) + pkt.EgressRoute = r + pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber pkts.PushBack(pkt) - if _, err := c.ep.WritePackets(r, pkts, header.IPv4ProtocolNumber); err != nil { + if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } } diff --git a/pkg/tcpip/link/sniffer/sniffer.go b/pkg/tcpip/link/sniffer/sniffer.go index 339992dce..4f4828d03 100644 --- a/pkg/tcpip/link/sniffer/sniffer.go +++ b/pkg/tcpip/link/sniffer/sniffer.go @@ -164,11 +164,11 @@ func (e *endpoint) dumpPacket(dir direction, protocol tcpip.NetworkProtocolNumbe // WritePackets implements the stack.LinkEndpoint interface. It is called by // higher-level protocols to write packets; it just logs the packet and // forwards the request to the lower endpoint. -func (e *endpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { e.dumpPacket(directionSend, pkt.NetworkProtocolNumber, pkt) } - return e.Endpoint.WritePackets(r, pkts, protocol) + return e.Endpoint.WritePackets(pkts) } func logPacket(prefix string, dir direction, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { diff --git a/pkg/tcpip/link/waitable/waitable.go b/pkg/tcpip/link/waitable/waitable.go index 8513054a0..449ba94b6 100644 --- a/pkg/tcpip/link/waitable/waitable.go +++ b/pkg/tcpip/link/waitable/waitable.go @@ -99,12 +99,12 @@ func (e *Endpoint) LinkAddress() tcpip.LinkAddress { // WritePackets implements stack.LinkEndpoint.WritePackets. It is called by // higher-level protocols to write packets. It only forwards packets to the // lower endpoint if Wait or WaitWrite haven't been called. -func (e *Endpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *Endpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { if !e.writeGate.Enter() { return pkts.Len(), nil } - n, err := e.lower.WritePackets(r, pkts, protocol) + n, err := e.lower.WritePackets(pkts) e.writeGate.Leave() return n, err } diff --git a/pkg/tcpip/link/waitable/waitable_test.go b/pkg/tcpip/link/waitable/waitable_test.go index 28dbd0f23..5a0afd9cd 100644 --- a/pkg/tcpip/link/waitable/waitable_test.go +++ b/pkg/tcpip/link/waitable/waitable_test.go @@ -72,7 +72,7 @@ func (e *countedEndpoint) LinkAddress() tcpip.LinkAddress { } // WritePackets implements stack.LinkEndpoint.WritePackets. -func (e *countedEndpoint) WritePackets(_ stack.RouteInfo, pkts stack.PacketBufferList, _ tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *countedEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { e.writeCount += pkts.Len() return pkts.Len(), nil } @@ -102,7 +102,11 @@ func TestWaitWrite(t *testing.T) { var pkts stack.PacketBufferList pkts.PushBack(stack.NewPacketBuffer(stack.PacketBufferOptions{})) // Write and check that it goes through. - wep.WritePackets(stack.RouteInfo{}, pkts, 0) + if n, err := wep.WritePackets(pkts); err != nil { + t.Fatalf("WritePackets(_): %s", err) + } else if n != 1 { + t.Fatalf("got WritePackets(_) = %d, want = 1", n) + } if want := 1; ep.writeCount != want { t.Fatalf("Unexpected writeCount: got=%v, want=%v", ep.writeCount, want) } @@ -112,7 +116,11 @@ func TestWaitWrite(t *testing.T) { pkts.PushBack(stack.NewPacketBuffer(stack.PacketBufferOptions{})) // Wait on dispatches, then try to write. It must go through. wep.WaitDispatch() - wep.WritePackets(stack.RouteInfo{}, pkts, 0) + if n, err := wep.WritePackets(pkts); err != nil { + t.Fatalf("WritePackets(_): %s", err) + } else if n != 1 { + t.Fatalf("got WritePackets(_) = %d, want = 1", n) + } if want := 2; ep.writeCount != want { t.Fatalf("Unexpected writeCount: got=%v, want=%v", ep.writeCount, want) } @@ -123,7 +131,11 @@ func TestWaitWrite(t *testing.T) { pkts.PushBack(stack.NewPacketBuffer(stack.PacketBufferOptions{})) // Wait on writes, then try to write. It must not go through. wep.WaitWrite() - wep.WritePackets(stack.RouteInfo{}, pkts, 0) + if n, err := wep.WritePackets(pkts); err != nil { + t.Fatalf("WritePackets(_): %s", err) + } else if n != 1 { + t.Fatalf("got WritePackets(_) = %d, want = 1", n) + } if want := 2; ep.writeCount != want { t.Fatalf("Unexpected writeCount: got=%v, want=%v", ep.writeCount, want) } diff --git a/pkg/tcpip/network/arp/arp.go b/pkg/tcpip/network/arp/arp.go index 9ba69ba41..c57d029cd 100644 --- a/pkg/tcpip/network/arp/arp.go +++ b/pkg/tcpip/network/arp/arp.go @@ -218,7 +218,7 @@ func (e *endpoint) HandlePacket(pkt *stack.PacketBuffer) { // // Send the packet to the (new) target hardware address on the same // hardware on which the request was received. - if err := e.nic.WritePacketToRemote(tcpip.LinkAddress(origSender), ProtocolNumber, respPkt); err != nil { + if err := e.nic.WritePacketToRemote(tcpip.LinkAddress(origSender), respPkt); err != nil { stats.outgoingRepliesDropped.Increment() } else { stats.outgoingRepliesSent.Increment() @@ -351,7 +351,7 @@ func (e *endpoint) sendARPRequest(localAddr, targetAddr tcpip.Address, remoteLin } stats := e.stats.arp - if err := e.nic.WritePacketToRemote(remoteLinkAddr, ProtocolNumber, pkt); err != nil { + if err := e.nic.WritePacketToRemote(remoteLinkAddr, pkt); err != nil { stats.outgoingRequestsDropped.Increment() return err } diff --git a/pkg/tcpip/network/arp/arp_test.go b/pkg/tcpip/network/arp/arp_test.go index 4e4d385b8..8abb65349 100644 --- a/pkg/tcpip/network/arp/arp_test.go +++ b/pkg/tcpip/network/arp/arp_test.go @@ -420,12 +420,12 @@ type testLinkEndpoint struct { writeErr tcpip.Error } -func (t *testLinkEndpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (t *testLinkEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { if t.writeErr != nil { return 0, t.writeErr } - return t.LinkEndpoint.WritePackets(r, pkts, protocol) + return t.LinkEndpoint.WritePackets(pkts) } func TestLinkAddressRequest(t *testing.T) { diff --git a/pkg/tcpip/network/internal/testutil/testutil.go b/pkg/tcpip/network/internal/testutil/testutil.go index 220aa766d..fb9089ef2 100644 --- a/pkg/tcpip/network/internal/testutil/testutil.go +++ b/pkg/tcpip/network/internal/testutil/testutil.go @@ -62,7 +62,7 @@ func (*MockLinkEndpoint) MaxHeaderLength() uint16 { return 0 } func (*MockLinkEndpoint) LinkAddress() tcpip.LinkAddress { return "" } // WritePackets implements LinkEndpoint.WritePackets. -func (ep *MockLinkEndpoint) WritePackets(r stack.RouteInfo, pkts stack.PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (ep *MockLinkEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { var n int for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { if ep.allowPackets == 0 { diff --git a/pkg/tcpip/network/ip_test.go b/pkg/tcpip/network/ip_test.go index a2fed55ac..eef3d66c5 100644 --- a/pkg/tcpip/network/ip_test.go +++ b/pkg/tcpip/network/ip_test.go @@ -191,7 +191,7 @@ func (*testObject) Wait() {} // WritePacket is called by network endpoints after producing a packet and // writing it to the link endpoint. This is used by the test object to verify // that the produced packet is as expected. -func (t *testObject) WritePacket(_ *stack.Route, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { +func (t *testObject) WritePacket(_ *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { var prot tcpip.TransportProtocolNumber var srcAddr tcpip.Address var dstAddr tcpip.Address @@ -341,7 +341,7 @@ func (t *testInterface) setEnabled(v bool) { t.mu.disabled = !v } -func (*testInterface) WritePacketToRemote(tcpip.LinkAddress, tcpip.NetworkProtocolNumber, *stack.PacketBuffer) tcpip.Error { +func (*testInterface) WritePacketToRemote(tcpip.LinkAddress, *stack.PacketBuffer) tcpip.Error { return &tcpip.ErrNotSupported{} } diff --git a/pkg/tcpip/network/ipv4/icmp.go b/pkg/tcpip/network/ipv4/icmp.go index 7260fa84e..e31a04624 100644 --- a/pkg/tcpip/network/ipv4/icmp.go +++ b/pkg/tcpip/network/ipv4/icmp.go @@ -284,10 +284,6 @@ func (e *endpoint) handleICMP(pkt *stack.PacketBuffer) { return } - // TODO(gvisor.dev/issue/3810:) When adding protocol numbers into the - // header information, we may have to change this code to handle the - // ICMP header no longer being in the data buffer. - // Because IP and ICMP are so closely intertwined, we need to handcraft our // IP header to be able to follow RFC 792. The wording on page 13 is as // follows: diff --git a/pkg/tcpip/network/ipv4/igmp.go b/pkg/tcpip/network/ipv4/igmp.go index 862b893df..91a85434a 100644 --- a/pkg/tcpip/network/ipv4/igmp.go +++ b/pkg/tcpip/network/ipv4/igmp.go @@ -342,7 +342,7 @@ func (igmp *igmpState) writePacket(destAddress tcpip.Address, groupAddress tcpip } sentStats := igmp.ep.stats.igmp.packetsSent - if err := igmp.ep.nic.WritePacketToRemote(header.EthernetAddressFromMulticastIPv4Address(destAddress), ProtocolNumber, pkt); err != nil { + if err := igmp.ep.nic.WritePacketToRemote(header.EthernetAddressFromMulticastIPv4Address(destAddress), pkt); err != nil { sentStats.dropped.Increment() return false, err } diff --git a/pkg/tcpip/network/ipv4/ipv4.go b/pkg/tcpip/network/ipv4/ipv4.go index 85413a868..54be0cbe3 100644 --- a/pkg/tcpip/network/ipv4/ipv4.go +++ b/pkg/tcpip/network/ipv4/ipv4.go @@ -499,14 +499,14 @@ func (e *endpoint) writePacketPostRouting(r *stack.Route, pkt *stack.PacketBuffe // fragment one by one using WritePacket() (current strategy) or if we // want to create a PacketBufferList from the fragments and feed it to // WritePackets(). It'll be faster but cost more memory. - return e.nic.WritePacket(r, ProtocolNumber, fragPkt) + return e.nic.WritePacket(r, fragPkt) }) stats.PacketsSent.IncrementBy(uint64(sent)) stats.OutgoingPacketErrors.IncrementBy(uint64(remain)) return err } - if err := e.nic.WritePacket(r, ProtocolNumber, pkt); err != nil { + if err := e.nic.WritePacket(r, pkt); err != nil { stats.OutgoingPacketErrors.Increment() return err } diff --git a/pkg/tcpip/network/ipv6/icmp_test.go b/pkg/tcpip/network/ipv6/icmp_test.go index 99f1acfb1..36a1d7051 100644 --- a/pkg/tcpip/network/ipv6/icmp_test.go +++ b/pkg/tcpip/network/ipv6/icmp_test.go @@ -77,8 +77,8 @@ func (*stubLinkEndpoint) LinkAddress() tcpip.LinkAddress { return "" } -func (*stubLinkEndpoint) WritePackets(stack.RouteInfo, stack.PacketBufferList, tcpip.NetworkProtocolNumber) (int, tcpip.Error) { - return 0, nil +func (*stubLinkEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.Error) { + return pkts.Len(), nil } func (*stubLinkEndpoint) Attach(stack.NetworkDispatcher) {} @@ -130,20 +130,20 @@ func (*testInterface) Spoofing() bool { return false } -func (t *testInterface) WritePacket(r *stack.Route, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { +func (t *testInterface) WritePacket(r *stack.Route, pkt *stack.PacketBuffer) tcpip.Error { + pkt.EgressRoute = r.Fields() var pkts stack.PacketBufferList pkts.PushBack(pkt) - _, err := t.LinkEndpoint.WritePackets(r.Fields(), pkts, protocol) + _, err := t.LinkEndpoint.WritePackets(pkts) return err } -func (t *testInterface) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { - var r stack.RouteInfo - r.NetProto = protocol - r.RemoteLinkAddress = remoteLinkAddr +func (t *testInterface) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt *stack.PacketBuffer) tcpip.Error { + pkt.EgressRoute.NetProto = pkt.NetworkProtocolNumber + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr var pkts stack.PacketBufferList pkts.PushBack(pkt) - _, err := t.LinkEndpoint.WritePackets(r, pkts, protocol) + _, err := t.LinkEndpoint.WritePackets(pkts) return err } diff --git a/pkg/tcpip/network/ipv6/ipv6.go b/pkg/tcpip/network/ipv6/ipv6.go index fa492c062..8ec8bf221 100644 --- a/pkg/tcpip/network/ipv6/ipv6.go +++ b/pkg/tcpip/network/ipv6/ipv6.go @@ -814,14 +814,14 @@ func (e *endpoint) writePacket(r *stack.Route, pkt *stack.PacketBuffer, protocol // fragment one by one using WritePacket() (current strategy) or if we // want to create a PacketBufferList from the fragments and feed it to // WritePackets(). It'll be faster but cost more memory. - return e.nic.WritePacket(r, ProtocolNumber, fragPkt) + return e.nic.WritePacket(r, fragPkt) }) stats.PacketsSent.IncrementBy(uint64(sent)) stats.OutgoingPacketErrors.IncrementBy(uint64(remain)) return err } - if err := e.nic.WritePacket(r, ProtocolNumber, pkt); err != nil { + if err := e.nic.WritePacket(r, pkt); err != nil { stats.OutgoingPacketErrors.Increment() return err } diff --git a/pkg/tcpip/network/ipv6/mld.go b/pkg/tcpip/network/ipv6/mld.go index 06a8e1b89..7238053c8 100644 --- a/pkg/tcpip/network/ipv6/mld.go +++ b/pkg/tcpip/network/ipv6/mld.go @@ -278,7 +278,7 @@ func (mld *mldState) writePacket(destAddress, groupAddress tcpip.Address, mldTyp }, extensionHeaders); err != nil { panic(fmt.Sprintf("failed to add IP header: %s", err)) } - if err := mld.ep.nic.WritePacketToRemote(header.EthernetAddressFromMulticastIPv6Address(destAddress), ProtocolNumber, pkt); err != nil { + if err := mld.ep.nic.WritePacketToRemote(header.EthernetAddressFromMulticastIPv6Address(destAddress), pkt); err != nil { sentStats.dropped.Increment() return false, err } diff --git a/pkg/tcpip/network/ipv6/ndp.go b/pkg/tcpip/network/ipv6/ndp.go index c363da890..59ed42c27 100644 --- a/pkg/tcpip/network/ipv6/ndp.go +++ b/pkg/tcpip/network/ipv6/ndp.go @@ -1817,7 +1817,7 @@ func (ndp *ndpState) startSolicitingRouters() { panic(fmt.Sprintf("failed to add IP header: %s", err)) } - if err := ndp.ep.nic.WritePacketToRemote(header.EthernetAddressFromMulticastIPv6Address(header.IPv6AllRoutersLinkLocalMulticastAddress), ProtocolNumber, pkt); err != nil { + if err := ndp.ep.nic.WritePacketToRemote(header.EthernetAddressFromMulticastIPv6Address(header.IPv6AllRoutersLinkLocalMulticastAddress), pkt); err != nil { sent.dropped.Increment() // Don't send any more messages if we had an error. remaining = 0 @@ -1935,7 +1935,7 @@ func (e *endpoint) sendNDPNS(srcAddr, dstAddr, targetAddr tcpip.Address, remoteL } sent := e.stats.icmp.packetsSent - err := e.nic.WritePacketToRemote(remoteLinkAddr, ProtocolNumber, pkt) + err := e.nic.WritePacketToRemote(remoteLinkAddr, pkt) if err != nil { sent.dropped.Increment() } else { diff --git a/pkg/tcpip/stack/forwarding_test.go b/pkg/tcpip/stack/forwarding_test.go index f6c2d4bc2..16254796b 100644 --- a/pkg/tcpip/stack/forwarding_test.go +++ b/pkg/tcpip/stack/forwarding_test.go @@ -126,13 +126,9 @@ func (f *fwdTestNetworkEndpoint) WritePacket(r *Route, params NetworkHeaderParam b[dstAddrOffset] = r.RemoteAddress()[0] b[srcAddrOffset] = r.LocalAddress()[0] b[protocolNumberOffset] = byte(params.Protocol) + pkt.NetworkProtocolNumber = fwdTestNetNumber - return f.nic.WritePacket(r, fwdTestNetNumber, pkt) -} - -// WritePackets implements LinkEndpoint.WritePackets. -func (*fwdTestNetworkEndpoint) WritePackets(*Route, PacketBufferList, NetworkHeaderParams) (int, tcpip.Error) { - panic("not implemented") + return f.nic.WritePacket(r, pkt) } func (f *fwdTestNetworkEndpoint) WriteHeaderIncludedPacket(r *Route, pkt *PacketBuffer) tcpip.Error { @@ -140,8 +136,9 @@ func (f *fwdTestNetworkEndpoint) WriteHeaderIncludedPacket(r *Route, pkt *Packet if _, ok := pkt.NetworkHeader().Consume(fwdTestNetHeaderLen); !ok { return &tcpip.ErrMalformedHeader{} } + pkt.NetworkProtocolNumber = fwdTestNetNumber - return f.nic.WritePacket(r, fwdTestNetNumber, pkt) + return f.nic.WritePacket(r, pkt) } func (f *fwdTestNetworkEndpoint) Close() { @@ -250,13 +247,6 @@ func (f *fwdTestNetworkEndpoint) SetForwarding(v bool) { f.mu.forwarding = v } -// fwdTestPacketInfo holds all the information about an outbound packet. -type fwdTestPacketInfo struct { - RemoteLinkAddress tcpip.LinkAddress - LocalLinkAddress tcpip.LinkAddress - Pkt *PacketBuffer -} - var _ LinkEndpoint = (*fwdTestLinkEndpoint)(nil) type fwdTestLinkEndpoint struct { @@ -265,7 +255,7 @@ type fwdTestLinkEndpoint struct { linkAddr tcpip.LinkAddress // C is where outbound packets are queued. - C chan fwdTestPacketInfo + C chan *PacketBuffer } // InjectInbound injects an inbound packet. @@ -313,17 +303,11 @@ func (e *fwdTestLinkEndpoint) LinkAddress() tcpip.LinkAddress { } // WritePackets stores outbound packets into the channel. -func (e *fwdTestLinkEndpoint) WritePackets(r RouteInfo, pkts PacketBufferList, protocol tcpip.NetworkProtocolNumber) (int, tcpip.Error) { +func (e *fwdTestLinkEndpoint) WritePackets(pkts PacketBufferList) (int, tcpip.Error) { n := 0 for pkt := pkts.Front(); pkt != nil; pkt = pkt.Next() { - p := fwdTestPacketInfo{ - RemoteLinkAddress: r.RemoteLinkAddress, - LocalLinkAddress: r.LocalLinkAddress, - Pkt: pkt, - } - select { - case e.C <- p: + case e.C <- pkt: default: } @@ -368,7 +352,7 @@ func fwdTestNetFactory(t *testing.T, proto *fwdTestNetworkProtocol) (*faketime.M // NIC 1 has the link address "a", and added the network address 1. ep1 := &fwdTestLinkEndpoint{ - C: make(chan fwdTestPacketInfo, 300), + C: make(chan *PacketBuffer, 300), mtu: fwdTestNetDefaultMTU, linkAddr: "a", } @@ -388,7 +372,7 @@ func fwdTestNetFactory(t *testing.T, proto *fwdTestNetworkProtocol) (*faketime.M // NIC 2 has the link address "b", and added the network address 2. ep2 := &fwdTestLinkEndpoint{ - C: make(chan fwdTestPacketInfo, 300), + C: make(chan *PacketBuffer, 300), mtu: fwdTestNetDefaultMTU, linkAddr: "b", } @@ -450,7 +434,7 @@ func TestForwardingWithStaticResolver(t *testing.T) { Data: buf.ToVectorisedView(), })) - var p fwdTestPacketInfo + var p *PacketBuffer clock.Advance(proto.addrResolveDelay) select { @@ -460,11 +444,11 @@ func TestForwardingWithStaticResolver(t *testing.T) { } // Test that the static address resolution happened correctly. - if p.RemoteLinkAddress != "c" { - t.Fatalf("got p.RemoteLinkAddress = %s, want = c", p.RemoteLinkAddress) + if p.EgressRoute.RemoteLinkAddress != "c" { + t.Fatalf("got p.EgressRoute.RemoteLinkAddress = %s, want = c", p.EgressRoute.RemoteLinkAddress) } - if p.LocalLinkAddress != "b" { - t.Fatalf("got p.LocalLinkAddress = %s, want = b", p.LocalLinkAddress) + if p.EgressRoute.LocalLinkAddress != "b" { + t.Fatalf("got p.EgressRoute.LocalLinkAddress = %s, want = b", p.EgressRoute.LocalLinkAddress) } } @@ -494,7 +478,7 @@ func TestForwardingWithFakeResolver(t *testing.T) { Data: buf.ToVectorisedView(), })) - var p fwdTestPacketInfo + var p *PacketBuffer clock.Advance(proto.addrResolveDelay) select { @@ -504,11 +488,11 @@ func TestForwardingWithFakeResolver(t *testing.T) { } // Test that the address resolution happened correctly. - if p.RemoteLinkAddress != "c" { - t.Fatalf("got p.RemoteLinkAddress = %s, want = c", p.RemoteLinkAddress) + if p.EgressRoute.RemoteLinkAddress != "c" { + t.Fatalf("got p.EgressRoute.RemoteLinkAddress = %s, want = c", p.EgressRoute.RemoteLinkAddress) } - if p.LocalLinkAddress != "b" { - t.Fatalf("got p.LocalLinkAddress = %s, want = b", p.LocalLinkAddress) + if p.EgressRoute.LocalLinkAddress != "b" { + t.Fatalf("got p.EgressRoute.LocalLinkAddress = %s, want = b", p.EgressRoute.LocalLinkAddress) } } @@ -605,7 +589,7 @@ func TestForwardingWithFakeResolverPartialTimeout(t *testing.T) { Data: buf.ToVectorisedView(), })) - var p fwdTestPacketInfo + var p *PacketBuffer clock.Advance(proto.addrResolveDelay) select { @@ -614,16 +598,16 @@ func TestForwardingWithFakeResolverPartialTimeout(t *testing.T) { t.Fatal("packet not forwarded") } - if nh := PayloadSince(p.Pkt.NetworkHeader()); nh[dstAddrOffset] != 3 { - t.Fatalf("got p.Pkt.NetworkHeader[dstAddrOffset] = %d, want = 3", nh[dstAddrOffset]) + if nh := PayloadSince(p.NetworkHeader()); nh[dstAddrOffset] != 3 { + t.Fatalf("got p.NetworkHeader[dstAddrOffset] = %d, want = 3", nh[dstAddrOffset]) } // Test that the address resolution happened correctly. - if p.RemoteLinkAddress != "c" { - t.Fatalf("got p.RemoteLinkAddress = %s, want = c", p.RemoteLinkAddress) + if p.EgressRoute.RemoteLinkAddress != "c" { + t.Fatalf("got p.EgressRoute.RemoteLinkAddress = %s, want = c", p.EgressRoute.RemoteLinkAddress) } - if p.LocalLinkAddress != "b" { - t.Fatalf("got p.LocalLinkAddress = %s, want = b", p.LocalLinkAddress) + if p.EgressRoute.LocalLinkAddress != "b" { + t.Fatalf("got p.EgressRoute.LocalLinkAddress = %s, want = b", p.EgressRoute.LocalLinkAddress) } } @@ -655,7 +639,7 @@ func TestForwardingWithFakeResolverTwoPackets(t *testing.T) { } for i := 0; i < 2; i++ { - var p fwdTestPacketInfo + var p *PacketBuffer clock.Advance(proto.addrResolveDelay) select { @@ -664,16 +648,16 @@ func TestForwardingWithFakeResolverTwoPackets(t *testing.T) { t.Fatal("packet not forwarded") } - if nh := PayloadSince(p.Pkt.NetworkHeader()); nh[dstAddrOffset] != 3 { - t.Fatalf("got p.Pkt.NetworkHeader[dstAddrOffset] = %d, want = 3", nh[dstAddrOffset]) + if nh := PayloadSince(p.NetworkHeader()); nh[dstAddrOffset] != 3 { + t.Fatalf("got p.NetworkHeader[dstAddrOffset] = %d, want = 3", nh[dstAddrOffset]) } // Test that the address resolution happened correctly. - if p.RemoteLinkAddress != "c" { - t.Fatalf("got p.RemoteLinkAddress = %s, want = c", p.RemoteLinkAddress) + if p.EgressRoute.RemoteLinkAddress != "c" { + t.Fatalf("got p.EgressRoute.RemoteLinkAddress = %s, want = c", p.EgressRoute.RemoteLinkAddress) } - if p.LocalLinkAddress != "b" { - t.Fatalf("got p.LocalLinkAddress = %s, want = b", p.LocalLinkAddress) + if p.EgressRoute.LocalLinkAddress != "b" { + t.Fatalf("got p.EgressRoute.LocalLinkAddress = %s, want = b", p.EgressRoute.LocalLinkAddress) } } } @@ -708,7 +692,7 @@ func TestForwardingWithFakeResolverManyPackets(t *testing.T) { } for i := 0; i < maxPendingPacketsPerResolution; i++ { - var p fwdTestPacketInfo + var p *PacketBuffer clock.Advance(proto.addrResolveDelay) select { @@ -717,7 +701,7 @@ func TestForwardingWithFakeResolverManyPackets(t *testing.T) { t.Fatal("packet not forwarded") } - b := PayloadSince(p.Pkt.NetworkHeader()) + b := PayloadSince(p.NetworkHeader()) if b[dstAddrOffset] != 3 { t.Fatalf("got b[dstAddrOffset] = %d, want = 3", b[dstAddrOffset]) } @@ -734,11 +718,11 @@ func TestForwardingWithFakeResolverManyPackets(t *testing.T) { } // Test that the address resolution happened correctly. - if p.RemoteLinkAddress != "c" { - t.Fatalf("got p.RemoteLinkAddress = %s, want = c", p.RemoteLinkAddress) + if p.EgressRoute.RemoteLinkAddress != "c" { + t.Fatalf("got p.EgressRoute.RemoteLinkAddress = %s, want = c", p.EgressRoute.RemoteLinkAddress) } - if p.LocalLinkAddress != "b" { - t.Fatalf("got p.LocalLinkAddress = %s, want = b", p.LocalLinkAddress) + if p.EgressRoute.LocalLinkAddress != "b" { + t.Fatalf("got p.EgressRoute.LocalLinkAddress = %s, want = b", p.EgressRoute.LocalLinkAddress) } } } @@ -773,7 +757,7 @@ func TestForwardingWithFakeResolverManyResolutions(t *testing.T) { } for i := 0; i < maxPendingResolutions; i++ { - var p fwdTestPacketInfo + var p *PacketBuffer clock.Advance(proto.addrResolveDelay) select { @@ -784,16 +768,16 @@ func TestForwardingWithFakeResolverManyResolutions(t *testing.T) { // The first 5 packets (address 3 to 7) should not be forwarded // because their address resolutions are interrupted. - if nh := PayloadSince(p.Pkt.NetworkHeader()); nh[dstAddrOffset] < 8 { - t.Fatalf("got p.Pkt.NetworkHeader[dstAddrOffset] = %d, want p.Pkt.NetworkHeader[dstAddrOffset] >= 8", nh[dstAddrOffset]) + if nh := PayloadSince(p.NetworkHeader()); nh[dstAddrOffset] < 8 { + t.Fatalf("got p.NetworkHeader[dstAddrOffset] = %d, want p.NetworkHeader[dstAddrOffset] >= 8", nh[dstAddrOffset]) } // Test that the address resolution happened correctly. - if p.RemoteLinkAddress != "c" { - t.Fatalf("got p.RemoteLinkAddress = %s, want = c", p.RemoteLinkAddress) + if p.EgressRoute.RemoteLinkAddress != "c" { + t.Fatalf("got p.EgressRoute.RemoteLinkAddress = %s, want = c", p.EgressRoute.RemoteLinkAddress) } - if p.LocalLinkAddress != "b" { - t.Fatalf("got p.LocalLinkAddress = %s, want = b", p.LocalLinkAddress) + if p.EgressRoute.LocalLinkAddress != "b" { + t.Fatalf("got p.EgressRoute.LocalLinkAddress = %s, want = b", p.EgressRoute.LocalLinkAddress) } } } diff --git a/pkg/tcpip/stack/nic.go b/pkg/tcpip/stack/nic.go index 93bc6541f..bea88be8f 100644 --- a/pkg/tcpip/stack/nic.go +++ b/pkg/tcpip/stack/nic.go @@ -146,10 +146,10 @@ 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 { +func (qDisc *delegatingQueueingDiscipline) WritePacket(pkt *PacketBuffer) tcpip.Error { var pkts PacketBufferList pkts.PushBack(pkt) - _, err := qDisc.LinkWriter.WritePackets(r, pkts, protocol) + _, err := qDisc.LinkWriter.WritePackets(pkts) return err } @@ -347,11 +347,12 @@ func (n *nic) WriteRawPacket(pkt *PacketBuffer) tcpip.Error { } // WritePacket implements NetworkEndpoint. -func (n *nic) WritePacket(r *Route, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) tcpip.Error { +func (n *nic) WritePacket(r *Route, pkt *PacketBuffer) tcpip.Error { routeInfo, _, err := r.resolvedFields(nil) switch err.(type) { case nil: - return n.writePacket(routeInfo, protocol, pkt) + pkt.EgressRoute = routeInfo + return n.writePacket(pkt) case *tcpip.ErrWouldBlock: // As per relevant RFCs, we should queue packets while we wait for link // resolution to complete. @@ -370,29 +371,25 @@ func (n *nic) WritePacket(r *Route, protocol tcpip.NetworkProtocolNumber, pkt *P // SHOULD be limited to some small value. When a queue overflows, the new // arrival SHOULD replace the oldest entry. Once address resolution // completes, the node transmits any queued packets. - return n.linkResQueue.enqueue(r, protocol, pkt) + return n.linkResQueue.enqueue(r, pkt) default: return err } } // WritePacketToRemote implements NetworkInterface. -func (n *nic) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) tcpip.Error { - var r RouteInfo - r.NetProto = protocol - r.RemoteLinkAddress = remoteLinkAddr - return n.writePacket(r, protocol, pkt) +func (n *nic) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt *PacketBuffer) tcpip.Error { + pkt.EgressRoute = RouteInfo{routeInfo: routeInfo{NetProto: pkt.NetworkProtocolNumber}, RemoteLinkAddress: remoteLinkAddr} + return n.writePacket(pkt) } -func (n *nic) writePacket(r RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) tcpip.Error { +func (n *nic) writePacket(pkt *PacketBuffer) tcpip.Error { // WritePacket modifies pkt, calculate numBytes first. numBytes := pkt.Size() - pkt.EgressRoute = r - pkt.NetworkProtocolNumber = protocol - n.deliverOutboundPacket(r.RemoteLinkAddress, pkt) + n.deliverOutboundPacket(pkt.EgressRoute.RemoteLinkAddress, pkt) - if err := n.qDisc.WritePacket(r, protocol, pkt); err != nil { + if err := n.qDisc.WritePacket(pkt); err != nil { return err } diff --git a/pkg/tcpip/stack/pending_packets.go b/pkg/tcpip/stack/pending_packets.go index 358e24de4..434355d00 100644 --- a/pkg/tcpip/stack/pending_packets.go +++ b/pkg/tcpip/stack/pending_packets.go @@ -30,7 +30,6 @@ const ( type pendingPacket struct { routeInfo RouteInfo - proto tcpip.NetworkProtocolNumber pkt *PacketBuffer } @@ -56,10 +55,10 @@ type packetsPendingLinkResolution struct { } } -func (f *packetsPendingLinkResolution) incrementOutgoingPacketErrors(proto tcpip.NetworkProtocolNumber, pkt *PacketBuffer) { +func (f *packetsPendingLinkResolution) incrementOutgoingPacketErrors(pkt *PacketBuffer) { f.nic.stack.stats.IP.OutgoingPacketErrors.Increment() - if ipEndpointStats, ok := f.nic.getNetworkEndpoint(proto).Stats().(IPNetworkEndpointStats); ok { + if ipEndpointStats, ok := f.nic.getNetworkEndpoint(pkt.NetworkProtocolNumber).Stats().(IPNetworkEndpointStats); ok { ipEndpointStats.IPStats().OutgoingPacketErrors.Increment() } } @@ -101,7 +100,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 *PacketBuffer) tcpip.Error { +func (f *packetsPendingLinkResolution) enqueue(r *Route, 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. @@ -119,7 +118,8 @@ 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.writePacket(routeInfo, proto, pkt) + pkt.EgressRoute = routeInfo + return f.nic.writePacket(pkt) case *tcpip.ErrWouldBlock: // We need to wait for link resolution to complete. default: @@ -132,13 +132,12 @@ func (f *packetsPendingLinkResolution) enqueue(r *Route, proto tcpip.NetworkProt packets, ok := f.mu.packets[ch] packets = append(packets, pendingPacket{ routeInfo: routeInfo, - proto: proto, pkt: pkt, }) pkt.IncRef() if len(packets) > maxPendingPacketsPerResolution { - f.incrementOutgoingPacketErrors(packets[0].proto, packets[0].pkt) + f.incrementOutgoingPacketErrors(packets[0].pkt) packets[0] = pendingPacket{} packets = packets[1:] @@ -193,11 +192,12 @@ func (f *packetsPendingLinkResolution) dequeuePackets(packets []pendingPacket, l for _, p := range packets { if err == nil { p.routeInfo.RemoteLinkAddress = linkAddr - _ = f.nic.writePacket(p.routeInfo, p.proto, p.pkt) + p.pkt.EgressRoute = p.routeInfo + _ = f.nic.writePacket(p.pkt) } else { - f.incrementOutgoingPacketErrors(p.proto, p.pkt) + f.incrementOutgoingPacketErrors(p.pkt) - if linkResolvableEP, ok := f.nic.getNetworkEndpoint(p.proto).(LinkResolvableNetworkEndpoint); ok { + if linkResolvableEP, ok := f.nic.getNetworkEndpoint(p.pkt.NetworkProtocolNumber).(LinkResolvableNetworkEndpoint); ok { linkResolvableEP.HandleLinkResolutionFailure(p.pkt) } } diff --git a/pkg/tcpip/stack/registration.go b/pkg/tcpip/stack/registration.go index 19ee098f8..34e0cc55c 100644 --- a/pkg/tcpip/stack/registration.go +++ b/pkg/tcpip/stack/registration.go @@ -567,14 +567,13 @@ type NetworkInterface interface { CheckLocalAddress(tcpip.NetworkProtocolNumber, tcpip.Address) bool // WritePacketToRemote writes the packet to the given remote link address. - WritePacketToRemote(tcpip.LinkAddress, tcpip.NetworkProtocolNumber, *PacketBuffer) tcpip.Error + WritePacketToRemote(tcpip.LinkAddress, *PacketBuffer) tcpip.Error - // WritePacket writes a packet with the given protocol through the given - // route. + // WritePacket writes a packet through the given route. // // WritePacket may modify the packet buffer. The packet buffer's // network and transport header must be set. - WritePacket(*Route, tcpip.NetworkProtocolNumber, *PacketBuffer) tcpip.Error + WritePacket(*Route, *PacketBuffer) tcpip.Error // HandleNeighborProbe processes an incoming neighbor probe (e.g. ARP // request or NDP Neighbor Solicitation). @@ -764,12 +763,12 @@ const ( // layer endpoint. It is used with QueueingDiscipline to batch writes from // upper layer endpoints. type LinkWriter interface { - // WritePackets writes packets with the given protocol and route. Must not be - // called with an empty list of packet buffers. + // WritePackets writes packets. Must not be called with an empty list of + // packet buffers. // // 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) + WritePackets(PacketBufferList) (int, tcpip.Error) } // LinkRawWriter is an interface that must be implemented by all Link endpoints @@ -840,15 +839,15 @@ type NetworkLinkEndpoint interface { // QueueingDiscipline provides a queueing strategy for outgoing packets (e.g // FIFO, LIFO, Random Early Drop etc). type QueueingDiscipline interface { - // WritePacket writes a packet with the given protocol and route. + // WritePacket writes a packet. // // 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 + // pkg.EgressRoute.LocalLinkAddress if it is provided. + WritePacket(*PacketBuffer) tcpip.Error Close() } diff --git a/pkg/tcpip/stack/stack.go b/pkg/tcpip/stack/stack.go index f751b484f..04bbce39e 100644 --- a/pkg/tcpip/stack/stack.go +++ b/pkg/tcpip/stack/stack.go @@ -1625,7 +1625,7 @@ func (s *Stack) WritePacketToRemote(nicID tcpip.NICID, remote tcpip.LinkAddress, }) defer pkt.DecRef() pkt.NetworkProtocolNumber = netProto - return nic.WritePacketToRemote(remote, netProto, pkt) + return nic.WritePacketToRemote(remote, pkt) } // WriteRawPacket writes data directly to the specified NIC without adding any diff --git a/pkg/tcpip/stack/stack_test.go b/pkg/tcpip/stack/stack_test.go index 1e1d5b66e..1e0c20701 100644 --- a/pkg/tcpip/stack/stack_test.go +++ b/pkg/tcpip/stack/stack_test.go @@ -196,7 +196,7 @@ func (f *fakeNetworkEndpoint) WritePacket(r *stack.Route, params stack.NetworkHe return nil } - return f.nic.WritePacket(r, fakeNetNumber, pkt) + return f.nic.WritePacket(r, pkt) } // WritePackets implements stack.LinkEndpoint.WritePackets.