From ad021f48c08e78a2dbd2223525fbb80990cfaf0f Mon Sep 17 00:00:00 2001 From: Ghanan Gowripalan Date: Wed, 26 Jan 2022 14:27:06 -0800 Subject: [PATCH] Add link-layer headers in nic This removes the need for the stack to add a link header out-of-line the write path when delivering outbound packets to a packet socket. PiperOrigin-RevId: 424444109 --- pkg/tcpip/link/ethernet/ethernet.go | 16 ---- pkg/tcpip/link/ethernet/ethernet_test.go | 24 ++---- pkg/tcpip/link/fdbased/endpoint.go | 12 +-- pkg/tcpip/link/fdbased/endpoint_test.go | 20 ++--- pkg/tcpip/link/sharedmem/sharedmem.go | 7 +- pkg/tcpip/link/sharedmem/sharedmem_server.go | 8 +- pkg/tcpip/link/sharedmem/sharedmem_test.go | 64 +++++++--------- pkg/tcpip/link/tun/device.go | 15 +--- pkg/tcpip/network/ipv6/icmp_test.go | 4 + pkg/tcpip/network/ipv6/ndp_test.go | 1 + pkg/tcpip/stack/forwarding_test.go | 1 - pkg/tcpip/stack/nic.go | 77 ++++++-------------- 12 files changed, 85 insertions(+), 164 deletions(-) diff --git a/pkg/tcpip/link/ethernet/ethernet.go b/pkg/tcpip/link/ethernet/ethernet.go index 7c4529ea1..715d45d37 100644 --- a/pkg/tcpip/link/ethernet/ethernet.go +++ b/pkg/tcpip/link/ethernet/ethernet.go @@ -79,17 +79,6 @@ func (e *Endpoint) Capabilities() stack.LinkEndpointCapabilities { return c } -// WritePackets implements stack.LinkEndpoint. -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(pkts) -} - // MaxHeaderLength implements stack.LinkEndpoint. func (e *Endpoint) MaxHeaderLength() uint16 { return header.EthernetMinimumSize + e.Endpoint.MaxHeaderLength() @@ -113,8 +102,3 @@ func (*Endpoint) AddHeader(local, remote tcpip.LinkAddress, proto tcpip.NetworkP } eth.Encode(&fields) } - -// WriteRawPacket implements stack.LinkEndpoint. -func (e *Endpoint) WriteRawPacket(pkt *stack.PacketBuffer) tcpip.Error { - return e.Endpoint.WriteRawPacket(pkt) -} diff --git a/pkg/tcpip/link/ethernet/ethernet_test.go b/pkg/tcpip/link/ethernet/ethernet_test.go index a2f124fbc..4ce308a5c 100644 --- a/pkg/tcpip/link/ethernet/ethernet_test.go +++ b/pkg/tcpip/link/ethernet/ethernet_test.go @@ -120,32 +120,24 @@ func TestMTU(t *testing.T) { } } -func TestWritePacketsAddHeader(t *testing.T) { +func TestWritePacketToRemoteAddHeader(t *testing.T) { const ( localLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x06") remoteLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x07") netProto = 55 + nicID = 1 ) c := channel.New(1, header.EthernetMinimumSize, localLinkAddr) - e := ethernet.New(c) - { - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ - ReserveHeaderBytes: int(e.MaxHeaderLength()), - }) - defer pkt.DecRef() - pkt.NetworkProtocolNumber = netProto - pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr + s := stack.New(stack.Options{}) + if err := s.CreateNIC(nicID, ethernet.New(c)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) + } - var pkts stack.PacketBufferList - pkts.PushFront(pkt) - if n, err := e.WritePackets(pkts); err != nil { - t.Fatalf("e.WritePackets(_): %s", err) - } else if n != 1 { - t.Fatalf("got e.WritePackets(_) = %d, want = 1", n) - } + if err := s.WritePacketToRemote(nicID, remoteLinkAddr, netProto, buffer.VectorisedView{}); err != nil { + t.Fatalf("s.WritePacketToRemote(%d, %s, _): %s", nicID, remoteLinkAddr, err) } { diff --git a/pkg/tcpip/link/fdbased/endpoint.go b/pkg/tcpip/link/fdbased/endpoint.go index 129397259..95769c67d 100644 --- a/pkg/tcpip/link/fdbased/endpoint.go +++ b/pkg/tcpip/link/fdbased/endpoint.go @@ -510,11 +510,7 @@ func (*endpoint) WriteRawPacket(*stack.PacketBuffer) tcpip.Error { return &tcpip // writePacket writes outbound packets to the file descriptor. If it is not // currently writable, the packet is dropped. -func (e *endpoint) writePacket(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { - if e.hdrSize > 0 { - e.AddHeader(r.LocalLinkAddress, r.RemoteLinkAddress, protocol, pkt) - } - +func (e *endpoint) writePacket(pkt *stack.PacketBuffer) tcpip.Error { fd := e.fds[pkt.Hash%uint32(len(e.fds))] var vnetHdrBuf []byte if e.gsoKind == stack.HWGSOSupported { @@ -572,10 +568,6 @@ func (e *endpoint) sendBatch(batchFD int, pkts []*stack.PacketBuffer) (int, tcpi batch := pkts[packets:] syscallHeaderBytes := uintptr(0) for _, pkt := range batch { - if e.hdrSize > 0 { - e.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) - } - var vnetHdrBuf []byte if e.gsoKind == stack.HWGSOSupported { vnetHdr := virtioNetHdr{} @@ -641,7 +633,7 @@ func (e *endpoint) sendBatch(batchFD int, pkts []*stack.PacketBuffer) (int, tcpi // if necessary (by using e.writevMaxIovs instead of // rawfile.MaxIovs). pkt := batch[0] - if err := e.writePacket(pkt.EgressRoute, pkt.NetworkProtocolNumber, pkt); err != nil { + if err := e.writePacket(pkt); err != nil { return packets, err } packets++ diff --git a/pkg/tcpip/link/fdbased/endpoint_test.go b/pkg/tcpip/link/fdbased/endpoint_test.go index 4586ec2db..f308af69d 100644 --- a/pkg/tcpip/link/fdbased/endpoint_test.go +++ b/pkg/tcpip/link/fdbased/endpoint_test.go @@ -181,9 +181,6 @@ func testWritePacket(t *testing.T, plen int, eth bool, gsoMaxSize uint32, hash u c := newContext(t, &Options{Address: laddr, MTU: mtu, EthernetHeader: eth, GSOMaxSize: gsoMaxSize}) defer c.cleanup() - var r stack.RouteInfo - r.RemoteLinkAddress = raddr - // Build payload. payload := buffer.NewView(plen) if _, err := rand.Read(payload); err != nil { @@ -199,7 +196,8 @@ func testWritePacket(t *testing.T, plen int, eth bool, gsoMaxSize uint32, hash u pkt.Hash = hash // Every PacketBuffer must have these set: // See nic.writePacket. - pkt.EgressRoute = r + pkt.EgressRoute.LocalLinkAddress = laddr + pkt.EgressRoute.RemoteLinkAddress = raddr pkt.NetworkProtocolNumber = proto defer pkt.DecRef() @@ -221,6 +219,9 @@ func testWritePacket(t *testing.T, plen int, eth bool, gsoMaxSize uint32, hash u L3HdrLen: l3HdrLen, } } + + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) + var pkts stack.PacketBufferList pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { @@ -326,11 +327,6 @@ func TestPreserveSrcAddress(t *testing.T) { c := newContext(t, &Options{Address: laddr, MTU: mtu, EthernetHeader: true}) defer c.cleanup() - // Set LocalLinkAddress in route to the value of the bridged address. - var r stack.RouteInfo - r.LocalLinkAddress = baddr - r.RemoteLinkAddress = raddr - pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ // WritePacket panics given a prependable with anything less than // the minimum size of the ethernet header. @@ -342,7 +338,11 @@ func TestPreserveSrcAddress(t *testing.T) { // Every PacketBuffer must have these set: // See nic.writePacket. pkt.NetworkProtocolNumber = proto - pkt.EgressRoute = r + // Set LocalLinkAddress in route to the value of the bridged address. + pkt.EgressRoute.LocalLinkAddress = baddr + pkt.EgressRoute.RemoteLinkAddress = raddr + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) + var pkts stack.PacketBufferList pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { diff --git a/pkg/tcpip/link/sharedmem/sharedmem.go b/pkg/tcpip/link/sharedmem/sharedmem.go index 43eb629b5..808e3ca79 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem.go +++ b/pkg/tcpip/link/sharedmem/sharedmem.go @@ -321,6 +321,10 @@ func (e *endpoint) LinkAddress() tcpip.LinkAddress { // AddHeader implements stack.LinkEndpoint.AddHeader. func (e *endpoint) AddHeader(local, remote tcpip.LinkAddress, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { // Add ethernet header if needed. + if len(e.addr) == 0 { + return + } + eth := header.Ethernet(pkt.LinkHeader().Push(header.EthernetMinimumSize)) ethHdr := &header.EthernetFields{ DstAddr: remote, @@ -346,9 +350,6 @@ func (*endpoint) WriteRawPacket(*stack.PacketBuffer) tcpip.Error { return &tcpip // +checklocks:e.mu func (e *endpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { - if e.addr != "" { - e.AddHeader(r.LocalLinkAddress, r.RemoteLinkAddress, protocol, pkt) - } if e.virtioNetHeaderRequired { e.AddVirtioNetHeader(pkt) } diff --git a/pkg/tcpip/link/sharedmem/sharedmem_server.go b/pkg/tcpip/link/sharedmem/sharedmem_server.go index 6df530fa9..84d1763d6 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_server.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_server.go @@ -207,6 +207,10 @@ func (e *serverEndpoint) LinkAddress() tcpip.LinkAddress { // AddHeader implements stack.LinkEndpoint.AddHeader. func (e *serverEndpoint) AddHeader(local, remote tcpip.LinkAddress, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) { // Add ethernet header if needed. + if len(e.addr) == 0 { + return + } + eth := header.Ethernet(pkt.LinkHeader().Push(header.EthernetMinimumSize)) ethHdr := &header.EthernetFields{ DstAddr: remote, @@ -242,10 +246,6 @@ func (e *serverEndpoint) WriteRawPacket(pkt *stack.PacketBuffer) tcpip.Error { // +checklocks:e.mu func (e *serverEndpoint) writePacketLocked(r stack.RouteInfo, protocol tcpip.NetworkProtocolNumber, pkt *stack.PacketBuffer) tcpip.Error { - if e.addr != "" { - e.AddHeader(r.LocalLinkAddress, r.RemoteLinkAddress, protocol, pkt) - } - if e.virtioNetHeaderRequired { e.AddVirtioNetHeader(pkt) } diff --git a/pkg/tcpip/link/sharedmem/sharedmem_test.go b/pkg/tcpip/link/sharedmem/sharedmem_test.go index f5f4521e6..4beb72e5a 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_test.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_test.go @@ -204,11 +204,6 @@ func TestSimpleSend(t *testing.T) { c := newTestContext(t, 20000, 1500, localLinkAddr) defer c.cleanup() - // Prepare route. - var r stack.RouteInfo - r.RemoteLinkAddress = remoteLinkAddr - r.LocalLinkAddress = localLinkAddr - for iters := 1000; iters > 0; iters-- { func() { hdrLen, dataLen := rand.Intn(10000), rand.Intn(10000) @@ -228,8 +223,10 @@ func TestSimpleSend(t *testing.T) { proto := tcpip.NetworkProtocolNumber(rand.Intn(0x10000)) // Every PacketBuffer must have these set: // See nic.writePacket. - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr + pkt.EgressRoute.LocalLinkAddress = localLinkAddr pkt.NetworkProtocolNumber = proto + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) var pkts stack.PacketBufferList pkts.PushBack(pkt) defer pkts.DecRef() @@ -291,10 +288,6 @@ func TestPreserveSrcAddressInSend(t *testing.T) { defer c.cleanup() newLocalLinkAddress := tcpip.LinkAddress(strings.Repeat("0xFE", 6)) - // Set both remote and local link address in route. - var r stack.RouteInfo - r.LocalLinkAddress = newLocalLinkAddress - r.RemoteLinkAddress = remoteLinkAddr pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ // WritePacket panics given a prependable with anything less than @@ -304,8 +297,10 @@ func TestPreserveSrcAddressInSend(t *testing.T) { proto := tcpip.NetworkProtocolNumber(rand.Intn(0x10000)) // Every PacketBuffer must have these set: // See nic.writePacket. - pkt.EgressRoute = r + pkt.EgressRoute.LocalLinkAddress = newLocalLinkAddress + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = proto + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) var pkts stack.PacketBufferList defer pkts.DecRef() @@ -351,10 +346,6 @@ func TestFillTxQueue(t *testing.T) { c := newTestContext(t, 20000, 1500, localLinkAddr) defer c.cleanup() - // Prepare to send a packet. - var r stack.RouteInfo - r.RemoteLinkAddress = remoteLinkAddr - buf := buffer.NewView(100) // Each packet is uses no more than 40 bytes, so write that many packets @@ -367,8 +358,9 @@ func TestFillTxQueue(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) var pkts stack.PacketBufferList pkts.PushBack(pkt) @@ -392,8 +384,9 @@ func TestFillTxQueue(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) var pkts stack.PacketBufferList pkts.PushBack(pkt) @@ -414,10 +407,6 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { queue.EncodeTxCompletion(c.txq.rx.Push(8), 1) c.txq.rx.Flush() - // Prepare to send a packet. - var r stack.RouteInfo - r.RemoteLinkAddress = remoteLinkAddr - buf := buffer.NewView(100) // Send two packets so that the id slice has at least two slots. @@ -429,8 +418,9 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { Data: buf.ToVectorisedView(), }) pkts.PushBack(pkt) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) } if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) @@ -456,8 +446,10 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) + var pkts stack.PacketBufferList pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { @@ -479,8 +471,10 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) + var pkts stack.PacketBufferList pkts.PushBack(pkt) _, err := c.ep.WritePackets(pkts) @@ -496,10 +490,6 @@ func TestFillTxMemory(t *testing.T) { c := newTestContext(t, 20000, bufferSize, localLinkAddr) defer c.cleanup() - // Prepare to send a packet. - var r stack.RouteInfo - r.RemoteLinkAddress = remoteLinkAddr - buf := buffer.NewView(100) // Each packet is uses up one buffer, so write as many as possible until @@ -510,8 +500,10 @@ func TestFillTxMemory(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) + var pkts stack.PacketBufferList pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { @@ -535,7 +527,7 @@ func TestFillTxMemory(t *testing.T) { Data: buf.ToVectorisedView(), }) pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr var pkts stack.PacketBufferList pkts.PushBack(pkt) _, err := c.ep.WritePackets(pkts) @@ -553,10 +545,6 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { c := newTestContext(t, 20000, bufferSize, localLinkAddr) defer c.cleanup() - // Prepare to send a packet. - var r stack.RouteInfo - r.RemoteLinkAddress = remoteLinkAddr - buf := buffer.NewView(100) // Each packet is uses up one buffer, so write as many as possible @@ -567,7 +555,7 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { Data: buf.ToVectorisedView(), }) var pkts stack.PacketBufferList - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { @@ -587,8 +575,10 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buffer.NewView(bufferSize).ToVectorisedView(), }) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber + c.ep.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) + pkts.PushBack(pkt) _, err := c.ep.WritePackets(pkts) if _, ok := err.(*tcpip.ErrWouldBlock); !ok { @@ -604,7 +594,7 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { ReserveHeaderBytes: int(c.ep.MaxHeaderLength()), Data: buf.ToVectorisedView(), }) - pkt.EgressRoute = r + pkt.EgressRoute.RemoteLinkAddress = remoteLinkAddr pkt.NetworkProtocolNumber = header.IPv4ProtocolNumber pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { diff --git a/pkg/tcpip/link/tun/device.go b/pkg/tcpip/link/tun/device.go index 3181a4070..83e785008 100644 --- a/pkg/tcpip/link/tun/device.go +++ b/pkg/tcpip/link/tun/device.go @@ -266,20 +266,7 @@ func (d *Device) encodePkt(pkt *stack.PacketBuffer) (buffer.View, bool) { vv.AppendView(buffer.View(hdr)) } - // Ethernet header (TAP only). - if d.flags.TAP { - // Add ethernet header if not provided. - if pkt.LinkHeader().View().IsEmpty() { - d.endpoint.AddHeader(pkt.EgressRoute.LocalLinkAddress, pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) - } - vv.AppendView(pkt.LinkHeader().View()) - } - - // Append upper headers. - vv.AppendView(pkt.NetworkHeader().View()) - vv.AppendView(pkt.TransportHeader().View()) - // Append data payload. - vv.Append(pkt.Data().ExtractVV()) + vv.AppendViews(pkt.Views()) return vv.ToView(), true } diff --git a/pkg/tcpip/network/ipv6/icmp_test.go b/pkg/tcpip/network/ipv6/icmp_test.go index 8a21874c3..7e2bd7e04 100644 --- a/pkg/tcpip/network/ipv6/icmp_test.go +++ b/pkg/tcpip/network/ipv6/icmp_test.go @@ -83,6 +83,9 @@ func (*stubLinkEndpoint) WritePackets(pkts stack.PacketBufferList) (int, tcpip.E func (*stubLinkEndpoint) Attach(stack.NetworkDispatcher) {} +func (*stubLinkEndpoint) AddHeader(_, _ tcpip.LinkAddress, _ tcpip.NetworkProtocolNumber, _ *stack.PacketBuffer) { +} + type stubDispatcher struct { stack.TransportDispatcher } @@ -1290,6 +1293,7 @@ func TestLinkAddressRequest(t *testing.T) { var want stack.RouteInfo want.NetProto = ProtocolNumber + want.LocalLinkAddress = linkAddr0 want.RemoteLinkAddress = test.expectedRemoteLinkAddr if diff := cmp.Diff(want, pkt.EgressRoute, cmp.AllowUnexported(want)); diff != "" { t.Errorf("route info mismatch (-want +got):\n%s", diff) diff --git a/pkg/tcpip/network/ipv6/ndp_test.go b/pkg/tcpip/network/ipv6/ndp_test.go index 78c23f900..be616773c 100644 --- a/pkg/tcpip/network/ipv6/ndp_test.go +++ b/pkg/tcpip/network/ipv6/ndp_test.go @@ -473,6 +473,7 @@ func TestNeighborSolicitationResponse(t *testing.T) { respNSDst := header.SolicitedNodeAddr(test.nsSrc) var want stack.RouteInfo want.NetProto = ProtocolNumber + want.LocalLinkAddress = nicLinkAddr want.RemoteLinkAddress = header.EthernetAddressFromMulticastIPv6Address(respNSDst) if diff := cmp.Diff(want, p.EgressRoute, cmp.AllowUnexported(want)); diff != "" { t.Errorf("route info mismatch (-want +got):\n%s", diff) diff --git a/pkg/tcpip/stack/forwarding_test.go b/pkg/tcpip/stack/forwarding_test.go index 5673bbf10..e984e9638 100644 --- a/pkg/tcpip/stack/forwarding_test.go +++ b/pkg/tcpip/stack/forwarding_test.go @@ -331,7 +331,6 @@ func (*fwdTestLinkEndpoint) ARPHardwareType() header.ARPHardwareType { // AddHeader implements stack.LinkEndpoint.AddHeader. func (e *fwdTestLinkEndpoint) AddHeader(tcpip.LinkAddress, tcpip.LinkAddress, tcpip.NetworkProtocolNumber, *PacketBuffer) { - panic("not implemented") } func fwdTestNetFactory(t *testing.T, proto *fwdTestNetworkProtocol) (*faketime.ManualClock, *fwdTestLinkEndpoint, *fwdTestLinkEndpoint) { diff --git a/pkg/tcpip/stack/nic.go b/pkg/tcpip/stack/nic.go index de0eab024..5c6943aac 100644 --- a/pkg/tcpip/stack/nic.go +++ b/pkg/tcpip/stack/nic.go @@ -381,7 +381,13 @@ func (n *nic) WritePacket(r *Route, pkt *PacketBuffer) tcpip.Error { // WritePacketToRemote implements NetworkInterface. func (n *nic) WritePacketToRemote(remoteLinkAddr tcpip.LinkAddress, pkt *PacketBuffer) tcpip.Error { - pkt.EgressRoute = RouteInfo{routeInfo: routeInfo{NetProto: pkt.NetworkProtocolNumber}, RemoteLinkAddress: remoteLinkAddr} + pkt.EgressRoute = RouteInfo{ + routeInfo: routeInfo{ + NetProto: pkt.NetworkProtocolNumber, + LocalLinkAddress: n.LinkAddress(), + }, + RemoteLinkAddress: remoteLinkAddr, + } return n.writePacket(pkt) } @@ -389,7 +395,9 @@ func (n *nic) writePacket(pkt *PacketBuffer) tcpip.Error { // WritePacket modifies pkt, calculate numBytes first. numBytes := pkt.Size() - n.deliverOutboundPacket(pkt.EgressRoute.RemoteLinkAddress, pkt) + n.NetworkLinkEndpoint.AddHeader(n.LinkAddress(), pkt.EgressRoute.RemoteLinkAddress, pkt.NetworkProtocolNumber, pkt) + + n.deliverLinkPacket(pkt.NetworkProtocolNumber, pkt, false /* incoming */) if err := n.qDisc.WritePacket(pkt); err != nil { return err @@ -722,6 +730,12 @@ func (n *nic) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *Pa pkt.RXTransportChecksumValidated = n.NetworkLinkEndpoint.Capabilities()&CapabilityRXChecksumOffload != 0 + n.deliverLinkPacket(protocol, pkt, true /* incoming */) + + networkEndpoint.HandlePacket(pkt) +} + +func (n *nic) deliverLinkPacket(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer, incoming bool) { // Deliver to interested packet endpoints without holding NIC lock. var packetEPPkt *PacketBuffer defer func() { @@ -747,7 +761,12 @@ func (n *nic) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *Pa // populate it in the packet buffer we provide to packet endpoints as // packet endpoints inspect link headers. packetEPPkt.LinkHeader().Consume(pkt.LinkHeader().View().Size()) - packetEPPkt.PktType = tcpip.PacketHost + + if incoming { + packetEPPkt.PktType = tcpip.PacketHost + } else { + packetEPPkt.PktType = tcpip.PacketOutgoing + } } clone := packetEPPkt.Clone() @@ -762,61 +781,13 @@ func (n *nic) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *Pa anyEPs, anyEPsOK := n.packetEPs[header.EthernetProtocolAll] n.packetEPsMu.Unlock() - if protoEPsOK { + // On Linux, only ETH_P_ALL endpoints get outbound packets. + if incoming && protoEPsOK { protoEPs.forEach(deliverPacketEPs) } if anyEPsOK { anyEPs.forEach(deliverPacketEPs) } - - networkEndpoint.HandlePacket(pkt) -} - -// deliverOutboundPacket delivers outgoing packets to interested endpoints. -func (n *nic) deliverOutboundPacket(remote tcpip.LinkAddress, pkt *PacketBuffer) { - n.packetEPsMu.RLock() - defer n.packetEPsMu.RUnlock() - // We do not deliver to protocol specific packet endpoints as on Linux - // only ETH_P_ALL endpoints get outbound packets. - // Add any other packet sockets that maybe listening for all protocols. - eps, ok := n.packetEPs[header.EthernetProtocolAll] - if !ok { - return - } - - local := n.LinkAddress() - - var packetEPPkt *PacketBuffer - defer func() { - if packetEPPkt != nil { - packetEPPkt.DecRef() - } - }() - eps.forEach(func(ep PacketEndpoint) { - if packetEPPkt == nil { - // Packet endpoints hold the full packet. - // - // We perform a deep copy because higher-level endpoints may point to - // the middle of a view that is held by a packet endpoint. Save/Restore - // does not support overlapping slices and will panic in this case. - // - // TODO(https://gvisor.dev/issue/6517): Avoid this copy once S/R supports - // overlapping slices (e.g. by passing a shallow copy of pkt to the packet - // endpoint). - packetEPPkt = NewPacketBuffer(PacketBufferOptions{ - ReserveHeaderBytes: pkt.AvailableHeaderBytes(), - Data: PayloadSince(pkt.NetworkHeader()).ToVectorisedView(), - }) - // Add the link layer header as outgoing packets are intercepted before - // the link layer header is created and packet endpoints are interested - // in the link header. - n.NetworkLinkEndpoint.AddHeader(local, remote, pkt.NetworkProtocolNumber, packetEPPkt) - packetEPPkt.PktType = tcpip.PacketOutgoing - } - clone := packetEPPkt.Clone() - defer clone.DecRef() - ep.HandlePacket(n.id, pkt.NetworkProtocolNumber, clone) - }) } // DeliverTransportPacket delivers the packets to the appropriate transport