passive MAC learning and ARP cache revamp (#96)

* begin working on arp cache

* fix little things

* rework arp cache priority

* fix test

* keep working on ARP

* push changes before thinking about ip prefix issue

* add passive peer MAC setting

* patch egress mac with correct ethernet CRCs

* consolidate subnet learning in subnetTable type

* add tests and fix subnet table bug
This commit is contained in:
Pat Whittingslow
2026-04-24 23:34:53 -03:00
committed by GitHub
parent 67c42aa136
commit aa77403a2b
10 changed files with 619 additions and 284 deletions
+5 -86
View File
@@ -2,8 +2,6 @@ package arp
import (
"bytes"
"log"
"slices"
"testing"
"github.com/soypat/lneto"
@@ -49,9 +47,9 @@ func TestHandler(t *testing.T) {
}
// Perform ARP exchange.
expectHWAddr := c2.ourHWAddr
expectHWAddr := c2.ourHWAddr[:]
queryAddr := c2.ourProtoAddr
err = c1.StartQuery(nil, queryAddr)
err = c1.StartQuery(queryAddr, false)
if err != nil {
t.Fatal(err)
}
@@ -85,11 +83,11 @@ func TestHandler(t *testing.T) {
if err != nil {
t.Fatal(err)
}
hwaddr, err := c1.QueryResult(queryAddr)
hwaddr, err := c1.CacheLookup(queryAddr)
if err != nil {
log.Fatal("expected query result:", err)
t.Fatal("expected query result:", err)
} else if !bytes.Equal(hwaddr, expectHWAddr) {
log.Fatalf("expected to get hwaddr %x!=%x", hwaddr, expectHWAddr)
t.Fatalf("expected to get hwaddr %x!=%x", hwaddr, expectHWAddr)
}
n, err = c1.Encapsulate(buf[:], -1, 0)
if err != nil {
@@ -105,85 +103,6 @@ func TestHandler(t *testing.T) {
}
}
func TestQueryCompaction(t *testing.T) {
var h Handler
startQuery := func(addr []byte) {
err := h.StartQuery(nil, addr)
if err != nil {
t.Fatal(err)
}
}
err := h.Reset(HandlerConfig{
HardwareAddr: []byte{0xde, 0xad, 0xbe, 0xef, 0x00, 0x00},
ProtocolAddr: []byte{192, 168, 1, 1},
MaxQueries: 5,
MaxPending: 1,
HardwareType: 1,
ProtocolType: ethernet.TypeIPv4,
})
if err != nil {
t.Fatal(err)
}
// Create multiple queries
addr1 := []byte{192, 168, 1, 10}
addr2 := []byte{192, 168, 1, 20}
addr3 := []byte{192, 168, 1, 30}
// Start 3 queries
startQuery(addr1)
startQuery(addr2)
startQuery(addr3)
if len(h.queries) != 3 {
t.Fatalf("expected 3 queries, got %d", len(h.queries))
}
// Discard the middle query (addr2)
if err := h.DiscardQuery(addr2); err != nil {
t.Fatal(err)
}
// Verify addr2 is marked as invalid
hasAddr2 := slices.ContainsFunc(h.queries, func(q queryResult) bool {
return bytes.Equal(q.protoaddr, addr2)
})
if hasAddr2 {
t.Fatal("addr2 query found after discard")
}
// Start new queries to trigger compaction
addr4 := []byte{192, 168, 1, 40}
addr5 := []byte{192, 168, 1, 50}
addr6 := []byte{192, 168, 1, 60}
startQuery(addr4)
startQuery(addr5)
startQuery(addr6)
// After compaction we should be left with 5 queries
expectedAddrs := [][]byte{addr1, addr3, addr4, addr5, addr6}
if len(h.queries) != len(expectedAddrs) {
t.Fatalf("after compaction: expected %d queries, got %d", len(expectedAddrs), len(h.queries))
}
gotAddrs := [][]byte{}
for _, q := range h.queries {
if !q.isInvalid() {
gotAddrs = append(gotAddrs, q.protoaddr)
} else {
t.Fatalf("invalid query %v should have been removed during compaction", q.protoaddr)
}
}
if !slices.EqualFunc(gotAddrs, expectedAddrs, bytes.Equal) {
t.Fatalf("expected %v, got %v", expectedAddrs, gotAddrs)
}
}
func validateARP(t *testing.T, buf []byte) {
t.Helper()
afrm, err := NewFrame(buf)
+167
View File
@@ -0,0 +1,167 @@
package arp
import (
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
)
type cache struct {
entries []entry
}
// entry is designed for compactness. size=class=24 bytes, same as a slice header on x86.
type entry struct {
addr [16]byte
mac [6]byte
age uint8
flags eflags
}
func (e *entry) use(mac [6]byte, proto []byte, flags eflags) {
e.flags = eflagInUse | flags
if copy(e.addr[:], proto) == 16 {
e.flags |= eflagIPv6
}
e.age = 0
e.mac = mac
}
func (e *entry) destroy() { *e = entry{} }
func (e *entry) put(frame, ourAddr []byte, ourMAC [6]byte, op Operation) (int, error) {
if len(ourMAC) != 6 {
return 0, lneto.ErrInvalidAddr
}
f, err := NewFrame(frame)
if err != nil {
return 0, err
}
f.SetHardware(1, 6)
f.SetOperation(op)
var n int
if e.flags&eflagIPv6 != 0 {
if len(ourAddr) != 16 {
return 0, lneto.ErrInvalidAddr
} else if len(frame) < sizeHeaderv6 {
return 0, lneto.ErrShortBuffer
}
f.SetProtocol(ethernet.TypeIPv6, 16)
hw, addr := f.Sender16()
*hw = ourMAC
copy(addr[:], ourAddr)
hw, addr = f.Target16()
copy(hw[:], e.mac[:])
copy(addr[:], e.addr[:])
n = sizeHeaderv6
} else {
if len(ourAddr) != 4 {
return 0, lneto.ErrInvalidAddr
}
f.SetProtocol(ethernet.TypeIPv4, 4)
hw, addr := f.Sender4()
*hw = ourMAC
copy(addr[:], ourAddr)
hw, addr = f.Target4()
copy(hw[:], e.mac[:])
copy(addr[:], e.addr[:])
n = sizeHeaderv4
}
return n, nil
}
func (flags eflags) hasAny(bits eflags) bool { return flags&bits != 0 }
type eflags uint8
// unset eflagInUse to signal the entry can be acquired for a new query.
const (
// eflagInUse set when in use. Discarded/unused entries have this bit unset.
eflagInUse eflags = 1 << iota
// set for IPv6 addressed entries.
eflagIPv6
// network device queried our address and we must respond to it.
// Both MAC and IP are valid in this case.
eflagPendingResponse
// user asked to query this address and query has yet to be answered. May or may not be sent.
eflagIncomplete
// user asked to query this address and the query has not been sent out yet.
// The MAC address is invalid in this case.
eflagIncompletePendingQuery
// eflagPriority set for prioritized cache entries. These entries are discarded last.
// i.e: set for user created queries, unset for external incoming network queries.
eflagPriority
// trigger callback, you know the drill.
eflagResolveTriggersCallback
)
func (c *cache) age() {
for i := range c.entries {
if c.entries[i].flags&eflagInUse != 0 && c.entries[i].age < 255 {
c.entries[i].age++
}
}
}
func (c *cache) reset(size int) {
internal.SliceReuse(&c.entries, size)
c.entries = c.entries[:cap(c.entries)] // maximize queries given allocation.
}
func (c *cache) getNextFlagged(entryHasFlags eflags) *entry {
for i := range c.entries {
flags := c.entries[i].flags
if flags&eflagInUse != 0 && flags.hasAny(entryHasFlags) {
return &c.entries[i]
}
}
return nil
}
func (c *cache) clearFlags(entryHasFlags, clrTheseFlagsIfMatch eflags) {
for i := range c.entries {
// Can clear flags on unused too, simpler.
if c.entries[i].flags&entryHasFlags != 0 {
c.entries[i].flags &^= clrTheseFlagsIfMatch
}
}
}
func (c *cache) Lookup(addr []byte) *entry {
n := len(addr)
for i := range c.entries {
if c.entries[i].flags&eflagInUse != 0 && internal.BytesEqual(c.entries[i].addr[:n], addr) {
return &c.entries[i]
}
}
return nil
}
// acquireNext gets next available entry for use. If all are in use evicts
// the oldest passive entry (learned from incoming requests) before touching
// active user queries or pending responses.
func (c *cache) acquireNext() *entry {
const priorityFlags = eflagPendingResponse | eflagIncomplete | eflagPriority
oldest, oldestPassive := 0, -1
for i := range c.entries {
if c.entries[i].flags&eflagInUse == 0 {
oldest = i
break
}
if !c.entries[i].flags.hasAny(priorityFlags) {
if oldestPassive < 0 || c.entries[i].age > c.entries[oldestPassive].age {
oldestPassive = i
}
}
if c.entries[i].age > c.entries[oldest].age {
oldest = i
}
}
if oldestPassive >= 0 && c.entries[oldest].flags&eflagInUse != 0 {
oldest = oldestPassive
}
c.age()
e := &c.entries[oldest]
e.destroy()
return e
}
+109 -180
View File
@@ -1,21 +1,21 @@
package arp
import (
"log/slog"
"github.com/soypat/lneto"
"github.com/soypat/lneto/ethernet"
"github.com/soypat/lneto/internal"
)
type Handler struct {
connID uint64
ourHWAddr []byte
ourProtoAddr []byte
htype uint16
protoType ethernet.Type
pendingResponse [][sizeHeaderv6]byte
queries []queryResult
connID uint64
cache cache
vld lneto.Validator
ourProtoAddr []byte
onresolve func(hw, proto []byte)
htype uint16
protoType ethernet.Type
ourHWAddr [6]byte
}
type HandlerConfig struct {
@@ -41,6 +41,10 @@ func (h *Handler) UpdateProtoAddr(protoAddr []byte) error {
return nil
}
func (h *Handler) SetOnResolveCallback(cb func(hwAddr, protoAddr []byte)) {
h.onresolve = cb
}
func (h *Handler) Reset(cfg HandlerConfig) error {
if len(cfg.HardwareAddr) == 0 || len(cfg.HardwareAddr) > 255 ||
len(cfg.ProtocolAddr) == 0 || len(cfg.ProtocolAddr) > 255 {
@@ -48,197 +52,120 @@ func (h *Handler) Reset(cfg HandlerConfig) error {
} else if cfg.MaxQueries <= 0 || cfg.MaxPending <= 0 {
return lneto.ErrInvalidConfig
}
if cfg.HardwareType != 1 || cfg.ProtocolType != ethernet.TypeIPv4 && cfg.ProtocolType != ethernet.TypeIPv6 {
return lneto.ErrUnsupported // We only support common for now.
}
*h = Handler{
connID: h.connID + 1,
ourHWAddr: h.ourHWAddr[:0],
ourProtoAddr: h.ourProtoAddr[:0],
htype: cfg.HardwareType,
protoType: cfg.ProtocolType,
pendingResponse: h.pendingResponse[:0],
queries: h.queries[:0],
connID: h.connID + 1,
ourHWAddr: h.ourHWAddr,
ourProtoAddr: h.ourProtoAddr[:0],
htype: cfg.HardwareType,
protoType: cfg.ProtocolType,
cache: h.cache,
}
h.ourHWAddr = append(h.ourHWAddr, cfg.HardwareAddr...)
h.cache.reset(cfg.MaxPending + cfg.MaxQueries)
h.ourHWAddr = [6]byte(cfg.HardwareAddr)
h.ourProtoAddr = append(h.ourProtoAddr, cfg.ProtocolAddr...)
if cap(h.pendingResponse) < cfg.MaxPending {
h.pendingResponse = make([][52]byte, cfg.MaxPending)[:0]
}
if cap(h.queries) < cfg.MaxQueries {
h.queries = make([]queryResult, cfg.MaxQueries)[:0]
}
return nil
}
type queryResult struct {
protoaddr []byte
hwaddr []byte
dstHw []byte
querysent bool
// inc tracks amount of times the query survived compaction.
// The queries with higher inc will be discarded first.
inc uint16
}
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]
h.queries = h.queries[:0]
h.cache.clearFlags(eflagPendingResponse|eflagIncomplete, eflagInUse)
}
func (h *Handler) expectSize() int {
return sizeHeader + 2*len(h.ourHWAddr) + 2*len(h.ourProtoAddr)
// CacheSeed pre-populates the cache with a known proto→hardware mapping, making it
// immediately resolvable via [Handler.CacheLookup] without an ARP exchange.
// Seeded entries are evicted before active user queries when the cache is full.
func (h *Handler) CacheSeed(protoAddr, hwAddr []byte) error {
if len(hwAddr) != 6 {
return lneto.ErrUnsupported
}
e := h.cache.acquireNext()
e.use([6]byte(hwAddr), protoAddr, 0)
return nil
}
func (h *Handler) QueryResult(protoAddr []byte) (hwAddr []byte, err error) {
for i := range h.queries {
if internal.BytesEqual(protoAddr, h.queries[i].protoaddr) {
if !h.queries[i].querysent {
return nil, errQueryPending
}
mac := h.queries[i].response()
if mac == nil {
return nil, errQueryPending
}
return mac, nil
}
// CacheLookup returns the hardware address for protoAddr if it is resolved in the cache.
// Returns [errQueryPending] if a query is in flight, [errQueryNotFound] if no entry exists.
func (h *Handler) CacheLookup(protoAddr []byte) (hwAddr []byte, err error) {
e := h.cache.Lookup(protoAddr)
if e == nil {
return nil, errQueryNotFound
} else if e.flags.hasAny(eflagIncomplete) {
return nil, errQueryPending
}
return nil, errQueryNotFound
return e.mac[:], nil
}
func (h *Handler) DiscardQuery(protoAddr []byte) error {
for i := range h.queries {
q := &h.queries[i]
if internal.BytesEqual(protoAddr, q.protoaddr) {
q.destroy()
return nil
}
// CacheRemove cancels a pending query or evicts a cached entry for protoAddr.
func (h *Handler) CacheRemove(protoAddr []byte) error {
e := h.cache.Lookup(protoAddr)
if e == nil {
return errQueryNotFound
}
return errQueryNotFound
}
func (h *Handler) compactQueries() {
validOff := 0
maxIdx := -1
maxInc := uint16(0)
for i := 0; i < len(h.queries); i++ {
discard := h.queries[i].isInvalid()
if !discard {
h.queries[i].inc++
if h.queries[i].inc > maxInc {
maxInc = h.queries[i].inc
maxIdx = i
}
if i != validOff {
// We swap the queries here so that when `StartQuery` extends
// queries slice, we don't have sharing of the internal structures.
// An alternative would be to zero things, however that would incur
// an allocation cost.
h.queries[validOff], h.queries[i] = h.queries[i], h.queries[validOff]
}
validOff++
}
}
if validOff == len(h.queries) && maxIdx >= 0 {
h.queries[maxIdx].destroy() // Destroy oldest query if unable to compact.
}
h.queries = h.queries[:validOff]
e.destroy()
return nil
}
// 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 lneto.ErrExhausted // Should never fail.
}
}
// Use [Handler.SetOnResolveCallback] to asynchronously set an ARP request result.
func (h *Handler) StartQuery(proto []byte, triggerCallback bool) error {
if len(proto) != len(h.ourProtoAddr) {
return lneto.ErrMismatchLen
} else if dstHWAddr != nil && len(dstHWAddr) != len(h.ourHWAddr) {
return lneto.ErrMismatchLen
} else if dstHWAddr != nil && !internal.IsZeroed(dstHWAddr...) {
return lneto.ErrInvalidConfig
}
q := internal.SliceReclaim(&h.queries)
*q = queryResult{
protoaddr: append(q.protoaddr[:0], proto...),
hwaddr: q.hwaddr[:0],
dstHw: dstHWAddr,
e := h.cache.acquireNext()
e.use([6]byte{}, proto, eflagIncomplete|eflagIncompletePendingQuery|eflagPriority)
if triggerCallback {
e.flags |= eflagResolveTriggersCallback
}
return nil
}
func (h *Handler) Encapsulate(carrierData []byte, offsetToIP, offsetToFrame int) (int, error) {
func (h *Handler) Encapsulate(carrierData []byte, _, offsetToFrame int) (int, error) {
b := carrierData[offsetToFrame:]
n := h.expectSize()
if len(b) < n {
return 0, errShortARP
afrm, err := h.newframe(b)
if err != nil {
return 0, err
}
if len(h.pendingResponse) > 0 {
// pop frame.
afrm, _ := NewFrame(h.pendingResponse[len(h.pendingResponse)-1][:])
h.pendingResponse = h.pendingResponse[:len(h.pendingResponse)-1]
afrm.SetOperation(OpReply)
afrm.SwapTargetSender()
hwsender, _ := afrm.Sender()
copy(hwsender, h.ourHWAddr)
n := copy(b, afrm.Clip().RawData())
tgt, _ := afrm.Target()
trySetEthernetDst(carrierData[:offsetToFrame], tgt)
return n, nil
op := OpReply
e := h.cache.getNextFlagged(eflagPendingResponse) // Prioritize responses.
if e == nil {
e = h.cache.getNextFlagged(eflagIncompletePendingQuery)
if e == nil {
return 0, nil // No action to perform
}
e.flags &^= eflagIncompletePendingQuery
op = OpRequest
} else {
e.flags &^= eflagPendingResponse
}
for i := range h.queries {
if h.queries[i].isInvalid() || h.queries[i].querysent {
continue
}
h.queries[i].querysent = true
afrm, _ := NewFrame(b)
afrm.SetHardware(h.htype, uint8(len(h.ourHWAddr)))
afrm.SetProtocol(h.protoType, uint8(len(h.ourProtoAddr)))
afrm.SetOperation(OpRequest)
hwSender, protoSender := afrm.Sender()
copy(hwSender, h.ourHWAddr)
copy(protoSender, h.ourProtoAddr)
hwTarget, protoTarget := afrm.Target()
copy(protoTarget, h.queries[i].protoaddr)
for j := range hwTarget {
hwTarget[j] = 0
}
// Write Request or Reply, depending on which entry we got.
n, err := e.put(b, h.ourProtoAddr, h.ourHWAddr, op)
if err != nil {
return 0, err
}
switch op {
case OpRequest:
broadcast := ethernet.BroadcastAddr()
trySetEthernetDst(carrierData[:offsetToFrame], broadcast[:])
return n, nil
case OpReply:
tgt, _ := afrm.Target()
trySetEthernetDst(carrierData[:offsetToFrame], tgt)
}
return 0, nil
return n, nil
}
func (h *Handler) Demux(ethFrame []byte, frameOffset int) error {
if len(h.pendingResponse) == cap(h.pendingResponse) {
return lneto.ErrExhausted
}
b := ethFrame[frameOffset:]
afrm, err := NewFrame(b)
afrm, err := h.newframe(b)
if err != nil {
return err
}
var vld lneto.Validator
afrm.ValidateSize(&vld)
if vld.HasError() {
return vld.ErrPop()
afrm.ValidateSize(&h.vld)
if h.vld.HasError() {
return h.vld.ErrPop()
}
htype, hlen := afrm.Hardware()
if htype != h.htype || int(hlen) != len(h.ourHWAddr) {
@@ -254,35 +181,37 @@ func (h *Handler) Demux(ethFrame []byte, frameOffset int) error {
if !internal.BytesEqual(protoaddr, h.ourProtoAddr) {
return nil // Not for us.
}
h.pendingResponse = h.pendingResponse[:len(h.pendingResponse)+1] // Extend pending buffer.
copy(h.pendingResponse[len(h.pendingResponse)-1][:], afrm.buf) // Set pending buffer.
hw, proto := afrm.Sender()
e := h.cache.acquireNext()
e.use([6]byte(hw), proto, eflagPendingResponse)
case OpReply:
hwaddr, protoaddr := afrm.Sender()
for i := range h.queries {
q := &h.queries[i]
mac := q.response()
if mac == nil && internal.BytesEqual(q.protoaddr, protoaddr) {
q.hwaddr = append(q.hwaddr, hwaddr...)
if q.dstHw != nil {
if !internal.IsZeroed(q.dstHw...) {
internal.LogAttrs(nil, slog.LevelError, "race-condition:ARP-reused-buffer")
}
// External write to user buffer.
// Copy data and free up this memory.
copy(q.dstHw, hwaddr)
q.inc = 10000
}
return nil
}
e := h.cache.Lookup(protoaddr)
if e == nil {
return nil
}
copy(e.mac[:], hwaddr)
e.flags &^= eflagIncomplete | eflagIncompletePendingQuery
if e.flags.hasAny(eflagResolveTriggersCallback) && h.onresolve != nil {
h.onresolve(e.mac[:], protoaddr)
}
default:
return errARPUnsupported
}
return nil
}
func (h *Handler) newframe(b []byte) (Frame, error) {
f, err := NewFrame(b)
if err != nil {
return f, err
} else if h.protoType == ethernet.TypeIPv6 && len(b) < sizeHeaderv6 {
return f, lneto.ErrMismatch
}
return f, nil
}
func trySetEthernetDst(ethFrame []byte, dst []byte) {
if len(ethFrame) >= 14 {
copy(ethFrame[:6], dst)