diff --git a/buffers_test.go b/buffers_test.go new file mode 100644 index 0000000..d1bd571 --- /dev/null +++ b/buffers_test.go @@ -0,0 +1,146 @@ +package extsort + +// 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" +) + +// 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) + } + }) + } +} + +// 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/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) + } + } +} 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/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) + } +} 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 87b3422..039cb70 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -15,6 +15,19 @@ import ( "golang.org/x/sync/errgroup" ) +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 + // 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 { @@ -48,10 +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 // []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 // []E slices } // GenericSorter implements external sorting for any type E using a divide-and-conquer approach. @@ -75,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. @@ -98,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, byte 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{} @@ -110,27 +141,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 - }, - } - - // 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 { - slice := make([]byte, binary.MaxVarintLen64) - return &slice + return new([]E) }, } @@ -269,21 +284,42 @@ 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: 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 + } + 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 @@ -303,6 +339,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 { @@ -460,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) + + 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 +} - for _, d := range b.data { - // binary encoding for size +// 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 } @@ -557,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 @@ -601,11 +669,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 @@ -630,15 +697,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 @@ -661,7 +728,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 } @@ -685,8 +752,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 { @@ -700,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 @@ -714,15 +816,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 { @@ -734,19 +830,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 { @@ -759,8 +859,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) } @@ -792,23 +892,51 @@ 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 + 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 +} + +// 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, } - c.hasNext = false - return false } // mergeFile represents each sorted chunk on disk and its next value @@ -816,6 +944,8 @@ 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. @@ -831,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 }