diff --git a/internet/definitions.go b/internet/definitions.go index 55a4fb9..b4aded0 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -4,6 +4,7 @@ import ( "errors" "math" "net" + "slices" ) // StackNode is an abstraction of a packet exchanging protocol controller. This is the building block for all protocols, @@ -39,6 +40,118 @@ type node struct { encapsulate func([]byte, int) (int, error) proto uint16 port uint16 + remoteAddr []byte +} + +type handlers struct { + nodes []node +} + +func (h *handlers) reset(maxNodes int) { + h.nodes = slices.Grow(h.nodes[:0], maxNodes) +} + +func (h *handlers) registerByProto(n node) error { + err := h.prepAdd() + if err != nil { + return err + } + if h.nodeByProto(n.proto) != nil { + return errProtoRegistered + } + h.nodes = append(h.nodes, n) + return nil +} + +func (h *handlers) registerByPortProto(n node) error { + err := h.prepAdd() + if err != nil { + return err + } + if h.nodeByPortProto(n.port, n.proto) != nil { + return errProtoRegistered + } + h.nodes = append(h.nodes, n) + return nil +} + +func (h *handlers) prepAdd() error { + if h.full() { + h.compact() + if h.full() { + return errNodesFull + } + } + return nil +} + +func (h *handlers) full() bool { return cap(h.nodes) == len(h.nodes) } + +func (h *handlers) compact() { + nilOff := 0 + for i := 0; i < len(h.nodes); i++ { + if !h.nodes[i].IsInvalid() { + h.nodes[nilOff] = h.nodes[i] + nilOff++ + } + } + h.nodes = h.nodes[:nilOff] +} + +func (h *handlers) tryHandleError(node *node, err error) (discardedGracefully bool) { + if err != nil && (err == net.ErrClosed || node.IsInvalid()) { + node.destroy() + discardedGracefully = true + } + return discardedGracefully +} + +func (h *handlers) nodeByProto(proto uint16) *node { + for i := range h.nodes { + node := &h.nodes[i] + if node.proto == proto { + return node + } + } + return nil +} + +func (h *handlers) nodeByPort(port uint16) *node { + for i := range h.nodes { + node := &h.nodes[i] + if node.port == port { + return node + } + } + return nil +} + +func (h *handlers) nodeByPortProto(port uint16, protocol uint16) *node { + for i := range h.nodes { + node := &h.nodes[i] + if node.port == port && node.proto == protocol { + return 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) { + for i := range h.nodes { + node := &h.nodes[i] + if node.IsInvalid() { + continue + } + 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 { + return node, n, err + } + } + return nil, 0, nil } var ( @@ -84,7 +197,7 @@ func checkNodeErr(node *node, err error) (discard bool) { return node.IsInvalid() || (err != nil && err == net.ErrClosed) } -func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node { +func nodeFromStackNode(s StackNode, port uint16, protocol uint64, remoteAddr []byte) node { if protocol > math.MaxUint16 { panic(">16bit protocol number unsupported") } @@ -100,6 +213,7 @@ func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node { encapsulate: s.Encapsulate, proto: uint16(protocol), port: port, + remoteAddr: append([]byte{}, remoteAddr...), } } @@ -113,28 +227,6 @@ func getNode(nodes []node, port uint16, protocol uint16) (node *node) { return nil } -func getEncapsulateNode(nodes *[]node, carrierData []byte, frameOffset int) (nodeIdx int, written int, err error) { - destroyed := false - for i := range *nodes { - node := &(*nodes)[i] - if node.IsInvalid() { - destroyed = true - node.destroy() - continue - } - written, err = node.encapsulate(carrierData, frameOffset) - if written > 0 { - return i, written, err - } else if err != nil { - - } - } - if destroyed { - *nodes = nodesCompact(*nodes) - } - return -1, 0, nil -} - // destroy removes all references to underlying StackNode. Allows garbage collection of node if possible. func (n *node) destroy() { *n = node{} diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index c2ebaad..bdc16e7 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -6,7 +6,6 @@ import ( "log/slog" "math" "net" - "slices" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -14,7 +13,7 @@ import ( type StackEthernet struct { connID uint64 - handlers []node + handlers handlers logger mac [6]byte gwmac [6]byte @@ -43,7 +42,7 @@ func (ls *StackEthernet) Reset6(mac, gateway [6]byte, mtu, maxNodes int) error { } else if maxNodes <= 0 { return errZeroMaxNodesArg } - ls.handlers = slices.Grow(ls.handlers[:0], maxNodes) + ls.handlers.reset(maxNodes) *ls = StackEthernet{ connID: ls.connID + 1, handlers: ls.handlers, @@ -68,18 +67,7 @@ func (ls *StackEthernet) Register(h StackNode) error { if proto > math.MaxUint16 || proto <= 1500 { return errInvalidProto } - eproto := uint16(proto) - for i := range ls.handlers { - hgot := &ls.handlers[i] - if hgot.proto == eproto { - return errProtoRegistered - } - } - return registerNode(&ls.handlers, node{ - demux: h.Demux, - encapsulate: h.Encapsulate, - proto: eproto, - }) + return ls.handlers.registerByProto(nodeFromStackNode(h, 0, proto, nil)) } func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) { @@ -98,11 +86,14 @@ func (ls *StackEthernet) Demux(carrierData []byte, frameOffset int) (err error) if vld.HasError() { return vld.ErrPop() } - - for i := range ls.handlers { - h := &ls.handlers[i] - if h.proto == uint16(etype) { - return h.demux(efrm.Payload(), 0) + { + 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 } } DROP: @@ -121,19 +112,16 @@ func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int return 0, err } *efrm.DestinationHardwareAddr() = ls.gwmac - for i := range ls.handlers { - h := &ls.handlers[i] - n, err = h.encapsulate(dst[:mtu], 14) + 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("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error())) - continue - } - if n > 0 { - // Found packet - *efrm.SourceHardwareAddr() = ls.mac - efrm.SetEtherType(ethernet.Type(h.proto)) - return n + 14, nil + ls.error("Ethernet:encapuslate", slog.String("err", err.Error())) } } - return 0, err + return n, err } diff --git a/internet/stack-ip.go b/internet/stack-ip.go index 5ac4549..2d129ea 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -5,7 +5,6 @@ import ( "io" "log/slog" "net/netip" - "slices" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -19,12 +18,11 @@ import ( var _ StackNode = (*StackIP)(nil) type StackIP struct { - connID uint64 - ipID uint16 - ip [4]byte - validator lneto.Validator - handlers []node - pendingICMP [][]byte + connID uint64 + ipID uint16 + ip [4]byte + validator lneto.Validator + handlers handlers logger } @@ -36,14 +34,13 @@ func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { if err != nil { return err } - sb.handlers = slices.Grow(sb.handlers[:0], maxNodes) + sb.handlers.reset(maxNodes) *sb = StackIP{ - connID: sb.connID + 1, - validator: sb.validator, - handlers: sb.handlers, - logger: sb.logger, - ip: sb.ip, - pendingICMP: make([][]byte, maxNodes*4), + connID: sb.connID + 1, + validator: sb.validator, + handlers: sb.handlers, + logger: sb.logger, + ip: sb.ip, } return nil } @@ -105,8 +102,9 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { if proto == lneto.IPProtoICMP { return sb.recvicmp(ifrm.RawData(), ifrm.HeaderLength()) } - nodeIdx := getNodeByProto(sb.handlers, uint16(proto)) - if nodeIdx < 0 { + node := sb.handlers.nodeByProto(uint16(proto)) + // 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())) return nil @@ -136,8 +134,8 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { } } sb.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen))) - err = sb.handlers[nodeIdx].demux(frame[:totalLen], off) - if handleNodeError(&sb.handlers, nodeIdx, err) { + err = node.demux(frame[:totalLen], off) + if sb.handlers.tryHandleError(node, err) { sb.info("ipclose", slog.String("proto", proto.String())) err = nil } @@ -162,14 +160,13 @@ 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 { - h := &sb.handlers[i] + 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 handleNodeError(&sb.handlers, i, err) { + if sb.handlers.tryHandleError(h, err) { println("IP NODE REMOVED", proto.String(), h.port) - h.destroy() } sb.error("StackIP:encapsulate", slog.String("proto", proto.String()), slog.String("err", err.Error())) continue @@ -209,19 +206,7 @@ func (sb *StackIP) Register(h StackNode) error { if proto > 255 { return errInvalidProto } - connID := h.ConnectionID() - var currConnID uint64 - if connID != nil { - currConnID = *connID - } - return registerNode(&sb.handlers, node{ - demux: h.Demux, - encapsulate: h.Encapsulate, - proto: uint16(proto), - port: h.LocalPort(), - currConnID: currConnID, - connID: connID, - }) + return sb.handlers.registerByPortProto(nodeFromStackNode(h, h.LocalPort(), proto, nil)) } func (sb *StackIP) recvicmp(carrierData []byte, offset int) error { diff --git a/internet/stack-ports.go b/internet/stack-ports.go index bd51284..d2f68f6 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -4,14 +4,13 @@ import ( "encoding/binary" "io" "math" - "slices" "github.com/soypat/lneto" ) type StackPorts struct { connID uint64 - handlers []node + handlers handlers dstPortOff uint16 protocol uint16 } @@ -30,7 +29,7 @@ func (ps *StackPorts) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) } else if maxNodes <= 0 { return errZeroMaxNodesArg } - ps.handlers = slices.Grow(ps.handlers[:0], maxNodes) + ps.handlers.reset(maxNodes) *ps = StackPorts{ connID: ps.connID + 1, handlers: ps.handlers, @@ -50,20 +49,7 @@ func (ps *StackPorts) Encapsulate(b []byte, offset int) (n int, err error) { if int(ps.dstPortOff)+offset+2 > len(b) { return 0, io.ErrShortBuffer } - var i int - for i = 0; i < len(ps.handlers); i++ { - if ps.handlers[i].IsInvalid() { - continue - } - n, err = ps.handlers[i].encapsulate(b, offset) - if err != nil || n > 0 { - if ps.handleResult(i, n, err) { - err = nil // Handler discarded. Keep looking for other handlers. - continue - } - break - } - } + _, n, err = ps.handlers.encapsulateAny(b, offset) return n, err } @@ -72,52 +58,25 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { return io.ErrShortBuffer } port := binary.BigEndian.Uint16(b[int(ps.dstPortOff)+offset:]) - var i int - for i = 0; i < len(ps.handlers); i++ { - if port != ps.handlers[i].port { - continue - } - err = ps.handlers[i].demux(b, offset) - if err != nil { - if ps.handleResult(i, 0, err) { - err = nil // Handler discarded. Keep looking for other maybe available handlers. - continue - } - break - } + 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 } - ps.handleResult(i, 0, err) return err } func (ps *StackPorts) Register(h StackNode) error { port := h.LocalPort() proto := h.Protocol() - if port <= 0 { return errZeroPort } else if proto != uint64(ps.protocol) { return errInvalidProto } - var cid uint64 - cidPtr := h.ConnectionID() - if cidPtr != nil { - cid = *cidPtr - } - return registerNode(&ps.handlers, node{ - demux: h.Demux, - encapsulate: h.Encapsulate, - port: port, - currConnID: cid, - connID: cidPtr, - proto: uint16(proto), - }) -} - -func (ps *StackPorts) handleResult(handlerIdx, n int, err error) (discarded bool) { - if handleNodeError(&ps.handlers, handlerIdx, err) { - discarded = true - println("DISCARD", handlerIdx, "witherr", err.Error()) - } - return discarded + return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil)) } diff --git a/internet/stack-udpport.go b/internet/stack-udpport.go index aaa4ae8..c63ee1b 100644 --- a/internet/stack-udpport.go +++ b/internet/stack-udpport.go @@ -17,7 +17,7 @@ type StackUDPPort struct { } func (sudp *StackUDPPort) SetStackNode(node StackNode, raddr []byte, rmport uint16) { - sudp.h = nodeFromStackNode(node, node.LocalPort(), node.Protocol()) + sudp.h = nodeFromStackNode(node, node.LocalPort(), node.Protocol(), raddr) sudp.rmport = rmport sudp.raddr = append(sudp.raddr[:0], raddr...) }