From eaa36a589aa90cf7ed1656030880bec67fed9757 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Thu, 18 Dec 2025 19:16:35 -0300 Subject: [PATCH] internet: further use handlers data structure --- internet/definitions.go | 102 ++++++++++++++----------------------- internet/stack-ethernet.go | 37 +++++--------- internet/stack-ip.go | 82 +++++++++++++---------------- internet/stack-ports.go | 13 ++--- 4 files changed, 89 insertions(+), 145 deletions(-) diff --git a/internet/definitions.go b/internet/definitions.go index b4aded0..683d34b 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -2,6 +2,7 @@ package internet import ( "errors" + "log/slog" "math" "net" "slices" @@ -44,11 +45,14 @@ type node struct { } type handlers struct { + context string + logger nodes []node } -func (h *handlers) reset(maxNodes int) { +func (h *handlers) reset(context string, maxNodes int) { h.nodes = slices.Grow(h.nodes[:0], maxNodes) + h.context = context } func (h *handlers) registerByProto(n node) error { @@ -136,22 +140,50 @@ func (h *handlers) nodeByPortProto(port uint16, protocol uint16) *node { return nil } -// encapsulateAny does not add the offset to the amount of bytes written. -func (h *handlers) encapsulateAny(buf []byte, offset int) (*node, int, error) { +func (h *handlers) demuxByProto(buf []byte, offset int, proto uint16) (*node, error) { + node := h.nodeByProto(proto) + if node == nil { + return nil, nil + } + err := node.demux(buf, offset) + if h.tryHandleError(node, err) { + err = nil + } + return node, err +} + +func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, error) { + node := h.nodeByPort(port) + if node == nil { + return nil, nil + } + err := node.demux(buf, offset) + if h.tryHandleError(node, err) { + err = nil + } + return node, err +} + +// encapsulateAny finds a node suitable to write and encapsulates the package. +// If no data is sent it returns the last error encountered. +func (h *handlers) encapsulateAny(buf []byte, offset int) (_ *node, n int, err error) { for i := range h.nodes { node := &h.nodes[i] if node.IsInvalid() { continue } - n, err := node.encapsulate(buf, offset) + n, err = node.encapsulate(buf, offset) if h.tryHandleError(node, err) { err = nil // CLOSE error handled gracefully by deleting node. } - if err != nil || n > 0 { + if n > 0 { return node, 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 nil, 0, nil + return nil, 0, err // Return last written error. } var ( @@ -163,32 +195,6 @@ var ( _ = net.ErrClosed ) -func registerNode(nodesPtr *[]node, h node) error { - if cap(*nodesPtr)-len(*nodesPtr) <= 0 { - *nodesPtr = nodesCompact(*nodesPtr) - } - if cap(*nodesPtr)-len(*nodesPtr) <= 0 { - return errNodesFull - } - *nodesPtr = append(*nodesPtr, h) - return nil -} - -func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) (discarded bool) { - if err != nil { - if nodeIdx >= len(*nodesPtr) { - panic("unreachable") - } - nodes := *nodesPtr - if checkNodeErr(&nodes[nodeIdx], err) { - // *nodesPtr = slices.Delete(nodes, nodeIdx, nodeIdx+1) - (*nodesPtr)[nodeIdx] = node{} // 'Delete' node without modifying slice length. - discarded = true - } - } - return discarded -} - func (node *node) IsInvalid() bool { return node.demux == nil || node.encapsulate == nil || (node.connID != nil && node.currConnID != *node.connID) } @@ -217,39 +223,7 @@ func nodeFromStackNode(s StackNode, port uint16, protocol uint64, remoteAddr []b } } -func getNode(nodes []node, port uint16, protocol uint16) (node *node) { - for i := range nodes { - node := &nodes[i] - if node.port == port && node.proto == protocol { - return node - } - } - return nil -} - // destroy removes all references to underlying StackNode. Allows garbage collection of node if possible. func (n *node) destroy() { *n = node{} } - -func getNodeByProto(nodes []node, protocol uint16) int { - for i := range nodes { - node := &nodes[i] - if node.proto == protocol { - return i - } - } - - return -1 -} - -func nodesCompact(nodes []node) []node { - nilOff := 0 - for i := 0; i < len(nodes); i++ { - if !nodes[i].IsInvalid() { - nodes[nilOff] = nodes[i] - nilOff++ - } - } - return nodes[:nilOff] -} diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index bdc16e7..46d6794 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -14,10 +14,9 @@ import ( type StackEthernet struct { connID uint64 handlers handlers - logger - mac [6]byte - gwmac [6]byte - mtu uint16 + mac [6]byte + gwmac [6]byte + mtu uint16 } func (ls *StackEthernet) SetGateway6(gw [6]byte) { @@ -42,11 +41,10 @@ func (ls *StackEthernet) Reset6(mac, gateway [6]byte, mtu, maxNodes int) error { } else if maxNodes <= 0 { return errZeroMaxNodesArg } - ls.handlers.reset(maxNodes) + ls.handlers.reset("StackEthernet", maxNodes) *ls = StackEthernet{ connID: ls.connID + 1, handlers: ls.handlers, - logger: ls.logger, mac: mac, gwmac: gateway, mtu: uint16(mtu), @@ -86,18 +84,11 @@ func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) if vld.HasError() { return vld.ErrPop() } - { - h := ls.handlers.nodeByProto(uint16(etype)) - if h != nil { - err := h.demux(efrm.Payload(), 0) - if ls.handlers.tryHandleError(h, err) { - err = nil - } - return err - } + if h, err := ls.handlers.demuxByProto(efrm.Payload(), 0, uint16(etype)); h != nil { + return err } DROP: - ls.info("LinkStack:drop-packet", slog.String("dsthw", net.HardwareAddr(dstaddr[:]).String()), slog.String("ethertype", efrm.EtherTypeOrSize().String())) + ls.handlers.info("LinkStack:drop-packet", slog.String("dsthw", net.HardwareAddr(dstaddr[:]).String()), slog.String("ethertype", efrm.EtherTypeOrSize().String())) return lneto.ErrPacketDrop } @@ -114,14 +105,12 @@ func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int *efrm.DestinationHardwareAddr() = ls.gwmac var h *node h, n, err = ls.handlers.encapsulateAny(dst[:mtu], 14) - if n > 0 { - // Found packet - *efrm.SourceHardwareAddr() = ls.mac - efrm.SetEtherType(ethernet.Type(h.proto)) - n += 14 - if err != nil { - ls.error("Ethernet:encapuslate", slog.String("err", err.Error())) - } + if n == 0 { + return n, err } + // Found packet + *efrm.SourceHardwareAddr() = ls.mac + efrm.SetEtherType(ethernet.Type(h.proto)) + n += 14 return n, err } diff --git a/internet/stack-ip.go b/internet/stack-ip.go index 2d129ea..e7748e9 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -23,7 +23,6 @@ type StackIP struct { ip [4]byte validator lneto.Validator handlers handlers - logger } func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { @@ -34,12 +33,11 @@ func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { if err != nil { return err } - sb.handlers.reset(maxNodes) + sb.handlers.reset("StackIP", maxNodes) *sb = StackIP{ connID: sb.connID + 1, validator: sb.validator, handlers: sb.handlers, - logger: sb.logger, ip: sb.ip, } return nil @@ -70,11 +68,11 @@ func (sb *StackIP) Addr() netip.Addr { } func (sb *StackIP) SetLogger(logger *slog.Logger) { - sb.logger.log = logger + sb.handlers.log = logger } func (sb *StackIP) Demux(carrierData []byte, offset int) error { - sb.info("StackIP.Demux:start") + sb.handlers.info("StackIP.Demux:start") frame := carrierData[offset:] // we don't care about carrier data in IP. ifrm, err := ipv4.NewFrame(frame) if err != nil { @@ -93,7 +91,7 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { gotCRC := ifrm.CRC() wantCRC := ifrm.CalculateHeaderCRC() if gotCRC != wantCRC { - sb.error("StackIP:Demux:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC))) + sb.handlers.error("StackIP:Demux:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC))) return errors.New("IPv4 CRC mismatch") } off := ifrm.HeaderLength() @@ -106,7 +104,7 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { // nodeIdx := getNodeByProto(sb.handlers, uint16(proto)) if node == nil { // Drop packet. - sb.info("iprecv:drop", slog.String("dstaddr", netip.AddrFrom4(*ifrm.DestinationAddr()).String()), slog.String("proto", ifrm.Protocol().String())) + sb.handlers.info("iprecv:drop", slog.String("dstaddr", netip.AddrFrom4(*ifrm.DestinationAddr()).String()), slog.String("proto", ifrm.Protocol().String())) return nil } // Incoming CRC Validation of common IP Protocols. @@ -133,10 +131,10 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { return errors.New("UDP CRC mismatch") } } - sb.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen))) + sb.handlers.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen))) err = node.demux(frame[:totalLen], off) if sb.handlers.tryHandleError(node, err) { - sb.info("ipclose", slog.String("proto", proto.String())) + sb.handlers.info("ipclose", slog.String("proto", proto.String())) err = nil } return err @@ -160,45 +158,35 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) ifrm.SetTTL(64) *ifrm.SourceAddr() = sb.ip sb.ipID = id - for i := range sb.handlers.nodes { - h := &sb.handlers.nodes[i] - proto := lneto.IPProto(h.proto) - n, err := h.encapsulate(frame[:], headerlen) - if err != nil { - if sb.handlers.tryHandleError(h, err) { - println("IP NODE REMOVED", proto.String(), h.port) - } - sb.error("StackIP:encapsulate", slog.String("proto", proto.String()), slog.String("err", err.Error())) - continue - } else if n == 0 { - continue - } - totalLen := n + headerlen - ifrm.SetTotalLength(uint16(totalLen)) - ifrm.SetProtocol(proto) - ifrm.SetCRC(ifrm.CalculateHeaderCRC()) - // Calculate CRC for our newly generated packet. - var crc lneto.CRC791 - switch proto { - case lneto.IPProtoTCP: - ifrm.CRCWriteTCPPseudo(&crc) - tfrm, _ := tcp.NewFrame(ifrm.Payload()) - tfrm.CRCWrite(&crc) - tfrm.SetCRC(crc.Sum16()) - case lneto.IPProtoUDP: - ifrm.CRCWriteUDPPseudo(&crc) - ufrm, _ := udp.NewFrame(ifrm.Payload()) - ufrm.SetLength(uint16(n)) - ufrm.CRCWriteIPv4(&crc) - ufrm.SetCRC(crc.Sum16()) - if n != int(ufrm.Length()) { - sb.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length()))) - return 0, errors.New("invalid UDP length") - } - } - return totalLen, nil + node, n, err := sb.handlers.encapsulateAny(frame, headerlen) + if n == 0 { + return n, err } - return 0, nil + proto := lneto.IPProto(node.proto) + totalLen := n + headerlen + ifrm.SetTotalLength(uint16(totalLen)) + ifrm.SetProtocol(proto) + ifrm.SetCRC(ifrm.CalculateHeaderCRC()) + // Calculate CRC for our newly generated packet. + var crc lneto.CRC791 + switch proto { + case lneto.IPProtoTCP: + ifrm.CRCWriteTCPPseudo(&crc) + tfrm, _ := tcp.NewFrame(ifrm.Payload()) + tfrm.CRCWrite(&crc) + tfrm.SetCRC(crc.Sum16()) + case lneto.IPProtoUDP: + ifrm.CRCWriteUDPPseudo(&crc) + ufrm, _ := udp.NewFrame(ifrm.Payload()) + ufrm.SetLength(uint16(n)) + ufrm.CRCWriteIPv4(&crc) + ufrm.SetCRC(crc.Sum16()) + if n != int(ufrm.Length()) { + sb.handlers.error("StackIP:encaps", slog.Int("n", n), slog.Int("un", int(ufrm.Length()))) + return 0, errors.New("invalid UDP length") + } + } + return totalLen, err } func (sb *StackIP) Register(h StackNode) error { diff --git a/internet/stack-ports.go b/internet/stack-ports.go index d2f68f6..87eece3 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -4,6 +4,7 @@ import ( "encoding/binary" "io" "math" + "strconv" "github.com/soypat/lneto" ) @@ -29,7 +30,7 @@ func (ps *StackPorts) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) } else if maxNodes <= 0 { return errZeroMaxNodesArg } - ps.handlers.reset(maxNodes) + ps.handlers.reset("StackPorts(proto="+strconv.Itoa(int(protocol))+")", maxNodes) *ps = StackPorts{ connID: ps.connID + 1, handlers: ps.handlers, @@ -58,15 +59,7 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { return io.ErrShortBuffer } port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:]) - node := ps.handlers.nodeByPort(port) - if node == nil { - return nil - } - err = node.demux(b, offset) - if ps.handlers.tryHandleError(node, err) { - // discarded handler gracefully. - err = nil - } + _, err = ps.handlers.demuxByPort(b, offset, port) return err }