From f6bc73ee6ca471c6d3330dbaa00f59e0cd05f3b8 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Mon, 22 Dec 2025 18:41:14 -0300 Subject: [PATCH] apply some of @ddirect suggestions --- arp/handler.go | 14 ++----- examples/bridge/main.go | 6 +-- examples/stack/main.go | 4 +- internal/ip.go | 11 ++++++ internet/definitions.go | 15 ++------ internet/stack-ports.go | 85 ++++++++++++++++++++++++++++++++++++----- x/xnet/stack-async.go | 10 ++--- 7 files changed, 104 insertions(+), 41 deletions(-) diff --git a/arp/handler.go b/arp/handler.go index aca916b..4796735 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -7,6 +7,7 @@ import ( "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" ) type Handler struct { @@ -143,7 +144,7 @@ func (h *Handler) StartQuery(dstHWAddr, proto []byte) error { return errors.New("bad protocol address length") } else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) { return errors.New("mismatch hardware size") - } else if dstHWAddr != nil && !allZeros(dstHWAddr) { + } else if dstHWAddr != nil && !internal.IsZeroed(dstHWAddr...) { return errors.New("write-to buffer must be zeroed out") } h.queries = h.queries[:len(h.queries)+1] @@ -238,7 +239,7 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error { if mac == nil && bytes.Equal(q.protoaddr, protoaddr) { q.hwaddr = append(q.hwaddr, hwaddr...) if q.dstHw != nil { - if !allZeros(q.dstHw) { + if !internal.IsZeroed(q.dstHw...) { slog.Error("race-condition:ARP-reused-buffer") } copy(q.dstHw, hwaddr) // External write to user buffer. @@ -258,12 +259,3 @@ func trySetEthernetDst(ethFrame []byte, dst []byte) { copy(ethFrame[:6], dst) } } - -func allZeros(b []byte) bool { - for i := range b { - if b[i] != 0 { - return false - } - } - return true -} diff --git a/examples/bridge/main.go b/examples/bridge/main.go index b8ed5d2..8a423d1 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -359,7 +359,7 @@ func (s *Stack) StartLookupIP(host string) error { var u internet.StackUDPPort dns4 := dnsSrvs.As4() u.SetStackNode(&s.dns, dns4[:], dns.ServerPort) - err = s.udps.Register(&u, nil) + err = s.udps.Register(&u) if err != nil { return err } @@ -405,7 +405,7 @@ func (s *Stack) BeginDHCPRequest(request [4]byte) error { } var u internet.StackUDPPort u.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort) - err = s.udps.Register(&u, nil) + err = s.udps.Register(&u) if err != nil { return err } @@ -417,7 +417,7 @@ func (s *Stack) StartNTP(addr netip.Addr) error { var u internet.StackUDPPort addr4 := addr.As4() u.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort) - err := s.udps.Register(&u, nil) + err := s.udps.Register(&u) return err } diff --git a/examples/stack/main.go b/examples/stack/main.go index b248202..9370c3e 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -239,7 +239,7 @@ func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, err if err != nil { return nil, err } - err = stack.tcpports.Register(&listener, nil) // Passive TCP requires no MAC setting. + err = stack.tcpports.Register(&listener) // Passive TCP requires no MAC setting. if err != nil { return nil, err } @@ -261,7 +261,7 @@ func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*tcp.Conn, error if err != nil { return nil, err } - err = stack.tcpports.Register(conn, nil) // Passive MAC with no listening. + err = stack.tcpports.Register(conn) // Passive MAC with no listening. if err != nil { return nil, err } diff --git a/internal/ip.go b/internal/ip.go index ac86109..77a9263 100644 --- a/internal/ip.go +++ b/internal/ip.go @@ -56,3 +56,14 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) { copy(dstaddr, dst) return nil } + +// IsZeroed returns true if all arguments are set to their zero value. +func IsZeroed[T comparable](a ...T) bool { + var z T + for i := range a { + if a[i] != z { + return false + } + } + return true +} diff --git a/internet/definitions.go b/internet/definitions.go index 47e0ee1..8246f97 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -165,6 +165,7 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro err := node.demux(buf, offset) if h.tryHandleError(node, err) { err = nil + node = nil // Node is destroyed in tryHandleError and invalidated. } return node, err } @@ -174,12 +175,13 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ *node, n int, err error) { for i := range h.nodes { node := &h.nodes[i] - if node.IsInvalid() || (len(node.remoteAddr) > 0 && isAllZeros(node.remoteAddr)) { + if node.IsInvalid() { continue } n, err = node.encapsulate(buf, offsetIP, offsetThisFrame) if h.tryHandleError(node, err) { - err = nil // CLOSE error handled gracefully by deleting node. + err = nil // CLOSE error handled gracefully by deleting node. + node = nil // Node is destroyed in tryHandleError and invalidated. } if n > 0 { return node, n, err @@ -191,15 +193,6 @@ func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ return nil, 0, err // Return last written error. } -func isAllZeros(b []byte) bool { - for i := range b { - if b[i] != 0 { - return false - } - } - return true -} - var ( errZeroMaxNodesArg = errors.New("zero max nodes arg") errZeroPort = errors.New("port must be greater than zero") diff --git a/internet/stack-ports.go b/internet/stack-ports.go index d98faa6..413cc09 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -4,11 +4,13 @@ import ( "encoding/binary" "errors" "io" + "log/slog" "math" "strconv" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" ) type StackPorts struct { @@ -16,6 +18,7 @@ type StackPorts struct { handlers handlers dstPortOff uint16 protocol uint16 + // stores last node to demux/encapsulate. } func (ps *StackPorts) ResetUDP(maxNodes int) error { @@ -52,12 +55,7 @@ func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) { return 0, io.ErrShortBuffer } - var node *node - node, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) - if n > 0 && len(node.remoteAddr) == 6 && offsetToIP >= 14 { - efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:]) - *efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr) - } + _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) return n, err } @@ -72,15 +70,84 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { // Register registers a port StackNode on StackPorts. // If dstMAC is set to non-nil, length six buffer then -func (ps *StackPorts) Register(h StackNode, dstMAC []byte) error { +func (ps *StackPorts) Register(h StackNode) error { port := h.LocalPort() proto := h.Protocol() if port <= 0 { return errZeroPort } else if proto != uint64(ps.protocol) { return errInvalidProto - } else if dstMAC != nil && len(dstMAC) != 6 { + } + return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil)) +} + +// StackPortsMACFiltered is a StackPorts implementation but that avoids calling encapsulate on nodes +// with a non-nil MAC address registered via Register method that is set to all zero values. +// If the address is set to nil no filtering occurs. MAC Address is set automatically on the ethernet frame by StackPortsMACFiltered when non-nil. +type StackPortsMACFiltered struct { + sp StackPorts +} + +func (mfsp *StackPortsMACFiltered) Register(h StackNode, addr []byte) error { + port := h.LocalPort() + proto := h.Protocol() + if port <= 0 { + return errZeroPort + } else if proto != uint64(mfsp.sp.protocol) { + return errInvalidProto + } else if addr != nil && len(addr) != 6 { return errors.New("invalid MAC") } - return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, dstMAC)) + return mfsp.sp.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, addr)) +} + +func (ps *StackPortsMACFiltered) ResetUDP(maxNodes int) error { + return ps.sp.ResetUDP(maxNodes) +} + +func (ps *StackPortsMACFiltered) ResetTCP(maxNodes int) error { + return ps.sp.ResetTCP(maxNodes) +} + +func (ps *StackPortsMACFiltered) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) error { + return ps.sp.Reset(protocol, dstPortOffset, maxNodes) +} + +func (ps *StackPortsMACFiltered) LocalPort() uint16 { return 0 } + +func (ps *StackPortsMACFiltered) Protocol() uint64 { return uint64(ps.sp.protocol) } + +func (ps *StackPortsMACFiltered) ConnectionID() *uint64 { return &ps.sp.connID } + +func (ps *StackPortsMACFiltered) Demux(b []byte, offset int) (err error) { + // No MAC Filtering on ingress. TODO? + return ps.sp.Demux(b, offset) +} + +func (ps *StackPortsMACFiltered) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { + if int(ps.sp.dstPortOff)+offsetToFrame+2 > len(carrierData) { + return 0, io.ErrShortBuffer + } + h := &ps.sp.handlers + for i := range h.nodes { + node := &h.nodes[i] + if node.IsInvalid() || (len(node.remoteAddr) > 0 && internal.IsZeroed(node.remoteAddr...)) { + continue + } + n, err = node.encapsulate(carrierData, offsetToIP, offsetToFrame) + if h.tryHandleError(node, err) { + err = nil // CLOSE error handled gracefully by deleting node. + } + if n > 0 { + if len(node.remoteAddr) == 6 && offsetToIP >= 14 { + efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:]) + *efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr) + } + return n, err + } else if err != nil { + // Make sure not to hang on one handler that keeps returning an error. + h.error("handlers:encapsulate", slog.String("func", "encapsulateAny"), slog.String("ctx", h.context), slog.String("err", err.Error())) + } + } + return 0, err // Return last written error. } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 13e604e..9a3c5b6 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -250,7 +250,7 @@ func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrP if err != nil { return err } - err = s.tcps.Register(conn, mac) // MAC is set later on by ARP response arriving to our network. + err = s.tcps.Register(conn) // MAC is set later on by ARP response arriving to our network. if err != nil { conn.Abort() return err @@ -265,7 +265,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) { if err != nil { return err } - err = s.tcps.Register(conn, nil) + err = s.tcps.Register(conn) if err != nil { conn.Abort() return err @@ -306,7 +306,7 @@ func (s *StackAsync) StartLookupIP(host string) error { } dns4 := s.dnssv.As4() s.dnsUDP.SetStackNode(&s.dns, dns4[:], dns.ServerPort) - err = s.udps.Register(&s.dnsUDP, nil) + err = s.udps.Register(&s.dnsUDP) return err } @@ -353,7 +353,7 @@ func (s *StackAsync) StartDHCPv4Request(request [4]byte) error { } s.dhcpUDP.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort) - err = s.udps.Register(&s.dhcpUDP, nil) + err = s.udps.Register(&s.dhcpUDP) if err != nil { return err } @@ -367,7 +367,7 @@ func (s *StackAsync) StartNTP(addr netip.Addr) error { addr4 := addr.As4() s.ntpUDP.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort) - err := s.udps.Register(&s.ntpUDP, nil) + err := s.udps.Register(&s.ntpUDP) return err }