Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 26 additions & 15 deletions binary_heap.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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)
}
}
Expand Down Expand Up @@ -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
Expand All @@ -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()
Expand All @@ -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

Expand Down
74 changes: 74 additions & 0 deletions binary_heap_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading