From 2129eb858f2c0dd4a0d83525d651a7cb52242600 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Wed, 30 Sep 2026 13:12:43 -0700 Subject: [PATCH 1/7] queue: keep elements in a heap of values Push wrapped each element in an item, boxed it for container/heap and stored a pointer to a copy: two allocations per element. It then ran heap.Fix on the index of the original, which was always 0, so it re-sifted the root for nothing. Keep the elements in a slice and sift them directly with the algorithm of container/heap, so Push and Pop no longer allocate and the comparisons are no longer interface calls. Pushing and popping through a 64-element queue drops from 66.2 ns and 2 allocations to 23.9 ns and none (BenchmarkPushPop); replacing the top of a 64-way merge drops from 38.9 ns to 22.2 ns (BenchmarkPeekUpdate). Co-Authored-By: Claude Opus 5.5 --- queue/priority_queue.go | 107 ++++++++++++++++++--------------------- queue/regression_test.go | 77 ++++++++++++++++++++++++++++ 2 files changed, 125 insertions(+), 59 deletions(-) create mode 100644 queue/regression_test.go diff --git a/queue/priority_queue.go b/queue/priority_queue.go index f2f10f2..9bb9d33 100644 --- a/queue/priority_queue.go +++ b/queue/priority_queue.go @@ -1,36 +1,21 @@ // Package queue provides a generic priority queue implementation optimized for external sorting. -// It uses Go's container/heap package internally to maintain heap properties efficiently. +// It is a binary heap of values that follows the same algorithm as Go's container/heap, +// without its interface: elements are not boxed, so Push and Pop do not allocate. // The priority queue supports any type E with a user-provided comparison function. package queue -// Implementation is based on the Go standard library example: -// https://golang.org/pkg/container/heap/#example__priorityQueue - import ( - "container/heap" "fmt" ) -// item is a container for holding values with a priority in the queue -type item[E any] struct { - value E - // The index is needed by update and is maintained by the heap.Interface methods. - index int // The index of the item in the heap. -} - -// innerPriorityQueue implements heap.Interface and holds Items -type innerPriorityQueue[E any] struct { - items []*item[E] - compareFunc func(E, E) int -} - // PriorityQueue is a generic priority queue that maintains elements in sorted order // according to a user-provided comparison function. It provides efficient O(log n) // insertion and removal of the minimum/maximum element. This implementation is // specifically optimized for the external sorting use case where elements need // to be efficiently merged from multiple sorted streams. type PriorityQueue[E any] struct { - ipq innerPriorityQueue[E] + items []E // a binary heap: items[i] is not after its children 2i+1 and 2i+2 + compareFunc func(E, E) int } // NewPriorityQueue creates a new priority queue with the given comparison function. @@ -39,42 +24,42 @@ type PriorityQueue[E any] struct { // if the first should appear later. For ascending order, use cmp.Compare(a, b). // The queue starts empty and elements can be added with Push(). func NewPriorityQueue[E any](cmpFunc func(E, E) int) *PriorityQueue[E] { - var pq PriorityQueue[E] - pq.ipq.items = make([]*(item[E]), 0) - pq.ipq.compareFunc = cmpFunc - heap.Init(&pq.ipq) - return &pq + return &PriorityQueue[E]{compareFunc: cmpFunc} } // Len returns the current number of elements in the priority queue. // This operation is O(1). func (pq *PriorityQueue[E]) Len() int { - return pq.ipq.Len() + return len(pq.items) } // Push adds a new element to the priority queue, maintaining heap properties. // The element will be positioned according to the comparison function provided // during queue creation. This operation is O(log n). func (pq *PriorityQueue[E]) Push(x E) { - var i item[E] - i.value = x - heap.Push(&pq.ipq, i) - heap.Fix(&pq.ipq, i.index) + pq.items = append(pq.items, x) + pq.up(len(pq.items) - 1) } // Pop removes and returns the highest priority element from the queue. // The returned element is the one that would be returned by Peek(). // This operation is O(log n). Panics if the queue is empty. func (pq *PriorityQueue[E]) Pop() E { - item := heap.Pop(&pq.ipq).(*item[E]) - return item.value + n := len(pq.items) - 1 + top := pq.items[0] + pq.items[0] = pq.items[n] + var zero E + pq.items[n] = zero // don't keep the popped element reachable + pq.items = pq.items[:n] + pq.down(0) + return top } // Peek returns the highest priority element without removing it from the queue. // This allows inspection of the next element that would be returned by Pop(). // This operation is O(1). Panics if the queue is empty. func (pq *PriorityQueue[E]) Peek() E { - return pq.ipq.items[0].value + return pq.items[0] } // PeekUpdate must be called after modifying the value returned by Peek() in-place. @@ -82,7 +67,7 @@ func (pq *PriorityQueue[E]) Peek() E { // This is more efficient than Pop() followed by Push() when updating the top element. // This operation is O(log n). func (pq *PriorityQueue[E]) PeekUpdate() { - heap.Fix(&pq.ipq, 0) + pq.down(0) } // Print outputs the current contents of the priority queue to stdout. @@ -90,40 +75,44 @@ func (pq *PriorityQueue[E]) PeekUpdate() { // This method is primarily intended for debugging purposes. func (pq *PriorityQueue[E]) Print() { fmt.Print("[") - for i := range pq.ipq.items { - fmt.Print(pq.ipq.items[i].value, ", ") + for i := range pq.items { + fmt.Print(pq.items[i], ", ") } fmt.Println("]") } -func (pq *innerPriorityQueue[E]) Len() int { - return len(pq.items) -} - -func (pq *innerPriorityQueue[E]) Less(i, j int) bool { - // TODO make full use of compareFunc returning an int - return pq.compareFunc(pq.items[i].value, pq.items[j].value) < 0 +func (pq *PriorityQueue[E]) less(i, j int) bool { + return pq.compareFunc(pq.items[i], pq.items[j]) < 0 } -func (pq *innerPriorityQueue[E]) Swap(i, j int) { - pq.items[i], pq.items[j] = pq.items[j], pq.items[i] - pq.items[i].index = i - pq.items[j].index = j +// up moves element j towards the root until it is not before its parent. +func (pq *PriorityQueue[E]) up(j int) { + for j > 0 { + i := (j - 1) / 2 // parent + if !pq.less(j, i) { + break + } + pq.items[i], pq.items[j] = pq.items[j], pq.items[i] + j = i + } } -func (pq *innerPriorityQueue[E]) Push(x any) { +// down moves element i towards the leaves until neither child is before it. +func (pq *PriorityQueue[E]) down(i int) { n := len(pq.items) - i := x.(item[E]) - i.index = n - pq.items = append(pq.items, &i) -} - -func (pq *innerPriorityQueue[E]) Pop() any { - old := pq.items - n := len(old) - item := old[n-1] - item.index = -1 // for safety - pq.items = old[0 : n-1] - return item + for { + j := 2*i + 1 // left child + if j >= n || j < 0 { // j < 0 after int overflow + break + } + if j2 := j + 1; j2 < n && pq.less(j2, j) { + j = j2 // right child + } + if !pq.less(j, i) { + break + } + pq.items[i], pq.items[j] = pq.items[j], pq.items[i] + i = j + } } diff --git a/queue/regression_test.go b/queue/regression_test.go new file mode 100644 index 0000000..ff759e8 --- /dev/null +++ b/queue/regression_test.go @@ -0,0 +1,77 @@ +package queue_test + +// Regression tests for PriorityQueue allocations and PeekUpdate. + +import ( + "cmp" + "slices" + "testing" + + "github.com/lanrat/extsort/queue" +) + +// Push used to allocate twice per element (boxing it for container/heap, then storing a +// pointer to a copy) and ran a redundant heap.Fix. Push and Pop now allocate nothing +// once the queue has grown. +func TestPushPopDoNotAllocate(t *testing.T) { + q := queue.NewPriorityQueue(cmp.Compare[int]) + for i := range 64 { + q.Push(i * 1000) + } + allocs := testing.AllocsPerRun(100, func() { + q.Push(q.Pop() + 1000) + }) + if allocs != 0 { + t.Errorf("Push and Pop allocated %v times per call, want 0", allocs) + } +} + +// The merge updates the top element in place and calls PeekUpdate, as tested here. +func TestPeekUpdate(t *testing.T) { + type source struct{ next int } + q := queue.NewPriorityQueue(func(a, b *source) int { return cmp.Compare(a.next, b.next) }) + for _, v := range []int{5, 1, 4, 2, 3} { + q.Push(&source{v}) + } + var got []int + for q.Len() > 0 { + top := q.Peek() + got = append(got, top.next) + if top.next < 10 { + top.next += 10 // each source yields its value, then that value plus 10 + q.PeekUpdate() + } else { + q.Pop() + } + } + want := []int{1, 2, 3, 4, 5, 11, 12, 13, 14, 15} + if !slices.Equal(got, want) { + t.Errorf("got %v, want %v", got, want) + } +} + +// BenchmarkPeekUpdate replaces the top of a 64-element queue, as a 64-way merge does. +func BenchmarkPeekUpdate(b *testing.B) { + type source struct{ next int } + q := queue.NewPriorityQueue(func(a, b *source) int { return cmp.Compare(a.next, b.next) }) + for i := range 64 { + q.Push(&source{i}) + } + b.ReportAllocs() + for b.Loop() { + q.Peek().next += 64 + q.PeekUpdate() + } +} + +// BenchmarkPushPop pushes and pops through a 64-element queue. +func BenchmarkPushPop(b *testing.B) { + q := queue.NewPriorityQueue(cmp.Compare[int]) + for i := range 64 { + q.Push(i * 1000) + } + b.ReportAllocs() + for b.Loop() { + q.Push(q.Pop() + 64000) + } +} From 584a942fb4e294e7bd542051ec2253d7ae17926a Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Wed, 30 Sep 2026 13:12:43 -0700 Subject: [PATCH 2/7] diff: receive without locking ctx's channel for every value Every read was a select on the data channel and ctx.Done(), and a select locks both channels, so readers of a busy stream contended on the context's channel for every value. Try a non-blocking receive first and fall back to the select only when the stream has nothing ready. ctx is then checked every 1024 reads, starting with the first, so a diff with values always ready still sees a cancellation, and a diff whose ctx is already cancelled returns before reading anything (the select used to pick a case at random). Diffing two streams of 1M ints drops from 179 ms to 70 ms (BenchmarkDiffOrdered). Co-Authored-By: Claude Opus 5.5 --- diff/diff_generic.go | 77 ++++++++++++++++++++++----------------- diff/regression_test.go | 79 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 124 insertions(+), 32 deletions(-) diff --git a/diff/diff_generic.go b/diff/diff_generic.go index 4c63541..681d549 100644 --- a/diff/diff_generic.go +++ b/diff/diff_generic.go @@ -5,6 +5,10 @@ import ( "fmt" ) +// ctxCheckInterval is how many values recv reads between checks of the context +// while its non-blocking receive keeps succeeding. +const ctxCheckInterval = 1024 + // differ is an internal struct that holds the state for performing diff operations // between two sorted channels of type T. It manages the comparison logic and // result reporting through callback functions. @@ -14,6 +18,7 @@ type differ[T any] struct { aErrChan, bErrChan <-chan error resultFunc ResultFunc[T] compare CompareFunc[T] + reads int // values read by recv } // Generic performs a diff operation on two sorted channels of any comparable type T. @@ -53,16 +58,12 @@ func (d *differ[T]) diff() (r Result, err error) { var okA, okB bool // read from channel A - select { - case dataA, okA = <-d.aChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataA, okA, err = d.recv(d.aChan); err != nil { + return } // read from channel B - select { - case dataB, okB = <-d.bChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataB, okB, err = d.recv(d.bChan); err != nil { + return } for okA && okB { c := d.compare(dataA, dataB) @@ -73,10 +74,8 @@ func (d *differ[T]) diff() (r Result, err error) { if err != nil { return } - select { - case dataB, okB = <-d.bChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataB, okB, err = d.recv(d.bChan); err != nil { + return } } else if c < 0 { r.TotalA++ @@ -85,25 +84,19 @@ func (d *differ[T]) diff() (r Result, err error) { if err != nil { return } - select { - case dataA, okA = <-d.aChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataA, okA, err = d.recv(d.aChan); err != nil { + return } } else { // common r.Common++ r.TotalA++ r.TotalB++ - select { - case dataA, okA = <-d.aChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataA, okA, err = d.recv(d.aChan); err != nil { + return } - select { - case dataB, okB = <-d.bChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataB, okB, err = d.recv(d.bChan); err != nil { + return } } } @@ -128,10 +121,8 @@ func (d *differ[T]) diff() (r Result, err error) { if err != nil { return } - select { - case dataA, okA = <-d.aChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataA, okA, err = d.recv(d.aChan); err != nil { + return } } // check for A errors if not read above @@ -148,10 +139,8 @@ func (d *differ[T]) diff() (r Result, err error) { if err != nil { return } - select { - case dataB, okB = <-d.bChan: - case <-d.ctx.Done(): - return r, d.ctx.Err() + if dataB, okB, err = d.recv(d.bChan); err != nil { + return } } // check for B errors if not read above @@ -163,6 +152,30 @@ func (d *differ[T]) diff() (r Result, err error) { return } +// recv reads the next value from ch. It tries a non-blocking receive first: unlike a +// select with ctx.Done(), that does not lock the context's channel, so a stream with a +// value ready costs one channel operation. ctx is then checked every ctxCheckInterval +// values, starting with the first, and whenever ch has nothing ready. +func (d *differ[T]) recv(ch <-chan T) (v T, ok bool, err error) { + if d.reads%ctxCheckInterval == 0 { + if err := d.ctx.Err(); err != nil { + return v, false, err + } + } + d.reads++ + select { + case v, ok = <-ch: + return v, ok, nil + default: + } + select { + case v, ok = <-ch: + return v, ok, nil + case <-d.ctx.Done(): + return v, false, d.ctx.Err() + } +} + // readErr waits for the error from a stream whose data channel has closed. // It gives up when ctx is done, so an error channel that is never closed cannot hang the diff. func (d *differ[T]) readErr(errChan <-chan error) error { diff --git a/diff/regression_test.go b/diff/regression_test.go index 27d7748..7a3e3f5 100644 --- a/diff/regression_test.go +++ b/diff/regression_test.go @@ -106,3 +106,82 @@ func TestStringResultChanContextStopsWaiting(t *testing.T) { t.Fatalf("got error %v, want %v", err, context.Canceled) } } + +// ints returns a closed channel holding n values of start, start+step, ... +func ints(n, start, step int) chan int { + ch := make(chan int, n) + for i := range n { + ch <- start + i*step + } + close(ch) + return ch +} + +// Reads try a non-blocking receive before selecting on ctx.Done(), so while values are +// ready ctx is only checked every so often. A cancellation must still stop the diff, +// including one that happened before the diff started. +func TestDiffStopsWhenCancelledWithValuesReady(t *testing.T) { + closedErr := func() chan error { + ch := make(chan error) + close(ch) + return ch + } + t.Run("cancelled before", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + calls := 0 + _, err := diff.Ordered(ctx, ints(10, 0, 2), ints(10, 1, 2), closedErr(), closedErr(), func(diff.Delta, int) error { + calls++ + return nil + }) + if !errors.Is(err, context.Canceled) || calls != 0 { + t.Fatalf("got error %v after %d results, want %v before any", err, calls, context.Canceled) + } + }) + t.Run("cancelled during", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + const n = 100_000 + calls := 0 + r, err := diff.Ordered(ctx, ints(n, 0, 2), ints(n, 1, 2), closedErr(), closedErr(), func(diff.Delta, int) error { + calls++ + if calls == 10 { + cancel() + } + return nil + }) + if !errors.Is(err, context.Canceled) { + t.Fatalf("got error %v, want %v", err, context.Canceled) + } + if r.TotalA+r.TotalB > 10+2*1024 { + t.Errorf("read %d values after the cancel, want at most about 1024 per stream", r.TotalA+r.TotalB-10) + } + }) +} + +// BenchmarkDiffOrdered diffs two streams of 1M ints fed by producer goroutines, with a +// cancellable context as callers use. +func BenchmarkDiffOrdered(b *testing.B) { + const n = 1_000_000 + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + produce := func(start, step int) (chan int, chan error) { + ch, errc := make(chan int, 1000), make(chan error) + go func() { + defer close(errc) + defer close(ch) + for i := range n { + ch <- start + i*step + } + }() + return ch, errc + } + for b.Loop() { + a, aErr := produce(0, 2) // even numbers + c, cErr := produce(0, 3) // multiples of 3 + r, err := diff.Ordered(ctx, a, c, aErr, cErr, func(diff.Delta, int) error { return nil }) + if err != nil || r.TotalA != n || r.TotalB != n { + b.Fatalf("got %s, %v", r.String(), err) + } + } +} From 60595a65c0724cc3527547854b2287a6c3a979aa Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Wed, 30 Sep 2026 13:12:43 -0700 Subject: [PATCH 3/7] Remove the unused byteSlicePool Every sorter created the pool, and nothing ever took a slice from it. Co-Authored-By: Claude Opus 5.5 --- sort_generic.go | 17 ++++------------- 1 file changed, 4 insertions(+), 13 deletions(-) diff --git a/sort_generic.go b/sort_generic.go index 87b3422..18e9feb 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -48,10 +48,9 @@ func (s *GenericSorter[E]) putChunk(c *genericChunk[E]) { // memoryPools holds sync.Pool instances for memory reuse type memoryPools struct { - chunkPool sync.Pool // *chunk objects - slicePool sync.Pool // []any slices - byteSlicePool sync.Pool // []byte slices for serialization - scratchPool sync.Pool // scratch buffers for binary encoding + chunkPool sync.Pool // *chunk objects + slicePool sync.Pool // []any slices + scratchPool sync.Pool // scratch buffers for binary encoding } // GenericSorter implements external sorting for any type E using a divide-and-conquer approach. @@ -98,7 +97,7 @@ func newSorter[E any](input <-chan E, fromBytes FromBytesGeneric[E], toBytes ToB } // initMemoryPools initializes sync.Pool instances for efficient memory reuse during sorting. -// Creates pools for chunks, slices, byte slices, and scratch buffers to reduce GC pressure +// Creates pools for chunks, slices, and scratch buffers to reduce GC pressure // and improve performance during high-frequency allocation/deallocation cycles. func (s *GenericSorter[E]) initMemoryPools() *memoryPools { pools := &memoryPools{} @@ -118,14 +117,6 @@ func (s *GenericSorter[E]) initMemoryPools() *memoryPools { }, } - // Pool for byte slices (for serialization) - store pointers to slices - pools.byteSlicePool = sync.Pool{ - New: func() any { - slice := make([]byte, 0, 1024) // Start with 1KB capacity - return &slice - }, - } - // Pool for scratch buffers (for binary encoding) - store pointers to slices pools.scratchPool = sync.Pool{ New: func() any { From 07d4eecb6b26c0ac70db7d834b3c590501b7b6c1 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Wed, 30 Sep 2026 13:12:43 -0700 Subject: [PATCH 4/7] Batch records between the parallel merge workers and the final merge Each merge worker sent its records to the final merge one per channel operation, every send a select on the same ctx.Done() channel as all the other workers, so the handoffs cost more than the merge: most CPU time went to the scheduler and to channel locks. Workers now send batches of 1024 records, with up to two queued per worker, and the final merge hands used-up batches back to the worker for reuse. Errors, cancellation and panic recovery take the same paths as before, and a worker checks ctx as it sends each batch, so it stops within one batch of a cancel. Sorting 1M ints with Generic drops from 350 ms to 134 ms (BenchmarkGenericVarintInts) and 1M strings from 377 ms to 183 ms (BenchmarkLegacyStringSort/size_1000000), and 2,000 chunks merge in 52 ms instead of 111 ms (BenchmarkSortManyChunks). New tests make a read error partway into a chunk, a compareFunc panic and a cancel each happen mid-merge, and check that the error reaches the error channel, the temp file is closed and no sorter goroutine is left running. Another covers record counts around the batch size. trackedTemp gains readErrAfter to fail a section partway, and the final merge panic test feeds the final merge through batch streams. Co-Authored-By: Claude Opus 5.5 --- merge_test.go | 177 +++++++++++++++++++++++++++++++++++++++++++++ regression_test.go | 20 ++--- sort_generic.go | 140 +++++++++++++++++++++++++---------- 3 files changed, 289 insertions(+), 48 deletions(-) create mode 100644 merge_test.go diff --git a/merge_test.go b/merge_test.go new file mode 100644 index 0000000..ede7e71 --- /dev/null +++ b/merge_test.go @@ -0,0 +1,177 @@ +package extsort + +// Tests for the parallel merge, whose workers hand records to the final merge in batches. + +import ( + "cmp" + "context" + "errors" + "math/rand" + "runtime" + "slices" + "strconv" + "strings" + "sync/atomic" + "testing" + "time" +) + +// sorterGoroutines counts the goroutines running a GenericSorter method. +func sorterGoroutines() (int, string) { + buf := make([]byte, 1<<22) + all := string(buf[:runtime.Stack(buf, true)]) + n := 0 + for _, g := range strings.Split(all, "\n\n") { + if strings.Contains(g, "extsort.(*GenericSorter") { + n++ + } + } + return n, all +} + +// checkSorterGoroutinesExit fails the test unless the goroutines running GenericSorter +// methods drop back to before within a few seconds. +func checkSorterGoroutinesExit(t *testing.T, before int) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for { + n, stacks := sorterGoroutines() + if n <= before { + return + } + if time.Now().After(deadline) { + t.Fatalf("%d sorter goroutine(s) still running after the sort ended:\n%s", n-before, stacks) + } + time.Sleep(10 * time.Millisecond) + } +} + +// intsChan returns a closed channel holding values. +func intsChan(values []int) chan int { + ch := make(chan int, len(values)) + for _, v := range values { + ch <- v + } + close(ch) + return ch +} + +// Batches must carry every record, in order, whatever the record count and chunk size. +func TestParallelMergeBatchBoundaries(t *testing.T) { + r := rand.New(rand.NewSource(1)) + for _, n := range []int{mergeBatchSize - 1, mergeBatchSize, mergeBatchSize + 1, 3*mergeBatchSize + 7} { + values := make([]int, n) + for i := range values { + values[i] = r.Intn(n / 2) // duplicates, so the merge sees ties + } + want := slices.Sorted(slices.Values(values)) + for _, chunkSize := range []int{7, mergeBatchSize - 1, mergeBatchSize + 1} { + for _, workers := range []int{2, 3} { + t.Run("n="+strconv.Itoa(n)+" chunk="+strconv.Itoa(chunkSize)+" workers="+strconv.Itoa(workers), func(t *testing.T) { + s, out, errc := MockGeneric(intsChan(values), atoiBytes, itoaBytes, cmp.Compare[int], &Config{ChunkSize: chunkSize, NumWorkers: workers}, 0) + s.Sort(context.Background()) + got, err := drainWithTimeout(t, out, errc) + if err != nil { + t.Fatal(err) + } + if !slices.Equal(got, want) { + t.Errorf("got %d records (sorted: %v), want the %d input records in order", len(got), slices.IsSorted(got), n) + } + }) + } + } + } +} + +// 50 chunks of 1000 records merged by 4 workers, so the parallel merge is used and the +// final merge sends records on while the workers still have batches to send. +const ( + midMergeRecords = 50_000 + midMergeChunkSize = 1000 + midMergeWorkers = 4 +) + +// A worker's read error in the middle of a chunk must reach the error channel, and every +// merge goroutine must stop, including workers blocked sending a batch. +func TestParallelMergeReadErrorMidChunk(t *testing.T) { + before, _ := sorterGoroutines() + errDisk := errors.New("disk read error") + // Section 30 fails a third of the way in. Its worker merges the chunks with + // smaller records first, so the error comes after many records were sent on. + tt := &trackedTemp{readErr: map[int]error{30: errDisk}, readErrAfter: map[int]int64{30: 2000}} + s := newTrackedSorter(descendingInts(midMergeRecords), itoaBytes, cmp.Compare[int], + &Config{ChunkSize: midMergeChunkSize, NumWorkers: midMergeWorkers}, tt) + got, err := runSort(t, context.Background(), s) + if !errors.Is(err, errDisk) { + t.Fatalf("got error %v, want %q", err, errDisk) + } + if len(got) == 0 || len(got) >= midMergeRecords || !slices.IsSorted(got) { + t.Errorf("got %d records (sorted: %v) before the error, want some but not all, in order", len(got), slices.IsSorted(got)) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + checkSorterGoroutinesExit(t, before) +} + +// A compareFunc panic in the middle of the merge, in a worker or the final merge, must +// reach the error channel as a ComparisonError, and every merge goroutine must stop. +func TestParallelMergeComparePanicMidMerge(t *testing.T) { + before, _ := sorterGoroutines() + var armed atomic.Bool + var calls atomic.Int64 + compare := func(a, b int) int { + if armed.Load() && calls.Add(1) == 1000 { + panic("compare failed") + } + return cmp.Compare(a, b) + } + tt := &trackedTemp{} + s := newTrackedSorter(descendingInts(midMergeRecords), itoaBytes, compare, + &Config{ChunkSize: midMergeChunkSize, NumWorkers: midMergeWorkers}, tt) + // Sort returns once every chunk is sorted and saved. The merge then blocks on the + // full output channel well before its end, so only the merge can panic. + s.Sort(context.Background()) + armed.Store(true) + got, err := drainWithTimeout(t, s.mergeChunkChan, s.mergeErrChan) + var cmpErr *ComparisonError + if !errors.As(err, &cmpErr) { + t.Fatalf("got error %v, want a ComparisonError", err) + } + if len(got) >= midMergeRecords || !slices.IsSorted(got) { + t.Errorf("got %d records (sorted: %v), want fewer than %d, in order", len(got), slices.IsSorted(got), midMergeRecords) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + checkSorterGoroutinesExit(t, before) +} + +// Cancelling in the middle of the merge must reach the error channel, stop the output +// soon after, and stop every merge goroutine. +func TestParallelMergeCancelMidMerge(t *testing.T) { + before, _ := sorterGoroutines() + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + tt := &trackedTemp{} + s := newTrackedSorter(descendingInts(midMergeRecords), itoaBytes, cmp.Compare[int], + &Config{ChunkSize: midMergeChunkSize, NumWorkers: midMergeWorkers}, tt) + s.Sort(ctx) + const read = 5000 + for range read { + <-s.mergeChunkChan + } + cancel() + got, err := drainWithTimeout(t, s.mergeChunkChan, s.mergeErrChan) + if !errors.Is(err, context.Canceled) { + t.Fatalf("got error %v, want %v", err, context.Canceled) + } + // At most the output channel's buffer, and the record being sent, follow the cancel + if after := len(got); after > cap(s.mergeChunkChan)+1 { + t.Errorf("got %d more records after the cancel, want at most %d", after, cap(s.mergeChunkChan)+1) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + checkSorterGoroutinesExit(t, before) +} diff --git a/regression_test.go b/regression_test.go index 7c8affa..9fa20d3 100644 --- a/regression_test.go +++ b/regression_test.go @@ -89,8 +89,9 @@ func drainWithTimeout[E any](t *testing.T, out <-chan E, errc <-chan error) ([]E // trackedTemp hands out in-memory temp files and counts how many the sorter // leaves open. readErr and closeErr inject failures. type trackedTemp struct { - readErr map[int]error // section index -> error returned when reading it - closeErr error // returned by the reader's Close + readErr map[int]error // section index -> error returned when reading it + readErrAfter map[int]int64 // section index -> bytes it reads before its readErr, 0 if unset + closeErr error // returned by the reader's Close mu sync.Mutex created int @@ -152,7 +153,8 @@ type trackedReader struct { func (r *trackedReader) Read(i int) *bufio.Reader { if err, ok := r.tt.readErr[i]; ok { - return bufio.NewReader(iotest.ErrReader(err)) + valid := io.LimitReader(r.TempReader.Read(i), r.tt.readErrAfter[i]) + return bufio.NewReader(io.MultiReader(valid, iotest.ErrReader(err))) } return r.TempReader.Read(i) } @@ -377,12 +379,12 @@ func TestCallbackPanicsBecomeErrors(t *testing.T) { t.Run("compare in final merge", func(t *testing.T) { s := newSorter(nil, atoiBytes, itoaBytes, panicCompare, nil) - a, b := make(chan int, 1), make(chan int, 1) - a <- 1 - b <- 2 - close(a) - close(b) - if err := s.finalMergeSimple(context.Background(), []chan int{a, b}); !errors.As(err, &cmpErr) { + a, b := newMergeStream[int](), newMergeStream[int]() + a.batches <- []int{1} + b.batches <- []int{2} + close(a.batches) + close(b.batches) + if err := s.finalMergeSimple(context.Background(), []mergeStream[int]{a, b}); !errors.As(err, &cmpErr) { t.Fatalf("got error %v, want a ComparisonError", err) } }) diff --git a/sort_generic.go b/sort_generic.go index 18e9feb..2229155 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -15,6 +15,14 @@ import ( "golang.org/x/sync/errgroup" ) +const ( + // mergeBatchSize is how many records a merge worker hands to the final merge at a time. + // Sending records one per channel operation made the handoffs cost more than the merge. + mergeBatchSize = 1024 + // mergeBatchBuffer is how many full batches a merge worker can queue for the final merge. + mergeBatchBuffer = 2 +) + // genericChunk represents a collection of any data that can be sorted. // It holds data in memory before being sorted using slices.SortFunc. type genericChunk[E any] struct { @@ -592,11 +600,10 @@ func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) error { mergeCtx, mergeCancel := context.WithCancel(ctx) defer mergeCancel() // Ensure all goroutines stop when this function returns - // Create intermediate channels for each worker - intermediateChanSize := s.config.SortedChanBuffSize - intermediateChans := make([]chan E, numWorkers) - for i := range intermediateChans { - intermediateChans[i] = make(chan E, intermediateChanSize) + // Create a stream for each worker to pass its merged records to the final merge in batches + streams := make([]mergeStream[E], numWorkers) + for i := range streams { + streams[i] = newMergeStream[E]() } // Error collection @@ -621,15 +628,15 @@ func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) error { workersStarted++ wg.Add(1) - go func(workerIdx, start, end int) { + go func(stream mergeStream[E], start, end int) { defer wg.Done() - defer close(intermediateChans[workerIdx]) // Each worker closes its own channel + defer close(stream.batches) // Each worker closes its own channel - if err := s.mergeWorkerSimple(mergeCtx, start, end, intermediateChans[workerIdx]); err != nil { + if err := s.mergeWorkerSimple(mergeCtx, start, end, stream); err != nil { errChan <- err mergeCancel() // Cancel all operations on error } - }(i, startChunk, endChunk) + }(streams[i], startChunk, endChunk) } // Start error collector with wait group for synchronization @@ -652,7 +659,7 @@ func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) error { finalMergeWg.Add(1) go func() { defer finalMergeWg.Done() - if err := s.finalMergeSimple(mergeCtx, intermediateChans[:workersStarted]); err != nil { + if err := s.finalMergeSimple(mergeCtx, streams[:workersStarted]); err != nil { errChan <- err mergeCancel() // Stop the workers, which may be blocked sending to the final merge } @@ -676,8 +683,46 @@ func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) error { return ctx.Err() } +// mergeStream carries one merge worker's records to the final merge in batches. +type mergeStream[E any] struct { + batches chan []E // batches of records in merge order, closed when the worker stops + free chan []E // used-up batches handed back to the worker for reuse +} + +func newMergeStream[E any]() mergeStream[E] { + return mergeStream[E]{ + batches: make(chan []E, mergeBatchBuffer), + // room for every batch a worker has: those queued, the one it fills and the one being merged + free: make(chan []E, mergeBatchBuffer+2), + } +} + +// emptyBatch returns a batch to fill, reusing one the final merge handed back if there is one. +func (m mergeStream[E]) emptyBatch() []E { + select { + case batch := <-m.free: + return batch[:0] + default: + return make([]E, 0, mergeBatchSize) + } +} + +// send queues a batch for the final merge. Once ctx is done it returns ctx's error instead, +// so a worker stops at its next batch after a cancellation. +func (m mergeStream[E]) send(ctx context.Context, batch []E) error { + if err := ctx.Err(); err != nil { + return err // the select below picks at random when the channel has room too + } + select { + case m.batches <- batch: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + // mergeWorkerSimple merges a subset of chunks with proper context handling -func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, endChunk int, output chan<- E) (err error) { +func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, endChunk int, output mergeStream[E]) (err error) { // A panicking compareFunc must not crash the process from this goroutine defer func() { if r := recover(); r != nil { @@ -705,15 +750,9 @@ func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, en pq.Push(merge) } - // Merge this worker's chunks + // Merge this worker's chunks, checking ctx as each batch is sent + batch := output.emptyBatch() for pq.Len() > 0 { - // Check context before processing - select { - case <-ctx.Done(): - return ctx.Err() - default: - } - merge := pq.Peek() rec, more, err := merge.getNext() if err != nil { @@ -725,19 +764,23 @@ func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, en pq.Pop() } - select { - case output <- rec: - case <-ctx.Done(): - return ctx.Err() + batch = append(batch, rec) + if len(batch) == mergeBatchSize { + if err := output.send(ctx, batch); err != nil { + return err + } + batch = output.emptyBatch() } } - + if len(batch) > 0 { + return output.send(ctx, batch) + } return nil } // finalMergeSimple performs streaming merge with simpler synchronization. // It returns nil when ctx is cancelled; the caller reports the cancellation. -func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, intermediateChans []chan E) (err error) { +func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, streams []mergeStream[E]) (err error) { // A panicking compareFunc must not crash the process from this goroutine defer func() { if r := recover(); r != nil { @@ -750,8 +793,8 @@ func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, intermediateCha }) // Initialize sources - for _, ch := range intermediateChans { - source := &channelMergeSource[E]{ch: ch} + for _, stream := range streams { + source := &channelMergeSource[E]{stream: stream} if source.getNextSimple() { pq.Push(source) } @@ -783,23 +826,42 @@ func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, intermediateCha return nil } -// channelMergeSource represents a source of sorted data from a channel +// channelMergeSource represents a source of sorted data from a merge worker's stream type channelMergeSource[E any] struct { - ch <-chan E + stream mergeStream[E] + batch []E // the batch being merged + pos int // index in batch of the record after nextRec nextRec E - hasNext bool } -// getNextSimple reads from channel without context (channel close handles cancellation) +// getNextSimple advances to the next record, receiving the next batch once the current one +// is used up. It reads without context: the worker closes its channel when it stops. func (c *channelMergeSource[E]) getNextSimple() bool { - rec, ok := <-c.ch - if ok { - c.nextRec = rec - c.hasNext = true - return true - } - c.hasNext = false - return false + for c.pos == len(c.batch) { + c.releaseBatch() + batch, ok := <-c.stream.batches + if !ok { + return false + } + c.batch, c.pos = batch, 0 + } + c.nextRec = c.batch[c.pos] + c.pos++ + return true +} + +// releaseBatch hands the used-up batch back to the worker, or drops it if the worker +// already has enough spare batches. +func (c *channelMergeSource[E]) releaseBatch() { + if c.batch == nil { + return + } + clear(c.batch) // the records were sent on; don't keep them reachable from a spare batch + select { + case c.stream.free <- c.batch: + default: + } + c.batch = nil } // mergeFile represents each sorted chunk on disk and its next value From 925c333dd163f8dda293660292d31918eadc8ef9 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Wed, 30 Sep 2026 13:12:43 -0700 Subject: [PATCH 5/7] Receive input without locking ctx's channel for every record buildChunks read every record with a select on the input and on ctx.Done(), which locks both channels. Try a non-blocking receive first and fall back to the select only when the input has nothing ready. A steady input then costs one channel operation per record, and ctx is checked every 1024 records so it still sees a cancellation. Sorting 1M ints fed by a goroutine drops from 134 ms to 109 ms (BenchmarkGenericVarintInts). The output send keeps its select: a non-blocking send first made the merge slower when measured. Co-Authored-By: Claude Opus 5.5 --- sort_generic.go | 31 +++++++++++++++++++++++-------- 1 file changed, 23 insertions(+), 8 deletions(-) diff --git a/sort_generic.go b/sort_generic.go index 2229155..e16492b 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -16,6 +16,9 @@ import ( ) const ( + // ctxCheckInterval is how many records a loop handles between checks of its context + // while a non-blocking receive keeps succeeding, since that path never selects on ctx.Done(). + ctxCheckInterval = 1024 // mergeBatchSize is how many records a merge worker hands to the final merge at a time. // Sending records one per channel operation made the handoffs cost more than the merge. mergeBatchSize = 1024 @@ -272,17 +275,29 @@ func (s *GenericSorter[E]) buildChunks() error { c := s.getChunk() fill: for i := 0; i < s.config.ChunkSize; i++ { + var rec E + var ok bool + // Try a non-blocking receive first: unlike the select below it does not lock + // the context's channel, so a steady input costs one channel operation per record select { - case rec, ok := <-s.input: - if !ok { - inputOpen = false - break fill // a plain break would only leave the select + case rec, ok = <-s.input: + if i%ctxCheckInterval == ctxCheckInterval-1 && s.sortCtx.Err() != nil { + s.putChunk(c) // Return unused chunk to pool + return s.sortCtx.Err() } - c.data = append(c.data, rec) - case <-s.sortCtx.Done(): - s.putChunk(c) // Return unused chunk to pool - return s.sortCtx.Err() + default: + select { + case rec, ok = <-s.input: + case <-s.sortCtx.Done(): + s.putChunk(c) // Return unused chunk to pool + return s.sortCtx.Err() + } + } + if !ok { + inputOpen = false + break fill } + c.data = append(c.data, rec) } if len(c.data) == 0 { // the chunk is empty, return it to pool From 9ce531aef082c54ba2b54b393e010a17ac216220 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Wed, 30 Sep 2026 13:12:43 -0700 Subject: [PATCH 6/7] Grow chunks instead of allocating ChunkSize up front MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The slice pool gave every chunk a slice of capacity ChunkSize, so sorting 10 ints with the default config allocated 8 MB. A chunk now starts at 1024 records and doubles, to exactly ChunkSize. Once a chunk has filled, the input spans several chunks, so later chunks take their full size at once instead of growing. Sorting 10 records drops from 47.7 µs and 7.6 MiB to 10.0 µs and 23 KiB (BenchmarkSortTenRecords). Large sorts pay for growing the first chunk once: BenchmarkGenericVarintInts allocates 3.5% more, in the same time. Co-Authored-By: Claude Opus 5.5 --- buffers_test.go | 46 ++++++++++++++++++++++++++++++++++++++++++++++ sort_generic.go | 32 ++++++++++++++++++++++++++++---- 2 files changed, 74 insertions(+), 4 deletions(-) create mode 100644 buffers_test.go diff --git a/buffers_test.go b/buffers_test.go new file mode 100644 index 0000000..0c735b3 --- /dev/null +++ b/buffers_test.go @@ -0,0 +1,46 @@ +package extsort + +// Tests for the buffers the sorter reuses: chunk slices. + +import ( + "cmp" + "context" + "slices" + "testing" +) + +// The first chunk used to allocate a full ChunkSize slice up front: 8 MB to sort 10 ints +// with the default config. A chunk now grows with its records, to exactly ChunkSize, and +// once a chunk has filled, the later ones take their full size at once. +func TestChunksGrowWithInput(t *testing.T) { + for _, tc := range []struct { + name string + n int + chunkSize int + wantCaps []int + }{ + {"small input", 10, 1 << 20, []int{firstChunkCap}}, + {"exactly one chunk", 5000, 5000, []int{5000}}, + {"several chunks", 12_000, 5000, []int{5000, 5000, 5000}}, + } { + t.Run(tc.name, func(t *testing.T) { + in := make(chan int, tc.n) + for i := range tc.n { + in <- i + } + close(in) + s := newSorter(in, atoiBytes, itoaBytes, cmp.Compare[int], &Config{ChunkSize: tc.chunkSize, ChanBuffSize: 8}) + s.sortCtx = context.Background() + if err := s.buildChunks(); err != nil { + t.Fatal(err) + } + var caps []int + for c := range s.chunkChan { + caps = append(caps, cap(c.data)) + } + if !slices.Equal(caps, tc.wantCaps) { + t.Errorf("chunk capacities %v, want %v", caps, tc.wantCaps) + } + }) + } +} diff --git a/sort_generic.go b/sort_generic.go index e16492b..f237cb8 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -19,6 +19,8 @@ const ( // ctxCheckInterval is how many records a loop handles between checks of its context // while a non-blocking receive keeps succeeding, since that path never selects on ctx.Done(). ctxCheckInterval = 1024 + // firstChunkCap is the capacity a chunk starts with before the input has filled a chunk. + firstChunkCap = 1024 // mergeBatchSize is how many records a merge worker hands to the final merge at a time. // Sending records one per channel operation made the handoffs cost more than the merge. mergeBatchSize = 1024 @@ -60,7 +62,7 @@ func (s *GenericSorter[E]) putChunk(c *genericChunk[E]) { // memoryPools holds sync.Pool instances for memory reuse type memoryPools struct { chunkPool sync.Pool // *chunk objects - slicePool sync.Pool // []any slices + slicePool sync.Pool // []E slices scratchPool sync.Pool // scratch buffers for binary encoding } @@ -120,11 +122,11 @@ func (s *GenericSorter[E]) initMemoryPools() *memoryPools { }, } - // Pool for slices - store pointers to slices + // Pool for slices - store pointers to slices. New slices start empty and + // buildChunks grows them, so a small input does not allocate a full ChunkSize slice. pools.slicePool = sync.Pool{ New: func() any { - slice := make([]E, 0, s.config.ChunkSize) - return &slice + return new([]E) }, } @@ -271,6 +273,9 @@ func (s *GenericSorter[E]) closeTempFiles() { func (s *GenericSorter[E]) buildChunks() error { defer close(s.chunkChan) // if this is not called on error, causes a deadlock + // Set once a chunk fills: the input spans several chunks, so new chunk slices + // are allocated at their full size instead of grown. + spansChunks := false for inputOpen := true; inputOpen; { c := s.getChunk() fill: @@ -297,8 +302,14 @@ func (s *GenericSorter[E]) buildChunks() error { inputOpen = false break fill } + if len(c.data) == cap(c.data) { + c.data = s.growChunk(c.data, spansChunks) + } c.data = append(c.data, rec) } + if len(c.data) == s.config.ChunkSize { + spansChunks = true + } if len(c.data) == 0 { // the chunk is empty, return it to pool s.putChunk(c) @@ -317,6 +328,19 @@ func (s *GenericSorter[E]) buildChunks() error { return nil } +// growChunk returns data with room for at least one more record, up to ChunkSize. The first +// chunk doubles from firstChunkCap, so a small input only allocates what it needs; once the +// input has filled a chunk (full), a chunk grows straight to ChunkSize. +func (s *GenericSorter[E]) growChunk(data []E, full bool) []E { + newCap := s.config.ChunkSize + if !full { + newCap = min(newCap, max(2*cap(data), firstChunkCap)) + } + grown := make([]E, len(data), newCap) // exact capacity: append could grow past ChunkSize + copy(grown, data) + return grown +} + // sortChunks is a worker for sorting the data stored in a chunk prior to save func (s *GenericSorter[E]) sortChunks() error { for { From dab2261d6af56574656dd3d3916fa07670bc4dff Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Wed, 30 Sep 2026 13:12:43 -0700 Subject: [PATCH 7/7] Reuse the encode and decode buffers of the built-in codecs Strings and Ordered allocated a slice to encode every record and another to read it back in the merge. Their codecs now append each record to one buffer reused across the chunk, written together with its length header in a single Write, and the merge reads all the records of a chunk into one reused buffer. That is only safe because these fromBytes functions never keep the slice they are given; Generic keeps allocating, since a user's fromBytes may. The scratch-buffer pool for the length header becomes a stack array. Sorting 1M ints with Ordered drops from 45.4 MiB and 3.0M allocations to 6.8 MiB and 277, and from 112 ms to 104 ms (BenchmarkOrderedInts). 1M strings take 34.9 MiB and 1M allocations instead of 51.4 MiB and 3M (BenchmarkLegacyStringSort/size_1000000). Co-Authored-By: Claude Opus 5.5 --- buffers_test.go | 102 ++++++++++++++++++++++++++++++++++++- sort_generic.go | 130 ++++++++++++++++++++++++++++++++---------------- sort_ordered.go | 58 +++++++++++---------- sort_strings.go | 7 +++ 4 files changed, 229 insertions(+), 68 deletions(-) diff --git a/buffers_test.go b/buffers_test.go index 0c735b3..d1bd571 100644 --- a/buffers_test.go +++ b/buffers_test.go @@ -1,11 +1,16 @@ package extsort -// Tests for the buffers the sorter reuses: chunk slices. +// Tests for the buffers the sorter reuses: chunk slices, and the encode and decode +// buffers of the built-in codecs. import ( + "bytes" "cmp" "context" + "math/rand" "slices" + "strconv" + "strings" "testing" ) @@ -44,3 +49,98 @@ func TestChunksGrowWithInput(t *testing.T) { }) } } + +// Strings and Ordered encode every record into one reused buffer, and decode all the +// records of a chunk from one reused buffer. Records of varying length must survive both. +func TestBuiltinCodecsReuseBuffers(t *testing.T) { + r := rand.New(rand.NewSource(1)) + words := make([]string, 5000) + for i := range words { + // 0 to 300 bytes, so a record's buffer is sometimes longer and sometimes shorter than the last + words[i] = strings.Repeat(string(rune('a'+r.Intn(26))), r.Intn(300)) + strconv.Itoa(r.Intn(100)) + } + nums := make([]int64, 5000) + for i := range nums { + nums[i] = r.Int63n(1<<62) - 1<<61 // varints of 1 to 9 bytes + } + config := func() *Config { return &Config{ChunkSize: 100, NumWorkers: 3, TempFilesDir: t.TempDir()} } + + t.Run("Strings", func(t *testing.T) { + in := make(chan string, len(words)) + for _, w := range words { + in <- w + } + close(in) + s, out, errc := Strings(in, config()) + if s.appendBytes == nil || !s.reuseReadBuffer { + t.Fatal("Strings does not reuse its buffers") + } + s.Sort(context.Background()) + checkSorted(t, out, errc, words) + }) + t.Run("Ordered string", func(t *testing.T) { + in := make(chan string, len(words)) + for _, w := range words { + in <- w + } + close(in) + s, out, errc := Ordered(in, config()) + if s.appendBytes == nil || !s.reuseReadBuffer { + t.Fatal("Ordered does not reuse its buffers") + } + s.Sort(context.Background()) + checkSorted(t, out, errc, words) + }) + t.Run("Ordered int64", func(t *testing.T) { + in := make(chan int64, len(nums)) + for _, v := range nums { + in <- v + } + close(in) + s, out, errc := Ordered(in, config()) + s.Sort(context.Background()) + checkSorted(t, out, errc, nums) + }) +} + +// checkSorted drains out and errc and checks that out delivered input in order. +func checkSorted[E cmp.Ordered](t *testing.T, out <-chan E, errc <-chan error, input []E) { + t.Helper() + got, err := drainWithTimeout(t, out, errc) + if err != nil { + t.Fatal(err) + } + if want := slices.Sorted(slices.Values(input)); !slices.Equal(got, want) { + t.Errorf("got %d records (sorted: %v), want the %d input records in order", len(got), slices.IsSorted(got), len(want)) + } +} + +// A user's fromBytes may keep the slice it is given, so Generic must not decode records +// from a reused buffer: here every record would end up aliasing the last one read. +func TestGenericDoesNotReuseReadBuffer(t *testing.T) { + r := rand.New(rand.NewSource(1)) + records := make([][]byte, 2000) + for i := range records { + records[i] = []byte(strconv.Itoa(r.Int())) + } + in := make(chan []byte, len(records)) + for _, rec := range records { + in <- rec + } + close(in) + keep := func(d []byte) ([]byte, error) { return d, nil } // the record is the slice itself + s, out, errc := Generic(in, keep, keep, bytes.Compare, &Config{ChunkSize: 50, NumWorkers: 3, TempFilesDir: t.TempDir()}) + if s.appendBytes != nil || s.reuseReadBuffer { + t.Fatal("Generic reuses buffers for user codecs") + } + s.Sort(context.Background()) + got, err := drainWithTimeout(t, out, errc) + if err != nil { + t.Fatal(err) + } + want := slices.Clone(records) + slices.SortFunc(want, bytes.Compare) + if !slices.EqualFunc(got, want, bytes.Equal) { + t.Error("records changed on their way through the sort") + } +} diff --git a/sort_generic.go b/sort_generic.go index f237cb8..039cb70 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -61,9 +61,8 @@ func (s *GenericSorter[E]) putChunk(c *genericChunk[E]) { // memoryPools holds sync.Pool instances for memory reuse type memoryPools struct { - chunkPool sync.Pool // *chunk objects - slicePool sync.Pool // []E slices - scratchPool sync.Pool // scratch buffers for binary encoding + chunkPool sync.Pool // *chunk objects + slicePool sync.Pool // []E slices } // GenericSorter implements external sorting for any type E using a divide-and-conquer approach. @@ -87,6 +86,26 @@ type GenericSorter[E any] struct { toBytes ToBytesGeneric[E] pools *memoryPools singleChunk *genericChunk[E] // Holds the single chunk for optimization + // Set by useBuiltinCodec for the package's own codecs only + appendBytes appendBytesFunc[E] // encodes into a reused buffer instead of toBytes + reuseReadBuffer bool // fromBytes never keeps its input, so the merge reuses one buffer +} + +// appendBytesFunc appends the encoding of v to dst and returns the extended slice. +type appendBytesFunc[E any] func(dst []byte, v E) []byte + +// useBuiltinCodec makes the sorter reuse buffers with one of the package's own codecs: +// records are encoded with appendBytes into one reused buffer, and the merge decodes all the +// records of a chunk from one reused buffer. The second is only safe because those fromBytes +// functions never keep the slice they are given; a user's FromBytesGeneric may. +func (s *GenericSorter[E]) useBuiltinCodec(appendBytes appendBytesFunc[E]) { + s.appendBytes = appendBytes + s.reuseReadBuffer = true +} + +// toBytesFunc returns the ToBytesGeneric form of appendBytes. +func toBytesFunc[E any](appendBytes appendBytesFunc[E]) ToBytesGeneric[E] { + return func(v E) ([]byte, error) { return appendBytes(nil, v), nil } } // newSorter creates a new GenericSorter instance with the given configuration. @@ -110,7 +129,7 @@ func newSorter[E any](input <-chan E, fromBytes FromBytesGeneric[E], toBytes ToB } // initMemoryPools initializes sync.Pool instances for efficient memory reuse during sorting. -// Creates pools for chunks, slices, and scratch buffers to reduce GC pressure +// Creates pools for chunks and their slices to reduce GC pressure // and improve performance during high-frequency allocation/deallocation cycles. func (s *GenericSorter[E]) initMemoryPools() *memoryPools { pools := &memoryPools{} @@ -130,14 +149,6 @@ func (s *GenericSorter[E]) initMemoryPools() *memoryPools { }, } - // Pool for scratch buffers (for binary encoding) - store pointers to slices - pools.scratchPool = sync.Pool{ - New: func() any { - slice := make([]byte, binary.MaxVarintLen64) - return &slice - }, - } - return pools } @@ -498,39 +509,61 @@ func (s *GenericSorter[E]) saveChunksOptimized() error { } } -// saveChunk processes a single chunk +// saveChunk writes a sorted chunk as the next section of the temp file, each record as a +// uvarint length followed by its encoding, and returns the chunk to the pool. func (s *GenericSorter[E]) saveChunk(b *genericChunk[E]) error { - scratchPtr := s.pools.scratchPool.Get().(*[]byte) - scratch := *scratchPtr - defer s.pools.scratchPool.Put(scratchPtr) + defer s.putChunk(b) - for _, d := range b.data { - // binary encoding for size + var err error + if s.appendBytes != nil { + err = s.writeAppended(b.data) + } else { + err = s.writeEncoded(b.data) + } + if err != nil { + return err + } + if _, err := s.tempWriter.Next(); err != nil { + return NewDiskError(err, "next chunk", "") + } + return nil +} + +// writeEncoded writes records encoded by toBytes. +func (s *GenericSorter[E]) writeEncoded(records []E) error { + var header [binary.MaxVarintLen64]byte + for _, d := range records { raw, err := s.encode(d) if err != nil { - s.putChunk(b) // Return chunk to pool on error return err } - n := binary.PutUvarint(scratch, uint64(len(raw))) - _, err = s.tempWriter.Write(scratch[:n]) - if err != nil { - s.putChunk(b) // Return chunk to pool on error + n := binary.PutUvarint(header[:], uint64(len(raw))) + if _, err := s.tempWriter.Write(header[:n]); err != nil { return NewDiskError(err, "write size header", "") } - // add data - _, err = s.tempWriter.Write(raw) - if err != nil { - s.putChunk(b) // Return chunk to pool on error + if _, err := s.tempWriter.Write(raw); err != nil { return NewDiskError(err, "write data", "") } } - _, err := s.tempWriter.Next() - if err != nil { - s.putChunk(b) // Return chunk to pool on error - return NewDiskError(err, "next chunk", "") + return nil +} + +// writeAppended writes records encoded by appendBytes into one reused buffer. Each record +// is appended after room for the longest length header, and the header is then put just +// before it, so a record and its header go out in a single Write. +func (s *GenericSorter[E]) writeAppended(records []E) error { + const room = binary.MaxVarintLen64 + buf := make([]byte, room, 256) + var header [room]byte + for _, d := range records { + buf = s.appendBytes(buf[:room], d) + n := binary.PutUvarint(header[:], uint64(len(buf)-room)) + start := room - n + copy(buf[start:], header[:n]) + if _, err := s.tempWriter.Write(buf[start:]); err != nil { + return NewDiskError(err, "write data", "") + } } - // Successfully processed chunk, return to pool - s.putChunk(b) return nil } @@ -595,10 +628,7 @@ func (s *GenericSorter[E]) mergeNChunksSingleThreaded(ctx context.Context) (err }) for i := 0; i < s.tempReader.Size(); i++ { - merge := &mergeFile[E]{ - fromBytes: s.fromBytes, - reader: s.tempReader.Read(i), - } + merge := s.newMergeFile(i) _, ok, err := merge.getNext() // start the merge by preloading the values if err != nil { return err @@ -775,10 +805,7 @@ func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, en // Initialize merge files for this worker's chunk range for i := startChunk; i < endChunk; i++ { - merge := &mergeFile[E]{ - fromBytes: s.fromBytes, - reader: s.tempReader.Read(i), - } + merge := s.newMergeFile(i) _, ok, err := merge.getNext() if err != nil { return err @@ -903,11 +930,22 @@ func (c *channelMergeSource[E]) releaseBatch() { c.batch = nil } +// newMergeFile returns a mergeFile reading section i of the temp file. +func (s *GenericSorter[E]) newMergeFile(i int) *mergeFile[E] { + return &mergeFile[E]{ + fromBytes: s.fromBytes, + reader: s.tempReader.Read(i), + reuseBuf: s.reuseReadBuffer, + } +} + // mergeFile represents each sorted chunk on disk and its next value type mergeFile[E any] struct { nextRec E fromBytes FromBytesGeneric[E] reader *bufio.Reader + reuseBuf bool // fromBytes never keeps its input, so every record can be read into buf + buf []byte // holds the last record read when reuseBuf is set } // getNext returns the next value from the sorted chunk on disk. @@ -923,7 +961,15 @@ func (m *mergeFile[E]) getNext() (E, bool, error) { if err != nil { return old, false, err } - newRecBytes := make([]byte, int(n)) + var newRecBytes []byte + if m.reuseBuf { + if uint64(cap(m.buf)) < n { + m.buf = make([]byte, int(n)) + } + newRecBytes = m.buf[:n] + } else { + newRecBytes = make([]byte, int(n)) + } if _, err := io.ReadFull(m.reader, newRecBytes); err != nil { if err == io.EOF { // a length header without its payload is a truncated record, not the end of the chunk diff --git a/sort_ordered.go b/sort_ordered.go index fa2ec91..5aabf6b 100644 --- a/sort_ordered.go +++ b/sort_ordered.go @@ -23,10 +23,16 @@ var errOrderedDecode = errors.New("extsort: invalid encoding of an ordered value // named types such as `type ID int64` use the codec of their underlying type. Integers are // stored as varints, floats as their exact IEEE 754 bits (so NaN and -0 survive), and // strings as their bytes. +func orderedCodec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { + fromBytes, appendBytes := orderedAppendCodec[T]() + return fromBytes, toBytesFunc(appendBytes) +} + +// orderedAppendCodec returns the codec of orderedCodec, with the encoder in append form. // // The codecs read and write T through a pointer to its underlying type. The kind check // guarantees the two types share a memory layout, which makes the conversion valid. -func orderedCodec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { +func orderedAppendCodec[T cmp.Ordered]() (FromBytesGeneric[T], appendBytesFunc[T]) { switch kind := reflect.TypeFor[T]().Kind(); kind { case reflect.Int: return signedCodec[T, int]() @@ -62,7 +68,7 @@ func orderedCodec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { } // signedCodec stores a T whose underlying type is I as a zig-zag varint. -func signedCodec[T cmp.Ordered, I int | int8 | int16 | int32 | int64]() (FromBytesGeneric[T], ToBytesGeneric[T]) { +func signedCodec[T cmp.Ordered, I int | int8 | int16 | int32 | int64]() (FromBytesGeneric[T], appendBytesFunc[T]) { fromBytes := func(d []byte) (T, error) { var v T x, n := binary.Varint(d) @@ -72,14 +78,14 @@ func signedCodec[T cmp.Ordered, I int | int8 | int16 | int32 | int64]() (FromByt *(*I)(unsafe.Pointer(&v)) = I(x) return v, nil } - toBytes := func(v T) ([]byte, error) { - return binary.AppendVarint(nil, int64(*(*I)(unsafe.Pointer(&v)))), nil + appendBytes := func(dst []byte, v T) []byte { + return binary.AppendVarint(dst, int64(*(*I)(unsafe.Pointer(&v)))) } - return fromBytes, toBytes + return fromBytes, appendBytes } // unsignedCodec stores a T whose underlying type is U as a varint. -func unsignedCodec[T cmp.Ordered, U uint | uint8 | uint16 | uint32 | uint64 | uintptr]() (FromBytesGeneric[T], ToBytesGeneric[T]) { +func unsignedCodec[T cmp.Ordered, U uint | uint8 | uint16 | uint32 | uint64 | uintptr]() (FromBytesGeneric[T], appendBytesFunc[T]) { fromBytes := func(d []byte) (T, error) { var v T x, n := binary.Uvarint(d) @@ -89,14 +95,14 @@ func unsignedCodec[T cmp.Ordered, U uint | uint8 | uint16 | uint32 | uint64 | ui *(*U)(unsafe.Pointer(&v)) = U(x) return v, nil } - toBytes := func(v T) ([]byte, error) { - return binary.AppendUvarint(nil, uint64(*(*U)(unsafe.Pointer(&v)))), nil + appendBytes := func(dst []byte, v T) []byte { + return binary.AppendUvarint(dst, uint64(*(*U)(unsafe.Pointer(&v)))) } - return fromBytes, toBytes + return fromBytes, appendBytes } // float32Codec stores a T whose underlying type is float32 as its 4 IEEE 754 bytes. -func float32Codec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { +func float32Codec[T cmp.Ordered]() (FromBytesGeneric[T], appendBytesFunc[T]) { fromBytes := func(d []byte) (T, error) { var v T if len(d) != 4 { @@ -105,14 +111,14 @@ func float32Codec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { *(*float32)(unsafe.Pointer(&v)) = math.Float32frombits(binary.BigEndian.Uint32(d)) return v, nil } - toBytes := func(v T) ([]byte, error) { - return binary.BigEndian.AppendUint32(nil, math.Float32bits(*(*float32)(unsafe.Pointer(&v)))), nil + appendBytes := func(dst []byte, v T) []byte { + return binary.BigEndian.AppendUint32(dst, math.Float32bits(*(*float32)(unsafe.Pointer(&v)))) } - return fromBytes, toBytes + return fromBytes, appendBytes } // float64Codec stores a T whose underlying type is float64 as its 8 IEEE 754 bytes. -func float64Codec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { +func float64Codec[T cmp.Ordered]() (FromBytesGeneric[T], appendBytesFunc[T]) { fromBytes := func(d []byte) (T, error) { var v T if len(d) != 8 { @@ -121,23 +127,23 @@ func float64Codec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { *(*float64)(unsafe.Pointer(&v)) = math.Float64frombits(binary.BigEndian.Uint64(d)) return v, nil } - toBytes := func(v T) ([]byte, error) { - return binary.BigEndian.AppendUint64(nil, math.Float64bits(*(*float64)(unsafe.Pointer(&v)))), nil + appendBytes := func(dst []byte, v T) []byte { + return binary.BigEndian.AppendUint64(dst, math.Float64bits(*(*float64)(unsafe.Pointer(&v)))) } - return fromBytes, toBytes + return fromBytes, appendBytes } // stringCodec stores a T whose underlying type is string as its bytes. -func stringCodec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { +func stringCodec[T cmp.Ordered]() (FromBytesGeneric[T], appendBytesFunc[T]) { fromBytes := func(d []byte) (T, error) { var v T *(*string)(unsafe.Pointer(&v)) = string(d) return v, nil } - toBytes := func(v T) ([]byte, error) { - return []byte(*(*string)(unsafe.Pointer(&v))), nil + appendBytes := func(dst []byte, v T) []byte { + return append(dst, *(*string)(unsafe.Pointer(&v))...) } - return fromBytes, toBytes + return fromBytes, appendBytes } // Ordered performs external sorting on a channel of cmp.Ordered types. @@ -148,8 +154,9 @@ func stringCodec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { // IMPORTANT: The input channel MUST be closed to signal the end of data. // Sort() will continue reading from the input channel until it is closed. func Ordered[T cmp.Ordered](input <-chan T, config *Config) (*OrderedSorter[T], <-chan T, <-chan error) { - fromBytes, toBytes := orderedCodec[T]() - s, output, errChan := Generic(input, fromBytes, toBytes, cmp.Compare, config) + fromBytes, appendBytes := orderedAppendCodec[T]() + s, output, errChan := Generic(input, fromBytes, toBytesFunc(appendBytes), cmp.Compare, config) + s.useBuiltinCodec(appendBytes) return &OrderedSorter[T]{GenericSorter: *s}, output, errChan } @@ -157,7 +164,8 @@ func Ordered[T cmp.Ordered](input <-chan T, config *Config) (*OrderedSorter[T], // the number of items to sort (useful for testing). Takes the same parameters as // Ordered plus n which limits the number of items processed. func OrderedMock[T cmp.Ordered](input <-chan T, config *Config, n int) (*OrderedSorter[T], <-chan T, <-chan error) { - fromBytes, toBytes := orderedCodec[T]() - s, output, errChan := MockGeneric(input, fromBytes, toBytes, cmp.Compare, config, n) + fromBytes, appendBytes := orderedAppendCodec[T]() + s, output, errChan := MockGeneric(input, fromBytes, toBytesFunc(appendBytes), cmp.Compare, config, n) + s.useBuiltinCodec(appendBytes) return &OrderedSorter[T]{GenericSorter: *s}, output, errChan } diff --git a/sort_strings.go b/sort_strings.go index b40feca..c55d192 100644 --- a/sort_strings.go +++ b/sort_strings.go @@ -23,6 +23,11 @@ func toBytesString(s string) ([]byte, error) { return []byte(s), nil } +// appendString is the append form of toBytesString. +func appendString(dst []byte, s string) []byte { + return append(dst, s...) +} + // Strings performs external sorting on a channel of strings using lexicographic ordering. // Returns the sorter instance, output channel with sorted strings, and error channel. // This function provides backward compatibility with the legacy string-specific API. @@ -31,6 +36,7 @@ func toBytesString(s string) ([]byte, error) { // Sort() will continue reading from the input channel until it is closed. func Strings(input <-chan string, config *Config) (*StringSorter, <-chan string, <-chan error) { genericSorter, output, errChan := Generic(input, fromBytesString, toBytesString, cmp.Compare, config) + genericSorter.useBuiltinCodec(appendString) s := &StringSorter{GenericSorter: *genericSorter} return s, output, errChan } @@ -40,6 +46,7 @@ func Strings(input <-chan string, config *Config) (*StringSorter, <-chan string, // The parameter n specifies the maximum number of strings to process. func StringsMock(input <-chan string, config *Config, n int) (*StringSorter, <-chan string, <-chan error) { genericSorter, output, errChan := MockGeneric(input, fromBytesString, toBytesString, cmp.Compare, config, n) + genericSorter.useBuiltinCodec(appendString) s := &StringSorter{GenericSorter: *genericSorter} return s, output, errChan }