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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
146 changes: 146 additions & 0 deletions buffers_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
}
77 changes: 45 additions & 32 deletions diff/diff_generic.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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.
Expand Down Expand Up @@ -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)
Expand All @@ -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++
Expand All @@ -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
}
}
}
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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 {
Expand Down
79 changes: 79 additions & 0 deletions diff/regression_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
}
Loading
Loading