diff --git a/binary_heap.go b/binary_heap.go index 83cb65d..2c71a95 100644 --- a/binary_heap.go +++ b/binary_heap.go @@ -14,8 +14,21 @@ type Item interface { GroupID() string } +// entry stores the insertion sequence to extract items with equal priority in FIFO order +type entry[T Item] struct { + item T + seq uint64 +} + +func (e entry[T]) less(o entry[T]) bool { + ep, op := e.item.Priority(), o.item.Priority() + return ep < op || (ep == op && e.seq < o.seq) +} + type BinHeap[T Item] struct { - items []T + items []entry[T] + // seq is the sequence number of the last inserted item + seq uint64 // exists used as a shadow structure to check if the item exists in the BinHeap exists map[string]struct{} st *stack @@ -26,7 +39,7 @@ type BinHeap[T Item] struct { func NewBinHeap[T Item](maxLen uint64) *BinHeap[T] { return &BinHeap[T]{ - items: make([]T, 0, 1000), + items: make([]entry[T], 0, 1000), exists: make(map[string]struct{}, 1000), st: newStack(), maxLen: maxLen, @@ -39,9 +52,7 @@ func (bh *BinHeap[T]) fixUp() { p := (k - 1) >> 1 // k-1 / 2 for k > 0 { - cur, par := (bh.items)[k], (bh.items)[p] - - if cur.Priority() < par.Priority() { + if bh.items[k].less(bh.items[p]) { bh.swap(k, p) k = p p = (k - 1) >> 1 @@ -64,10 +75,10 @@ func (bh *BinHeap[T]) fixDown(curr, end int) { } idxToSwap := cOneIdx - if cTwoIdx > -1 && (bh.items)[cTwoIdx].Priority() < (bh.items)[cOneIdx].Priority() { + if cTwoIdx > -1 && bh.items[cTwoIdx].less(bh.items[cOneIdx]) { idxToSwap = cTwoIdx } - if (bh.items)[idxToSwap].Priority() < (bh.items)[curr].Priority() { + if bh.items[idxToSwap].less(bh.items[curr]) { bh.swap(uint64(curr), uint64(idxToSwap)) //nolint:gosec curr = idxToSwap cOneIdx = (curr << 1) + 1 @@ -93,10 +104,10 @@ func (bh *BinHeap[T]) Remove(groupID string) []T { out := make([]T, 0, 10) for i := range bh.items { - if bh.items[i].GroupID() == groupID { + if bh.items[i].item.GroupID() == groupID { // delete element - delete(bh.exists, bh.items[i].ID()) - out = append(out, bh.items[i]) + delete(bh.exists, bh.items[i].item.ID()) + out = append(out, bh.items[i].item) bh.st.Add(i) } } @@ -134,7 +145,7 @@ func (bh *BinHeap[T]) PeekPriority() int64 { defer bh.cond.L.Unlock() if len(bh.items) > 0 { - return bh.items[0].Priority() + return bh.items[0].item.Priority() } return -1 @@ -153,7 +164,8 @@ func (bh *BinHeap[T]) Insert(item T) { bh.cond.Wait() } - bh.items = append(bh.items, item) + bh.seq++ + bh.items = append(bh.items, entry[T]{item: item, seq: bh.seq}) // fix binary heap up bh.fixUp() @@ -177,9 +189,8 @@ func (bh *BinHeap[T]) ExtractMin() T { n := uint64(len(bh.items)) bh.swap(0, n-1) - item := bh.items[n-1] - var zero T - bh.items[n-1] = zero + item := bh.items[n-1].item + bh.items[n-1] = entry[T]{} bh.items = bh.items[:n-1] bh.fixDown(0, int(n)-2) //nolint:gosec diff --git a/binary_heap_test.go b/binary_heap_test.go index 2b42e51..df3a803 100644 --- a/binary_heap_test.go +++ b/binary_heap_test.go @@ -849,6 +849,80 @@ func TestBinHeap_BroadcastPreventsDeadlock(t *testing.T) { }) } +func TestBinHeap_ExtractOrder(t *testing.T) { + // step inserts item, or calls ExtractMin when extract is true + type step struct { + item Test + extract bool + } + ins := func(priority int64, groupID, id string) step { + return step{item: NewTest(priority, groupID, id)} + } + ext := step{extract: true} + + tests := []struct { + name string + maxLen uint64 + steps []step + // remove is the group to remove after the steps + remove string + want []string + }{ + { + name: "equal priority, fill then drain", + maxLen: 10, + steps: []step{ins(10, "g", "1"), ins(10, "g", "2"), ins(10, "g", "3"), ins(10, "g", "4")}, + want: []string{"1", "2", "3", "4"}, + }, + { + name: "equal priority, full heap, one extract and one insert per round", + maxLen: 4, + steps: []step{ + ins(10, "g", "1"), ins(10, "g", "2"), ins(10, "g", "3"), ins(10, "g", "4"), + ext, ins(10, "g", "5"), + ext, ins(10, "g", "6"), + ext, ins(10, "g", "7"), + ext, ins(10, "g", "8"), + }, + want: []string{"1", "2", "3", "4", "5", "6", "7", "8"}, + }, + { + name: "mixed priority, equal priorities keep insertion order", + maxLen: 10, + steps: []step{ins(2, "g", "a"), ins(1, "g", "b"), ins(2, "g", "c"), ins(1, "g", "d")}, + want: []string{"b", "d", "a", "c"}, + }, + { + name: "equal priority after Remove of another group", + maxLen: 10, + steps: []step{ins(10, "g1", "a1"), ins(10, "g2", "b1"), ins(10, "g1", "a2"), ins(10, "g1", "a3"), ext}, + remove: "g2", + want: []string{"a1", "a2", "a3"}, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + bh := NewBinHeap[Test](tc.maxLen) + got := make([]string, 0, len(tc.want)) + for _, s := range tc.steps { + if s.extract { + got = append(got, bh.ExtractMin().ID()) + continue + } + bh.Insert(s.item) + } + if tc.remove != "" { + bh.Remove(tc.remove) + } + for bh.Len() > 0 { + got = append(got, bh.ExtractMin().ID()) + } + require.Equal(t, tc.want, got) + }) + } +} + func BenchmarkInsert(b *testing.B) { bh := NewBinHeap[Item](1 << 30) b.ReportAllocs()