diff --git a/arp/arp_test.go b/arp/arp_test.go index 59bb875..a1e69bf 100644 --- a/arp/arp_test.go +++ b/arp/arp_test.go @@ -3,6 +3,7 @@ package arp import ( "bytes" "log" + "slices" "testing" "github.com/soypat/lneto" @@ -104,6 +105,85 @@ 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) diff --git a/arp/definitions.go b/arp/definitions.go index f6925ae..62b81f7 100644 --- a/arp/definitions.go +++ b/arp/definitions.go @@ -13,7 +13,7 @@ const ( var ( errARPBufferFull = errors.New("ARP client need handling:too many ops pending") errShortARP = errors.New("packet too short to be ARP") - errARPUnsupported = errors.New("ARP not supprortedf") + errARPUnsupported = errors.New("ARP not supported") errLargeSizes = errors.New("size of ARP protocol+hardware is unusually large") ) diff --git a/arp/handler.go b/arp/handler.go index cce5e52..d6cb33c 100644 --- a/arp/handler.go +++ b/arp/handler.go @@ -129,8 +129,14 @@ func (h *Handler) DiscardQuery(protoAddr []byte) error { 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] + if !h.queries[i].isInvalid() { + 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++ } }