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 (
"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]
}
+13 -24
View File
@@ -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
}
+35 -47
View File
@@ -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 {
+3 -10
View File
@@ -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
}