internet: further use handlers data structure

This commit is contained in:
Patricio Whittingslow
2025-12-18 19:16:35 -03:00
parent ac7f752447
commit eaa36a589a
4 changed files with 89 additions and 145 deletions
+38 -64
View File
@@ -2,6 +2,7 @@ package internet
import ( import (
"errors" "errors"
"log/slog"
"math" "math"
"net" "net"
"slices" "slices"
@@ -44,11 +45,14 @@ type node struct {
} }
type handlers struct { type handlers struct {
context string
logger
nodes []node 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.nodes = slices.Grow(h.nodes[:0], maxNodes)
h.context = context
} }
func (h *handlers) registerByProto(n node) error { func (h *handlers) registerByProto(n node) error {
@@ -136,22 +140,50 @@ func (h *handlers) nodeByPortProto(port uint16, protocol uint16) *node {
return nil return nil
} }
// encapsulateAny does not add the offset to the amount of bytes written. func (h *handlers) demuxByProto(buf []byte, offset int, proto uint16) (*node, error) {
func (h *handlers) encapsulateAny(buf []byte, offset int) (*node, int, 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 { for i := range h.nodes {
node := &h.nodes[i] node := &h.nodes[i]
if node.IsInvalid() { if node.IsInvalid() {
continue continue
} }
n, err := node.encapsulate(buf, offset) n, err = node.encapsulate(buf, offset)
if h.tryHandleError(node, err) { if h.tryHandleError(node, err) {
err = nil // CLOSE error handled gracefully by deleting node. err = nil // CLOSE error handled gracefully by deleting node.
} }
if err != nil || n > 0 { if n > 0 {
return node, n, err 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 ( var (
@@ -163,32 +195,6 @@ var (
_ = net.ErrClosed _ = 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 { func (node *node) IsInvalid() bool {
return node.demux == nil || node.encapsulate == nil || (node.connID != nil && node.currConnID != *node.connID) 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. // destroy removes all references to underlying StackNode. Allows garbage collection of node if possible.
func (n *node) destroy() { func (n *node) destroy() {
*n = node{} *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]
}
+13 -24
View File
@@ -14,10 +14,9 @@ import (
type StackEthernet struct { type StackEthernet struct {
connID uint64 connID uint64
handlers handlers handlers handlers
logger mac [6]byte
mac [6]byte gwmac [6]byte
gwmac [6]byte mtu uint16
mtu uint16
} }
func (ls *StackEthernet) SetGateway6(gw [6]byte) { 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 { } else if maxNodes <= 0 {
return errZeroMaxNodesArg return errZeroMaxNodesArg
} }
ls.handlers.reset(maxNodes) ls.handlers.reset("StackEthernet", maxNodes)
*ls = StackEthernet{ *ls = StackEthernet{
connID: ls.connID + 1, connID: ls.connID + 1,
handlers: ls.handlers, handlers: ls.handlers,
logger: ls.logger,
mac: mac, mac: mac,
gwmac: gateway, gwmac: gateway,
mtu: uint16(mtu), mtu: uint16(mtu),
@@ -86,18 +84,11 @@ func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error)
if vld.HasError() { if vld.HasError() {
return vld.ErrPop() return vld.ErrPop()
} }
{ if h, err := ls.handlers.demuxByProto(efrm.Payload(), 0, uint16(etype)); h != nil {
h := ls.handlers.nodeByProto(uint16(etype)) return err
if h != nil {
err := h.demux(efrm.Payload(), 0)
if ls.handlers.tryHandleError(h, err) {
err = nil
}
return err
}
} }
DROP: 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 return lneto.ErrPacketDrop
} }
@@ -114,14 +105,12 @@ func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int
*efrm.DestinationHardwareAddr() = ls.gwmac *efrm.DestinationHardwareAddr() = ls.gwmac
var h *node var h *node
h, n, err = ls.handlers.encapsulateAny(dst[:mtu], 14) h, n, err = ls.handlers.encapsulateAny(dst[:mtu], 14)
if n > 0 { if n == 0 {
// Found packet return n, err
*efrm.SourceHardwareAddr() = ls.mac
efrm.SetEtherType(ethernet.Type(h.proto))
n += 14
if err != nil {
ls.error("Ethernet:encapuslate", slog.String("err", err.Error()))
}
} }
// Found packet
*efrm.SourceHardwareAddr() = ls.mac
efrm.SetEtherType(ethernet.Type(h.proto))
n += 14
return n, err return n, err
} }
+35 -47
View File
@@ -23,7 +23,6 @@ type StackIP struct {
ip [4]byte ip [4]byte
validator lneto.Validator validator lneto.Validator
handlers handlers handlers handlers
logger
} }
func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { 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 { if err != nil {
return err return err
} }
sb.handlers.reset(maxNodes) sb.handlers.reset("StackIP", maxNodes)
*sb = StackIP{ *sb = StackIP{
connID: sb.connID + 1, connID: sb.connID + 1,
validator: sb.validator, validator: sb.validator,
handlers: sb.handlers, handlers: sb.handlers,
logger: sb.logger,
ip: sb.ip, ip: sb.ip,
} }
return nil return nil
@@ -70,11 +68,11 @@ func (sb *StackIP) Addr() netip.Addr {
} }
func (sb *StackIP) SetLogger(logger *slog.Logger) { func (sb *StackIP) SetLogger(logger *slog.Logger) {
sb.logger.log = logger sb.handlers.log = logger
} }
func (sb *StackIP) Demux(carrierData []byte, offset int) error { 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. frame := carrierData[offset:] // we don't care about carrier data in IP.
ifrm, err := ipv4.NewFrame(frame) ifrm, err := ipv4.NewFrame(frame)
if err != nil { if err != nil {
@@ -93,7 +91,7 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
gotCRC := ifrm.CRC() gotCRC := ifrm.CRC()
wantCRC := ifrm.CalculateHeaderCRC() wantCRC := ifrm.CalculateHeaderCRC()
if gotCRC != wantCRC { 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") return errors.New("IPv4 CRC mismatch")
} }
off := ifrm.HeaderLength() off := ifrm.HeaderLength()
@@ -106,7 +104,7 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error {
// nodeIdx := getNodeByProto(sb.handlers, uint16(proto)) // nodeIdx := getNodeByProto(sb.handlers, uint16(proto))
if node == nil { if node == nil {
// Drop packet. // 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 return nil
} }
// Incoming CRC Validation of common IP Protocols. // 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") 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) err = node.demux(frame[:totalLen], off)
if sb.handlers.tryHandleError(node, err) { 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 err = nil
} }
return err return err
@@ -160,45 +158,35 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error)
ifrm.SetTTL(64) ifrm.SetTTL(64)
*ifrm.SourceAddr() = sb.ip *ifrm.SourceAddr() = sb.ip
sb.ipID = id sb.ipID = id
for i := range sb.handlers.nodes { node, n, err := sb.handlers.encapsulateAny(frame, headerlen)
h := &sb.handlers.nodes[i] if n == 0 {
proto := lneto.IPProto(h.proto) return n, err
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
} }
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 { func (sb *StackIP) Register(h StackNode) error {
+3 -10
View File
@@ -4,6 +4,7 @@ import (
"encoding/binary" "encoding/binary"
"io" "io"
"math" "math"
"strconv"
"github.com/soypat/lneto" "github.com/soypat/lneto"
) )
@@ -29,7 +30,7 @@ func (ps *StackPorts) Reset(protocol uint64, dstPortOffset uint16, maxNodes int)
} else if maxNodes <= 0 { } else if maxNodes <= 0 {
return errZeroMaxNodesArg return errZeroMaxNodesArg
} }
ps.handlers.reset(maxNodes) ps.handlers.reset("StackPorts(proto="+strconv.Itoa(int(protocol))+")", maxNodes)
*ps = StackPorts{ *ps = StackPorts{
connID: ps.connID + 1, connID: ps.connID + 1,
handlers: ps.handlers, handlers: ps.handlers,
@@ -58,15 +59,7 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) {
return io.ErrShortBuffer return io.ErrShortBuffer
} }
port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:]) port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:])
node := ps.handlers.nodeByPort(port) _, err = ps.handlers.demuxByPort(b, offset, port)
if node == nil {
return nil
}
err = node.demux(b, offset)
if ps.handlers.tryHandleError(node, err) {
// discarded handler gracefully.
err = nil
}
return err return err
} }