all in on StackNode refactor: Encapsulate/Demux on all stacks+ConnID+more abstraction

This commit is contained in:
soypat
2025-06-06 01:23:44 -03:00
parent 6ba727a0ef
commit c635bde512
11 changed files with 390 additions and 73 deletions
+1 -1
View File
@@ -39,7 +39,7 @@ func (efrm Frame) HeaderLength() int {
return sizeHeaderNoVLAN
}
// Payload returns the data portion of the ethernet packet with handling of VLAN packets.
// Payload returns the data portion of the ethernet packet with correct handling of VLAN packets.
func (efrm Frame) Payload() []byte {
hl := efrm.HeaderLength()
et := efrm.EtherTypeOrSize()
+6 -10
View File
@@ -169,21 +169,17 @@ func NewEthernetTCPStack(ourMAC, gwMAC [6]byte, ip netip.AddrPort, mtu uint16, s
gwmac: gwMAC,
}
var ipStack internet.StackBasic
var ipStack internet.StackIP
addr := ip.Addr()
addr4 := addr.As4()
_ = addr4
ipStack.SetAddr(addr)
lStack.Register(handler{
raddr: nil, //addr4[:],
recv: func(b []byte, i int) error {
return ipStack.Recv(b[i:])
},
handle: func(b []byte, i int) (int, error) {
return ipStack.Handle(b[i:])
},
proto: ethernet.TypeIPv4,
lport: 0,
raddr: nil, //addr4[:],
recv: ipStack.Demux,
handle: ipStack.Encapsulate,
proto: ethernet.TypeIPv4,
lport: 0,
})
var conn internet.TCPConn
err = conn.Configure(&internet.TCPConnConfig{
+83
View File
@@ -0,0 +1,83 @@
package internet
import (
"errors"
"math"
"net"
"slices"
)
// 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:]
// 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
// 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]).
//
// 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)
// 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
LocalPort() uint16
Protocol() uint64
ConnectionID() *uint64
// SetFlagPending(flagPending func(numPendingEncapsulations int))
}
// node is a concrete StackNode as stored in Stacks. Methods are devirtualized for performance benefits, especially on TinyGo.
type node struct {
currConnID uint64
connID *uint64
demux func([]byte, int) error
encapsulate func([]byte, int) (int, error)
lastErrs [2]error
proto uint16
port uint16
}
var (
errZeroPort = errors.New("port must be greater than zero")
errInvalidProto = errors.New("invalid protocol")
errProtoRegistered = errors.New("protocol already registered")
_ = net.ErrClosed
)
func handleNodeError(nodes *[]node, nodeIdx int, err error) {
if err != nil {
badConnID := (*nodes)[nodeIdx].connID != nil && *(*nodes)[nodeIdx].connID != (*nodes)[nodeIdx].currConnID
if err == net.ErrClosed || (*nodes)[nodeIdx].lastErrs[0] == err || (*nodes)[nodeIdx].lastErrs[1] == err || badConnID {
*nodes = slices.Delete(*nodes, nodeIdx, nodeIdx+1)
} else {
// Advance Queue of errors
(*nodes)[nodeIdx].lastErrs[1] = (*nodes)[nodeIdx].lastErrs[0]
(*nodes)[nodeIdx].lastErrs[0] = err
}
}
}
func addNode(nodes *[]node, h StackNode, port uint16, protocol uint64) {
if protocol > math.MaxUint16 {
panic(">16bit protocol number unsupported")
}
var currConnID uint64
connIDPtr := h.ConnectionID()
if connIDPtr != nil {
currConnID = *connIDPtr
}
*nodes = append(*nodes, node{
currConnID: currConnID,
connID: connIDPtr,
demux: h.Demux,
encapsulate: h.Encapsulate,
proto: uint16(protocol),
port: port,
})
}
+4
View File
@@ -0,0 +1,4 @@
package internet
type ARPEndpoint struct {
}
-14
View File
@@ -1,14 +0,0 @@
package internet
import "github.com/soypat/lneto"
type PortStack struct {
handlers []porthandler
proto lneto.IPProto
}
type porthandler struct {
recv func([]byte, int) error
handle func([]byte, int) (int, error)
port uint16
}
+76 -31
View File
@@ -9,38 +9,64 @@ import (
"slices"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
"github.com/soypat/lneto/ipv4"
"github.com/soypat/lneto/tcp"
)
type StackBasic struct {
var _ StackNode = (*StackIP)(nil)
type StackIP struct {
connID uint64
ip [4]byte
validator lneto.Validator
handlers []handler
handlers []node
logger
}
type handler struct {
recv func([]byte, int) error
handle func([]byte, int) (int, error)
proto lneto.IPProto
port uint16
func (sb *StackIP) Reset(addr netip.Addr) error {
err := sb.SetAddr(addr)
if err != nil {
return err
}
*sb = StackIP{
connID: sb.connID + 1,
validator: sb.validator,
handlers: sb.handlers[:0],
logger: sb.logger,
ip: sb.ip,
}
return nil
}
func (sb *StackBasic) SetAddr(addr netip.Addr) {
if !addr.Is4() {
panic("only support IPv4")
func (sb *StackIP) SetAddr(addr netip.Addr) error {
if !addr.IsValid() {
return errors.New("invalid IP")
} else if !addr.Is4() {
return errors.New("require IPv4")
}
sb.ip = addr.As4()
return nil
}
func (sb *StackBasic) Addr() netip.Addr {
func (sb *StackIP) ConnectionID() *uint64 {
return &sb.connID
}
func (sb *StackIP) Protocol() uint64 {
return uint64(ethernet.TypeIPv4) // Only support ipv4 for now.
}
func (sb *StackIP) LocalPort() uint16 { return 0 }
func (sb *StackIP) Addr() netip.Addr {
return netip.AddrFrom4(sb.ip)
}
func (sb *StackBasic) Recv(frame []byte) error {
sb.info("StackBasic.Recv:start")
func (sb *StackIP) Demux(carrierData []byte, offset int) error {
sb.info("StackIP.Demux:start")
frame := carrierData[offset:] // we don't care about carrier data in IP.
ifrm, err := ipv4.NewFrame(frame)
if err != nil {
return err
@@ -58,7 +84,7 @@ func (sb *StackBasic) Recv(frame []byte) error {
gotCRC := ifrm.CRC()
wantCRC := ifrm.CalculateHeaderCRC()
if gotCRC != wantCRC {
sb.error("IPv4Stack:Recv:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC)))
sb.error("StackIP:Demux:crc-mismatch", slog.Uint64("want", uint64(wantCRC)), slog.Uint64("got", uint64(gotCRC)))
return errors.New("IPv4 CRC mismatch")
}
off := ifrm.HeaderLength()
@@ -66,9 +92,9 @@ func (sb *StackBasic) Recv(frame []byte) error {
for i := range sb.handlers {
h := &sb.handlers[i]
proto := ifrm.Protocol()
if h.proto == proto {
sb.info("iprecv", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen)))
err = h.recv(frame[:totalLen], off)
if h.proto == uint16(proto) {
sb.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen)))
err = h.demux(frame[:totalLen], off)
if err == net.ErrClosed {
sb.info("ipclose", slog.String("proto", proto.String()))
sb.handlers = slices.Delete(sb.handlers, i, i+1)
@@ -83,7 +109,8 @@ DROP:
return nil
}
func (sb *StackBasic) Handle(frame []byte) (int, error) {
func (sb *StackIP) Encapsulate(carrierData []byte, frameOffset int) (int, error) {
frame := carrierData[frameOffset:]
if len(frame) < 256 {
return 0, io.ErrShortBuffer
}
@@ -91,14 +118,15 @@ func (sb *StackBasic) Handle(frame []byte) (int, error) {
const ihl = 5
const headerlen = ihl * 4
ifrm.SetVersionAndIHL(4, 5)
*ifrm.SourceAddr() = sb.ip
ifrm.SetToS(0)
ifrm.SetID(0)
*ifrm.SourceAddr() = sb.ip
for i := range sb.handlers {
h := &sb.handlers[i]
n, err := h.handle(frame[:], headerlen)
proto := lneto.IPProto(h.proto)
n, err := h.encapsulate(frame[:], headerlen)
if err != nil {
sb.error("IPv4Stack:handle", slog.String("proto", h.proto.String()), slog.String("err", err.Error()))
sb.error("StackIP:handle", slog.String("proto", proto.String()), slog.String("err", err.Error()))
continue
}
if n > 0 {
@@ -107,7 +135,7 @@ func (sb *StackBasic) Handle(frame []byte) (int, error) {
ifrm.SetTotalLength(uint16(totalLen))
ifrm.SetFlags(dontFrag)
ifrm.SetTTL(64)
ifrm.SetProtocol(h.proto)
ifrm.SetProtocol(proto)
ifrm.SetCRC(ifrm.CalculateHeaderCRC())
if ifrm.Protocol() == lneto.IPProtoTCP {
var crc lneto.CRC791
@@ -115,7 +143,7 @@ func (sb *StackBasic) Handle(frame []byte) (int, error) {
tfrm, _ := tcp.NewFrame(ifrm.Payload())
tfrm.CRCWrite(&crc)
tfrm.SetCRC(crc.Sum16())
sb.info("IPv4Stack:send", slog.String("ip", ifrm.String()), slog.String("tcp", tfrm.String()))
sb.info("StackIP:send", slog.String("ip", ifrm.String()), slog.String("tcp", tfrm.String()))
}
return totalLen, nil
}
@@ -123,15 +151,32 @@ func (sb *StackBasic) Handle(frame []byte) (int, error) {
return 0, nil
}
func (sb *StackBasic) RegisterTCPConn(conn *TCPConn) error {
if conn.LocalPort() == 0 {
return errors.New("undefined local port")
func (sb *StackIP) Register(h StackNode) error {
port := h.LocalPort()
proto := h.Protocol()
if port <= 0 {
return errZeroPort
} else if proto > 255 {
return errInvalidProto
}
sb.handlers = append(sb.handlers, handler{
recv: conn.RecvIP,
handle: conn.HandleIP,
proto: lneto.IPProtoTCP,
port: conn.LocalPort(),
sb.handlers = append(sb.handlers, node{
demux: h.Demux,
encapsulate: h.Encapsulate,
proto: uint16(proto),
port: port,
})
return nil
}
func (sb *StackIP) RegisterTCPConn(conn *TCPConn) error {
if conn.LocalPort() == 0 {
return errZeroPort
}
sb.handlers = append(sb.handlers, node{
demux: conn.RecvIP,
encapsulate: conn.HandleIP,
proto: uint16(lneto.IPProtoTCP),
port: conn.LocalPort(),
})
return nil
}
+118
View File
@@ -0,0 +1,118 @@
package internet
import (
"errors"
"io"
"log/slog"
"math"
"net"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
)
type StackLinkLayer struct {
connID uint64
handlers []node
logger
mac [6]byte
gwmac [6]byte
mtu uint16
}
func (ls *StackLinkLayer) Reset6(mac, gateway [6]byte, mtu int) error {
if mtu > math.MaxUint16 || mtu < 256 {
return errors.New("invalid MTU")
}
*ls = StackLinkLayer{
connID: ls.connID + 1,
handlers: ls.handlers[:0],
logger: ls.logger,
mac: mac,
gwmac: gateway,
mtu: uint16(mtu),
}
return nil
}
func (ls *StackLinkLayer) ConnectionID() *uint64 { return &ls.connID }
func (ls *StackLinkLayer) LocalPort() uint16 { return 0 }
func (ls *StackLinkLayer) Protocol() uint64 { return 1 }
func (ls *StackLinkLayer) Register(h StackNode) error {
proto := h.Protocol()
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
}
}
ls.handlers = append(ls.handlers, node{
demux: h.Demux,
encapsulate: h.Encapsulate,
proto: eproto,
})
return nil
}
func (ls *StackLinkLayer) Demux(carrierData []byte, frameOffset int) (err error) {
pkt := carrierData[frameOffset:]
efrm, err := ethernet.NewFrame(pkt)
if err != nil {
return err
}
etype := efrm.EtherTypeOrSize()
dstaddr := efrm.DestinationHardwareAddr()
var vld lneto.Validator
if !efrm.IsBroadcast() && ls.mac != *dstaddr {
goto DROP
}
efrm.ValidateSize(&vld)
if vld.HasError() {
return vld.Err()
}
for i := range ls.handlers {
h := &ls.handlers[i]
if h.proto == uint16(etype) {
return h.demux(efrm.Payload(), 0)
}
}
DROP:
ls.info("LinkStack:drop-packet", slog.String("dsthw", net.HardwareAddr(dstaddr[:]).String()), slog.String("ethertype", efrm.EtherTypeOrSize().String()))
return nil
}
func (ls *StackLinkLayer) Encapsulate(carrierData []byte, frameOffset int) (n int, err error) {
mtu := ls.mtu
dst := carrierData[frameOffset:]
if len(dst) < int(mtu) {
return 0, io.ErrShortBuffer
}
efrm, err := ethernet.NewFrame(dst)
if err != nil {
return 0, err
}
*efrm.DestinationHardwareAddr() = ls.gwmac
for i := range ls.handlers {
h := &ls.handlers[i]
n, err = h.encapsulate(dst[:mtu], 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
}
}
return 0, err
}
+81
View File
@@ -0,0 +1,81 @@
package internet
import (
"encoding/binary"
"io"
)
type StackPort struct {
handlers []node
dstPortOff int
protocol uint64
connID uint64
}
func (ps *StackPort) Reset(protocol uint64, dstPortOffset int) {
*ps = StackPort{
connID: ps.connID + 1,
handlers: ps.handlers[:0],
dstPortOff: dstPortOffset,
protocol: protocol,
}
}
func (ps *StackPort) LocalPort() uint16 { return 0 }
func (ps *StackPort) Protocol() uint64 { return ps.protocol }
func (ps *StackPort) ConnectionID() *uint64 { return &ps.connID }
func (ps *StackPort) Encapsulate(b []byte, offset int) (n int, err error) {
if ps.dstPortOff+offset+2 > len(b) {
return 0, io.ErrShortBuffer
}
var i int
for i = 0; i < len(ps.handlers); i++ {
n, err = ps.handlers[i].encapsulate(b, offset)
if err != nil || n > 0 {
break
}
}
ps.handleResult(i, n, err)
return n, err
}
func (ps *StackPort) Demux(b []byte, offset int) (err error) {
if ps.dstPortOff+offset+2 > len(b) {
return io.ErrShortBuffer
}
port := binary.BigEndian.Uint16(b[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 {
break
}
}
ps.handleResult(i, 0, err)
return err
}
func (ps *StackPort) Register(h StackNode) error {
port := h.LocalPort()
proto := h.Protocol()
if port <= 0 {
return errZeroPort
} else if proto != ps.protocol {
return errInvalidProto
}
ps.handlers = append(ps.handlers, node{
demux: h.Demux,
encapsulate: h.Encapsulate,
port: uint16(port),
})
return nil
}
func (ps *StackPort) handleResult(handlerIdx, n int, err error) {
handleNodeError(&ps.handlers, handlerIdx, err)
}
+7 -7
View File
@@ -10,7 +10,7 @@ import (
func TestBasicStack(t *testing.T) {
rng := rand.New(rand.NewSource(1))
var sbCl, sbSv StackBasic
var sbCl, sbSv StackIP
var connCl, connSv TCPConn
setupClientServer(t, rng, &sbCl, &sbSv, &connCl, &connSv)
var buf [2048]byte
@@ -36,28 +36,28 @@ func TestBasicStack(t *testing.T) {
func TestBasicStack2(t *testing.T) {
rng := rand.New(rand.NewSource(1))
var sbCl, sbSv StackBasic
var sbCl, sbSv StackIP
var connCl, connSv TCPConn
setupClientServerEstablished(t, rng, &sbCl, &sbSv, &connCl, &connSv)
}
func expectExchange(t *testing.T, from, to *StackBasic, buf []byte) {
func expectExchange(t *testing.T, from, to *StackIP, buf []byte) {
t.Helper()
n, err := from.Handle(buf)
n, err := from.Encapsulate(buf, 0)
if err != nil {
t.Error("expectExchange:Handle:", err)
} else if n == 0 {
t.Error("expected data exchange")
return
}
err = to.Recv(buf[:n])
err = to.Demux(buf[:n], 0)
if err != nil {
t.Error("expectExchange:Recv:", err)
}
}
func setupClientServerEstablished(t *testing.T, rng *rand.Rand, client, server *StackBasic, connClient, connServer *TCPConn) {
func setupClientServerEstablished(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *TCPConn) {
t.Helper()
setupClientServer(t, rng, client, server, connClient, connServer)
var buf [2048]byte
@@ -85,7 +85,7 @@ func setupClientServerEstablished(t *testing.T, rng *rand.Rand, client, server *
}
}
func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackBasic, connClient, connServer *TCPConn) {
func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *TCPConn) {
bufsize := 2048
// Ensure buffer sizes are OK with reused buffers.
svip := netip.AddrPortFrom(netip.AddrFrom4([4]byte{192, 168, 1, 0}), 80)
+7 -4
View File
@@ -10,6 +10,7 @@ import (
"runtime"
"time"
"github.com/soypat/lneto"
"github.com/soypat/lneto/internal"
"github.com/soypat/lneto/ipv4"
"github.com/soypat/lneto/ipv6"
@@ -21,7 +22,6 @@ var (
)
type TCPConn struct {
// deprecated: here for debugging purposes only.
h tcp.Handler
remoteAddr []byte
@@ -221,9 +221,8 @@ func (conn *TCPConn) HandleIP(buf []byte, off int) (n int, err error) {
return n, nil
}
func (conn *TCPConn) Send(response []byte) (n int, err error) {
conn.trace("tcpconn.Send:start")
return conn.h.Send(response)
func (conn *TCPConn) Protocol() uint64 {
return uint64(lneto.IPProtoTCP)
}
func getIPAddr(buf []byte) (addr []byte, id uint16, err error) {
@@ -324,3 +323,7 @@ func (conn *TCPConn) SetWriteDeadline(t time.Time) error {
func (conn *TCPConn) deadlineExceeded(deadline time.Time) bool {
return !deadline.IsZero() && time.Since(deadline) > 0
}
func (conn *TCPConn) ConnectionID() *uint64 {
return conn.h.ConnectionID()
}
+7 -6
View File
@@ -21,9 +21,10 @@ var (
// Does NOT implement IP related logic, so no CRC calculation/validation or pseudo header logic.
// Does NOT implement connection lifetime handling, so NO deadlines, keepalives, backoffs or anything that requires use of time package.
type Handler struct {
scb ControlBlock
bufTx ringTx
bufRx internal.Ring
connid uint64
scb ControlBlock
bufTx ringTx
bufRx internal.Ring
logger
validator lneto.Validator
localPort uint16
@@ -31,7 +32,7 @@ type Handler struct {
// connid is a conenction counter that is incremented each time a new
// connection is established via Open calls. This disambiguate's whether
// Read and Write calls belong to the current connection.
connid uint16
optcodec OptionCodec
closing bool
}
@@ -42,8 +43,8 @@ func (h *Handler) SetLoggers(handler, scb *slog.Logger) {
}
// ConnectionID returns the connection identifier which is incremented every time the connection is closed or open.
func (h *Handler) ConnectionID() int {
return int(h.connid)
func (h *Handler) ConnectionID() *uint64 {
return &h.connid
}
// State returns the state of the TCP state machine as per RFC9293. See [State].