mirror of
https://github.com/soypat/lneto.git
synced 2026-08-10 09:53:44 +00:00
internet: further use handlers data structure
This commit is contained in:
+38
-64
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user