diff --git a/examples/bridge/main.go b/examples/bridge/main.go index 43c7037..daae2c7 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -14,7 +14,6 @@ import ( "strings" "time" - "github.com/soypat/lneto" "github.com/soypat/lneto/arp" "github.com/soypat/lneto/dhcpv4" "github.com/soypat/lneto/dns" @@ -82,7 +81,6 @@ func run() (err error) { return err } brHW := nicHW - brHW[4]++ mtu, err := iface.MTU() if err != nil { return err @@ -138,8 +136,6 @@ func run() (err error) { hw, err := stack.ResultResolveHardwareAddress6(netip.AddrFrom4(router)) if err == nil { stack.link.SetGateway6(hw) - stack.link.SetHardwareAddr6([6]byte{0xd8, 0x5e, 0xd3, 0x43, 0x03, 0xeb}) - stack.ip.SetAddr(netip.AddrFrom4([4]byte{192, 168, 1, 53})) if flagDoNTP { state = stateDNSNTP err = stack.StartLookupIP("pool.ntp.org") @@ -274,11 +270,12 @@ func (s *Stack) Encapsulate(b []byte, _ int) (int, error) { } func (s *Stack) Reset(mac [6]byte, addr netip.Addr, mtu uint16) error { - err := s.link.Reset6(mac, ethernet.BroadcastAddr(), int(mtu)) + const maxNodes = 8 + err := s.link.Reset6(mac, ethernet.BroadcastAddr(), int(mtu), maxNodes) if err != nil { return err } - err = s.ip.Reset(addr) + err = s.ip.Reset(addr, maxNodes) if err != nil { return err } @@ -298,7 +295,7 @@ func (s *Stack) Reset(mac [6]byte, addr netip.Addr, mtu uint16) error { if err != nil { return err } - err = s.udps.Reset(uint64(lneto.IPProtoUDP), 2) + err = s.udps.ResetUDP(maxNodes) if err != nil { return err } @@ -355,6 +352,7 @@ func (s *Stack) StartLookupIP(host string) error { if err != nil { return err } + fmt.Println("START LOOKUP", host, dns4[:]) return nil } @@ -404,7 +402,7 @@ func (s *Stack) BeginDHCPRequest(request [4]byte) error { } func (s *Stack) StartNTP(addr netip.Addr) error { - s.ntp.Reset(time.Now, s.sysprec) + s.ntp.Reset(s.sysprec, time.Now) var u internet.StackUDPPort addr4 := addr.As4() u.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort) diff --git a/examples/stack/main.go b/examples/stack/main.go index 8bec0dc..1b50a61 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -185,15 +185,16 @@ type Stack struct { } func (stack *Stack) Reset(ourMAC, gwMAC [6]byte, ip netip.Addr, mtu int) (err error) { - err = stack.ethernet.Reset6(ourMAC, gwMAC, mtu) + const maxNodes = 8 + err = stack.ethernet.Reset6(ourMAC, gwMAC, mtu, maxNodes) if err != nil { return err } - err = stack.ip.Reset(ip) + err = stack.ip.Reset(ip, maxNodes) if err != nil { return err } - stack.tcpports.Reset(uint64(lneto.IPProtoTCP), 2) + stack.tcpports.ResetTCP(maxNodes) ipaddr := ip.As4() err = stack.arp.Reset(arp.HandlerConfig{ HardwareAddr: ourMAC[:], diff --git a/internet/definitions.go b/internet/definitions.go index e14e153..ce9d4e6 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -4,7 +4,6 @@ import ( "errors" "math" "net" - "slices" ) // StackNode is an abstraction of a packet exchanging protocol controller. This is the building block for all protocols, @@ -46,9 +45,21 @@ var ( errZeroPort = errors.New("port must be greater than zero") errInvalidProto = errors.New("invalid protocol") errProtoRegistered = errors.New("protocol already registered") + errNodesFull = errors.New("no more room for new nodes") _ = 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) { @@ -56,7 +67,8 @@ func handleNodeError(nodesPtr *[]node, nodeIdx int, err error) (discarded bool) } nodes := *nodesPtr if checkNodeErr(&nodes[nodeIdx], err) { - *nodesPtr = slices.Delete(nodes, nodeIdx, nodeIdx+1) + // *nodesPtr = slices.Delete(nodes, nodeIdx, nodeIdx+1) + (*nodesPtr)[nodeIdx] = node{} // 'Delete' node without modifying slice length. discarded = true } } @@ -71,10 +83,6 @@ func checkNodeErr(node *node, err error) (discard bool) { return checkNode(node) || (err != nil && err == net.ErrClosed) } -func addNode(nodes *[]node, h StackNode, port uint16, protocol uint64) { - *nodes = append(*nodes, nodeFromStackNode(h, port, protocol)) -} - func nodeFromStackNode(s StackNode, port uint16, protocol uint64) node { if protocol > math.MaxUint16 { panic(">16bit protocol number unsupported") diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index b6f8f45..a3b9ef4 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -6,6 +6,7 @@ import ( "log/slog" "math" "net" + "slices" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -32,13 +33,14 @@ func (ls *StackEthernet) HardwareAddr6() [6]byte { return ls.mac } -func (ls *StackEthernet) Reset6(mac, gateway [6]byte, mtu int) error { +func (ls *StackEthernet) Reset6(mac, gateway [6]byte, mtu, maxNodes int) error { if mtu > math.MaxUint16 || mtu < 256 { return errors.New("invalid MTU") } + ls.handlers = slices.Grow(ls.handlers[:0], maxNodes) *ls = StackEthernet{ connID: ls.connID + 1, - handlers: ls.handlers[:0], + handlers: ls.handlers, logger: ls.logger, mac: mac, gwmac: gateway, diff --git a/internet/stack-ip.go b/internet/stack-ip.go index 34315b8..609e879 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -5,6 +5,7 @@ import ( "io" "log/slog" "net/netip" + "slices" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -25,15 +26,16 @@ type StackIP struct { logger } -func (sb *StackIP) Reset(addr netip.Addr) error { +func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error { err := sb.SetAddr(addr) if err != nil { return err } + sb.handlers = slices.Grow(sb.handlers[:0], maxNodes) *sb = StackIP{ connID: sb.connID + 1, validator: sb.validator, - handlers: sb.handlers[:0], + handlers: sb.handlers, logger: sb.logger, ip: sb.ip, } @@ -157,7 +159,7 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) n, err := h.encapsulate(frame[:], headerlen) if err != nil { if handleNodeError(&sb.handlers, i, err) { - println("NODE REMOVED", proto.String(), h.port) + println("IP NODE REMOVED", proto.String(), h.port) h.destroy() } sb.error("StackIP:encapsulate", slog.String("proto", proto.String()), slog.String("err", err.Error())) diff --git a/internet/stack-ports.go b/internet/stack-ports.go index 2cb5082..b9e8bfa 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -4,22 +4,34 @@ import ( "encoding/binary" "io" "math" + "slices" + + "github.com/soypat/lneto" ) type StackPorts struct { connID uint64 handlers []node - dstPortOff int + dstPortOff uint16 protocol uint16 } -func (ps *StackPorts) Reset(protocol uint64, dstPortOffset int) error { +func (ps *StackPorts) ResetUDP(maxNodes int) error { + return ps.Reset(uint64(lneto.IPProtoUDP), 2, maxNodes) +} + +func (ps *StackPorts) ResetTCP(maxNodes int) error { + return ps.Reset(uint64(lneto.IPProtoTCP), 2, maxNodes) +} + +func (ps *StackPorts) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) error { if protocol > math.MaxUint16 { return errInvalidProto } + ps.handlers = slices.Grow(ps.handlers[:0], maxNodes) *ps = StackPorts{ connID: ps.connID + 1, - handlers: ps.handlers[:0], + handlers: ps.handlers, dstPortOff: dstPortOffset, protocol: uint16(protocol), } @@ -33,25 +45,28 @@ func (ps *StackPorts) Protocol() uint64 { return uint64(ps.protocol) } func (ps *StackPorts) ConnectionID() *uint64 { return &ps.connID } func (ps *StackPorts) Encapsulate(b []byte, offset int) (n int, err error) { - if ps.dstPortOff+offset+2 > len(b) { + if int(ps.dstPortOff)+offset+2 > len(b) { return 0, io.ErrShortBuffer } var i int for i = 0; i < len(ps.handlers); i++ { 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 } } - ps.handleResult(i, n, err) return n, err } func (ps *StackPorts) Demux(b []byte, offset int) (err error) { - if ps.dstPortOff+offset+2 > len(b) { + if int(ps.dstPortOff)+offset+2 > len(b) { return io.ErrShortBuffer } - port := binary.BigEndian.Uint16(b[ps.dstPortOff+offset:]) + 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 { @@ -59,6 +74,10 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { } 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 } } @@ -74,16 +93,17 @@ func (ps *StackPorts) Register(h StackNode) error { } else if proto != uint64(ps.protocol) { return errInvalidProto } - ps.handlers = append(ps.handlers, node{ + return registerNode(&ps.handlers, node{ demux: h.Demux, encapsulate: h.Encapsulate, port: port, }) - return nil } -func (ps *StackPorts) handleResult(handlerIdx, n int, err error) { +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 } diff --git a/ntp/client.go b/ntp/client.go index 3d3c2da..8c65018 100644 --- a/ntp/client.go +++ b/ntp/client.go @@ -35,7 +35,7 @@ type Client struct { sysprec int8 } -func (c *Client) Reset(now func() time.Time, sysprec int8) { +func (c *Client) Reset(sysprec int8, now func() time.Time) { *c = Client{ connID: c.connID + 1, _now: now,