diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 9998da0..14f37a2 100644 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -9,8 +9,12 @@ on: jobs: build: - name: Build - runs-on: ubuntu-latest + name: Build (${{ matrix.os }}) + strategy: + fail-fast: false + matrix: + os: [ ubuntu-latest, macos-latest, windows-latest ] + runs-on: ${{ matrix.os }} steps: - name: Check out code into the Go module directory uses: actions/checkout@v7 @@ -24,10 +28,24 @@ jobs: run: go mod download - name: golangci-lint + if: runner.os == 'Linux' uses: golangci/golangci-lint-action@v9 with: version: v2.12 - name: Test + if: runner.os != 'Windows' run: make test + # make is not reliably available on Windows runners + - name: Test (Windows) + if: runner.os == 'Windows' + run: go test -timeout=5m ./... + + - name: Test with race detector + if: runner.os == 'Linux' + run: make test-race + + - name: Run examples + if: runner.os == 'Linux' + run: make examples diff --git a/Makefile b/Makefile index ba8e9ee..f632193 100644 --- a/Makefile +++ b/Makefile @@ -6,12 +6,17 @@ include release.mk ALL_SOURCES := $(shell find . -type f -name '*.go') -.PHONY: fmt lint test cover coverhtml examples readme +.PHONY: fmt lint test test-race cover coverhtml examples readme test: go test -timeout=60s $(shell go list ./... | grep -v "/examples") @echo "< ALL TESTS PASS >" +# the race detector slows the root package to about a minute, so allow longer than test +test-race: + go test -race -timeout=10m $(shell go list ./... | grep -v "/examples") + @echo "< ALL RACE TESTS PASS >" + update-deps: go.mod GOPROXY=direct go get -u ./... go mod tidy @@ -39,9 +44,9 @@ benchmark: examples: @for dir in examples/*/; do \ - if [ -f "$$dir"*.go ]; then \ + if ls "$$dir"*.go > /dev/null 2>&1; then \ echo "Running example in $$dir"; \ - (cd "$$dir" && go run *.go > /dev/null); \ + (cd "$$dir" && go run . > /dev/null) || exit 1; \ fi; \ done diff --git a/config.go b/config.go index d12d580..7976c82 100644 --- a/config.go +++ b/config.go @@ -1,26 +1,29 @@ package extsort // Config holds configuration settings for external sorting operations. -// All fields have sensible defaults and can be left as zero values to use defaults. +// Pass a nil *Config to use DefaultConfig(). In a non-nil Config, a ChunkSize or +// NumWorkers below 1 and a negative ChanBuffSize or SortedChanBuffSize are replaced by +// their defaults, but a zero ChanBuffSize or SortedChanBuffSize means an unbuffered channel. +// The sorter works on its own copy, so one Config can be shared by several sorters. type Config struct { // ChunkSize specifies the maximum number of records to store in each chunk // before writing to disk. Larger chunks use more memory but reduce I/O operations. - // Default: 1,000,000 records. Must be > 1. + // Default: 1,000,000 records. Values below 1 use the default. ChunkSize int // NumWorkers controls the maximum number of goroutines used for parallel // chunk sorting and merging. More workers can improve CPU utilization on multi-core systems. - // Default: 2 workers. Must be > 1. + // Default: 2 workers. Values below 1 use the default. NumWorkers int - // ChanBuffSize sets the buffer size for internal channels used during chunk merging. - // Larger buffers can improve throughput but use more memory. - // Default: 1. Must be >= 0. + // ChanBuffSize sets how many whole chunks can wait between reading the input and + // sorting them. Each buffered chunk holds up to ChunkSize records in memory. + // Default: 1. Zero means unbuffered; negative values use the default. ChanBuffSize int // SortedChanBuffSize sets the buffer size for the output channel that delivers // sorted results. Larger buffers allow more decoupling between sorting and consumption. - // Default: 1000. Must be >= 0. + // Default: 1000. Zero means unbuffered; negative values use the default. SortedChanBuffSize int // TempFilesDir specifies the directory for temporary files during sorting. @@ -45,31 +48,32 @@ func DefaultConfig() *Config { return &Config{ ChunkSize: int(1e6), // 1M NumWorkers: 2, - ChanBuffSize: 16, + ChanBuffSize: 1, SortedChanBuffSize: 1000, TempFilesDir: "", } } -// mergeConfig validates and normalizes a Config by replacing zero/invalid values -// with defaults. If config is nil, returns DefaultConfig(). -// This ensures all sorter instances have valid configuration values. +// mergeConfig returns a validated and normalized copy of c, replacing invalid values +// with defaults. If c is nil, returns DefaultConfig(). +// The caller's Config is never modified, since it may be shared by other sorters. func mergeConfig(c *Config) *Config { d := DefaultConfig() if c == nil { return d } - if c.ChunkSize < 1 { - c.ChunkSize = d.ChunkSize + merged := *c + if merged.ChunkSize < 1 { + merged.ChunkSize = d.ChunkSize } - if c.NumWorkers < 1 { - c.NumWorkers = d.NumWorkers + if merged.NumWorkers < 1 { + merged.NumWorkers = d.NumWorkers } - if c.ChanBuffSize < 0 { - c.ChanBuffSize = d.ChanBuffSize + if merged.ChanBuffSize < 0 { + merged.ChanBuffSize = d.ChanBuffSize } - if c.SortedChanBuffSize < 0 { - c.SortedChanBuffSize = d.SortedChanBuffSize + if merged.SortedChanBuffSize < 0 { + merged.SortedChanBuffSize = d.SortedChanBuffSize } - return c + return &merged } diff --git a/diff/diff_generic.go b/diff/diff_generic.go index 171761d..4c63541 100644 --- a/diff/diff_generic.go +++ b/diff/diff_generic.go @@ -107,14 +107,16 @@ func (d *differ[T]) diff() (r Result, err error) { } } } - // check for errors just in case + // check for errors just in case. Each error channel is read once: here if its + // stream has ended, otherwise after the stream is drained below. + aErrPending, bErrPending := okA, okB if !okA { - if err = <-d.aErrChan; err != nil { + if err = d.readErr(d.aErrChan); err != nil { return } } if !okB { - if err = <-d.bErrChan; err != nil { + if err = d.readErr(d.bErrChan); err != nil { return } } @@ -132,9 +134,11 @@ func (d *differ[T]) diff() (r Result, err error) { return r, d.ctx.Err() } } - // check for A errors once again - if err = <-d.aErrChan; err != nil { - return + // check for A errors if not read above + if aErrPending { + if err = d.readErr(d.aErrChan); err != nil { + return + } } // if only B has data left for okB { @@ -150,13 +154,26 @@ func (d *differ[T]) diff() (r Result, err error) { return r, d.ctx.Err() } } - // check for B errors once again - if err = <-d.bErrChan; err != nil { - return + // check for B errors if not read above + if bErrPending { + if err = d.readErr(d.bErrChan); err != nil { + return + } } return } +// 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 { + select { + case err := <-errChan: + return err + case <-d.ctx.Done(): + return d.ctx.Err() + } +} + // PrintDiff is a utility function that can be used as a ResultFunc to print // differences to stdout. It formats each difference with the Delta symbol // (< for OLD, > for NEW) followed by the item value. diff --git a/diff/diff_result_chan.go b/diff/diff_result_chan.go index 8994236..1cbe8b8 100644 --- a/diff/diff_result_chan.go +++ b/diff/diff_result_chan.go @@ -1,5 +1,7 @@ package diff +import "context" + // StringChanResult holds a single diff result from a string comparison. // It contains both the difference type (NEW/OLD) and the actual string value. // This type is used with StringResultChan to enable parallel processing of diff results. @@ -20,11 +22,24 @@ type StringChanResult struct { // - chan *StringChanResult: A channel to receive diff results from // // The caller is responsible for closing the returned channel when done. +// The returned function blocks until each result is received; use StringResultChanContext +// to stop waiting when a context is done. func StringResultChan() (StringResultFunc, chan *StringChanResult) { + return StringResultChanContext(context.Background()) +} + +// StringResultChanContext is like StringResultChan, but the returned function stops waiting +// for the receiver once ctx is done and returns ctx.Err(), which ends the diff with that error. +// Pass the same context to the diff. +func StringResultChanContext(ctx context.Context) (StringResultFunc, chan *StringChanResult) { c := make(chan *StringChanResult, 1) f := func(d Delta, s string) error { - c <- &StringChanResult{D: d, S: s} - return nil + select { + case c <- &StringChanResult{D: d, S: s}: + return nil + case <-ctx.Done(): + return ctx.Err() + } } return f, c } diff --git a/diff/diff_types.go b/diff/diff_types.go index f7e778b..1d88071 100644 --- a/diff/diff_types.go +++ b/diff/diff_types.go @@ -40,7 +40,7 @@ type Delta int const ( // NEW indicates an item that exists only in the second stream (B). // This represents a "new" or "added" item when comparing A to B. - NEW = iota // + + NEW Delta = iota // + // OLD indicates an item that exists only in the first stream (A). // This represents an "old" or "removed" item when comparing A to B. diff --git a/diff/regression_test.go b/diff/regression_test.go new file mode 100644 index 0000000..27d7748 --- /dev/null +++ b/diff/regression_test.go @@ -0,0 +1,108 @@ +package diff_test + +// Regression tests for diff hangs and the Delta constants. + +import ( + "context" + "errors" + "fmt" + "testing" + "time" + + "github.com/lanrat/extsort/diff" +) + +// stream returns a closed channel holding items. +func stream(items ...string) chan string { + ch := make(chan string, len(items)) + for _, s := range items { + ch <- s + } + close(ch) + return ch +} + +func ignoreResult(diff.Delta, string) error { return nil } + +// runDiff runs diff.Strings and fails the test if it does not return in time. +func runDiff(t *testing.T, ctx context.Context, a, b <-chan string, aErr, bErr <-chan error, f diff.StringResultFunc) (diff.Result, error) { + t.Helper() + type result struct { + r diff.Result + err error + } + done := make(chan result, 1) + go func() { + r, err := diff.Strings(ctx, a, b, aErr, bErr, f) + done <- result{r, err} + }() + select { + case res := <-done: + return res.r, res.err + case <-time.After(5 * time.Second): + t.Fatal("diff did not return within 5s") + return diff.Result{}, nil + } +} + +// The error channel of the stream that ended first used to be read twice, so a +// caller that sent one value without closing the channel hung the diff. +func TestErrChanOfShorterStreamIsReadOnce(t *testing.T) { + for _, tc := range []struct { + name string + a, b []string + }{ + {"A ends first", []string{"a"}, []string{"a", "b", "c"}}, + {"B ends first", []string{"a", "b", "c"}, []string{"a"}}, + } { + t.Run(tc.name, func(t *testing.T) { + aErr, bErr := make(chan error, 1), make(chan error, 1) + aErr <- nil // one value each, never closed + bErr <- nil + r, err := runDiff(t, context.Background(), stream(tc.a...), stream(tc.b...), aErr, bErr, ignoreResult) + if err != nil { + t.Fatal(err) + } + if r.Common != 1 || r.ExtraA+r.ExtraB != 2 { + t.Errorf("unexpected result %s", r.String()) + } + }) + } +} + +// Reading an error channel ignored ctx, so a channel that is never sent to or closed +// hung the diff even past the ctx deadline. +func TestErrChanReadRespectsContext(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + never := make(chan error) + _, err := runDiff(t, ctx, stream("a"), stream("b"), never, never, ignoreResult) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("got error %v, want %v", err, context.DeadlineExceeded) + } +} + +// NEW and OLD used to be untyped integer constants, so they printed as 0 and 1 +// instead of using Delta's String method. +func TestDeltaConstantsAreTyped(t *testing.T) { + if got := fmt.Sprint(diff.NEW, diff.OLD); got != "> <" { + t.Errorf("fmt.Sprint(NEW, OLD) = %q, want %q", got, "> <") + } +} + +// The function returned by StringResultChan blocks on its send even after ctx is done, +// so a diff whose results are no longer read hangs. StringResultChanContext stops waiting. +func TestStringResultChanContextStopsWaiting(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + resultFunc, results := diff.StringResultChanContext(ctx) + defer close(results) + noErr := make(chan error) + close(noErr) + time.AfterFunc(50*time.Millisecond, cancel) + // three differences, and nobody reads the results channel + _, err := runDiff(t, ctx, stream("a1", "a2", "a3"), stream(), noErr, noErr, resultFunc) + if !errors.Is(err, context.Canceled) { + t.Fatalf("got error %v, want %v", err, context.Canceled) + } +} diff --git a/error_scenarios_test.go b/error_scenarios_test.go index 4072f62..86a9bc4 100644 --- a/error_scenarios_test.go +++ b/error_scenarios_test.go @@ -153,7 +153,9 @@ func TestLargeDataElements(t *testing.T) { inputChan <- &largeVal{Key: 2, Data: largeString} close(inputChan) - sort, outChan, errChan := extsort.New(inputChan, fromBytesForLargeVal, largeLessThan, nil) + // One element per chunk, so the elements are written to disk and read back + config := &extsort.Config{ChunkSize: 1} + sort, outChan, errChan := extsort.New(inputChan, fromBytesForLargeVal, largeLessThan, config) sort.Sort(context.Background()) var results []*largeVal @@ -304,7 +306,9 @@ func TestMixedTypeComparison(t *testing.T) { panic("unknown type in deserialization") } - sort, outChan, errChan := extsort.New(inputChan, mixedFromBytes, mixedLessFunc, nil) + // One element per chunk, so both types are written to disk and read back with mixedFromBytes + config := &extsort.Config{ChunkSize: 1} + sort, outChan, errChan := extsort.New(inputChan, mixedFromBytes, mixedLessFunc, config) sort.Sort(context.Background()) var results []extsort.SortType diff --git a/generic_error_test.go b/generic_error_test.go index 4ef4b67..dcae7b9 100644 --- a/generic_error_test.go +++ b/generic_error_test.go @@ -106,16 +106,17 @@ func TestGenericDeserializationError(t *testing.T) { t.Logf("Successfully caught generic deserialization error: %v", err) } -// TestOrderedSerializationError tests gob encoding errors in Ordered API +// TestOrderedSerializationError tests the error path of the Ordered API func TestOrderedSerializationError(t *testing.T) { - // Note: It's hard to make gob encoding fail for basic types, - // so this test demonstrates the error handling path exists + // Note: the binary codec cannot fail to encode basic types, + // so this test checks that a sort through the codec reports no error inputChan := make(chan int, 2) inputChan <- 1 inputChan <- 2 close(inputChan) - sort, outChan, errChan := extsort.Ordered(inputChan, nil) + // One record per chunk, so the records go through the codec and disk + sort, outChan, errChan := extsort.Ordered(inputChan, &extsort.Config{ChunkSize: 1}) sort.Sort(context.Background()) // Should complete successfully since int is easily serializable diff --git a/ordered_codec_test.go b/ordered_codec_test.go new file mode 100644 index 0000000..98a6037 --- /dev/null +++ b/ordered_codec_test.go @@ -0,0 +1,150 @@ +package extsort + +import ( + "cmp" + "context" + "encoding/binary" + "math" + "slices" + "strings" + "testing" +) + +type ( + namedInt int + namedInt8 int8 + namedUint16 uint16 + namedFloat float64 + namedString string +) + +// roundTrip encodes and decodes each value with the Ordered codec for T and checks +// that the value and its exact encoding (including float bits) survive. +func roundTrip[T cmp.Ordered](t *testing.T, values ...T) { + t.Helper() + fromBytes, toBytes := orderedCodec[T]() + for _, v := range values { + d, err := toBytes(v) + if err != nil { + t.Fatalf("toBytes(%v): %v", v, err) + } + got, err := fromBytes(d) + if err != nil { + t.Fatalf("fromBytes(toBytes(%v)): %v", v, err) + } + if cmp.Compare(got, v) != 0 { + t.Errorf("round trip of %v (%T) gave %v", v, v, got) + } + if again, _ := toBytes(got); string(again) != string(d) { + t.Errorf("round trip of %v (%T) changed its encoding from %x to %x", v, v, d, again) + } + } +} + +func TestOrderedCodecRoundTrip(t *testing.T) { + negZero := math.Copysign(0, -1) + t.Run("int", func(t *testing.T) { roundTrip(t, 0, 1, -1, math.MaxInt, math.MinInt) }) + t.Run("int8", func(t *testing.T) { roundTrip[int8](t, 0, -1, math.MaxInt8, math.MinInt8) }) + t.Run("int16", func(t *testing.T) { roundTrip[int16](t, 0, math.MaxInt16, math.MinInt16) }) + t.Run("int32", func(t *testing.T) { roundTrip[int32](t, 0, math.MaxInt32, math.MinInt32) }) + t.Run("int64", func(t *testing.T) { roundTrip[int64](t, 0, math.MaxInt64, math.MinInt64) }) + t.Run("uint", func(t *testing.T) { roundTrip[uint](t, 0, 1, math.MaxUint) }) + t.Run("uint8", func(t *testing.T) { roundTrip[uint8](t, 0, math.MaxUint8) }) + t.Run("uint16", func(t *testing.T) { roundTrip[uint16](t, 0, math.MaxUint16) }) + t.Run("uint32", func(t *testing.T) { roundTrip[uint32](t, 0, math.MaxUint32) }) + t.Run("uint64", func(t *testing.T) { roundTrip[uint64](t, 0, math.MaxUint64) }) + t.Run("uintptr", func(t *testing.T) { roundTrip[uintptr](t, 0, ^uintptr(0)) }) + t.Run("float32", func(t *testing.T) { + roundTrip[float32](t, 0, float32(negZero), 1.5, -2.25, math.MaxFloat32, math.SmallestNonzeroFloat32, + float32(math.Inf(1)), float32(math.Inf(-1)), float32(math.NaN())) + }) + t.Run("float64", func(t *testing.T) { + roundTrip(t, 0, negZero, 1.5, -2.25, math.MaxFloat64, math.SmallestNonzeroFloat64, + math.Inf(1), math.Inf(-1), math.NaN()) + }) + t.Run("string", func(t *testing.T) { + roundTrip(t, "", "a", "héllo, 世界", "\x00\xff not UTF-8", strings.Repeat("x", 1<<16)) + }) + t.Run("named types", func(t *testing.T) { + roundTrip[namedInt](t, -42, math.MaxInt, math.MinInt) + roundTrip[namedInt8](t, math.MinInt8, math.MaxInt8) + roundTrip[namedUint16](t, 0, math.MaxUint16) + roundTrip(t, namedFloat(negZero), namedFloat(math.NaN()), namedFloat(math.Inf(-1))) + roundTrip[namedString](t, "", "named") + }) +} + +// decodeErr adapts a FromBytesGeneric to return only its error. +func decodeErr[T any](fromBytes FromBytesGeneric[T]) func([]byte) error { + return func(d []byte) error { + _, err := fromBytes(d) + return err + } +} + +func TestOrderedCodecRejectsInvalidData(t *testing.T) { + intFromBytes, _ := orderedCodec[int]() + int8FromBytes, _ := orderedCodec[int8]() + uint8FromBytes, _ := orderedCodec[uint8]() + float32FromBytes, _ := orderedCodec[float32]() + float64FromBytes, _ := orderedCodec[float64]() + for _, tc := range []struct { + name string + decode func([]byte) error + data []byte + }{ + {"int: empty", decodeErr(intFromBytes), nil}, + {"int: truncated varint", decodeErr(intFromBytes), []byte{0x80}}, + {"int: trailing bytes", decodeErr(intFromBytes), []byte{0x02, 0x00}}, + {"int8: out of range", decodeErr(int8FromBytes), binary.AppendVarint(nil, 300)}, + {"uint8: out of range", decodeErr(uint8FromBytes), binary.AppendUvarint(nil, 256)}, + {"float32: wrong length", decodeErr(float32FromBytes), []byte{1, 2, 3}}, + {"float64: wrong length", decodeErr(float64FromBytes), []byte{1, 2, 3, 4}}, + } { + if err := tc.decode(tc.data); err == nil { + t.Errorf("%s: decoding %x succeeded, want an error", tc.name, tc.data) + } + } +} + +// checkOrderedSort sorts values with Ordered through temp files on disk and compares +// the output with slices.SortFunc and cmp.Compare. +func checkOrderedSort[T cmp.Ordered](t *testing.T, values []T) { + t.Helper() + in := make(chan T, len(values)) + for _, v := range values { + in <- v + } + close(in) + sorter, out, errc := Ordered(in, &Config{ChunkSize: 2, TempFilesDir: t.TempDir()}) + sorter.Sort(context.Background()) + got, err := drainWithTimeout(t, out, errc) + if err != nil { + t.Fatal(err) + } + want := slices.Clone(values) + slices.SortFunc(want, cmp.Compare[T]) + if slices.CompareFunc(got, want, cmp.Compare[T]) != 0 { + t.Errorf("got %v, want %v", got, want) + } +} + +// Ordered used gob, which handled these types but allocated a new encoder and decoder +// per record. The binary codec must sort them the same way. +func TestOrderedSortsThroughDisk(t *testing.T) { + t.Run("float64 with NaN and -0", func(t *testing.T) { + checkOrderedSort(t, []float64{3, math.NaN(), -1, math.Inf(1), math.Copysign(0, -1), 0, math.Inf(-1), math.NaN(), 2.5}) + }) + t.Run("float32", func(t *testing.T) { + checkOrderedSort(t, []float32{2.5, -1, float32(math.NaN()), 0, 1e30}) + }) + t.Run("named int", func(t *testing.T) { + checkOrderedSort(t, []namedInt{5, -3, math.MaxInt, math.MinInt, 0, 5}) + }) + t.Run("uint8", func(t *testing.T) { + checkOrderedSort(t, []uint8{255, 0, 7, 128, 7}) + }) + t.Run("named string", func(t *testing.T) { + checkOrderedSort(t, []namedString{"pear", "", "apple", "fig", "apple", "\xff"}) + }) +} diff --git a/regression_test.go b/regression_test.go index d791b23..7c8affa 100644 --- a/regression_test.go +++ b/regression_test.go @@ -8,8 +8,10 @@ import ( "bytes" "cmp" "context" + "encoding/binary" "errors" "io" + "math/rand" "os" "path/filepath" "runtime" @@ -17,6 +19,7 @@ import ( "strconv" "strings" "sync" + "sync/atomic" "testing" "testing/iotest" "time" @@ -89,9 +92,10 @@ type trackedTemp struct { readErr map[int]error // section index -> error returned when reading it closeErr error // returned by the reader's Close - mu sync.Mutex - created int - closed int + mu sync.Mutex + created int + closed int + sections int // Size of the last reader returned by Save } func (tt *trackedTemp) newWriter() (tempfile.TempWriter, error) { @@ -135,6 +139,9 @@ func (w *trackedWriter) Save() (tempfile.TempReader, error) { if err != nil { return nil, err } + w.tt.mu.Lock() + w.tt.sections = r.Size() + w.tt.mu.Unlock() return &trackedReader{TempReader: r, tt: w.tt}, nil } @@ -243,16 +250,16 @@ func TestMergeReportsChunkReadErrors(t *testing.T) { workers int section int }{ - // 10 records in chunks of 5: two chunks plus the empty section Save appends. - // NumWorkers 4 >= 3 sections merges single-threaded; NumWorkers 2 merges in parallel. - {"single-threaded first chunk", 4, 0}, - {"single-threaded second chunk", 4, 1}, + // 15 records in chunks of 5 make three chunks. NumWorkers 3 merges them + // single-threaded; NumWorkers 2 merges them in parallel. + {"single-threaded first chunk", 3, 0}, + {"single-threaded second chunk", 3, 1}, {"parallel first chunk", 2, 0}, {"parallel second chunk", 2, 1}, } { t.Run(tc.name, func(t *testing.T) { tt := &trackedTemp{readErr: map[int]error{tc.section: errDisk}} - s := newTrackedSorter(descendingInts(10), itoaBytes, cmp.Compare[int], &Config{ChunkSize: 5, NumWorkers: tc.workers}, tt) + s := newTrackedSorter(descendingInts(15), itoaBytes, cmp.Compare[int], &Config{ChunkSize: 5, NumWorkers: tc.workers}, tt) got, err := runSort(t, context.Background(), s) if !errors.Is(err, errDisk) { t.Fatalf("got %d records and error %v, want error %q", len(got), err, errDisk) @@ -385,16 +392,16 @@ func TestCallbackPanicsBecomeErrors(t *testing.T) { // the deferred Close ran after the deferred close(mergeErrChan). func TestTempFileCloseErrorIsReported(t *testing.T) { errClose := errors.New("close failed") - for _, workers := range []int{4, 2} { // single-threaded and parallel merge + for _, workers := range []int{3, 2} { // three chunks: single-threaded and parallel merge t.Run("NumWorkers "+strconv.Itoa(workers), func(t *testing.T) { tt := &trackedTemp{closeErr: errClose} - s := newTrackedSorter(descendingInts(10), itoaBytes, cmp.Compare[int], &Config{ChunkSize: 5, NumWorkers: workers}, tt) + s := newTrackedSorter(descendingInts(15), itoaBytes, cmp.Compare[int], &Config{ChunkSize: 5, NumWorkers: workers}, tt) got, err := runSort(t, context.Background(), s) if !errors.Is(err, errClose) { t.Fatalf("got error %v, want %q", err, errClose) } - if len(got) != 10 || !slices.IsSorted(got) { - t.Errorf("got %v, want 1..10 in order", got) + if len(got) != 15 || !slices.IsSorted(got) { + t.Errorf("got %v, want 1..15 in order", got) } }) } @@ -628,3 +635,237 @@ func TestTempFileCreationErrorIsReportedOnErrChan(t *testing.T) { } }) } + +// buildChunks used to keep looping up to ChunkSize times on the closed input, +// because a plain break inside the select only left the select. +func TestBuildChunksStopsWhenInputCloses(t *testing.T) { + in := make(chan struct{}, 1) + in <- struct{}{} + close(in) + // Chunks of struct{} need no memory, so a huge ChunkSize only costs loop iterations + s := newSorter(in, + func([]byte) (struct{}, error) { return struct{}{}, nil }, + func(struct{}) ([]byte, error) { return nil, nil }, + func(a, b struct{}) int { return 0 }, + &Config{ChunkSize: 1 << 30, ChanBuffSize: 1}) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() // stops a spinning buildChunks once the test is over + s.sortCtx = ctx + + done := make(chan error, 1) + go func() { done <- s.buildChunks() }() + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("buildChunks still running 1s after the input closed") + } + if c := <-s.chunkChan; c == nil || len(c.data) != 1 { + t.Fatalf("got chunk %v, want one chunk with the record", c) + } + if _, ok := <-s.chunkChan; ok { + t.Error("buildChunks sent more than one chunk") + } +} + +// BenchmarkSortTenRecords measures a small sort with the default config. The loop +// on the closed input used to add about 2x ChunkSize (1M) iterations to every sort. +func BenchmarkSortTenRecords(b *testing.B) { + for b.Loop() { + in := make(chan int, 10) + for i := 10; i > 0; i-- { + in <- i + } + close(in) + s, out, errc := Generic(in, atoiBytes, itoaBytes, cmp.Compare[int], nil) + s.Sort(context.Background()) + for range out { + } + if err := <-errc; err != nil { + b.Fatal(err) + } + } +} + +// mergeConfig used to write defaults into the caller's Config, which raced when one +// Config was shared by several sorters (run with -race) and surprised callers. +func TestConfigIsNotModified(t *testing.T) { + shared := &Config{} // every zero field used to be overwritten + var wg sync.WaitGroup + for i := 0; i < 4; i++ { + wg.Add(1) + go func() { + defer wg.Done() + s, out, errc := MockGeneric(descendingInts(3), atoiBytes, itoaBytes, cmp.Compare[int], shared, 0) + s.Sort(context.Background()) + if _, err := drainWithTimeout(t, out, errc); err != nil { + t.Error(err) + } + }() + } + wg.Wait() + if *shared != (Config{}) { + t.Errorf("caller's Config was modified: %+v", *shared) + } +} + +// Zero buffer sizes mean unbuffered channels; nil and negative values mean the defaults. +func TestConfigBufferSizes(t *testing.T) { + for _, tc := range []struct { + name string + config *Config + wantChunk, wantOutput int + }{ + {"nil config", nil, 1, 1000}, + {"zero values", &Config{}, 0, 0}, + {"negative values", &Config{ChanBuffSize: -1, SortedChanBuffSize: -1}, 1, 1000}, + {"explicit values", &Config{ChanBuffSize: 3, SortedChanBuffSize: 7}, 3, 7}, + } { + t.Run(tc.name, func(t *testing.T) { + s := newSorter[int](nil, atoiBytes, itoaBytes, cmp.Compare[int], tc.config) + if got := cap(s.chunkChan); got != tc.wantChunk { + t.Errorf("chunk channel buffer = %d, want %d", got, tc.wantChunk) + } + if got := cap(s.mergeChunkChan); got != tc.wantOutput { + t.Errorf("output channel buffer = %d, want %d", got, tc.wantOutput) + } + }) + } +} + +// The parallel merge (v1.1.0) calls FromBytes from several goroutines at once, which +// broke legacy FromBytes functions written for v1.0 that are not safe for concurrent +// use. New and NewMock now serialize the calls. +func TestLegacyFromBytesIsNotCalledConcurrently(t *testing.T) { + const n = 200 + in := make(chan SortType, n) + for i := n; i > 0; i-- { + in <- legacyInt(i) + } + close(in) + var inFlight, maxInFlight atomic.Int32 + fromBytes := func(b []byte) SortType { + now := inFlight.Add(1) + defer inFlight.Add(-1) + for { + seen := maxInFlight.Load() + if now <= seen || maxInFlight.CompareAndSwap(seen, now) { + break + } + } + time.Sleep(20 * time.Microsecond) // widen the window for overlapping calls + v, _ := strconv.Atoi(string(b)) + return legacyInt(v) + } + less := func(a, b SortType) bool { return a.(legacyInt) < b.(legacyInt) } + // One record per chunk and 4 workers: the parallel merge reads 4 chunks at once + sorter, out, errc := NewMock(in, fromBytes, less, &Config{ChunkSize: 1, NumWorkers: 4}, 0) + sorter.Sort(context.Background()) + got, err := drainWithTimeout(t, out, errc) + if err != nil { + t.Fatal(err) + } + if len(got) != n { + t.Fatalf("got %d records, want %d", len(got), n) + } + if m := maxInFlight.Load(); m != 1 { + t.Errorf("FromBytes ran %d calls at once, want 1", m) + } +} + +// Save used to append an empty section after the last chunk, so the merge saw one more +// chunk than was written and chose the parallel merge for exactly NumWorkers chunks. +func TestTempFileHasOneSectionPerChunk(t *testing.T) { + for _, chunks := range []int{2, 3, 10} { + t.Run(strconv.Itoa(chunks)+" chunks", func(t *testing.T) { + tt := &trackedTemp{} + s := newTrackedSorter(descendingInts(chunks*5), itoaBytes, cmp.Compare[int], &Config{ChunkSize: 5}, tt) + got, err := runSort(t, context.Background(), s) + if err != nil || len(got) != chunks*5 { + t.Fatalf("got %d records, error %v", len(got), err) + } + tt.mu.Lock() + defer tt.mu.Unlock() + if tt.sections != chunks { + t.Errorf("temp file has %d sections, want %d", tt.sections, chunks) + } + }) + } +} + +// BenchmarkSortManyChunks sorts 200k records in 2,000 chunks. Save used to allocate a +// 64 KiB read buffer for every chunk up front (125 MiB here); they are now sized to the chunk. +func BenchmarkSortManyChunks(b *testing.B) { + const n = 200_000 + dir := b.TempDir() + for b.Loop() { + in := make(chan int, 1000) + go func() { + defer close(in) + for i := n; i > 0; i-- { + in <- i + } + }() + s, out, errc := Generic(in, atoiBytes, itoaBytes, cmp.Compare[int], &Config{ChunkSize: 100, TempFilesDir: dir}) + s.Sort(context.Background()) + for range out { + } + if err := <-errc; err != nil { + b.Fatal(err) + } + } +} + +// benchmarkInts sorts 1M random ints in chunks of 100k with the sorter that newSorter returns. +func benchmarkInts(b *testing.B, newSorter func(in chan int, config *Config) (Sorter, <-chan int, <-chan error)) { + r := rand.New(rand.NewSource(1)) + data := make([]int, 1_000_000) + for i := range data { + data[i] = r.Int() + } + config := DefaultConfig() + config.ChunkSize = 100_000 + config.TempFilesDir = b.TempDir() + b.ResetTimer() + for b.Loop() { + in := make(chan int, 1000) + go func() { + defer close(in) + for _, v := range data { + in <- v + } + }() + s, out, errc := newSorter(in, config) + s.Sort(context.Background()) + for range out { + } + if err := <-errc; err != nil { + b.Fatal(err) + } + } +} + +// BenchmarkOrderedInts sorts 1M ints with Ordered. Ordered used to gob-encode every record +// with a new encoder and decoder, which made it much slower than Generic. +func BenchmarkOrderedInts(b *testing.B) { + benchmarkInts(b, func(in chan int, config *Config) (Sorter, <-chan int, <-chan error) { + return Ordered(in, config) + }) +} + +// BenchmarkGenericVarintInts is BenchmarkOrderedInts using Generic with a varint codec. +func BenchmarkGenericVarintInts(b *testing.B) { + toBytes := func(v int) ([]byte, error) { return binary.AppendVarint(nil, int64(v)), nil } + fromBytes := func(d []byte) (int, error) { + v, n := binary.Varint(d) + if n <= 0 { + return 0, errors.New("bad varint") + } + return int(v), nil + } + benchmarkInts(b, func(in chan int, config *Config) (Sorter, <-chan int, <-chan error) { + return Generic(in, fromBytes, toBytes, cmp.Compare[int], config) + }) +} diff --git a/sort_generic.go b/sort_generic.go index 0a895be..87b3422 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -153,6 +153,9 @@ func (s *GenericSorter[E]) initMemoryPools() *memoryPools { // 3. Saves sorted chunks to temporary files using toBytes serialization // 4. Merges all chunks back into sorted order using fromBytes deserialization // +// fromBytes, toBytes and compareFunc are called from several goroutines at once, +// so they must be safe for concurrent use. +// // Call Sort() on the returned sorter to begin the sorting process. // Results are delivered via the output channel, errors via the error channel. // The temporary file is only created once the input spans more than one chunk; @@ -177,7 +180,8 @@ func Generic[E any](input <-chan E, fromBytes FromBytesGeneric[E], toBytes ToByt // MockGeneric creates an external sorter that uses in-memory storage instead of disk files. // This is primarily useful for testing and benchmarking without filesystem I/O overhead. // The parameter n specifies the initial capacity of the in-memory buffer. -// All other behavior is identical to Generic(). +// All other behavior is identical to Generic(), including that fromBytes, toBytes and +// compareFunc must be safe for concurrent use. func MockGeneric[E any](input <-chan E, fromBytes FromBytesGeneric[E], toBytes ToBytesGeneric[E], compareFunc CompareGeneric[E], config *Config, n int) (*GenericSorter[E], <-chan E, <-chan error) { s := newSorter(input, fromBytes, toBytes, compareFunc, config) s.newTempWriter = func() (tempfile.TempWriter, error) { @@ -265,13 +269,15 @@ 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 - for { + for inputOpen := true; inputOpen; { c := s.getChunk() + fill: for i := 0; i < s.config.ChunkSize; i++ { select { case rec, ok := <-s.input: if !ok { - break + inputOpen = false + break fill // a plain break would only leave the select } c.data = append(c.data, rec) case <-s.sortCtx.Done(): @@ -286,7 +292,7 @@ func (s *GenericSorter[E]) buildChunks() error { } select { - // chunk is now full + // chunk is now full, or holds the last records case s.chunkChan <- c: case <-s.sortCtx.Done(): s.putChunk(c) // Return unused chunk to pool diff --git a/sort_ordered.go b/sort_ordered.go index cf729cb..fa2ec91 100644 --- a/sort_ordered.go +++ b/sort_ordered.go @@ -1,87 +1,163 @@ package extsort import ( - "bytes" "cmp" - "encoding/gob" - "sync" + "encoding/binary" + "errors" + "math" + "reflect" + "unsafe" ) // OrderedSorter provides external sorting for types that implement cmp.Ordered. -// It embeds GenericSorter and adds optimized byte serialization using gob encoding -// with a sync.Pool for buffer reuse to reduce allocations. +// It embeds GenericSorter and serializes records with a compact binary codec +// chosen once for the underlying kind of T. type OrderedSorter[T cmp.Ordered] struct { GenericSorter[T] - bufferPool sync.Pool } -// newOrderedSorter creates a new OrderedSorter with an initialized buffer pool -// for efficient gob encoding/decoding operations. -func newOrderedSorter[T cmp.Ordered]() *OrderedSorter[T] { - s := &OrderedSorter[T]{ - bufferPool: sync.Pool{ - New: func() any { - return &bytes.Buffer{} - }, - }, +// errOrderedDecode reports a record that is not a valid encoding of the sorted type. +var errOrderedDecode = errors.New("extsort: invalid encoding of an ordered value") + +// orderedCodec returns serialization functions for T, chosen by its underlying kind, so +// 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. +// +// 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]) { + switch kind := reflect.TypeFor[T]().Kind(); kind { + case reflect.Int: + return signedCodec[T, int]() + case reflect.Int8: + return signedCodec[T, int8]() + case reflect.Int16: + return signedCodec[T, int16]() + case reflect.Int32: + return signedCodec[T, int32]() + case reflect.Int64: + return signedCodec[T, int64]() + case reflect.Uint: + return unsignedCodec[T, uint]() + case reflect.Uint8: + return unsignedCodec[T, uint8]() + case reflect.Uint16: + return unsignedCodec[T, uint16]() + case reflect.Uint32: + return unsignedCodec[T, uint32]() + case reflect.Uint64: + return unsignedCodec[T, uint64]() + case reflect.Uintptr: + return unsignedCodec[T, uintptr]() + case reflect.Float32: + return float32Codec[T]() + case reflect.Float64: + return float64Codec[T]() + case reflect.String: + return stringCodec[T]() + default: + panic("extsort: unsupported cmp.Ordered kind " + kind.String()) } - return s } -// fromBytesOrdered deserializes a byte slice back to the original type T -// using gob decoding. It reuses buffers from the pool for efficiency. -// Returns an error if decoding fails. -func (s *OrderedSorter[T]) fromBytesOrdered(d []byte) (T, error) { - var v T - buf := s.bufferPool.Get().(*bytes.Buffer) - buf.Reset() - buf.Write(d) - defer s.bufferPool.Put(buf) +// 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]) { + fromBytes := func(d []byte) (T, error) { + var v T + x, n := binary.Varint(d) + if n <= 0 || n != len(d) || int64(I(x)) != x { + return v, errOrderedDecode + } + *(*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 + } + return fromBytes, toBytes +} - dec := gob.NewDecoder(buf) - err := dec.Decode(&v) - return v, err +// 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]) { + fromBytes := func(d []byte) (T, error) { + var v T + x, n := binary.Uvarint(d) + if n <= 0 || n != len(d) || uint64(U(x)) != x { + return v, errOrderedDecode + } + *(*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 + } + return fromBytes, toBytes } -// toBytesOrdered serializes a value of type T to bytes using gob encoding. -// It reuses buffers from the pool and returns a copy of the serialized data. -// Returns an error if encoding fails. -func (s *OrderedSorter[T]) toBytesOrdered(d T) ([]byte, error) { - buf := s.bufferPool.Get().(*bytes.Buffer) - buf.Reset() - defer s.bufferPool.Put(buf) +// float32Codec stores a T whose underlying type is float32 as its 4 IEEE 754 bytes. +func float32Codec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { + fromBytes := func(d []byte) (T, error) { + var v T + if len(d) != 4 { + return v, errOrderedDecode + } + *(*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 + } + return fromBytes, toBytes +} - enc := gob.NewEncoder(buf) - err := enc.Encode(d) - if err != nil { - return nil, err +// float64Codec stores a T whose underlying type is float64 as its 8 IEEE 754 bytes. +func float64Codec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[T]) { + fromBytes := func(d []byte) (T, error) { + var v T + if len(d) != 8 { + return v, errOrderedDecode + } + *(*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 } + return fromBytes, toBytes +} - // Need to copy the bytes since we're returning the buffer to the pool - result := make([]byte, buf.Len()) - copy(result, buf.Bytes()) - return result, nil +// stringCodec stores a T whose underlying type is string as its bytes. +func stringCodec[T cmp.Ordered]() (FromBytesGeneric[T], ToBytesGeneric[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 + } + return fromBytes, toBytes } // Ordered performs external sorting on a channel of cmp.Ordered types. // It returns the sorter instance, output channel with sorted results, and error channel. -// Uses gob encoding for serialization and the < operator for comparison. +// Records are serialized with a compact binary codec for T's underlying kind and +// compared with cmp.Compare, which orders NaN before all other floats. // // 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) { - orderedSorter := newOrderedSorter[T]() - s, output, errChan := Generic(input, orderedSorter.fromBytesOrdered, orderedSorter.toBytesOrdered, cmp.Compare, config) - orderedSorter.GenericSorter = *s - return orderedSorter, output, errChan + fromBytes, toBytes := orderedCodec[T]() + s, output, errChan := Generic(input, fromBytes, toBytes, cmp.Compare, config) + return &OrderedSorter[T]{GenericSorter: *s}, output, errChan } // OrderedMock performs external sorting with a mock implementation that limits // 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) { - orderedSorter := newOrderedSorter[T]() - s, output, errChan := MockGeneric(input, orderedSorter.fromBytesOrdered, orderedSorter.toBytesOrdered, cmp.Compare, config, n) - orderedSorter.GenericSorter = *s - return orderedSorter, output, errChan + fromBytes, toBytes := orderedCodec[T]() + s, output, errChan := MockGeneric(input, fromBytes, toBytes, cmp.Compare, config, n) + return &OrderedSorter[T]{GenericSorter: *s}, output, errChan } diff --git a/sort_sorttype_legacy.go b/sort_sorttype_legacy.go index 8109e85..5561f06 100644 --- a/sort_sorttype_legacy.go +++ b/sort_sorttype_legacy.go @@ -1,5 +1,7 @@ package extsort +import "sync" + // SortType defines the interface required by the extsort library to be able to sort the items // // Deprecated: Use Generic() with custom types instead for new code. This interface is maintained for backward compatibility. @@ -44,7 +46,10 @@ func sortTypeToBytes(a SortType) (result []byte, err error) { // makeSortTypeFromBytes creates a generic-compatible deserialization function from a legacy FromBytes function. // It wraps the legacy function to catch any panics and convert them to DeserializationError instances, // enabling graceful error handling during the merge phase of external sorting. +// Calls are serialized with a mutex: FromBytes functions written for v1.0 were only called from one +// goroutine, and the parallel merge added in v1.1.0 broke those that are not safe for concurrent use. func makeSortTypeFromBytes(fromBytes FromBytes) func([]byte) (SortType, error) { + var mu sync.Mutex return func(d []byte) (result SortType, err error) { // named results: the deferred recover must be able to set the returned error defer func() { @@ -53,6 +58,8 @@ func makeSortTypeFromBytes(fromBytes FromBytes) func([]byte) (SortType, error) { err = NewDeserializationError(r, len(d), "FromBytes") } }() + mu.Lock() + defer mu.Unlock() return fromBytes(d), nil } } @@ -70,6 +77,8 @@ func makeCompareSortType(lessFunc CompareLessFunc) func(a, b SortType) int { // It takes a FromBytes function for deserialization and a CompareLessFunc for comparison. // Returns the sorter instance, output channel with sorted items, and error channel. // This function provides backward compatibility with the original extsort API. +// Calls to fromBytes are serialized, so it need not be safe for concurrent use, +// but lessFunc is called from several goroutines at once and must be. // // 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. @@ -88,7 +97,8 @@ func New(input <-chan SortType, fromBytes FromBytes, lessFunc CompareLessFunc, c // NewMock performs external sorting on SortType items with a mock implementation that limits // the number of items to sort. Useful for testing with a controlled dataset size. // The parameter n specifies the maximum number of items to process. -// Uses the same interface-based API as New for backward compatibility. +// Uses the same interface-based API as New for backward compatibility, with the same +// concurrency rules: fromBytes calls are serialized, lessFunc must be safe for concurrent use. // // Deprecated: Use MockGeneric() instead for new code. This function is maintained for backward compatibility. func NewMock(input <-chan SortType, fromBytes FromBytes, lessFunc CompareLessFunc, config *Config, n int) (*SortTypeSorter, <-chan SortType, <-chan error) { diff --git a/tempfile/mockfile.go b/tempfile/mockfile.go index 645d2cf..2fd3e8e 100644 --- a/tempfile/mockfile.go +++ b/tempfile/mockfile.go @@ -3,7 +3,6 @@ package tempfile import ( "bufio" "bytes" - "io" ) // MockFileWriter provides an in-memory implementation of the TempWriter interface. @@ -12,6 +11,7 @@ import ( type MockFileWriter struct { data *bytes.Buffer sections []int + pending bool // data was written since the last Next } // mockFileReader provides an in-memory implementation of the TempReader interface. @@ -33,11 +33,14 @@ func Mock(n int) *MockFileWriter { return &m } -// Size returns the total number of virtual file sections that have been created. -// This includes the current section being written plus all completed sections. +// Size returns the number of virtual file sections Save will produce: all completed +// sections, plus the current one if anything was written to it. A writer with no data +// has one empty section. func (w *MockFileWriter) Size() int { - // we add one because we only write to the sections when we are done - return len(w.sections) + 1 + if w.pending || len(w.sections) == 0 { + return len(w.sections) + 1 + } + return len(w.sections) } // Close terminates the MockFileWriter and releases all memory. @@ -52,12 +55,20 @@ func (w *MockFileWriter) Close() error { // Write appends data to the current virtual file section in memory. func (w *MockFileWriter) Write(p []byte) (int, error) { - return w.data.Write(p) + n, err := w.data.Write(p) + if n > 0 { + w.pending = true + } + return n, err } // WriteString appends string data to the current virtual file section in memory. func (w *MockFileWriter) WriteString(s string) (int, error) { - return w.data.WriteString(s) + n, err := w.data.WriteString(s) + if n > 0 { + w.pending = true + } + return n, err } // Next finalizes the current virtual file section and prepares for writing the next section. @@ -66,16 +77,19 @@ func (w *MockFileWriter) Next() (int64, error) { // save offsets pos := w.data.Len() w.sections = append(w.sections, pos) + w.pending = false return int64(pos), nil } // Save finalizes all virtual file sections and returns a TempReader for accessing the data. +// The current section becomes the last one only if anything was written to it, as with FileWriter. // After calling Save(), the MockFileWriter can no longer be used for writing. // The returned TempReader allows concurrent access to any virtual file section. func (w *MockFileWriter) Save() (TempReader, error) { - _, err := w.Next() - if err != nil { - return nil, err + if w.pending || len(w.sections) == 0 { + if _, err := w.Next(); err != nil { + return nil, err + } } return newMockTempReader(w.sections, w.data.Bytes()) } @@ -89,13 +103,6 @@ func newMockTempReader(sections []int, data []byte) (*mockFileReader, error) { r.sections = sections r.readers = make([]*bufio.Reader, len(r.sections)) - offset := 0 - for i, end := range r.sections { - section := io.NewSectionReader(r.data, int64(offset), int64(end-offset)) - offset = end - r.readers[i] = bufio.NewReaderSize(section, fileBufferSize) - } - return &r, nil } @@ -112,11 +119,19 @@ func (r *mockFileReader) Size() int { return len(r.readers) } -// Read returns a buffered reader for the specified virtual file section. +// Read returns a buffered reader for the specified virtual file section, created on +// first use like the disk-based reader. Read may be called concurrently for different sections. // Panics if the section index is out of range. func (r *mockFileReader) Read(i int) *bufio.Reader { if i < 0 || i >= len(r.readers) { panic("tempfile: read request out of range") } + if r.readers[i] == nil { + start := 0 + if i > 0 { + start = r.sections[i-1] + } + r.readers[i] = newSectionReader(r.data, int64(start), int64(r.sections[i]), len(r.sections)) + } return r.readers[i] } diff --git a/tempfile/regression_test.go b/tempfile/regression_test.go index fcc09a8..666cd59 100644 --- a/tempfile/regression_test.go +++ b/tempfile/regression_test.go @@ -3,12 +3,25 @@ package tempfile // Regression tests for temp directory selection and cleanup. import ( + "io" "os" "path/filepath" "runtime" "testing" ) +// testWriters creates each TempWriter implementation. +var testWriters = map[string]func(t *testing.T) TempWriter{ + "FileWriter": func(t *testing.T) TempWriter { + w, err := New(t.TempDir(), true) + if err != nil { + t.Fatal(err) + } + return w + }, + "MockFileWriter": func(*testing.T) TempWriter { return Mock(0) }, +} + // canTestPermissions reports whether chmod restrictions apply to this process. func canTestPermissions() bool { return runtime.GOOS != "windows" && os.Geteuid() != 0 @@ -119,3 +132,150 @@ func TestReaderCloseRemovesExtsortDir(t *testing.T) { t.Errorf("%s still exists after the reader was closed", dir) } } + +// BenchmarkWriteAndSave writes 64 MiB in 1 MiB sections, then saves and closes the file. +// Save used to fsync the file, which is already unlinked on Unix and only read back by +// this process. +func BenchmarkWriteAndSave(b *testing.B) { + data := make([]byte, 1<<20) + dir := b.TempDir() + for b.Loop() { + w, err := New(dir, true) + if err != nil { + b.Fatal(err) + } + for range 64 { + if _, err := w.Write(data); err != nil { + b.Fatal(err) + } + if _, err := w.Next(); err != nil { + b.Fatal(err) + } + } + r, err := w.Save() + if err != nil { + b.Fatal(err) + } + if err := r.Close(); err != nil { + b.Fatal(err) + } + } +} + +// Save used to end the current section even when nothing was written since the last +// Next, adding an empty section after the last one. +func TestSaveAddsNoEmptySection(t *testing.T) { + for name, newWriter := range testWriters { + t.Run(name, func(t *testing.T) { + for _, tc := range []struct { + name string + sections []string // each ended by Next + trailing string // written after the last Next + want []string + }{ + {"Next after every section", []string{"a", "b"}, "", []string{"a", "b"}}, + {"data after the last Next", []string{"a", "b"}, "c", []string{"a", "b", "c"}}, + {"nothing written", nil, "", []string{""}}, + } { + t.Run(tc.name, func(t *testing.T) { + w := newWriter(t) + for _, data := range tc.sections { + if _, err := w.WriteString(data); err != nil { + t.Fatal(err) + } + if _, err := w.Next(); err != nil { + t.Fatal(err) + } + } + if _, err := w.WriteString(tc.trailing); err != nil { + t.Fatal(err) + } + if got := w.Size(); got != len(tc.want) { + t.Errorf("writer Size() = %d, want %d", got, len(tc.want)) + } + r, err := w.Save() + if err != nil { + t.Fatal(err) + } + defer func() { _ = r.Close() }() + if got := r.Size(); got != len(tc.want) { + t.Fatalf("reader Size() = %d, want %d", got, len(tc.want)) + } + for i, want := range tc.want { + got, err := io.ReadAll(r.Read(i)) + if err != nil || string(got) != want { + t.Errorf("section %d = %q, %v; want %q", i, got, err, want) + } + } + }) + } + }) + } +} + +func TestReadBufferSize(t *testing.T) { + for _, tc := range []struct { + sections int + sectionLen int64 + want int + }{ + {1, 1 << 20, fileBufferSize}, + {1000, 1 << 20, fileBufferSize}, // a 64 MiB share per 1000 sections is still above 64 KiB + {10_000, 1 << 20, readBufferBudget / 10_000}, + {100_000, 1 << 20, minReadBufferSize}, + {1, 100, 100}, + {10_000, 100, 100}, + {0, 0, 0}, + } { + if got := readBufferSize(tc.sections, tc.sectionLen); got != tc.want { + t.Errorf("readBufferSize(%d, %d) = %d, want %d", tc.sections, tc.sectionLen, got, tc.want) + } + } +} + +// allocatedBytes returns how many bytes f allocated. +func allocatedBytes(f func()) uint64 { + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + f() + runtime.ReadMemStats(&after) + return after.TotalAlloc - before.TotalAlloc +} + +// Save used to allocate a 64 KiB read buffer for every section up front, so merge memory +// grew by 64 KiB per chunk. Buffers are now created on first Read and sized to the section. +func TestReadBuffersAreLazyAndSized(t *testing.T) { + const sections = 1000 // 62.5 MiB of read buffers before + for name, newWriter := range testWriters { + t.Run(name, func(t *testing.T) { + w := newWriter(t) + for range sections { + if _, err := w.WriteString("0123456789"); err != nil { + t.Fatal(err) + } + if _, err := w.Next(); err != nil { + t.Fatal(err) + } + } + var r TempReader + var err error + if n := allocatedBytes(func() { r, err = w.Save() }); n > 1<<20 { + t.Errorf("Save allocated %d bytes for %d sections", n, sections) + } + if err != nil { + t.Fatal(err) + } + defer func() { _ = r.Close() }() + if n := allocatedBytes(func() { + for i := range sections { + r.Read(i) + } + }); n > 1<<20 { + t.Errorf("reading %d 10-byte sections allocated %d bytes", sections, n) + } + if got, err := io.ReadAll(r.Read(sections - 1)); err != nil || string(got) != "0123456789" { + t.Errorf("last section = %q, %v", got, err) + } + }) + } +} diff --git a/tempfile/tempdir.go b/tempfile/tempdir.go index 3536392..7564731 100644 --- a/tempfile/tempdir.go +++ b/tempfile/tempdir.go @@ -16,7 +16,8 @@ var ( memoryAllowedDir string dirDiscoveryOnce sync.Once - // Cached expensive operations + // Cached expensive operations, filled on first use by cacheOnce + cacheOnce sync.Once cachedHomeDir string cachedWorkDir string cachedOSTemp string @@ -46,9 +47,6 @@ func GetTempDir(dir string, preferDiskBacked bool) string { // This runs once and caches expensive operations like os.UserHomeDir() and os.Getwd(). // Called by sync.Once to ensure thread-safe initialization. func discoverOptimalDirectories() { - // Cache expensive operations once - cacheExpensiveOperations() - // Find optimal directory for disk-preferred usage diskPreferredDir = findBestDirectory(true) @@ -106,6 +104,7 @@ func firstUsableDir(candidates []string) string { // When preferDiskBacked is true, disk-preferred candidates are prioritized first. // Uses cached values for performance. The order depends on the OS and preferDiskBacked setting. func buildCandidateList(preferDiskBacked bool) []string { + cacheOnce.Do(cacheExpensiveOperations) var candidates []string if preferDiskBacked { @@ -155,6 +154,7 @@ func buildDiskPreferredCandidates() []string { // Creates process-specific subdirectories in the user's home directory and current // working directory. Uses cached directory values for performance. func buildAdditionalFallbacks() []string { + cacheOnce.Do(cacheExpensiveOperations) var candidates []string // Try user home directory with subdirectory (using cached value) diff --git a/tempfile/tempfile.go b/tempfile/tempfile.go index 78c029e..63a7ac0 100644 --- a/tempfile/tempfile.go +++ b/tempfile/tempfile.go @@ -29,6 +29,33 @@ import ( // file IO buffer size for each file const fileBufferSize = 1 << 16 // 64k +const ( + // readBufferBudget caps the combined read buffers of all sections. A merge reads every + // section at once, so with many sections each gets an equal share instead of fileBufferSize. + readBufferBudget = 64 << 20 // 64 MiB + // minReadBufferSize is the smallest share a section larger than it gets. + minReadBufferSize = 4 << 10 // 4 KiB +) + +// readBufferSize returns the read buffer size for one of n sections of length sectionLen: +// fileBufferSize, shrunk to an equal share of readBufferBudget when there are many +// sections, and to the section length when the section is smaller. +func readBufferSize(n int, sectionLen int64) int { + size := fileBufferSize + if share := readBufferBudget / max(n, 1); share < size { + size = max(share, minReadBufferSize) + } + if sectionLen < int64(size) { + size = int(sectionLen) // bufio raises tiny sizes to its minimum + } + return size +} + +// newSectionReader returns a buffered reader for bytes [start, end) of ra, one of n sections. +func newSectionReader(ra io.ReaderAt, start, end int64, n int) *bufio.Reader { + return bufio.NewReaderSize(io.NewSectionReader(ra, start, end-start), readBufferSize(n, end-start)) +} + // filename prefix for files put in temp directory var mergeFilenamePrefix = fmt.Sprintf("extsort_%d_", os.Getpid()) @@ -49,6 +76,7 @@ type FileWriter struct { file *os.File bufWriter *bufio.Writer sections []int64 + pending bool // data was written since the last Next needsCleanup bool // true if manual cleanup is needed (Windows) createdDir string // directory we created (for cleanup) } @@ -56,10 +84,10 @@ type FileWriter struct { type fileReader struct { file *os.File sections []int64 - readers []*bufio.Reader - needsCleanup bool // true if manual cleanup is needed (Windows) - filename string // filename for cleanup - createdDir string // directory we created (for cleanup), taken over from the FileWriter + readers []*bufio.Reader // created by Read on first use + needsCleanup bool // true if manual cleanup is needed (Windows) + filename string // filename for cleanup + createdDir string // directory we created (for cleanup), taken over from the FileWriter } // New creates a new FileWriter for virtual temporary files in the specified directory. @@ -110,11 +138,14 @@ func New(dir string, preferDiskBacked bool) (*FileWriter, error) { return &w, nil } -// Size returns the total number of virtual file sections created. -// This includes the current section being written plus all completed sections. +// Size returns the number of virtual file sections Save will produce: all completed +// sections, plus the current one if anything was written to it. A writer with no data +// has one empty section. func (w *FileWriter) Size() int { - // we add one because we only write to the sections when we are done - return len(w.sections) + 1 + if w.pending || len(w.sections) == 0 { + return len(w.sections) + 1 + } + return len(w.sections) } // Name returns the full filesystem path of the underlying physical temporary file. @@ -152,13 +183,21 @@ func (w *FileWriter) Close() error { // Write appends data to the current virtual file section. // Data is buffered for efficiency and will be flushed when Next() or Save() is called. func (w *FileWriter) Write(p []byte) (int, error) { - return w.bufWriter.Write(p) + n, err := w.bufWriter.Write(p) + if n > 0 { + w.pending = true + } + return n, err } // WriteString appends a string to the current virtual file section. // This is more efficient than Write() for string data as it avoids byte slice conversion. func (w *FileWriter) WriteString(s string) (int, error) { - return w.bufWriter.WriteString(s) + n, err := w.bufWriter.WriteString(s) + if n > 0 { + w.pending = true + } + return n, err } // Next finalizes the current virtual file section and prepares for writing the next section. @@ -175,21 +214,25 @@ func (w *FileWriter) Next() (int64, error) { return 0, err } w.sections = append(w.sections, pos) + w.pending = false return pos, nil } // Save finalizes all virtual file sections and returns a TempReader for accessing the data. +// The current section becomes the last one only if anything was written to it, so a +// Next after the final section does not add an empty section. // After calling Save(), the FileWriter can no longer be used for writing. // The returned TempReader allows concurrent access to any virtual file section. func (w *FileWriter) Save() (TempReader, error) { - _, err := w.Next() - if err != nil { - return nil, err - } - err = w.file.Sync() - if err != nil { - return nil, err + // No Sync: the file is only read back by this process, and on Unix it is already + // unlinked, so flushing it to stable storage only costs time. + var err error + if w.pending || len(w.sections) == 0 { + // Next also flushes; otherwise the last Next already did + if _, err = w.Next(); err != nil { + return nil, err + } } var r *fileReader @@ -230,13 +273,6 @@ func newTempReader(filename string, sections []int64, needsCleanup bool) (*fileR r.needsCleanup = needsCleanup r.filename = filename - offset := int64(0) - for i, end := range r.sections { - section := io.NewSectionReader(r.file, offset, end-offset) - offset = end - r.readers[i] = bufio.NewReaderSize(section, fileBufferSize) - } - return &r, nil } @@ -251,13 +287,6 @@ func newTempReaderFromFile(file *os.File, sections []int64, needsCleanup bool) ( r.needsCleanup = needsCleanup r.filename = file.Name() - offset := int64(0) - for i, end := range r.sections { - section := io.NewSectionReader(r.file, offset, end-offset) - offset = end - r.readers[i] = bufio.NewReaderSize(section, fileBufferSize) - } - return &r, nil } @@ -289,11 +318,21 @@ func (r *fileReader) Size() int { } // Read returns a buffered reader for the specified virtual file section. +// The reader is created on first use, with a buffer sized by readBufferSize, so the +// memory used grows with the sections read rather than 64 KiB per section up front. +// Read may be called concurrently for different sections. // Panics if the section index is out of range. func (r *fileReader) Read(i int) *bufio.Reader { if i < 0 || i >= len(r.readers) { panic("tempfile: read request out of range") } + if r.readers[i] == nil { + var start int64 + if i > 0 { + start = r.sections[i-1] + } + r.readers[i] = newSectionReader(r.file, start, r.sections[i], len(r.sections)) + } return r.readers[i] } diff --git a/tempfile/tempfile_test.go b/tempfile/tempfile_test.go index 1bd98d2..edf93bd 100644 --- a/tempfile/tempfile_test.go +++ b/tempfile/tempfile_test.go @@ -87,7 +87,7 @@ func TestTempFileRepeat(t *testing.T) { // via the file handle, so we don't check for file existence here s := tempReader.Size() - if s != iterations+1 { + if s != iterations { t.Fatalf("tempReader.Size returned %d, expected %d", s, iterations) } @@ -175,7 +175,7 @@ func TestMockFileMultiSection(t *testing.T) { } s := tempReader.Size() - if s != iterations+1 { + if s != iterations { t.Fatalf("tempReader.Size returned %d, expected %d", s, iterations) } diff --git a/tempfile_error_test.go b/tempfile_error_test.go index b110d27..a460cb8 100644 --- a/tempfile_error_test.go +++ b/tempfile_error_test.go @@ -4,6 +4,7 @@ import ( "context" "os" "path/filepath" + "runtime" "testing" "github.com/lanrat/extsort" @@ -13,6 +14,8 @@ import ( // failures gracefully without causing segmentation faults. This test addresses // the bug reported in issue #10 where a full filesystem causes a segfault. func TestTempFileCreationFailure(t *testing.T) { + skipUnlessDirPermissionsEnforced(t) + // Create a test directory that will be automatically cleaned up testDir := t.TempDir() @@ -91,6 +94,8 @@ func TestTempFileCreationFailure(t *testing.T) { // TestTempFileCreationFailureStrings tests the same scenario with string sorting func TestTempFileCreationFailureStrings(t *testing.T) { + skipUnlessDirPermissionsEnforced(t) + testDir := t.TempDir() readOnlyDir := filepath.Join(testDir, "readonly") err := os.Mkdir(readOnlyDir, 0555) @@ -156,6 +161,18 @@ func TestTempFileCreationFailureStrings(t *testing.T) { } } +// skipUnlessDirPermissionsEnforced skips tests that rely on a 0555 directory +// rejecting new files, which does not hold on Windows or when running as root. +func skipUnlessDirPermissionsEnforced(t *testing.T) { + t.Helper() + if runtime.GOOS == "windows" { + t.Skip("directory permission bits do not prevent file creation on Windows") + } + if os.Geteuid() == 0 { + t.Skip("root can create files in a read-only directory") + } +} + // Helper types and functions for testing type testData struct { Key int diff --git a/types.go b/types.go index 2943003..7bc5725 100644 --- a/types.go +++ b/types.go @@ -18,6 +18,8 @@ type Sorter interface { // The function should be the inverse of the corresponding ToBytesGeneric function. // It returns an error for any deserialization failures, which will be wrapped // in a DeserializationError by the external sorter. +// The parallel merge calls it from several goroutines at once, so it must be safe +// for concurrent use. type FromBytesGeneric[E any] func([]byte) (E, error) // ToBytesGeneric is a function type for serializing type E to bytes. @@ -25,6 +27,7 @@ type FromBytesGeneric[E any] func([]byte) (E, error) // The function should produce deterministic output that can be read back // by the corresponding FromBytesGeneric function. It returns an error for any // serialization failures, which will be wrapped in a SerializationError by the external sorter. +// It may be called from several goroutines at once, so it must be safe for concurrent use. type ToBytesGeneric[E any] func(E) ([]byte, error) // CompareGeneric is a function type for comparing two items of type E. @@ -33,4 +36,6 @@ type ToBytesGeneric[E any] func(E) ([]byte, error) // and a positive integer if a should be ordered after b in the final sorted output. // The function must be consistent and must handle any errors by panicking. // This follows the same semantics as cmp.Compare and can be implemented using cmp.Compare[T] for ordered types. +// The sort and merge workers call it from several goroutines at once, so it must be +// safe for concurrent use. type CompareGeneric[E any] func(a, b E) int diff --git a/uniq.go b/uniq.go index 49b8f6e..367ac45 100644 --- a/uniq.go +++ b/uniq.go @@ -1,5 +1,8 @@ package extsort +// uniqChanBuffSize matches the default SortedChanBuffSize of the sorter output it usually reads. +const uniqChanBuffSize = 1000 + // UniqStringChan returns a channel that filters out consecutive duplicate strings from the input. // This function assumes the input channel provides strings in sorted order and uses string equality // to detect duplicates. It preserves the first occurrence of each unique string while filtering @@ -8,8 +11,9 @@ package extsort // // The returned channel will be closed when the input channel is closed. // This function spawns a goroutine that will terminate when the input channel is closed. -func UniqStringChan(in chan string) chan string { - out := make(chan string) +// It accepts the receive-only output channel of Strings; a bidirectional chan string works too. +func UniqStringChan(in <-chan string) chan string { + out := make(chan string, uniqChanBuffSize) go func() { var prior string priorSet := false diff --git a/uniq_test.go b/uniq_test.go index 9e0775b..43c88d6 100644 --- a/uniq_test.go +++ b/uniq_test.go @@ -1,7 +1,9 @@ package extsort_test import ( + "context" "fmt" + "slices" "testing" "github.com/lanrat/extsort" @@ -30,3 +32,30 @@ func TestUniqString(t *testing.T) { past = u } } + +// UniqStringChan used to take a bidirectional chan string, so it could not accept the +// <-chan string returned by Strings, and its output channel was unbuffered. +func TestUniqStringChanAcceptsSorterOutput(t *testing.T) { + in := make(chan string, 6) + for _, s := range []string{"b", "a", "b", "c", "a", "b"} { + in <- s + } + close(in) + sorter, sorted, errc := extsort.Strings(in, nil) + sorter.Sort(context.Background()) + + uniq := extsort.UniqStringChan(sorted) + if cap(uniq) == 0 { + t.Error("output channel is unbuffered") + } + var got []string + for s := range uniq { + got = append(got, s) + } + if err := <-errc; err != nil { + t.Fatal(err) + } + if want := []string{"a", "b", "c"}; !slices.Equal(got, want) { + t.Errorf("got %v, want %v", got, want) + } +}