From c635bde51247c222ccf3eba92c6bb7e3448b2808 Mon Sep 17 00:00:00 2001 From: soypat Date: Fri, 6 Jun 2025 01:23:44 -0300 Subject: [PATCH] all in on StackNode refactor: Encapsulate/Demux on all stacks+ConnID+more abstraction --- ethernet/frame.go | 2 +- examples/stackbasic/main.go | 16 ++-- internet/definitions.go | 83 +++++++++++++++++ internet/endpoint-arp.go | 4 + internet/portstack.go | 14 --- internet/{basicstack.go => stack-ip.go} | 107 ++++++++++++++------- internet/stack-linklayer.go | 118 ++++++++++++++++++++++++ internet/stack-port.go | 81 ++++++++++++++++ internet/stackbasic_test.go | 14 +-- internet/tcpconn.go | 11 ++- tcp/handler.go | 13 +-- 11 files changed, 390 insertions(+), 73 deletions(-) create mode 100644 internet/definitions.go create mode 100644 internet/endpoint-arp.go delete mode 100644 internet/portstack.go rename internet/{basicstack.go => stack-ip.go} (53%) create mode 100644 internet/stack-linklayer.go create mode 100644 internet/stack-port.go diff --git a/ethernet/frame.go b/ethernet/frame.go index a2ca59f..062be85 100644 --- a/ethernet/frame.go +++ b/ethernet/frame.go @@ -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() diff --git a/examples/stackbasic/main.go b/examples/stackbasic/main.go index fec7a14..6e35dc9 100644 --- a/examples/stackbasic/main.go +++ b/examples/stackbasic/main.go @@ -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{ diff --git a/internet/definitions.go b/internet/definitions.go new file mode 100644 index 0000000..16794f7 --- /dev/null +++ b/internet/definitions.go @@ -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, + }) +} diff --git a/internet/endpoint-arp.go b/internet/endpoint-arp.go new file mode 100644 index 0000000..a948bcc --- /dev/null +++ b/internet/endpoint-arp.go @@ -0,0 +1,4 @@ +package internet + +type ARPEndpoint struct { +} diff --git a/internet/portstack.go b/internet/portstack.go deleted file mode 100644 index a77e373..0000000 --- a/internet/portstack.go +++ /dev/null @@ -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 -} diff --git a/internet/basicstack.go b/internet/stack-ip.go similarity index 53% rename from internet/basicstack.go rename to internet/stack-ip.go index efd4b5e..c0adb3f 100644 --- a/internet/basicstack.go +++ b/internet/stack-ip.go @@ -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 } diff --git a/internet/stack-linklayer.go b/internet/stack-linklayer.go new file mode 100644 index 0000000..489e2eb --- /dev/null +++ b/internet/stack-linklayer.go @@ -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 +} diff --git a/internet/stack-port.go b/internet/stack-port.go new file mode 100644 index 0000000..c592451 --- /dev/null +++ b/internet/stack-port.go @@ -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) +} diff --git a/internet/stackbasic_test.go b/internet/stackbasic_test.go index fdbfa56..7733cbe 100644 --- a/internet/stackbasic_test.go +++ b/internet/stackbasic_test.go @@ -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) diff --git a/internet/tcpconn.go b/internet/tcpconn.go index 112b2af..27499bd 100644 --- a/internet/tcpconn.go +++ b/internet/tcpconn.go @@ -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() +} diff --git a/tcp/handler.go b/tcp/handler.go index ee1a052..0b31240 100644 --- a/tcp/handler.go +++ b/tcp/handler.go @@ -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].