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