ipv6: StackIP and StackAsync.Addr refactor (#105)

* apply StackIP changes and internet package test passing

* fix tests and examples

* remove old Reset method on StackIP

* use encapsulate for ipv6

* add TCP over IPv6 tests
This commit is contained in:
Pat Whittingslow
2026-05-09 16:11:31 -03:00
committed by GitHub
parent a430f6c40a
commit bdbd38ab44
22 changed files with 609 additions and 357 deletions
+29
View File
@@ -7,6 +7,7 @@ import (
"slices"
"github.com/soypat/lneto"
"github.com/soypat/lneto/internal"
)
// node is a concrete StackNode as stored in Stacks. Methods are devirtualized for performance benefits, especially on TinyGo.
@@ -231,3 +232,31 @@ func incLim(v, max int) int {
}
return v
}
type logger struct {
log *slog.Logger
}
func (l logger) error(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelError, msg, attrs...)
}
func (l logger) info(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelInfo, msg, attrs...)
}
func (l logger) warn(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelWarn, msg, attrs...)
}
func (l logger) debug(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelDebug, msg, attrs...)
}
func (l logger) trace(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, internal.LevelTrace, msg, attrs...)
}
const enableAllocLog = internal.HeapAllocDebugging
func debugLog(msg string) {
if enableAllocLog {
internal.LogAllocs(msg)
}
}
+33 -216
View File
@@ -1,251 +1,68 @@
package internet
import (
"io"
"log/slog"
"net/netip"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
"github.com/soypat/lneto/ipv4"
"github.com/soypat/lneto/tcp"
"github.com/soypat/lneto/udp"
)
var _ lneto.StackNode = (*StackIP)(nil)
type StackIP struct {
connID uint64
ipID uint16
ip [4]byte
acceptMulticast bool
validator lneto.Validator
handlers handlers
connID uint64
stackip4
stackip6
}
func (sb *StackIP) Reset(addr netip.Addr, maxNodes int) error {
if maxNodes <= 0 {
func (stackip *StackIP) Reset(vld *lneto.Validator, maxNodes4, maxNodes6 int) error {
if maxNodes4 <= 0 && maxNodes6 <= 0 || vld == nil {
return lneto.ErrInvalidConfig
}
err := sb.SetAddr(addr)
if err != nil {
return err
}
sb.handlers.reset("StackIP", maxNodes)
*sb = StackIP{
connID: sb.connID + 1,
validator: sb.validator,
handlers: sb.handlers,
ip: sb.ip,
acceptMulticast: sb.acceptMulticast,
}
stackip.connID++
stackip.reset4(vld, maxNodes4)
stackip.reset6(vld, maxNodes6)
return nil
}
func (sb *StackIP) SetAddr(addr netip.Addr) error {
if !addr.IsValid() {
return lneto.ErrInvalidAddr
} else if !addr.Is4() {
return lneto.ErrUnsupported
}
sb.ip = addr.As4()
return nil
func (stackip *StackIP) ConnectionID() *uint64 {
return &stackip.connID
}
func (sb *StackIP) ConnectionID() *uint64 {
return &sb.connID
}
func (sb *StackIP) Protocol() uint64 {
func (stackip *StackIP) Protocol() uint64 {
return uint64(ethernet.TypeIPv4) // Only support ipv4 for now.
}
func (sb *StackIP) LocalPort() uint16 { return 0 }
func (stackip *StackIP) LocalPort() uint16 { return 0 }
func (sb *StackIP) Addr() netip.Addr {
return netip.AddrFrom4(sb.ip)
func (stackip *StackIP) SetLogger(logger *slog.Logger) {
stackip.stackip4.handlers.log = logger
stackip.stackip6.handlers.log = logger
}
func (sb *StackIP) SetAcceptMulticast(accept bool) {
sb.acceptMulticast = accept
}
func (sb *StackIP) SetLogger(logger *slog.Logger) {
sb.handlers.log = logger
}
func (sb *StackIP) Demux(carrierData []byte, offset int) error {
func (stackip *StackIP) Demux(carrierData []byte, offset int) error {
debugLog("ip:demux")
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 {
return err
if len(carrierData) < 1 {
return lneto.ErrTruncatedFrame
}
dst := ifrm.DestinationAddr()
if sb.ip != ([4]byte{}) && *dst != sb.ip {
if !sb.acceptMulticast || dst[0]&0xF0 != 0xE0 {
sb.handlers.debug("ip:not-for-us")
return lneto.ErrPacketDrop // Not meant for us.
}
version := carrierData[offset] >> 4
switch version {
case 4:
return stackip.stackip4.demux4(carrierData, offset)
case 6:
return stackip.stackip6.demux6(carrierData, offset)
default:
return lneto.ErrUnsupported
}
sb.validator.ResetErr()
ifrm.ValidateExceptCRC(&sb.validator)
if err = sb.validator.ErrPop(); err != nil {
sb.handlers.error("ip:Demux.validate")
return err
}
if ifrm.CalculateHeaderCRC() != 0 {
sb.handlers.error("ip:demux.crc")
return lneto.ErrBadCRC
}
off := ifrm.HeaderLength()
totalLen := ifrm.TotalLength()
proto := ifrm.Protocol()
node := sb.handlers.nodeByProto(uint16(proto))
// nodeIdx := getNodeByProto(sb.handlers, uint16(proto))
if node == nil {
// Drop packet.
sb.handlers.info("ip:demux.drop", internal.SlogAddr4("dstaddr", ifrm.DestinationAddr()), slog.String("proto", ifrm.Protocol().String()))
return lneto.ErrPacketDrop
}
// Incoming CRC Validation of common IP Protocols.
var crc lneto.CRC791
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
if crc.PayloadSum16(ifrm.Payload()) != 0 {
sb.handlers.error("ip:demux.tcpcrc")
return lneto.ErrBadCRC
}
case lneto.IPProtoUDP:
ufrm, err := udp.NewFrame(ifrm.Payload())
if err != nil {
return err
}
ufrm.ValidateSize(&sb.validator)
if err = sb.validator.ErrPop(); err != nil {
sb.handlers.error("ip:demux.udpvalidatesize")
return err
}
frameLen := ufrm.Length()
ifrm.CRCWriteUDPPseudo(&crc, frameLen)
if crc.PayloadSum16(ufrm.RawData()[:frameLen]) != 0 {
sb.handlers.error("ip:demux.udpcrc")
return lneto.ErrBadCRC
}
}
sb.handlers.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("plen", int(totalLen)))
err = node.callbacks.Demux(frame[:totalLen], off)
if sb.handlers.tryHandleError(node, err) {
sb.handlers.info("ipclose", slog.String("proto", proto.String()))
err = nil
}
return err
}
func (sb *StackIP) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
frame := carrierData[offsetToFrame:]
if len(frame) < ipv4.MinimumMTU {
return 0, io.ErrShortBuffer
func (stackip *StackIP) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (n int, err error) {
if offsetToFrame != offsetToIP {
return 0, lneto.ErrBug
}
ifrm, _ := ipv4.NewFrame(frame)
const ihl = 5
const headerlen = ihl * 4
const dontFrag = 0x4000
ifrm.SetVersionAndIHL(4, ihl)
ifrm.SetToS(0)
seed := sb.ipID + uint16(sb.connID)
id := internal.Prand16(seed)
ifrm.SetID(id)
ifrm.SetFlags(dontFrag)
ifrm.SetTTL(64)
*ifrm.SourceAddr() = sb.ip
sb.ipID = id
// 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
}
proto := lneto.IPProto(node.proto)
totalLen := n + headerlen
ifrm.SetTotalLength(uint16(totalLen))
ifrm.SetProtocol(proto)
// Zero the CRC field so its value does not add to the final result.
ifrm.SetCRC(0)
crcValue := ifrm.CalculateHeaderCRC()
ifrm.SetCRC(crcValue)
// Calculate CRC for our newly generated packet.
var crc lneto.CRC791
payload := ifrm.Payload()
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
tfrm, _ := tcp.NewFrame(payload)
// Zero the CRC field so its value does not add to the final result.
tfrm.SetCRC(0)
crcValue = crc.PayloadSum16(payload)
tfrm.SetCRC(crcValue)
case lneto.IPProtoUDP:
ufrm, _ := udp.NewFrame(payload)
ifrm.CRCWriteUDPPseudo(&crc, uint16(n))
ufrm.SetLength(uint16(n))
// Zero the CRC field so its value does not add to the final result.
ufrm.SetCRC(0)
crcValue = lneto.NeverZeroSum(crc.PayloadSum16(payload))
ufrm.SetCRC(crcValue)
}
return totalLen, err
}
func (sb *StackIP) Register(h lneto.StackNode) error {
proto := h.Protocol()
if proto > 255 {
return lneto.ErrInvalidConfig
}
return sb.handlers.registerByPortProto(nodeFromStackNode(h, h.LocalPort(), proto, nil))
}
func (sb *StackIP) IsRegistered(proto lneto.IPProto) bool {
return sb.handlers.nodeByProto(uint16(proto)) != nil
}
func (sb *StackIP) recvicmp(icmpData []byte) error {
var crc lneto.CRC791
if crc.PayloadSum16(icmpData) != 0 {
return lneto.ErrBadCRC
}
return nil
}
type logger struct {
log *slog.Logger
}
func (l logger) error(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelError, msg, attrs...)
}
func (l logger) info(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelInfo, msg, attrs...)
}
func (l logger) warn(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelWarn, msg, attrs...)
}
func (l logger) debug(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, slog.LevelDebug, msg, attrs...)
}
func (l logger) trace(msg string, attrs ...slog.Attr) {
internal.LogAttrs(l.log, internal.LevelTrace, msg, attrs...)
}
const enableAllocLog = internal.HeapAllocDebugging
func debugLog(msg string) {
if enableAllocLog {
internal.LogAllocs(msg)
n, err = stackip.stackip4.encapsulate4(carrierData, offsetToIP)
if len(stackip.stackip6.handlers.nodes) > 0 && n == 0 {
n, err = stackip.stackip6.encapsulate6(carrierData, offsetToIP)
}
return n, err
}
+183
View File
@@ -0,0 +1,183 @@
package internet
import (
"io"
"log/slog"
"github.com/soypat/lneto"
"github.com/soypat/lneto/internal"
"github.com/soypat/lneto/ipv4"
"github.com/soypat/lneto/tcp"
"github.com/soypat/lneto/udp"
)
// stackip4 is NOT a StackNode implementation.
// It is meant to be embedded within StackNodes.
// var _ lneto.StackNode = (*stackip4)(nil)
type stackip4 struct {
handlers handlers
vld *lneto.Validator
ipID uint16
ip4 [4]byte
acceptMulticast bool
}
func (si4 *stackip4) reset4(vld *lneto.Validator, maxNodes int) {
*si4 = stackip4{
ip4: [4]byte{},
ipID: 1,
acceptMulticast: false,
handlers: si4.handlers,
vld: vld,
}
si4.handlers.reset("stackip4", maxNodes)
}
func (si4 *stackip4) Register4(h lneto.StackNode) error {
proto := h.Protocol()
if proto > 255 {
return lneto.ErrInvalidConfig
}
return si4.handlers.registerByPortProto(nodeFromStackNode(h, h.LocalPort(), proto, nil))
}
func (si4 *stackip4) IsRegistered4(proto lneto.IPProto) bool {
return si4.handlers.nodeByProto(uint16(proto)) != nil
}
func (si4 *stackip4) SetAcceptMulticast4(accept bool) {
si4.acceptMulticast = accept
}
func (si4 *stackip4) Addr4() [4]byte { return si4.ip4 }
func (si4 *stackip4) SetAddr4(ip4 [4]byte) {
si4.ip4 = ip4
}
func (si4 *stackip4) demux4(carrierData []byte, offset int) error {
debugLog("ip4:demux")
si4.handlers.info("demux:start")
frame := carrierData[offset:] // we don't care about carrier data in IP.
ifrm, err := ipv4.NewFrame(frame)
if err != nil {
return err
}
dst := ifrm.DestinationAddr()
if si4.ip4 != ([4]byte{}) && *dst != si4.ip4 {
if !si4.acceptMulticast || dst[0]&0xF0 != 0xE0 {
si4.handlers.debug("ip:not-for-us")
return lneto.ErrPacketDrop // Not meant for us.
}
}
si4.vld.ResetErr()
ifrm.ValidateExceptCRC(si4.vld)
if err = si4.vld.ErrPop(); err != nil {
si4.handlers.error("ip:Demux.validate")
return err
}
if ifrm.CalculateHeaderCRC() != 0 {
si4.handlers.error("ip:demux.crc")
return lneto.ErrBadCRC
}
off := ifrm.HeaderLength()
proto := ifrm.Protocol()
node := si4.handlers.nodeByProto(uint16(proto))
// nodeIdx := getNodeByProto(sb.handlers, uint16(proto))
if node == nil {
// Drop packet.
si4.handlers.info("ip:demux.drop", internal.SlogAddr4("dstaddr", ifrm.DestinationAddr()), slog.String("proto", ifrm.Protocol().String()))
return lneto.ErrPacketDrop
}
// Incoming CRC Validation of common IP Protocols.
var crc lneto.CRC791
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
if crc.PayloadSum16(ifrm.Payload()) != 0 {
si4.handlers.error("ip:demux.tcpcrc")
return lneto.ErrBadCRC
}
case lneto.IPProtoUDP:
ufrm, err := udp.NewFrame(ifrm.Payload())
if err != nil {
return err
}
ufrm.ValidateSize(si4.vld)
if err = si4.vld.ErrPop(); err != nil {
si4.handlers.error("ip:demux.udpvalidatesize")
return err
}
frameLen := ufrm.Length()
ifrm.CRCWriteUDPPseudo(&crc, frameLen)
if crc.PayloadSum16(ufrm.RawData()[:frameLen]) != 0 {
si4.handlers.error("ip:demux.udpcrc")
return lneto.ErrBadCRC
}
}
totalLen := ifrm.TotalLength()
si4.handlers.info("ipDemux", slog.String("ipproto", proto.String()), slog.Int("tlen", int(totalLen)))
err = node.callbacks.Demux(frame[:totalLen], off)
if si4.handlers.tryHandleError(node, err) {
si4.handlers.info("ipclose", slog.String("proto", proto.String()))
err = nil
}
return err
}
func (si4 *stackip4) encapsulate4(carrierData []byte, offsetToIP int) (int, error) {
frame := carrierData[offsetToIP:]
if len(frame) < ipv4.MinimumMTU {
return 0, io.ErrShortBuffer
}
ifrm, _ := ipv4.NewFrame(frame)
const ihl = 5
const headerlen = ihl * 4
const dontFrag = 0x4000
ifrm.SetVersionAndIHL(4, ihl)
ifrm.SetToS(0)
seed := (si4.ipID + 1) ^ uint16(si4.ip4[0])
id := internal.Prand16(seed)
ifrm.SetID(id)
ifrm.SetFlags(dontFrag)
ifrm.SetTTL(64)
*ifrm.SourceAddr() = si4.ip4
si4.ipID = id
// 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 := si4.handlers.encapsulateAny(carrierData, offsetToIP, offsetToIP+headerlen)
if n == 0 {
return n, err
}
proto := lneto.IPProto(node.proto)
totalLen := n + headerlen
ifrm.SetTotalLength(uint16(totalLen))
ifrm.SetProtocol(proto)
// Zero the CRC field so its value does not add to the final result.
ifrm.SetCRC(0)
crcValue := ifrm.CalculateHeaderCRC()
ifrm.SetCRC(crcValue)
// Calculate CRC for our newly generated packet.
var crc lneto.CRC791
payload := ifrm.Payload()
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWriteTCPPseudo(&crc)
tfrm, _ := tcp.NewFrame(payload)
// Zero the CRC field so its value does not add to the final result.
tfrm.SetCRC(0)
crcValue = crc.PayloadSum16(payload)
tfrm.SetCRC(crcValue)
case lneto.IPProtoUDP:
ufrm, _ := udp.NewFrame(payload)
ifrm.CRCWriteUDPPseudo(&crc, uint16(n))
ufrm.SetLength(uint16(n))
// Zero the CRC field so its value does not add to the final result.
ufrm.SetCRC(0)
crcValue = lneto.NeverZeroSum(crc.PayloadSum16(payload))
ufrm.SetCRC(crcValue)
}
return totalLen, err
}
+144
View File
@@ -0,0 +1,144 @@
package internet
import (
"log/slog"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ipv6"
"github.com/soypat/lneto/tcp"
"github.com/soypat/lneto/udp"
)
// stackip6 is NOT a StackNode implementation.
// It is meant to be embedded within StackNodes.
// var _ lneto.StackNode = (*stackip6)(nil)
type stackip6 struct {
handlers handlers
vld *lneto.Validator
ip6 [16]byte
acceptMulticast bool
}
func (si6 *stackip6) Register6(h lneto.StackNode) error {
proto := h.Protocol()
if proto > 255 {
return lneto.ErrInvalidConfig
}
return si6.handlers.registerByPortProto(nodeFromStackNode(h, h.LocalPort(), proto, nil))
}
func (si6 *stackip6) IsRegistered6(proto lneto.IPProto) bool {
return si6.handlers.nodeByProto(uint16(proto)) != nil
}
func (si6 *stackip6) SetAcceptMulticast6(accept bool) { si6.acceptMulticast = accept }
func (si6 *stackip6) Addr6() [16]byte { return si6.ip6 }
func (si6 *stackip6) SetAddr6(ip6 [16]byte) { si6.ip6 = ip6 }
func (si6 *stackip6) reset6(vld *lneto.Validator, maxNodes int) {
*si6 = stackip6{
handlers: si6.handlers,
vld: vld,
}
si6.handlers.reset("stackip6", maxNodes)
}
func (si6 *stackip6) demux6(carrierData []byte, offset int) error {
debugLog("ip6:demux")
si6.handlers.info("StackIP6.Demux:start")
ifrm, err := ipv6.NewFrame(carrierData[offset:])
if err != nil {
return err
}
dst := ifrm.DestinationAddr()
if si6.ip6 != ([16]byte{}) && *dst != si6.ip6 {
if !si6.acceptMulticast || dst[0] != 0xFF {
si6.handlers.debug("ip6:not-for-us")
return lneto.ErrPacketDrop
}
}
si6.vld.ResetErr()
ifrm.ValidateSize(si6.vld)
if err = si6.vld.ErrPop(); err != nil {
si6.handlers.error("ip6:Demux.validate")
return err
}
proto := ifrm.NextHeader()
node := si6.handlers.nodeByProto(uint16(proto))
if node == nil {
si6.handlers.info("ip6:demux.drop", slog.String("proto", proto.String()))
return lneto.ErrPacketDrop
}
payload := ifrm.Payload()
var crc lneto.CRC791
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWritePseudo(&crc)
if crc.PayloadSum16(payload) != 0 {
si6.handlers.error("ip6:demux.tcpcrc")
return lneto.ErrBadCRC
}
case lneto.IPProtoUDP:
ufrm, err := udp.NewFrame(payload)
if err != nil {
return err
}
ufrm.ValidateSize(si6.vld)
if err = si6.vld.ErrPop(); err != nil {
si6.handlers.error("ip6:demux.udpvalidatesize")
return err
}
ifrm.CRCWritePseudo(&crc)
if crc.PayloadSum16(payload) != 0 {
si6.handlers.error("ip6:demux.udpcrc")
return lneto.ErrBadCRC
}
}
const headerlen = 40
plen := ifrm.PayloadLength()
si6.handlers.info("ip6Demux", slog.String("ipproto", proto.String()), slog.Int("plen", int(plen)))
err = node.callbacks.Demux(carrierData[offset:offset+headerlen+int(plen)], headerlen)
if si6.handlers.tryHandleError(node, err) {
si6.handlers.info("ip6close", slog.String("proto", proto.String()))
err = nil
}
return err
}
func (si6 *stackip6) encapsulate6(carrierData []byte, offsetToIP int) (int, error) {
ifrm, err := ipv6.NewFrame(carrierData[offsetToIP:])
if err != nil {
return 0, err
}
// Set default parameters which node is free to change.
ifrm.SetVersionTrafficAndFlow(6, 0, 0)
ifrm.SetHopLimit(64)
*ifrm.SourceAddr() = si6.ip6
const headerlen = 40
node, n, err := si6.handlers.encapsulateAny(carrierData, offsetToIP, offsetToIP+headerlen)
if n == 0 {
return n, err
}
proto := lneto.IPProto(node.proto)
ifrm.SetNextHeader(proto)
ifrm.SetPayloadLength(uint16(n))
var crc lneto.CRC791
payload := ifrm.Payload()
switch proto {
case lneto.IPProtoTCP:
ifrm.CRCWritePseudo(&crc)
tfrm, _ := tcp.NewFrame(payload)
tfrm.SetCRC(0)
tfrm.SetCRC(crc.PayloadSum16(payload))
case lneto.IPProtoUDP:
ufrm, _ := udp.NewFrame(payload)
ufrm.SetLength(uint16(n))
ifrm.CRCWritePseudo(&crc)
ufrm.SetCRC(0)
ufrm.SetCRC(lneto.NeverZeroSum(crc.PayloadSum16(payload)))
}
return headerlen + n, err
}
+89 -6
View File
@@ -5,6 +5,7 @@ import (
"net/netip"
"testing"
"github.com/soypat/lneto"
"github.com/soypat/lneto/tcp"
)
@@ -44,7 +45,7 @@ func TestBasicStack2(t *testing.T) {
func expectExchange(t *testing.T, from, to *StackIP, buf []byte) {
t.Helper()
n, err := from.Encapsulate(buf, -1, 0)
n, err := from.Encapsulate(buf, 0, 0)
if err != nil {
t.Error("expectExchange:encapsulate:", err)
} else if n == 0 {
@@ -90,15 +91,97 @@ func testClientServerEstablish(t *testing.T, client, server *StackIP, connClient
}
}
func TestBasicStack6(t *testing.T) {
rng := rand.New(rand.NewSource(1))
var sbCl, sbSv StackIP
var connCl, connSv tcp.Conn
setupClientServer6(t, rng, &sbCl, &sbSv, &connCl, &connSv)
var buf [2048]byte
nextToSend := &sbCl
nextToRecv := &sbSv
exchangeAndExpectStates := func(clState, svState tcp.State) {
t.Helper()
expectExchange(t, nextToSend, nextToRecv, buf[:])
gotCl := connCl.State()
gotSv := connSv.State()
if gotCl != clState {
t.Errorf("want client state %s, got %s", clState, gotCl)
}
if gotSv != svState {
t.Errorf("want server state %s, got %s", svState, gotSv)
}
nextToSend, nextToRecv = nextToRecv, nextToSend
}
exchangeAndExpectStates(tcp.StateSynSent, tcp.StateSynRcvd)
exchangeAndExpectStates(tcp.StateEstablished, tcp.StateSynRcvd)
exchangeAndExpectStates(tcp.StateEstablished, tcp.StateEstablished)
}
func TestBasicStack6Established(t *testing.T) {
rng := rand.New(rand.NewSource(1))
var sbCl, sbSv StackIP
var connCl, connSv tcp.Conn
setupClientServer6(t, rng, &sbCl, &sbSv, &connCl, &connSv)
testClientServerEstablish(t, &sbCl, &sbSv, &connCl, &connSv)
}
func setupClientServer6(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *tcp.Conn) {
t.Helper()
_ = rng
const maxNodes = 1
bufsize := 2048
svip6 := netip.AddrFrom16([16]byte{0x20, 0x01, 0x0d, 0xb8, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}) // 2001:db8::1
clip6 := netip.AddrFrom16([16]byte{0x20, 0x01, 0x0d, 0xb8, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}) // 2001:db8::2
svip := netip.AddrPortFrom(svip6, 80)
clip := netip.AddrPortFrom(clip6, 1337)
if err := server.Reset(new(lneto.Validator), 0, maxNodes); err != nil {
t.Fatal(err)
}
if err := client.Reset(new(lneto.Validator), 0, maxNodes); err != nil {
t.Fatal(err)
}
server.SetAddr6(svip6.As16())
client.SetAddr6(clip6.As16())
err := connServer.Configure(tcp.ConnConfig{
RxBuf: make([]byte, bufsize),
TxBuf: make([]byte, bufsize),
TxPacketQueueSize: 3,
})
if err != nil {
t.Fatal(err)
}
err = connClient.Configure(tcp.ConnConfig{
RxBuf: make([]byte, bufsize),
TxBuf: make([]byte, bufsize),
TxPacketQueueSize: 3,
})
if err != nil {
t.Fatal(err)
}
if err = connServer.OpenListen(svip.Port(), 200); err != nil {
t.Fatal(err)
}
if err = connClient.OpenActive(clip.Port(), svip, 100); err != nil {
t.Fatal(err)
}
if err = server.Register6(connServer); err != nil {
t.Fatal(err)
}
if err = client.Register6(connClient); err != nil {
t.Fatal(err)
}
}
func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackIP, connClient, connServer *tcp.Conn) {
const maxNodes = 1
bufsize := 2048
// Ensure buffer sizes are OK with reused buffers.
svip := netip.AddrPortFrom(netip.AddrFrom4([4]byte{192, 168, 1, 0}), 80)
clip := netip.AddrPortFrom(netip.AddrFrom4([4]byte{192, 168, 1, 1}), 1337)
server.Reset(svip.Addr(), maxNodes)
client.Reset(clip.Addr(), maxNodes)
server.Reset(new(lneto.Validator), maxNodes, 0)
client.Reset(new(lneto.Validator), maxNodes, 0)
server.SetAddr4(svip.Addr().As4())
client.SetAddr4(clip.Addr().As4())
err := connServer.Configure(tcp.ConnConfig{
RxBuf: make([]byte, bufsize),
TxBuf: make([]byte, bufsize),
@@ -127,11 +210,11 @@ func setupClientServer(t *testing.T, rng *rand.Rand, client, server *StackIP, co
t.Fatal(err)
}
err = server.Register(connServer)
err = server.Register4(connServer)
if err != nil {
t.Fatal(err)
}
err = client.Register(connClient)
err = client.Register4(connClient)
if err != nil {
t.Fatal(err)
}
+13 -11
View File
@@ -6,6 +6,7 @@ import (
"net/netip"
"testing"
"github.com/soypat/lneto"
"github.com/soypat/lneto/tcp"
)
@@ -24,7 +25,7 @@ func TestListener_SingleConnection(t *testing.T) {
if err := listener.Reset(serverPort, pool); err != nil {
t.Fatal(err)
}
if err := serverStack.Register(&listener); err != nil {
if err := serverStack.Register4(&listener); err != nil {
t.Fatal(err)
}
@@ -77,7 +78,7 @@ func TestListener_AcceptAfterEstablished(t *testing.T) {
if err := listener.Reset(serverPort, pool); err != nil {
t.Fatal(err)
}
if err := serverStack.Register(&listener); err != nil {
if err := serverStack.Register4(&listener); err != nil {
t.Fatal(err)
}
@@ -105,7 +106,7 @@ func TestListener_AcceptAfterEstablished(t *testing.T) {
// Setup second client and verify we can still accept.
var client2Stack StackIP
var client2Conn tcp.Conn
setupClient(t, &client2Stack, &client2Conn, serverStack.Addr(), serverPort, 1338)
setupClient(t, &client2Stack, &client2Conn, netip.AddrFrom4(serverStack.Addr4()), serverPort, 1338)
// Complete full handshake for client2.
expectExchange(t, &client2Stack, &serverStack, buf[:]) // SYN
@@ -147,14 +148,14 @@ func TestListener_MultiConn(t *testing.T) {
if err := listener.Reset(serverPort, pool); err != nil {
t.Fatal(err)
}
if err := serverStack.Register(&listener); err != nil {
if err := serverStack.Register4(&listener); err != nil {
t.Fatal(err)
}
// Setup remaining clients.
for i := 1; i < numClients; i++ {
clientPort := uint16(1337 + i)
setupClient(t, &clientStacks[i], &clientConns[i], serverStack.Addr(), serverPort, clientPort)
setupClient(t, &clientStacks[i], &clientConns[i], netip.AddrFrom4(serverStack.Addr4()), serverPort, clientPort)
}
var buf [2048]byte
@@ -324,7 +325,7 @@ func TestListener_RSTOnPoolExhaustion(t *testing.T) {
if err := listener.Reset(serverPort, pool); err != nil {
t.Fatal(err)
}
if err := serverStack.Register(&listener); err != nil {
if err := serverStack.Register4(&listener); err != nil {
t.Fatal(err)
}
@@ -340,10 +341,10 @@ func TestListener_RSTOnPoolExhaustion(t *testing.T) {
// Setup client2 and send its SYN — pool is full, server should queue RST.
const client2Port = uint16(1338)
setupClient(t, &client2Stack, &client2Conn, serverStack.Addr(), serverPort, client2Port)
setupClient(t, &client2Stack, &client2Conn, netip.AddrFrom4(serverStack.Addr4()), serverPort, client2Port)
// Client2 sends SYN.
n, err := client2Stack.Encapsulate(buf[:], -1, 0)
n, err := client2Stack.Encapsulate(buf[:], 0, 0)
if err != nil {
t.Fatal("client2 encapsulate:", err)
} else if n == 0 {
@@ -356,7 +357,7 @@ func TestListener_RSTOnPoolExhaustion(t *testing.T) {
}
// Server encapsulates — should produce RST (no connection data pending).
n, err = serverStack.Encapsulate(buf[:], -1, 0)
n, err = serverStack.Encapsulate(buf[:], 0, 0)
if err != nil {
t.Fatal("server encapsulate RST:", err)
} else if n == 0 {
@@ -619,7 +620,8 @@ func setupClient(t *testing.T, client *StackIP, conn *tcp.Conn, serverAddr netip
t.Helper()
bufsize := 2048
clientIP := netip.AddrFrom4([4]byte{192, 168, 1, byte(clientPort % 256)})
client.Reset(clientIP, 1)
client.Reset(new(lneto.Validator), 1, 0)
client.SetAddr4(clientIP.As4())
err := conn.Configure(tcp.ConnConfig{
RxBuf: make([]byte, bufsize),
TxBuf: make([]byte, bufsize),
@@ -633,7 +635,7 @@ func setupClient(t *testing.T, client *StackIP, conn *tcp.Conn, serverAddr netip
if err != nil {
t.Fatal(err)
}
err = client.Register(conn)
err = client.Register4(conn)
if err != nil {
t.Fatal(err)
}