From 6f3da983b0a9fe22e29c288bb2a2583d4b903e22 Mon Sep 17 00:00:00 2001 From: ddirect <62487612+ddirect@users.noreply.github.com> Date: Sun, 30 Nov 2025 16:51:42 +0200 Subject: [PATCH 01/11] Conn: fixed double mutex lock in SetDeadline --- tcp/conn.go | 18 ++++++++++++++---- 1 file changed, 14 insertions(+), 4 deletions(-) diff --git a/tcp/conn.go b/tcp/conn.go index 3dff2db..c7577d6 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -246,7 +246,7 @@ func (conn *Conn) checkPipeOpen() error { if conn.abortErr != nil { return conn.abortErr } - state := conn.State() + state := conn.h.State() if state.IsClosed() { return net.ErrClosed } @@ -332,11 +332,11 @@ func (conn *Conn) reset(h Handler) { func (conn *Conn) SetDeadline(t time.Time) error { conn.mu.Lock() defer conn.mu.Unlock() - err := conn.SetReadDeadline(t) + err := conn.setReadDeadline(t) if err != nil { return err } - return conn.SetWriteDeadline(t) + return conn.setWriteDeadline(t) } // SetReadDeadline sets the deadline for future Read calls @@ -344,7 +344,11 @@ func (conn *Conn) SetDeadline(t time.Time) error { func (conn *Conn) SetReadDeadline(t time.Time) error { conn.mu.Lock() defer conn.mu.Unlock() - conn.trace("TCPConn.SetReadDeadline:start") + return conn.setReadDeadline(t) +} + +func (conn *Conn) setReadDeadline(t time.Time) error { + conn.trace("TCPConn.setReadDeadline:start") err := conn.checkPipeOpen() if err == nil { conn.rdead = t @@ -358,6 +362,12 @@ func (conn *Conn) SetReadDeadline(t time.Time) error { // some of the data was successfully written. // A zero value for t means Write will not time out. func (conn *Conn) SetWriteDeadline(t time.Time) error { + conn.mu.Lock() + defer conn.mu.Unlock() + return conn.setWriteDeadline(t) +} + +func (conn *Conn) setWriteDeadline(t time.Time) error { conn.trace("TCPConn.SetWriteDeadline:start") err := conn.checkPipeOpen() if err == nil { From ac7f752447f6e05d8902bd381169e32af62f87e9 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Thu, 18 Dec 2025 18:44:50 -0300 Subject: [PATCH 02/11] internet: add handlers data structure --- internet/definitions.go | 138 ++++++++++++++++++++++++++++++------- internet/stack-ethernet.go | 52 ++++++-------- internet/stack-ip.go | 55 ++++++--------- internet/stack-ports.go | 65 ++++------------- internet/stack-udpport.go | 2 +- 5 files changed, 168 insertions(+), 144 deletions(-) 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...) } From eaa36a589aa90cf7ed1656030880bec67fed9757 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Thu, 18 Dec 2025 19:16:35 -0300 Subject: [PATCH 03/11] internet: further use handlers data structure --- internet/definitions.go | 102 ++++++++++++++----------------------- internet/stack-ethernet.go | 37 +++++--------- internet/stack-ip.go | 82 +++++++++++++---------------- internet/stack-ports.go | 13 ++--- 4 files changed, 89 insertions(+), 145 deletions(-) diff --git a/internet/definitions.go b/internet/definitions.go index b4aded0..683d34b 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -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] -} diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index bdc16e7..46d6794 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -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 } diff --git a/internet/stack-ip.go b/internet/stack-ip.go index 2d129ea..e7748e9 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -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 { diff --git a/internet/stack-ports.go b/internet/stack-ports.go index d2f68f6..87eece3 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -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 } From 1554b89a080ad770afbf1702ada72e6210120934 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Fri, 19 Dec 2025 00:38:28 -0300 Subject: [PATCH 04/11] StackNode refactor: add offsetToIP argument to Encapsulate --- README.md | 8 +-- arp/arp_test.go | 16 +++--- arp/handler.go | 97 ++++++++++++++++++++++++++++++------ dhcpv4/client.go | 14 +++--- dhcpv4/dhcp_test.go | 16 +++--- dhcpv4/server.go | 8 +-- dns/client.go | 4 +- examples/bridge/main.go | 10 ++-- examples/stack/main.go | 4 +- examples/stackbasic/main.go | 8 +-- examples/xnet/main.go | 2 +- internet/definitions.go | 21 +++++--- internet/node-tcplistener.go | 4 +- internet/stack-ethernet.go | 11 ++-- internet/stack-ip.go | 8 +-- internet/stack-ports.go | 6 +-- internet/stack-udpport.go | 11 ++-- internet/stackbasic_test.go | 2 +- ntp/client.go | 4 +- tcp/conn.go | 12 +++-- x/xnet/stack-async.go | 6 +-- x/xnet/xnet_test.go | 2 +- 22 files changed, 180 insertions(+), 94 deletions(-) diff --git a/README.md b/README.md index 68e38fb..77cc069 100644 --- a/README.md +++ b/README.md @@ -57,11 +57,11 @@ The following interface is implemented by networking stack nodes and the stack t ```go type StackNode interface { // Encapsulate receives a buffer the receiver must fill with data. - // The receiver's start byte is at carrierData[frameOffset]. - Encapsulate(carrierData []byte, frameOffset int) (int, error) + // The receiver's start byte is at carrierData[offsetToFrame]. + Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) // Demux receives a buffer the receiver must decode and pass on to corresponding child StackNode(s). - // The receiver's start byte is at carrierData[frameOffset]. - Demux(carrierData []byte, frameOffset int) error + // The receiver's start byte is at carrierData[offsetToFrame]. + Demux(carrierData []byte, offsetToFrame int) error // LocalPort returns the port of the node if applicable or zero. Used for UDP/TCP nodes. LocalPort() uint16 // Protocol returns the protocol of this node if applicable or zero. Usually either a ethernet.Type (EtherType) or lneto.IPProto (IP Protocol number). diff --git a/arp/arp_test.go b/arp/arp_test.go index 6b10a33..59bb875 100644 --- a/arp/arp_test.go +++ b/arp/arp_test.go @@ -34,13 +34,13 @@ func TestHandler(t *testing.T) { t.Fatal(err) } var buf, discard [64]byte - n, err := c1.Encapsulate(buf[:], 0) + n, err := c1.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal("error on should be nop send:", err) } else if n > 0 { t.Fatal("should not send if no query") } - n, err = c2.Encapsulate(buf[:], 0) + n, err = c2.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal("error on should be nop send:", err) } else if n > 0 { @@ -50,11 +50,11 @@ func TestHandler(t *testing.T) { // Perform ARP exchange. expectHWAddr := c2.ourHWAddr queryAddr := c2.ourProtoAddr - err = c1.StartQuery(queryAddr) + err = c1.StartQuery(nil, queryAddr) if err != nil { t.Fatal(err) } - n, err = c1.Encapsulate(buf[:], 0) // Send Request. + n, err = c1.Encapsulate(buf[:], -1, 0) // Send Request. if err != nil { t.Fatal(err) } else if n == 0 { @@ -66,14 +66,14 @@ func TestHandler(t *testing.T) { t.Fatal(err) } - n, err = c2.Encapsulate(buf[:], 0) // Send response. + n, err = c2.Encapsulate(buf[:], -1, 0) // Send response. if err != nil { t.Fatal(err) } else if n == 0 { t.Fatal("got no response to request") } validateARP(t, buf[:]) - n, err = c2.Encapsulate(discard[:], 0) // Double tap check, should send nothing. + n, err = c2.Encapsulate(discard[:], -1, 0) // Double tap check, should send nothing. if err != nil { t.Fatal("double tap send error:", err) } else if n > 0 { @@ -90,13 +90,13 @@ func TestHandler(t *testing.T) { } else if !bytes.Equal(hwaddr, expectHWAddr) { log.Fatalf("expected to get hwaddr %x!=%x", hwaddr, expectHWAddr) } - n, err = c1.Encapsulate(buf[:], 0) + n, err = c1.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n > 0 { t.Fatal("expected no data") } - n, err = c2.Encapsulate(buf[:], 0) + n, err = c2.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n > 0 { diff --git a/arp/handler.go b/arp/handler.go index e761ad1..aca916b 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -3,6 +3,7 @@ package arp import ( "bytes" "errors" + "log/slog" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" @@ -63,9 +64,22 @@ func (h *Handler) Reset(cfg HandlerConfig) error { type queryResult struct { protoaddr []byte hwaddr []byte + dstHw []byte querysent bool } +func (qr *queryResult) destroy() { + *qr = queryResult{protoaddr: qr.protoaddr[:0], hwaddr: qr.hwaddr[:0]} +} + +func (qr *queryResult) response() []byte { + if len(qr.hwaddr) == 0 { + return nil + } + return qr.hwaddr[:] +} +func (qr *queryResult) isInvalid() bool { return len(qr.protoaddr) == 0 } + // AbortPending drops pending queries and incoming requests. func (h *Handler) AbortPending() { h.pendingResponse = h.pendingResponse[:0] @@ -81,31 +95,69 @@ func (h *Handler) QueryResult(protoAddr []byte) (hwAddr []byte, err error) { if bytes.Equal(protoAddr, h.queries[i].protoaddr) { if !h.queries[i].querysent { return nil, errors.New("query not yet sent") - } else if len(h.queries[i].hwaddr) == 0 { + } + mac := h.queries[i].response() + if mac == nil { return nil, errors.New("no response yet") } - return h.queries[i].hwaddr, nil + return mac, nil } } return nil, errors.New("query not exist or dropped") } -func (h *Handler) StartQuery(proto []byte) error { +func (h *Handler) DiscardQuery(protoAddr []byte) error { + for i := range h.queries { + q := &h.queries[i] + if bytes.Equal(protoAddr, q.protoaddr) { + q.destroy() + return nil + } + } + return errors.New("query not found") +} + +func (h *Handler) compactQueries() { + validOff := 0 + for i := 0; i < len(h.queries); i++ { + if h.queries[i].isInvalid() { + h.queries[validOff] = h.queries[i] + validOff++ + } + } + h.queries = h.queries[:validOff] +} + +// StartQuery queues a query to perform over ARP for the protocol address `proto`. +// The user can additionally specify an dstHWAddr to write query result to on completion. +// If dstHWAddr is nil then query still occurs but no external buffer is written on query completion. +// dstHWAddr must be zeroed out (invalid MAC). +func (h *Handler) StartQuery(dstHWAddr, proto []byte) error { + if len(h.queries) == cap(h.queries) { + h.compactQueries() + if len(h.queries) == cap(h.queries) { + return errors.New("too many ongoing queries") + } + } if len(proto) != len(h.ourProtoAddr) { return errors.New("bad protocol address length") - } else if len(h.queries) == cap(h.queries) { - return errors.New("too many ongoing queries") + } else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) { + return errors.New("mismatch hardware size") + } else if dstHWAddr != nil && !allZeros(dstHWAddr) { + return errors.New("write-to buffer must be zeroed out") } h.queries = h.queries[:len(h.queries)+1] q := &h.queries[len(h.queries)-1] - q.hwaddr = q.hwaddr[:0] - q.querysent = false - q.protoaddr = append(q.protoaddr[:0], proto...) + *q = queryResult{ + protoaddr: append(q.protoaddr[:0], proto...), + hwaddr: q.hwaddr[:0], + dstHw: dstHWAddr, + } return nil } -func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) { - b := eth[frameOffset:] +func (h *Handler) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + b := carrierData[offsetToFrame:] n := h.expectSize() if len(b) < n { return 0, errShortARP @@ -120,7 +172,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) { copy(hwsender, h.ourHWAddr) n := copy(b, afrm.Clip().RawData()) tgt, _ := afrm.Target() - trySetEthernetDst(eth[:frameOffset], tgt) + trySetEthernetDst(carrierData[:offsetToFrame], tgt) return n, nil } for i := range h.queries { @@ -139,7 +191,7 @@ func (h *Handler) Encapsulate(eth []byte, frameOffset int) (int, error) { hwTarget[j] = 0 } broadcast := ethernet.BroadcastAddr() - trySetEthernetDst(eth[:frameOffset], broadcast[:]) + trySetEthernetDst(carrierData[:offsetToFrame], broadcast[:]) return n, nil } } @@ -181,8 +233,16 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error { case OpReply: hwaddr, protoaddr := afrm.Sender() for i := range h.queries { - if len(h.queries[i].hwaddr) == 0 && bytes.Equal(h.queries[i].protoaddr, protoaddr) { - h.queries[i].hwaddr = append(h.queries[i].hwaddr[:0], hwaddr...) + q := &h.queries[i] + mac := q.response() + if mac == nil && bytes.Equal(q.protoaddr, protoaddr) { + q.hwaddr = append(q.hwaddr, hwaddr...) + if q.dstHw != nil { + if !allZeros(q.dstHw) { + slog.Error("race-condition:ARP-reused-buffer") + } + copy(q.dstHw, hwaddr) // External write to user buffer. + } return nil } } @@ -198,3 +258,12 @@ func trySetEthernetDst(ethFrame []byte, dst []byte) { copy(ethFrame[:6], dst) } } + +func allZeros(b []byte) bool { + for i := range b { + if b[i] != 0 { + return false + } + } + return true +} diff --git a/dhcpv4/client.go b/dhcpv4/client.go index eb6ab5f..2c73312 100644 --- a/dhcpv4/client.go +++ b/dhcpv4/client.go @@ -104,11 +104,11 @@ func (c *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) } func (c *Client) LocalPort() uint16 { return DefaultClientPort } func (c *Client) ConnectionID() *uint64 { return &c.connID } -func (c *Client) setIP(b []byte, frameOffset int) { - if frameOffset < 28 { - return // Not an IP/UDP frame. +func (c *Client) setIP(carrierFrame []byte, offsetToIP int) { + if offsetToIP < 0 { + return // No IP layer present. } - ifrm, _ := ipv4.NewFrame(b) + ifrm, _ := ipv4.NewFrame(carrierFrame[offsetToIP:]) ifrm.SetID((uint16(c.currentXID) ^ uint16(c.currentXID>>16)) + uint16(c.state)) if c.state > StateInit { // Match server ToS since some routers drop DHCP requests if no ToS set apparently? @@ -124,7 +124,7 @@ func (c *Client) setIP(b []byte, frameOffset int) { } } -func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) { +func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if c.isClosed() { return 0, net.ErrClosed } else if c.state == StateSelecting && !c.offer.valid { @@ -134,7 +134,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) } else if c.state == StateRequesting { return 0, nil // Currently awaiting ACK. } - dst := carrierFrame[frameOffset:] + dst := carrierData[offsetToFrame:] frm, err := NewFrame(dst) if err != nil { return 0, err @@ -194,7 +194,7 @@ func (c *Client) Encapsulate(carrierFrame []byte, frameOffset int) (int, error) opts[numOpts] = byte(OptEnd) numOpts++ c.setHeader(frm) - c.setIP(carrierFrame, frameOffset) + c.setIP(carrierData, offsetToIP) c.state = nextState return OptionsOffset + numOpts, nil } diff --git a/dhcpv4/dhcp_test.go b/dhcpv4/dhcp_test.go index 66f0990..fe69f7a 100644 --- a/dhcpv4/dhcp_test.go +++ b/dhcpv4/dhcp_test.go @@ -28,7 +28,7 @@ func TestClientServer(t *testing.T) { // CLIENT DISCOVER. assertClState(StateInit) var buf [1024]byte - n, err := cl.Encapsulate(buf[:], 0) + n, err := cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -40,7 +40,7 @@ func TestClientServer(t *testing.T) { t.Fatal(err) } // SERVER REPLY OFFER - n, err = sv.Encapsulate(buf[:], 0) + n, err = sv.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -53,7 +53,7 @@ func TestClientServer(t *testing.T) { assertClState(StateSelecting) // CLIENT SEND OUT REQUEST. - n, err = cl.Encapsulate(buf[:], 0) + n, err = cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -66,7 +66,7 @@ func TestClientServer(t *testing.T) { } // SERVER REPLIES WITH ACK. - n, err = sv.Encapsulate(buf[:], 0) + n, err = sv.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { @@ -99,13 +99,13 @@ func TestExample(t *testing.T) { }) buf := make([]byte, 2048) buf2 := make([]byte, len(buf)) - n, err := cl.Encapsulate(buf, 0) + n, err := cl.Encapsulate(buf, -1, 0) if err != nil { t.Fatal(err) } else if n <= 0 { t.Fatal("no data sent out by client after starting request") } - n, err = cl.Encapsulate(buf2, 0) + n, err = cl.Encapsulate(buf2, -1, 0) if err != nil { t.Error("client encaps double tap after discover:", err) } @@ -141,13 +141,13 @@ func TestExample(t *testing.T) { t.Fatal(err) } - n, err = cl.Encapsulate(buf[:], 0) + n, err = cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n <= 0 { t.Fatal("no data written from client in response to offer") } - n, err = cl.Encapsulate(buf[:], 0) + n, err = cl.Encapsulate(buf[:], -1, 0) if err != nil { t.Error("encapsulate double tap after request:", err) } else if n > 0 { diff --git a/dhcpv4/server.go b/dhcpv4/server.go index df2a143..88bceb0 100644 --- a/dhcpv4/server.go +++ b/dhcpv4/server.go @@ -159,9 +159,9 @@ func (sv *Server) Demux(carrierData []byte, frameOffset int) error { return nil } -func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error) { - carrierIsIP := frameOffset >= 28 - dfrm, err := NewFrame(carrierData[frameOffset:]) +func (sv *Server) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + carrierIsIP := offsetToIP >= 0 + dfrm, err := NewFrame(carrierData[offsetToFrame:]) optBuf := dfrm.OptionsPayload()[:] if err != nil { return 0, err @@ -220,7 +220,7 @@ func (sv *Server) Encapsulate(carrierData []byte, frameOffset int) (int, error) copy(dfrm.CHAddrAs6()[:], client.hwaddr[:]) dfrm.SetMagicCookie(MagicCookie) if carrierIsIP { - err = internal.SetIPAddrs(carrierData, 0, sv.siaddr[:], client.addr[:]) + err = internal.SetIPAddrs(carrierData[offsetToIP:], 0, sv.siaddr[:], client.addr[:]) if err != nil { return 0, err } diff --git a/dns/client.go b/dns/client.go index c143112..802e143 100644 --- a/dns/client.go +++ b/dns/client.go @@ -43,7 +43,7 @@ func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error { return nil } -func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { +func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if c.isClosed() { return 0, net.ErrClosed } else if c.state != dnsSendQuery { @@ -51,7 +51,7 @@ func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { } msg := &c.msg - frame := carrierData[frameOffset:] + frame := carrierData[offsetToFrame:] msglen := msg.Len() if msglen > uint16(len(frame)) { return 0, errCalcLen diff --git a/examples/bridge/main.go b/examples/bridge/main.go index 3b00d6e..8a423d1 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -199,7 +199,7 @@ func run() (err error) { prevState = state clear(buf) - nwrite, err := stack.Encapsulate(buf[:], 0) + nwrite, err := stack.Encapsulate(buf[:], -1, 0) if err != nil { fmt.Println("ERR:ENCAPSULATE", err) } else if nwrite > 0 { @@ -267,10 +267,10 @@ func (s *Stack) Demux(b []byte, _ int) (err error) { return s.link.Demux(b, 0) } -func (s *Stack) Encapsulate(b []byte, _ int) (int, error) { - n, err := s.link.Encapsulate(b, 0) +func (s *Stack) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + n, err := s.link.Encapsulate(carrierData, offsetToIP, offsetToFrame) if n > 0 { - iframes, errpcap := s.shark.CaptureEthernet(s.aux[:0], b[:n], 0) + iframes, errpcap := s.shark.CaptureEthernet(s.aux[:0], carrierData[:n], 0) if errpcap != nil { fmt.Println("OU", iframes, errpcap.Error()) } else { @@ -426,7 +426,7 @@ func (s *Stack) StartResolveHardwareAddress6(ip netip.Addr) error { return errors.New("unsupported or invalid IP address") } addr := ip.As4() - return s.arp.StartQuery(addr[:]) + return s.arp.StartQuery(nil, addr[:]) } func (s *Stack) ResultResolveHardwareAddress6(ip netip.Addr) (hw [6]byte, err error) { diff --git a/examples/stack/main.go b/examples/stack/main.go index 8d84603..f81c4c2 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -107,7 +107,7 @@ func main() { } } - nw, err := stack.ethernet.Encapsulate(buf[:], 0) + nw, err := stack.ethernet.Encapsulate(buf[:], -1, 0) if err != nil { lg.Error("handle", slog.String("err", err.Error())) } else if nw > 0 { @@ -230,7 +230,7 @@ func (stack *Stack) Recv(b []byte) error { } func (stack *Stack) Send(b []byte) (int, error) { - return stack.ethernet.Encapsulate(b, 0) + return stack.ethernet.Encapsulate(b, -1, 0) } func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, error) { diff --git a/examples/stackbasic/main.go b/examples/stackbasic/main.go index 8e124bb..39ace54 100644 --- a/examples/stackbasic/main.go +++ b/examples/stackbasic/main.go @@ -232,7 +232,7 @@ func NewEthernetTCPStack(ourMAC, gwMAC [6]byte, ip netip.AddrPort, mtu uint16, s type handler struct { raddr []byte recv func([]byte, int) error - handle func([]byte, int) (int, error) + handle func([]byte, int, int) (int, error) proto ethernet.Type lport uint16 } @@ -295,7 +295,7 @@ func (ls *LinkStack) HandleEth(dst []byte) (n int, err error) { copy(efrm.DestinationHardwareAddr()[:], ls.gwmac[:]) // default set the gateway. for i := range ls.handlers { h := &ls.handlers[i] - n, err = h.handle(dst[:mtu], 14) + n, err = h.handle(dst[:mtu], 14, 14) if err != nil { ls.error("handling", slog.String("proto", ethernet.Type(h.proto).String()), slog.String("err", err.Error())) continue @@ -322,8 +322,8 @@ func (as *ARPStack) Recv(EtherFrame []byte, arpOff int) error { return as.handler.Demux(EtherFrame, arpOff) } -func (as *ARPStack) Handle(EtherFrame []byte, arpOff int) (int, error) { - n, err := as.handler.Encapsulate(EtherFrame, arpOff) +func (as *ARPStack) Handle(EtherFrame []byte, offsetToIP, arpOff int) (int, error) { + n, err := as.handler.Encapsulate(EtherFrame, offsetToIP, arpOff) if err != nil || n == 0 { return 0, err } diff --git a/examples/xnet/main.go b/examples/xnet/main.go index 1408ab2..2061de0 100644 --- a/examples/xnet/main.go +++ b/examples/xnet/main.go @@ -110,7 +110,7 @@ func run() (err error) { var frames []pcap.Frame for { clear(buf) - nwrite, err := stack.Encapsulate(buf[:], 0) + nwrite, err := stack.Encapsulate(buf[:], -1, 0) if err != nil { fmt.Println("ERR:ENCAPSULATE", err) } else if nwrite > 0 { diff --git a/internet/definitions.go b/internet/definitions.go index 683d34b..5b47ee0 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -11,18 +11,21 @@ import ( // StackNode is an abstraction of a packet exchanging protocol controller. This is the building block for all protocols, // from Ethernet to IP to TCP, practically any protocol can be expressed as a StackNode and function completely. type StackNode interface { - // Encapsulate writes the stack node's frame into carrierData[frameOffset:] + // Encapsulate writes the stack node's frame into carrierData[offsetToFrame:] // along with any other frame or payload the stack node encapsulates. - // The returned integer is amount of bytes written such that carrierData[frameOffset:frameOffset+n] - // contains written data. Data inside carrierData[:frameOffset] usually contains data necessary for + // The returned integer is amount of bytes written such that carrierData[offsetToFrame:offsetToFrame+n] + // contains written data. Data inside carrierData[:offsetToFrame] usually contains data necessary for // a StackNode to correctly emit valid frame data: such is the case for TCP packets which require IP // frame data for checksum calculation. Thus StackNodes must provide fields in their own frame // required by sub-stacknodes for correct encapsulation; in the case of IPv4/6 this means including fields // used in pseudo-header checksum like local IP (see [ipv4.CRCWriteUDPPseudo]). // + // offsetToIP is the offset to the IP frame, if present, else its value should be -1. + // The relation offsetToIP<=offsetToFrame should always hold. + // // When [net.ErrClosed] is returned the StackNode should be discarded and any written data passed up normally. // Errors returned by Encapsulate are "extraordinary" and should not be returned unless the StackNode is receiving invalid carrierData/frameOffset. - Encapsulate(carrierData []byte, frameOffset int) (int, error) + Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) // Demux reads from the argument buffer where frameOffset is the offset of this StackNode's frame first byte. // The stack node then dispatches(demuxes) the encapsulated frames to its corresponding sub-stack-node(s). Demux(carrierData []byte, frameOffset int) error @@ -38,10 +41,12 @@ type node struct { currConnID uint64 connID *uint64 demux func([]byte, int) error - encapsulate func([]byte, int) (int, error) + encapsulate func([]byte, int, int) (int, error) proto uint16 port uint16 - remoteAddr []byte + // remoteAddr will be set on active(outbound) port connections + // that require an ARP to set the remoteAddr beforehand. + remoteAddr []byte } type handlers struct { @@ -166,13 +171,13 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro // 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) { +func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame 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, offsetIP, offsetThisFrame) if h.tryHandleError(node, err) { err = nil // CLOSE error handled gracefully by deleting node. } diff --git a/internet/node-tcplistener.go b/internet/node-tcplistener.go index 8cac79e..9b1aa14 100644 --- a/internet/node-tcplistener.go +++ b/internet/node-tcplistener.go @@ -94,7 +94,7 @@ func (listener *NodeTCPListener) TryAccept() (*tcp.Conn, error) { } // Encapsulate implements [StackNode]. -func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset int) (int, error) { +func (listener *NodeTCPListener) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if listener.isClosed() { return 0, net.ErrClosed } @@ -102,7 +102,7 @@ func (listener *NodeTCPListener) Encapsulate(carrierData []byte, tcpFrameOffset if conn == nil { continue } - n, err := conn.Encapsulate(carrierData, tcpFrameOffset) + n, err := conn.Encapsulate(carrierData, offsetToIP, offsetToFrame) if err != nil { err = listener.maintainConn(listener.accepted, i, err) } diff --git a/internet/stack-ethernet.go b/internet/stack-ethernet.go index 46d6794..fd1edcc 100644 --- a/internet/stack-ethernet.go +++ b/internet/stack-ethernet.go @@ -92,9 +92,9 @@ DROP: return lneto.ErrPacketDrop } -func (ls *StackEthernet) Encapsulate(carrierData []byte, frameOffset int) (n int, err error) { +func (ls *StackEthernet) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { mtu := ls.mtu - dst := carrierData[frameOffset:] + dst := carrierData[offsetToFrame:] if len(dst) < int(mtu) { return 0, io.ErrShortBuffer } @@ -104,7 +104,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) + // Children (IP/ARP) start at offset 14 (after ethernet header). + // For IP: offsetToIP=14, offsetToFrame=14 + // For ARP: offsetToIP=-1, offsetToFrame=14 (but ARP ignores offsetToIP) + // Clip carrierData to MTU to prevent writes beyond MTU limit. + mtuLimit := offsetToFrame + int(mtu) + h, n, err = ls.handlers.encapsulateAny(carrierData[:mtuLimit], offsetToFrame+14, offsetToFrame+14) if n == 0 { return n, err } diff --git a/internet/stack-ip.go b/internet/stack-ip.go index e7748e9..a4ec5ce 100644 --- a/internet/stack-ip.go +++ b/internet/stack-ip.go @@ -140,8 +140,8 @@ func (sb *StackIP) Demux(carrierData []byte, offset int) error { return err } -func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) { - frame := carrierData[frameOffset:] +func (sb *StackIP) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { + frame := carrierData[offsetToFrame:] if len(frame) < 256 { return 0, io.ErrShortBuffer } @@ -158,7 +158,9 @@ func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) ifrm.SetTTL(64) *ifrm.SourceAddr() = sb.ip sb.ipID = id - node, n, err := sb.handlers.encapsulateAny(frame, headerlen) + // Children (TCP/UDP) start at offset headerlen (20 bytes after IP header start). + // offsetToIP is 0 relative to this slice (frame), children's frame starts at headerlen. + node, n, err := sb.handlers.encapsulateAny(carrierData, offsetToFrame, offsetToFrame+headerlen) if n == 0 { return n, err } diff --git a/internet/stack-ports.go b/internet/stack-ports.go index 87eece3..dd5e304 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -46,11 +46,11 @@ 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 int(ps.dstPortOff)+offset+2 > len(b) { +func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { + if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) { return 0, io.ErrShortBuffer } - _, n, err = ps.handlers.encapsulateAny(b, offset) + _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) return n, err } diff --git a/internet/stack-udpport.go b/internet/stack-udpport.go index c63ee1b..e64615f 100644 --- a/internet/stack-udpport.go +++ b/internet/stack-udpport.go @@ -61,24 +61,25 @@ func (sudp *StackUDPPort) Demux(carrierData []byte, frameOffset int) error { return err } -func (sudp *StackUDPPort) Encapsulate(carrierData []byte, frameOffset int) (int, error) { +func (sudp *StackUDPPort) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if sudp.h.IsInvalid() { sudp.h.destroy() return 0, net.ErrClosed } - ufrm, err := udp.NewFrame(carrierData[frameOffset:]) + ufrm, err := udp.NewFrame(carrierData[offsetToFrame:]) if err != nil { return 0, err } ufrm.SetSourcePort(sudp.h.port) ufrm.SetDestinationPort(sudp.rmport) - if len(sudp.raddr) > 0 && frameOffset >= 20 { - err = internal.SetIPAddrs(carrierData, 0, nil, sudp.raddr) + if len(sudp.raddr) > 0 && offsetToIP >= 0 { + err = internal.SetIPAddrs(carrierData[offsetToIP:], 0, nil, sudp.raddr) if err != nil { return 0, err } } - n, err := sudp.h.encapsulate(carrierData, frameOffset+8) + // Child payload starts 8 bytes after UDP header start. + n, err := sudp.h.encapsulate(carrierData, offsetToIP, offsetToFrame+8) if n == 0 { if err != nil { slog.Error("stackudp:encapsulate", slog.String("err", err.Error())) diff --git a/internet/stackbasic_test.go b/internet/stackbasic_test.go index 39a704a..2e93702 100644 --- a/internet/stackbasic_test.go +++ b/internet/stackbasic_test.go @@ -44,7 +44,7 @@ func TestBasicStack2(t *testing.T) { func expectExchange(t *testing.T, from, to *StackIP, buf []byte) { t.Helper() - n, err := from.Encapsulate(buf, 0) + n, err := from.Encapsulate(buf, -1, 0) if err != nil { t.Error("expectExchange:encapsulate:", err) } else if n == 0 { diff --git a/ntp/client.go b/ntp/client.go index 8c65018..2d6c387 100644 --- a/ntp/client.go +++ b/ntp/client.go @@ -50,11 +50,11 @@ func (c *Client) ConnectionID() *uint64 { return &c.connID } -func (c *Client) Encapsulate(carrierData []byte, frameOffset int) (int, error) { +func (c *Client) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { if c.IsDone() { return 0, nil } - payload := carrierData[frameOffset:] + payload := carrierData[offsetToFrame:] frm, err := NewFrame(payload) if err != nil { return 0, err diff --git a/tcp/conn.go b/tcp/conn.go index 62c7d7a..3dff2db 100644 --- a/tcp/conn.go +++ b/tcp/conn.go @@ -278,23 +278,27 @@ func (conn *Conn) Demux(buf []byte, off int) (err error) { return nil } -func (conn *Conn) Encapsulate(buf []byte, off int) (n int, err error) { +func (conn *Conn) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { conn.mu.Lock() defer conn.mu.Unlock() if len(conn.remoteAddr) == 0 { return 0, errNoRemoteAddr } - raddr, _, _, _, err := internal.GetIPAddr(buf[:off]) + if offsetToIP < 0 { + return 0, errNoRemoteAddr // No IP layer present. + } + ipFrame := carrierData[offsetToIP:offsetToFrame] + raddr, _, _, _, err := internal.GetIPAddr(ipFrame) if err != nil { return 0, err } else if len(raddr) != len(conn.remoteAddr) { return 0, errMismatchedIPVersion } - n, err = conn.h.Send(buf[off:]) + n, err = conn.h.Send(carrierData[offsetToFrame:]) if err != nil { return 0, err } - err = internal.SetIPAddrs(buf[:off], conn.ipID, nil, conn.remoteAddr) + err = internal.SetIPAddrs(ipFrame, conn.ipID, nil, conn.remoteAddr) if err != nil { return 0, err } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 6e4e09b..7946153 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -73,11 +73,11 @@ func (s *StackAsync) Demux(carrierData []byte, etherOff int) error { return s.link.Demux(carrierData, etherOff) } -func (s *StackAsync) Encapsulate(carrierData []byte, etherOff int) (int, error) { +func (s *StackAsync) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) { s.mu.Lock() defer s.mu.Unlock() - n, err := s.link.Encapsulate(carrierData, etherOff) + n, err := s.link.Encapsulate(carrierData, offsetToIP, offsetToFrame) s.totalsent += uint64(n) return n, err } @@ -376,7 +376,7 @@ func (s *StackAsync) StartResolveHardwareAddress6(ip netip.Addr) error { return errors.New("unsupported or invalid IP address") } addr := ip.As4() - return s.arp.StartQuery(addr[:]) + return s.arp.StartQuery(nil, addr[:]) } // ResultResolveHardwareAddress6 diff --git a/x/xnet/xnet_test.go b/x/xnet/xnet_test.go index 086bf99..d066224 100644 --- a/x/xnet/xnet_test.go +++ b/x/xnet/xnet_test.go @@ -322,7 +322,7 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs default: panic("OOB") } - n, err := src.Encapsulate(buf[:], 0) + n, err := src.Encapsulate(buf[:], -1, 0) if err != nil { t.Fatal(err) } else if n == 0 { From 5e631a1b3dd313c866ec7cc514fe567acef9af94 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Fri, 19 Dec 2025 01:04:18 -0300 Subject: [PATCH 05/11] StackPorts: set MAC destination when node has MAC --- examples/bridge/main.go | 6 +++--- examples/stack/main.go | 4 ++-- internet/definitions.go | 13 +++++++++++-- internet/stack-ports.go | 17 ++++++++++++++--- x/xnet/stack-async.go | 22 +++++++++++++++++----- 5 files changed, 47 insertions(+), 15 deletions(-) diff --git a/examples/bridge/main.go b/examples/bridge/main.go index 8a423d1..b8ed5d2 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -359,7 +359,7 @@ func (s *Stack) StartLookupIP(host string) error { var u internet.StackUDPPort dns4 := dnsSrvs.As4() u.SetStackNode(&s.dns, dns4[:], dns.ServerPort) - err = s.udps.Register(&u) + err = s.udps.Register(&u, nil) if err != nil { return err } @@ -405,7 +405,7 @@ func (s *Stack) BeginDHCPRequest(request [4]byte) error { } var u internet.StackUDPPort u.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort) - err = s.udps.Register(&u) + err = s.udps.Register(&u, nil) if err != nil { return err } @@ -417,7 +417,7 @@ func (s *Stack) StartNTP(addr netip.Addr) error { var u internet.StackUDPPort addr4 := addr.As4() u.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort) - err := s.udps.Register(&u) + err := s.udps.Register(&u, nil) return err } diff --git a/examples/stack/main.go b/examples/stack/main.go index f81c4c2..b248202 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -239,7 +239,7 @@ func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, err if err != nil { return nil, err } - err = stack.tcpports.Register(&listener) + err = stack.tcpports.Register(&listener, nil) // Passive TCP requires no MAC setting. if err != nil { return nil, err } @@ -261,7 +261,7 @@ func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*tcp.Conn, error if err != nil { return nil, err } - err = stack.tcpports.Register(conn) + err = stack.tcpports.Register(conn, nil) // Passive MAC with no listening. if err != nil { return nil, err } diff --git a/internet/definitions.go b/internet/definitions.go index 5b47ee0..47e0ee1 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -174,7 +174,7 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ *node, n int, err error) { for i := range h.nodes { node := &h.nodes[i] - if node.IsInvalid() { + if node.IsInvalid() || (len(node.remoteAddr) > 0 && isAllZeros(node.remoteAddr)) { continue } n, err = node.encapsulate(buf, offsetIP, offsetThisFrame) @@ -191,6 +191,15 @@ func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ return nil, 0, err // Return last written error. } +func isAllZeros(b []byte) bool { + for i := range b { + if b[i] != 0 { + return false + } + } + return true +} + var ( errZeroMaxNodesArg = errors.New("zero max nodes arg") errZeroPort = errors.New("port must be greater than zero") @@ -224,7 +233,7 @@ func nodeFromStackNode(s StackNode, port uint16, protocol uint64, remoteAddr []b encapsulate: s.Encapsulate, proto: uint16(protocol), port: port, - remoteAddr: append([]byte{}, remoteAddr...), + remoteAddr: remoteAddr, // SHARED MEMORY- used to signal. } } diff --git a/internet/stack-ports.go b/internet/stack-ports.go index dd5e304..d98faa6 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -2,11 +2,13 @@ package internet import ( "encoding/binary" + "errors" "io" "math" "strconv" "github.com/soypat/lneto" + "github.com/soypat/lneto/ethernet" ) type StackPorts struct { @@ -50,7 +52,12 @@ func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) { return 0, io.ErrShortBuffer } - _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) + var node *node + node, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) + if n > 0 && len(node.remoteAddr) == 6 && offsetToIP >= 14 { + efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:]) + *efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr) + } return n, err } @@ -63,13 +70,17 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { return err } -func (ps *StackPorts) Register(h StackNode) error { +// Register registers a port StackNode on StackPorts. +// If dstMAC is set to non-nil, length six buffer then +func (ps *StackPorts) Register(h StackNode, dstMAC []byte) error { port := h.LocalPort() proto := h.Protocol() if port <= 0 { return errZeroPort } else if proto != uint64(ps.protocol) { return errInvalidProto + } else if dstMAC != nil && len(dstMAC) != 6 { + return errors.New("invalid MAC") } - return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil)) + return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, dstMAC)) } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 7946153..13e604e 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -234,11 +234,23 @@ func (s *StackAsync) Gateway6() [6]byte { func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrPort) (err error) { s.mu.Lock() defer s.mu.Unlock() + var mac []byte + if s.dhcpResults.Subnet.Contains(addrp.Addr()) { + mac = make([]byte, 6) + ip := addrp.Addr().As4() + // StartQuery starts an ARP query for addresses in this network. + // On finishing query MAC is set and thus the StackPort will allow encapsulating + // data on that connection. + err = s.arp.StartQuery(mac, ip[:]) + if err != nil { + return err + } + } err = conn.OpenActive(localPort, addrp, tcp.Value(s.Prand32())) if err != nil { return err } - err = s.tcps.Register(conn) + err = s.tcps.Register(conn, mac) // MAC is set later on by ARP response arriving to our network. if err != nil { conn.Abort() return err @@ -253,7 +265,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) { if err != nil { return err } - err = s.tcps.Register(conn) + err = s.tcps.Register(conn, nil) if err != nil { conn.Abort() return err @@ -294,7 +306,7 @@ func (s *StackAsync) StartLookupIP(host string) error { } dns4 := s.dnssv.As4() s.dnsUDP.SetStackNode(&s.dns, dns4[:], dns.ServerPort) - err = s.udps.Register(&s.dnsUDP) + err = s.udps.Register(&s.dnsUDP, nil) return err } @@ -341,7 +353,7 @@ func (s *StackAsync) StartDHCPv4Request(request [4]byte) error { } s.dhcpUDP.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort) - err = s.udps.Register(&s.dhcpUDP) + err = s.udps.Register(&s.dhcpUDP, nil) if err != nil { return err } @@ -355,7 +367,7 @@ func (s *StackAsync) StartNTP(addr netip.Addr) error { addr4 := addr.As4() s.ntpUDP.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort) - err := s.udps.Register(&s.ntpUDP) + err := s.udps.Register(&s.ntpUDP, nil) return err } From f6bc73ee6ca471c6d3330dbaa00f59e0cd05f3b8 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Mon, 22 Dec 2025 18:41:14 -0300 Subject: [PATCH 06/11] apply some of @ddirect suggestions --- arp/handler.go | 14 ++----- examples/bridge/main.go | 6 +-- examples/stack/main.go | 4 +- internal/ip.go | 11 ++++++ internet/definitions.go | 15 ++------ internet/stack-ports.go | 85 ++++++++++++++++++++++++++++++++++++----- x/xnet/stack-async.go | 10 ++--- 7 files changed, 104 insertions(+), 41 deletions(-) diff --git a/arp/handler.go b/arp/handler.go index aca916b..4796735 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -7,6 +7,7 @@ import ( "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" ) type Handler struct { @@ -143,7 +144,7 @@ func (h *Handler) StartQuery(dstHWAddr, proto []byte) error { return errors.New("bad protocol address length") } else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) { return errors.New("mismatch hardware size") - } else if dstHWAddr != nil && !allZeros(dstHWAddr) { + } else if dstHWAddr != nil && !internal.IsZeroed(dstHWAddr...) { return errors.New("write-to buffer must be zeroed out") } h.queries = h.queries[:len(h.queries)+1] @@ -238,7 +239,7 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error { if mac == nil && bytes.Equal(q.protoaddr, protoaddr) { q.hwaddr = append(q.hwaddr, hwaddr...) if q.dstHw != nil { - if !allZeros(q.dstHw) { + if !internal.IsZeroed(q.dstHw...) { slog.Error("race-condition:ARP-reused-buffer") } copy(q.dstHw, hwaddr) // External write to user buffer. @@ -258,12 +259,3 @@ func trySetEthernetDst(ethFrame []byte, dst []byte) { copy(ethFrame[:6], dst) } } - -func allZeros(b []byte) bool { - for i := range b { - if b[i] != 0 { - return false - } - } - return true -} diff --git a/examples/bridge/main.go b/examples/bridge/main.go index b8ed5d2..8a423d1 100644 --- a/examples/bridge/main.go +++ b/examples/bridge/main.go @@ -359,7 +359,7 @@ func (s *Stack) StartLookupIP(host string) error { var u internet.StackUDPPort dns4 := dnsSrvs.As4() u.SetStackNode(&s.dns, dns4[:], dns.ServerPort) - err = s.udps.Register(&u, nil) + err = s.udps.Register(&u) if err != nil { return err } @@ -405,7 +405,7 @@ func (s *Stack) BeginDHCPRequest(request [4]byte) error { } var u internet.StackUDPPort u.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort) - err = s.udps.Register(&u, nil) + err = s.udps.Register(&u) if err != nil { return err } @@ -417,7 +417,7 @@ func (s *Stack) StartNTP(addr netip.Addr) error { var u internet.StackUDPPort addr4 := addr.As4() u.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort) - err := s.udps.Register(&u, nil) + err := s.udps.Register(&u) return err } diff --git a/examples/stack/main.go b/examples/stack/main.go index b248202..9370c3e 100644 --- a/examples/stack/main.go +++ b/examples/stack/main.go @@ -239,7 +239,7 @@ func (stack *Stack) OpenTCPListener(port uint16) (*internet.NodeTCPListener, err if err != nil { return nil, err } - err = stack.tcpports.Register(&listener, nil) // Passive TCP requires no MAC setting. + err = stack.tcpports.Register(&listener) // Passive TCP requires no MAC setting. if err != nil { return nil, err } @@ -261,7 +261,7 @@ func (stack *Stack) OpenPassiveTCP(port uint16, iss tcp.Value) (*tcp.Conn, error if err != nil { return nil, err } - err = stack.tcpports.Register(conn, nil) // Passive MAC with no listening. + err = stack.tcpports.Register(conn) // Passive MAC with no listening. if err != nil { return nil, err } diff --git a/internal/ip.go b/internal/ip.go index ac86109..77a9263 100644 --- a/internal/ip.go +++ b/internal/ip.go @@ -56,3 +56,14 @@ func SetIPAddrs(buf []byte, id uint16, src, dst []byte) (err error) { copy(dstaddr, dst) return nil } + +// IsZeroed returns true if all arguments are set to their zero value. +func IsZeroed[T comparable](a ...T) bool { + var z T + for i := range a { + if a[i] != z { + return false + } + } + return true +} diff --git a/internet/definitions.go b/internet/definitions.go index 47e0ee1..8246f97 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -165,6 +165,7 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro err := node.demux(buf, offset) if h.tryHandleError(node, err) { err = nil + node = nil // Node is destroyed in tryHandleError and invalidated. } return node, err } @@ -174,12 +175,13 @@ func (h *handlers) demuxByPort(buf []byte, offset int, port uint16) (*node, erro func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ *node, n int, err error) { for i := range h.nodes { node := &h.nodes[i] - if node.IsInvalid() || (len(node.remoteAddr) > 0 && isAllZeros(node.remoteAddr)) { + if node.IsInvalid() { continue } n, err = node.encapsulate(buf, offsetIP, offsetThisFrame) if h.tryHandleError(node, err) { - err = nil // CLOSE error handled gracefully by deleting node. + err = nil // CLOSE error handled gracefully by deleting node. + node = nil // Node is destroyed in tryHandleError and invalidated. } if n > 0 { return node, n, err @@ -191,15 +193,6 @@ func (h *handlers) encapsulateAny(buf []byte, offsetIP, offsetThisFrame int) (_ return nil, 0, err // Return last written error. } -func isAllZeros(b []byte) bool { - for i := range b { - if b[i] != 0 { - return false - } - } - return true -} - var ( errZeroMaxNodesArg = errors.New("zero max nodes arg") errZeroPort = errors.New("port must be greater than zero") diff --git a/internet/stack-ports.go b/internet/stack-ports.go index d98faa6..413cc09 100644 --- a/internet/stack-ports.go +++ b/internet/stack-ports.go @@ -4,11 +4,13 @@ import ( "encoding/binary" "errors" "io" + "log/slog" "math" "strconv" "github.com/soypat/lneto" "github.com/soypat/lneto/ethernet" + "github.com/soypat/lneto/internal" ) type StackPorts struct { @@ -16,6 +18,7 @@ type StackPorts struct { handlers handlers dstPortOff uint16 protocol uint16 + // stores last node to demux/encapsulate. } func (ps *StackPorts) ResetUDP(maxNodes int) error { @@ -52,12 +55,7 @@ func (ps *StackPorts) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame if int(ps.dstPortOff)+offsetToFrame+2 > len(carrierData) { return 0, io.ErrShortBuffer } - var node *node - node, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) - if n > 0 && len(node.remoteAddr) == 6 && offsetToIP >= 14 { - efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:]) - *efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr) - } + _, n, err = ps.handlers.encapsulateAny(carrierData, offsetToIP, offsetToFrame) return n, err } @@ -72,15 +70,84 @@ func (ps *StackPorts) Demux(b []byte, offset int) (err error) { // Register registers a port StackNode on StackPorts. // If dstMAC is set to non-nil, length six buffer then -func (ps *StackPorts) Register(h StackNode, dstMAC []byte) error { +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 - } else if dstMAC != nil && len(dstMAC) != 6 { + } + return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, nil)) +} + +// StackPortsMACFiltered is a StackPorts implementation but that avoids calling encapsulate on nodes +// with a non-nil MAC address registered via Register method that is set to all zero values. +// If the address is set to nil no filtering occurs. MAC Address is set automatically on the ethernet frame by StackPortsMACFiltered when non-nil. +type StackPortsMACFiltered struct { + sp StackPorts +} + +func (mfsp *StackPortsMACFiltered) Register(h StackNode, addr []byte) error { + port := h.LocalPort() + proto := h.Protocol() + if port <= 0 { + return errZeroPort + } else if proto != uint64(mfsp.sp.protocol) { + return errInvalidProto + } else if addr != nil && len(addr) != 6 { return errors.New("invalid MAC") } - return ps.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, dstMAC)) + return mfsp.sp.handlers.registerByPortProto(nodeFromStackNode(h, port, proto, addr)) +} + +func (ps *StackPortsMACFiltered) ResetUDP(maxNodes int) error { + return ps.sp.ResetUDP(maxNodes) +} + +func (ps *StackPortsMACFiltered) ResetTCP(maxNodes int) error { + return ps.sp.ResetTCP(maxNodes) +} + +func (ps *StackPortsMACFiltered) Reset(protocol uint64, dstPortOffset uint16, maxNodes int) error { + return ps.sp.Reset(protocol, dstPortOffset, maxNodes) +} + +func (ps *StackPortsMACFiltered) LocalPort() uint16 { return 0 } + +func (ps *StackPortsMACFiltered) Protocol() uint64 { return uint64(ps.sp.protocol) } + +func (ps *StackPortsMACFiltered) ConnectionID() *uint64 { return &ps.sp.connID } + +func (ps *StackPortsMACFiltered) Demux(b []byte, offset int) (err error) { + // No MAC Filtering on ingress. TODO? + return ps.sp.Demux(b, offset) +} + +func (ps *StackPortsMACFiltered) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) { + if int(ps.sp.dstPortOff)+offsetToFrame+2 > len(carrierData) { + return 0, io.ErrShortBuffer + } + h := &ps.sp.handlers + for i := range h.nodes { + node := &h.nodes[i] + if node.IsInvalid() || (len(node.remoteAddr) > 0 && internal.IsZeroed(node.remoteAddr...)) { + continue + } + n, err = node.encapsulate(carrierData, offsetToIP, offsetToFrame) + if h.tryHandleError(node, err) { + err = nil // CLOSE error handled gracefully by deleting node. + } + if n > 0 { + if len(node.remoteAddr) == 6 && offsetToIP >= 14 { + efrm, _ := ethernet.NewFrame(carrierData[offsetToIP-14:]) + *efrm.DestinationHardwareAddr() = [6]byte(node.remoteAddr) + } + return 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 0, err // Return last written error. } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 13e604e..9a3c5b6 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -250,7 +250,7 @@ func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrP if err != nil { return err } - err = s.tcps.Register(conn, mac) // MAC is set later on by ARP response arriving to our network. + err = s.tcps.Register(conn) // MAC is set later on by ARP response arriving to our network. if err != nil { conn.Abort() return err @@ -265,7 +265,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) { if err != nil { return err } - err = s.tcps.Register(conn, nil) + err = s.tcps.Register(conn) if err != nil { conn.Abort() return err @@ -306,7 +306,7 @@ func (s *StackAsync) StartLookupIP(host string) error { } dns4 := s.dnssv.As4() s.dnsUDP.SetStackNode(&s.dns, dns4[:], dns.ServerPort) - err = s.udps.Register(&s.dnsUDP, nil) + err = s.udps.Register(&s.dnsUDP) return err } @@ -353,7 +353,7 @@ func (s *StackAsync) StartDHCPv4Request(request [4]byte) error { } s.dhcpUDP.SetStackNode(&s.dhcp, nil, dhcpv4.DefaultServerPort) - err = s.udps.Register(&s.dhcpUDP, nil) + err = s.udps.Register(&s.dhcpUDP) if err != nil { return err } @@ -367,7 +367,7 @@ func (s *StackAsync) StartNTP(addr netip.Addr) error { addr4 := addr.As4() s.ntpUDP.SetStackNode(&s.ntp, addr4[:], ntp.ServerPort) - err := s.udps.Register(&s.ntpUDP, nil) + err := s.udps.Register(&s.ntpUDP) return err } From 48847af1e4a39e1b6738ed8c4418bde2b0a45d78 Mon Sep 17 00:00:00 2001 From: ddirect <62487612+ddirect@users.noreply.github.com> Date: Sun, 28 Dec 2025 18:13:40 +0200 Subject: [PATCH 07/11] Preliminary patch to avoid ARP to stop working after resetting it. --- arp/handler.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/arp/handler.go b/arp/handler.go index 4796735..4ba5a86 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -43,7 +43,7 @@ func (h *Handler) Reset(cfg HandlerConfig) error { return errors.New("invalid Handler query or pending config") } *h = Handler{ - connID: h.connID + 1, + connID: h.connID, ourHWAddr: h.ourHWAddr[:0], ourProtoAddr: h.ourProtoAddr[:0], htype: cfg.HardwareType, From b9fbe0d25ddd3e67c9230080fd29c6de817e365b Mon Sep 17 00:00:00 2001 From: ddirect <62487612+ddirect@users.noreply.github.com> Date: Sun, 28 Dec 2025 18:30:47 +0200 Subject: [PATCH 08/11] Minor fixes --- arp/handler.go | 2 +- x/xnet/stack-async.go | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/arp/handler.go b/arp/handler.go index 4ba5a86..45748bf 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -255,7 +255,7 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error { } func trySetEthernetDst(ethFrame []byte, dst []byte) { - if len(ethFrame) > 14 { + if len(ethFrame) >= 14 { copy(ethFrame[:6], dst) } } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 9a3c5b6..6be8eaa 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -28,7 +28,7 @@ type StackAsync struct { ip internet.StackIP arp arp.Handler udps internet.StackPorts - tcps internet.StackPorts + tcps internet.StackPortsMACFiltered dhcpUDP internet.StackUDPPort dhcp dhcpv4.Client @@ -250,7 +250,7 @@ func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrP if err != nil { return err } - err = s.tcps.Register(conn) // MAC is set later on by ARP response arriving to our network. + err = s.tcps.Register(conn, mac) // MAC is set later on by ARP response arriving to our network. if err != nil { conn.Abort() return err @@ -265,7 +265,7 @@ func (s *StackAsync) ListenTCP(conn *tcp.Conn, localPort uint16) (err error) { if err != nil { return err } - err = s.tcps.Register(conn) + err = s.tcps.Register(conn, nil) if err != nil { conn.Abort() return err From a2886423a2fd4668811328588aad05a402c66321 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Sun, 28 Dec 2025 15:33:10 -0300 Subject: [PATCH 09/11] fix StackAsync losing ARP handler on reset during DHCP result assimilation --- x/xnet/stack-async.go | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 9a3c5b6..016d883 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -169,7 +169,7 @@ func (s *StackAsync) resetARP() error { if addr.Is6() { proto = ethernet.TypeIPv6 } - return s.arp.Reset(arp.HandlerConfig{ + err := s.arp.Reset(arp.HandlerConfig{ HardwareAddr: mac[:], ProtocolAddr: addr.AsSlice(), MaxQueries: 3, @@ -177,6 +177,14 @@ func (s *StackAsync) resetARP() error { HardwareType: 1, ProtocolType: proto, }) + if err != nil { + return err + } + err = s.link.Register(&s.arp) + if err != nil { + return err + } + return nil } // Prand32 generates a pseudo random 32-bit unsigned integer from the internal state and advances the seed. From 890fd33b22b98c80b3da9c4140b689dae8c3a60c Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Sun, 28 Dec 2025 15:42:10 -0300 Subject: [PATCH 10/11] do not increment ARP connectionID on IP addr update --- arp/handler.go | 8 ++++++++ x/xnet/stack-async.go | 11 ++++------- 2 files changed, 12 insertions(+), 7 deletions(-) diff --git a/arp/handler.go b/arp/handler.go index 4796735..adaae36 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -35,6 +35,14 @@ func (h *Handler) Protocol() uint64 { return uint64(ethernet.TypeARP) } func (h *Handler) ConnectionID() *uint64 { return &h.connID } +func (h *Handler) UpdateProtoAddr(protoAddr []byte) error { + if len(protoAddr) != len(h.ourProtoAddr) { + return errors.New("mismatch ARP proto size") + } + copy(h.ourProtoAddr, protoAddr) + return nil +} + func (h *Handler) Reset(cfg HandlerConfig) error { if len(cfg.HardwareAddr) == 0 || len(cfg.HardwareAddr) > 255 || len(cfg.ProtocolAddr) == 0 || len(cfg.ProtocolAddr) > 255 { diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 016d883..6e89363 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -205,7 +205,9 @@ func (s *StackAsync) SetIPAddr(addr netip.Addr) error { if err != nil { return err } - return s.resetARP() + ip := addr.As4() + err = s.arp.UpdateProtoAddr(ip[:]) + return err } func (s *StackAsync) Addr() netip.Addr { @@ -456,12 +458,7 @@ func (stack *StackAsync) AssimilateDHCPResults(results *DHCPResults) error { stack.mu.Lock() defer stack.mu.Unlock() if results.AssignedAddr.IsValid() { - err := stack.ip.SetAddr(results.AssignedAddr) - if err != nil { - return err - } - // Reset ARP handler with new IP address so it can respond to ARP requests. - err = stack.resetARP() + err := stack.SetIPAddr(results.AssignedAddr) if err != nil { return err } From 8307403c5c5814815e9c39fa12f80dba4bbc86e1 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Tue, 30 Dec 2025 18:00:50 -0300 Subject: [PATCH 11/11] add subnet to StackAsync and add test for local arp resolving --- arp/handler.go | 2 +- internet/definitions.go | 6 +- x/xnet/stack-async.go | 19 +++-- x/xnet/xnet_arp_test.go | 49 ++++++++++++ x/xnet/xnet_test.go | 166 ++++++++++++++++++++++++++++++++++++++-- 5 files changed, 224 insertions(+), 18 deletions(-) create mode 100644 x/xnet/xnet_arp_test.go diff --git a/arp/handler.go b/arp/handler.go index 897e5d1..cce5e52 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -51,7 +51,7 @@ func (h *Handler) Reset(cfg HandlerConfig) error { return errors.New("invalid Handler query or pending config") } *h = Handler{ - connID: h.connID, + connID: h.connID + 1, ourHWAddr: h.ourHWAddr[:0], ourProtoAddr: h.ourProtoAddr[:0], htype: cfg.HardwareType, diff --git a/internet/definitions.go b/internet/definitions.go index 8246f97..010a896 100644 --- a/internet/definitions.go +++ b/internet/definitions.go @@ -118,7 +118,7 @@ func (h *handlers) tryHandleError(node *node, err error) (discardedGracefully bo func (h *handlers) nodeByProto(proto uint16) *node { for i := range h.nodes { node := &h.nodes[i] - if node.proto == proto { + if node.proto == proto && !node.IsInvalid() { return node } } @@ -128,7 +128,7 @@ func (h *handlers) nodeByProto(proto uint16) *node { func (h *handlers) nodeByPort(port uint16) *node { for i := range h.nodes { node := &h.nodes[i] - if node.port == port { + if node.port == port && !node.IsInvalid() { return node } } @@ -138,7 +138,7 @@ func (h *handlers) nodeByPort(port uint16) *node { 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 { + if node.port == port && node.proto == protocol && !node.IsInvalid() { return node } } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 5533736..84430c9 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -33,6 +33,7 @@ type StackAsync struct { dhcpUDP internet.StackUDPPort dhcp dhcpv4.Client dhcpResults DHCPResults + subnet netip.Prefix // Local subnet for ARP resolution. dnsUDP internet.StackUDPPort dns dns.Client @@ -112,6 +113,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error { if err != nil { return err } + // err = s.resetARP() if err != nil { return err @@ -135,10 +137,7 @@ func (s *StackAsync) Reset(cfg StackConfig) error { } // Now setup stacks. - err = s.link.Register(&s.arp) // ARP. - if err != nil { - return err - } + // ARP registered in resetARP. err = s.link.Register(&s.ip) // IPv4 | IPv6 if err != nil { return err @@ -201,6 +200,10 @@ func (s *StackAsync) Prand32() uint32 { func (s *StackAsync) SetIPAddr(addr netip.Addr) error { s.mu.Lock() defer s.mu.Unlock() + return s.setIPAddr(addr) +} + +func (s *StackAsync) setIPAddr(addr netip.Addr) error { err := s.ip.SetAddr(addr) if err != nil { return err @@ -245,7 +248,7 @@ func (s *StackAsync) DialTCP(conn *tcp.Conn, localPort uint16, addrp netip.AddrP s.mu.Lock() defer s.mu.Unlock() var mac []byte - if s.dhcpResults.Subnet.Contains(addrp.Addr()) { + if s.subnet.Contains(addrp.Addr()) { mac = make([]byte, 6) ip := addrp.Addr().As4() // StartQuery starts an ARP query for addresses in this network. @@ -454,11 +457,15 @@ func (s *StackAsync) ReadStatistics(stats *Statistics) { // AssimilateDHCPResults sets the stack's following parameters: // - IPv4 address. // - DNS server. +// - Subnet (for ARP resolution of local addresses). func (stack *StackAsync) AssimilateDHCPResults(results *DHCPResults) error { stack.mu.Lock() defer stack.mu.Unlock() + if results.Subnet.IsValid() { + stack.subnet = results.Subnet + } if results.AssignedAddr.IsValid() { - err := stack.SetIPAddr(results.AssignedAddr) + err := stack.setIPAddr(results.AssignedAddr) if err != nil { return err } diff --git a/x/xnet/xnet_arp_test.go b/x/xnet/xnet_arp_test.go new file mode 100644 index 0000000..859135e --- /dev/null +++ b/x/xnet/xnet_arp_test.go @@ -0,0 +1,49 @@ +package xnet + +import ( + "bytes" + "net/netip" + "testing" +) + +func TestARPLocal(t *testing.T) { + const mtu = 1500 + const seed = 1 + s1, s2, c1, c2 := newTCPStacks(t, seed, mtu) + routerHw := [6]byte{1, 2, 3, 4, 5, 6} + // Most common case: we have a router in between computers. + s1.SetGateway6(routerHw) + s2.SetGateway6(routerHw) + addr1 := netip.AddrPortFrom(s1.Addr(), 1024) // dialer, client. + addr2 := netip.AddrPortFrom(s2.Addr(), 80) // listener, server. + err := s1.AssimilateDHCPResults(&DHCPResults{ + Router: netip.AddrFrom4([4]byte{10, 0, 0, 255}), + BroadcastAddr: netip.AddrFrom4([4]byte{255, 255, 255, 255}), + AssignedAddr: s1.Addr(), + Subnet: netip.PrefixFrom(s2.Addr(), 24), // Subnet containing s2 will force an ARP on s1. + TRenewal: 1000, + TRebind: 1000, + TLease: 1000, + }) + if err != nil { + t.Fatal(err) + } + hw2 := s2.HardwareAddress() + err = s1.DialTCP(c1, addr1.Port(), addr2) // addr2 MAC address is unknown and must be resolved by stack. + if err != nil { + t.Fatal(err) + } + err = s2.ListenTCP(c2, addr2.Port()) + if err != nil { + t.Fatal(err) + } + tst := testerFrom(t, mtu) + _ = tst + tst.ARPExchangeOnly(s1, s2) + hwaddr, err := s1.arp.QueryResult(addr2.Addr().AsSlice()) + if err != nil { + t.Fatal(err) + } else if !bytes.Equal(hwaddr[:], hw2[:]) { + t.Errorf("expected hardware address %x, got %x", hw2, hwaddr) + } +} diff --git a/x/xnet/xnet_test.go b/x/xnet/xnet_test.go index d066224..3e19ddc 100644 --- a/x/xnet/xnet_test.go +++ b/x/xnet/xnet_test.go @@ -2,11 +2,13 @@ package xnet import ( "bytes" + "encoding/binary" "errors" "math/rand" "net/netip" "testing" + "github.com/soypat/lneto/arp" "github.com/soypat/lneto/ethernet" "github.com/soypat/lneto/internet/pcap" "github.com/soypat/lneto/tcp" @@ -24,9 +26,7 @@ func TestStackAsyncTCP_multipacket(t *testing.T) { const svPort = 8080 const maxPktLen = 30 client, sv, clconn, svconn := newTCPStacks(t, seed, MTU) - tst := tester{ - t: t, buf: make([]byte, MTU), - } + tst := testerFrom(t, MTU) rng := rand.New(rand.NewSource(seed)) client2, sv2, clconn2, svconn2 := newTCPStacks(t, seed, MTU) _, _, _, _ = client2, sv2, clconn2, svconn2 @@ -58,10 +58,7 @@ func TestStackAsyncTCP_singlepacket(t *testing.T) { const MTU = 1500 const svPort = 80 client, sv, clconn, svconn := newTCPStacks(t, seed, MTU) - - tst := tester{ - t: t, buf: make([]byte, MTU), - } + tst := testerFrom(t, MTU) tst.TestTCPSetupAndEstablish(sv, client, svconn, clconn, svPort, 1337) sendData := []byte("hello") @@ -80,7 +77,7 @@ func TestStackAsyncTCP_singlepacket(t *testing.T) { func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1, c2 *tcp.Conn) { s1, s2 = new(StackAsync), new(StackAsync) c1, c2 = new(tcp.Conn), new(tcp.Conn) - byte1 := byte(randSeed) / 4 + byte1 := byte(randSeed)/4 - 1 err := s1.Reset(StackConfig{ Hostname: "Stack1", RandSeed: randSeed, @@ -127,6 +124,13 @@ func newTCPStacks(t *testing.T, randSeed int64, mtu int) (s1, s2 *StackAsync, c1 return s1, s2, c1, c2 } +func testerFrom(t *testing.T, mtu int) *tester { + return &tester{ + t: t, + buf: make([]byte, mtu), + } +} + type tester struct { t *testing.T cap pcap.PacketBreakdown @@ -334,6 +338,7 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs t.Error("expected no data sent and got data") return } + defer setzero(buf[:n]) tst.buf = tst.buf[:n] tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0) @@ -374,7 +379,121 @@ func (tst *tester) TCPExchange(expect tcpExpectExchange, stack1, stack2 *StackAs if err != nil { t.Fatal(err) } +} + +func (tst *tester) ARPExchangeOnly(querying, target *StackAsync) { + t := tst.t + t.Helper() + buf := tst.buf[:cap(tst.buf)] + + // === PHASE 1: ARP Request from querying stack === + n, err := querying.Encapsulate(buf[:], -1, 0) + if err != nil { + t.Fatal(err) + } else if n == 0 { + t.Error("zero bits sent by ARP querying stack") + return + } + + tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0) + if err != nil { + t.Fatal(err) + } + tst.buf = tst.buf[:n] + + qHw := querying.HardwareAddress() + tgtHw := target.HardwareAddress() + broadcast := ethernet.BroadcastAddr() + qIP := querying.Addr() + tgtIP := target.Addr() + + // Validate Ethernet layer (request is broadcast) + if !bytes.Equal(qHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) { + t.Errorf("request: mismatched ethernet src addr %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) + } + if !bytes.Equal(broadcast[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) { + t.Errorf("request: expected broadcast ethernet dst addr, got %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) + } + + // Validate ARP request fields + // ARP fields: FieldClassSrc with 6 octets = HW addr, 4 octets = proto addr + // occurrence 0 = sender, occurrence 1 = target + if tst.getARPOperation() != arp.OpRequest { + t.Errorf("request: expected ARP OpRequest, got %d", tst.getARPOperation()) + } + if !bytes.Equal(qHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 0)) { + t.Errorf("request: mismatched ARP sender HW") + } + if !bytes.Equal(qIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 0)) { + t.Errorf("request: mismatched ARP sender proto") + } + if !bytes.Equal(tgtIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 1)) { + t.Errorf("request: mismatched ARP target proto") + } + + // Deliver request to target + err = target.Demux(buf[:n], 0) + if err != nil { + t.Fatal("target demux request:", err) + } setzero(buf[:n]) + + // === PHASE 2: ARP Reply from target stack === + buf = tst.buf[:cap(tst.buf)] + n, err = target.Encapsulate(buf[:], -1, 0) + if err != nil { + t.Fatal(err) + } else if n == 0 { + t.Error("zero bits sent by ARP target stack (no reply)") + return + } + + tst.frmbuf, err = tst.cap.CaptureEthernet(tst.frmbuf[:0], buf[:n], 0) + if err != nil { + t.Fatal(err) + } + tst.buf = tst.buf[:n] + + // Validate Ethernet layer (reply is unicast to querying) + if !bytes.Equal(tgtHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) { + t.Errorf("reply: mismatched ethernet src addr %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassSrc)) + } + if !bytes.Equal(qHw[:], tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) { + t.Errorf("reply: expected unicast to querying, got %x", tst.getData(pcap.ProtoEthernet, pcap.FieldClassDst)) + } + + // Validate ARP reply fields + if tst.getARPOperation() != arp.OpReply { + t.Errorf("reply: expected ARP OpReply, got %d", tst.getARPOperation()) + } + if !bytes.Equal(tgtHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 0)) { + t.Errorf("reply: mismatched ARP sender HW (should be target's MAC)") + } + if !bytes.Equal(tgtIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 0)) { + t.Errorf("reply: mismatched ARP sender proto (should be target's IP)") + } + if !bytes.Equal(qHw[:], tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 6, 1)) { + t.Errorf("reply: mismatched ARP target HW (should be querying's MAC)") + } + if !bytes.Equal(qIP.AsSlice(), tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassSrc, 4, 1)) { + t.Errorf("reply: mismatched ARP target proto (should be querying's IP)") + } + + // Deliver reply to querying stack + err = querying.Demux(buf[:n], 0) + if err != nil { + t.Fatal("querying demux reply:", err) + } + setzero(buf[:n]) + + // === PHASE 3: Verify querying stack learned target's MAC === + resolvedHw, err := querying.ResultResolveHardwareAddress6(tgtIP) + if err != nil { + t.Fatalf("ARP query result failed: %v", err) + } + if resolvedHw != tgtHw { + t.Errorf("ARP resolved wrong MAC: got %x, want %x", resolvedHw, tgtHw) + } } func (tst *tester) getTCPFrame() tcp.Frame { @@ -457,3 +576,34 @@ func setzero[T ~[]E, E any](s T) { s[i] = zero } } + +// getFieldByClassLen finds a field by protocol, class, and octet length. +// occurrence specifies which match to return (0 = first, 1 = second, etc.) +// This is needed for ARP where sender and target fields share the same class. +func (tst *tester) getFieldByClassLen(proto any, class pcap.FieldClass, octetLen, occurrence int) []byte { + tst.t.Helper() + frm := getProtoFrame(tst.frmbuf, proto) + if frm == nil { + tst.t.Fatalf("no frame for proto %v found", proto) + } + count := 0 + for _, field := range frm.Fields { + if field.Class == class && field.BitLength == octetLen*8 { + if count == occurrence { + bitoff := frm.PacketBitOffset + field.FrameBitOffset + return tst.buf[bitoff/8 : bitoff/8+field.BitLength/8] + } + count++ + } + } + tst.t.Fatalf("field (proto=%v, class=%v, octets=%d, occurrence=%d) not found", proto, class, octetLen, occurrence) + return nil +} + +func (tst *tester) getARPOperation() arp.Operation { + tst.t.Helper() + // ARP has 3 FieldClassType fields: Hardware type (0), Protocol type (1), Opcode (2) + // All are 2 bytes, so we need occurrence=2 to get Opcode. + data := tst.getFieldByClassLen(ethernet.TypeARP, pcap.FieldClassType, 2, 2) + return arp.Operation(binary.BigEndian.Uint16(data)) +}