From 35e71769c307e41a44d95d6e9742dd77445b6aa2 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 18:57:21 -0700 Subject: [PATCH 01/11] ci: test on Linux, macOS and Windows; add race and examples steps - Run the tests on ubuntu-latest, macos-latest and windows-latest. golangci-lint stays on ubuntu. Windows calls go test directly because make is not reliably available there. - Add a make test-race target with a 10 minute timeout (the root package takes over a minute under -race) and run it on ubuntu. - Fix make examples: stop at the first failing example instead of reporting only the last one, and detect directories with more than one .go file (the old [ -f dir*.go ] test broke on them). Run it on ubuntu. - Skip TestTempFileCreationFailure and TestTempFileCreationFailureStrings on Windows and as root, where a 0555 directory does not block file creation. - Fill the temp dir cache on first use in buildCandidateList and buildAdditionalFallbacks, so TestGetAdditionalFallbacks no longer depends on another test having called GetTempDir first. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/go.yml | 22 ++++++++++++++++++++-- Makefile | 11 ++++++++--- tempfile/tempdir.go | 8 ++++---- tempfile_error_test.go | 17 +++++++++++++++++ 4 files changed, 49 insertions(+), 9 deletions(-) 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/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_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 From f520dd226514344096b13aa7f281a0e5de1ffbfe Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 18:58:34 -0700 Subject: [PATCH 02/11] Stop reading input as soon as it closes buildChunks used a plain break when the input channel closed, which only left the select, so it kept polling the closed channel up to ChunkSize times, then allocated and polled one more empty chunk. With the default ChunkSize of 1M this added about 50 ms to every sort. Use a labeled break and stop after the last chunk. A 10-record sort drops from 51.8 ms to 0.05 ms (BenchmarkSortTenRecords) and allocates one chunk instead of two. Co-Authored-By: Claude Opus 5.5 --- regression_test.go | 53 ++++++++++++++++++++++++++++++++++++++++++++++ sort_generic.go | 8 ++++--- 2 files changed, 58 insertions(+), 3 deletions(-) diff --git a/regression_test.go b/regression_test.go index d791b23..57ff4e8 100644 --- a/regression_test.go +++ b/regression_test.go @@ -628,3 +628,56 @@ 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) + } + } +} diff --git a/sort_generic.go b/sort_generic.go index 0a895be..04dbddc 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -265,13 +265,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 +288,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 From 4f508b619f4c14cbd17ff9c3fd042830896210bc Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 18:59:33 -0700 Subject: [PATCH 03/11] Normalize a copy of Config and fix its docs - mergeConfig wrote defaults into the caller's *Config, which raced when one Config was shared by several sorters. It now normalizes a copy and never modifies the caller's struct. - Document the zero-value rules, which were misdescribed: a nil Config means DefaultConfig(), a ChunkSize or NumWorkers below 1 and a negative buffer size use the default, and a zero ChanBuffSize or SortedChanBuffSize means an unbuffered channel. The "Must be > 1" notes on ChunkSize and NumWorkers were also wrong: 1 is valid. - Change the ChanBuffSize default from 16 to 1, matching the docs and README. It only sizes the channel between reading the input and the sort workers, which holds whole chunks: 16 let up to 25 chunks sit in memory at once for no gain in wall time. Co-Authored-By: Claude Opus 5.5 --- config.go | 44 ++++++++++++++++++++++++-------------------- regression_test.go | 46 ++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 70 insertions(+), 20 deletions(-) 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/regression_test.go b/regression_test.go index 57ff4e8..4e1c343 100644 --- a/regression_test.go +++ b/regression_test.go @@ -681,3 +681,49 @@ func BenchmarkSortTenRecords(b *testing.B) { } } } + +// 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) + } + }) + } +} From ef9bf912bb6860442788d464a53c3c95b773d618 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:01:08 -0700 Subject: [PATCH 04/11] Document the concurrency contract; serialize legacy FromBytes compareFunc is called from several sort and merge goroutines at once, and since the parallel merge (v1.1.0) so is fromBytes. None of this was documented. Say so on the function types in types.go and on Generic, MockGeneric, New and NewMock. toBytes is documented the same way so that serializing records in the sort workers stays possible. The parallel merge broke legacy FromBytes functions written for v1.0, which were only ever called from one goroutine. New and NewMock now guard FromBytes with a per-sorter mutex. Generic stays lock-free. Co-Authored-By: Claude Opus 5.5 --- regression_test.go | 41 +++++++++++++++++++++++++++++++++++++++++ sort_generic.go | 6 +++++- sort_sorttype_legacy.go | 12 +++++++++++- types.go | 5 +++++ 4 files changed, 62 insertions(+), 2 deletions(-) diff --git a/regression_test.go b/regression_test.go index 4e1c343..607b0cf 100644 --- a/regression_test.go +++ b/regression_test.go @@ -17,6 +17,7 @@ import ( "strconv" "strings" "sync" + "sync/atomic" "testing" "testing/iotest" "time" @@ -727,3 +728,43 @@ func TestConfigBufferSizes(t *testing.T) { }) } } + +// 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) + } +} diff --git a/sort_generic.go b/sort_generic.go index 04dbddc..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) { 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/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 From 362f01f807802c5e545740806d6d2e5e5e73015e Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:03:01 -0700 Subject: [PATCH 05/11] diff: read each error channel once and respect ctx; type NEW and OLD - diff read the error channel of the stream that ended first twice, so it hung if the caller never closed that channel, and error reads ignored ctx, so the hang outlasted the ctx deadline. Each error channel is now read once, and the read gives up when ctx is done. - StringResultChan's send ignored ctx, so a diff whose results were no longer read hung. Add StringResultChanContext, whose result function returns ctx.Err() once ctx is done. StringResultChan keeps its signature and behavior. - BREAKING: NEW and OLD are now typed Delta constants instead of untyped integers, so they print as ">" and "<". Code that uses them as plain ints no longer compiles. Co-Authored-By: Claude Opus 5.5 --- diff/diff_generic.go | 35 +++++++++---- diff/diff_result_chan.go | 19 ++++++- diff/diff_types.go | 2 +- diff/regression_test.go | 108 +++++++++++++++++++++++++++++++++++++++ 4 files changed, 152 insertions(+), 12 deletions(-) create mode 100644 diff/regression_test.go 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) + } +} From f7ae0fc1780ac91a9713c1ee57cebb3c90cf82df Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:03:35 -0700 Subject: [PATCH 06/11] Let UniqStringChan accept the sorter's output and buffer it UniqStringChan took a bidirectional chan string, so it could not be passed the <-chan string that Strings returns, its main use. Take <-chan string instead; callers passing a chan string still compile. The output channel is now buffered like the sorter's output (1000). Co-Authored-By: Claude Opus 5.5 --- uniq.go | 8 ++++++-- uniq_test.go | 29 +++++++++++++++++++++++++++++ 2 files changed, 35 insertions(+), 2 deletions(-) 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) + } +} From 986d53a57ed8ec6fdbe2088d223446277dc3016b Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:04:54 -0700 Subject: [PATCH 07/11] tempfile: drop the fsync in Save Save called Sync on a file that is only read back by this process and, on Unix, is already unlinked, so durability buys nothing. Writing and saving 64 MiB drops from 20.4 ms to 7.3 ms on macOS (BenchmarkWriteAndSave). Co-Authored-By: Claude Opus 5.5 --- tempfile/regression_test.go | 29 +++++++++++++++++++++++++++++ tempfile/tempfile.go | 6 ++---- 2 files changed, 31 insertions(+), 4 deletions(-) diff --git a/tempfile/regression_test.go b/tempfile/regression_test.go index fcc09a8..d1df033 100644 --- a/tempfile/regression_test.go +++ b/tempfile/regression_test.go @@ -119,3 +119,32 @@ 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) + } + } +} diff --git a/tempfile/tempfile.go b/tempfile/tempfile.go index 78c029e..6ce2aac 100644 --- a/tempfile/tempfile.go +++ b/tempfile/tempfile.go @@ -183,14 +183,12 @@ func (w *FileWriter) Next() (int64, error) { // 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) { + // 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. _, err := w.Next() if err != nil { return nil, err } - err = w.file.Sync() - if err != nil { - return nil, err - } var r *fileReader if w.needsCleanup { From b9f109f209d936dd6036599ec3f9bd9f5a2634f8 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:06:42 -0700 Subject: [PATCH 08/11] tempfile: don't end an empty section in Save Save always called Next, so after the sorter's Next for its last chunk it appended an empty section. The merge then counted one chunk more than was written, which chose the parallel merge for exactly NumWorkers chunks and gave one worker an empty range. Save now ends the current section only if anything was written to it (or nothing was written at all, which keeps one empty section), and Size reports the count Save will produce. Same for MockFileWriter. TestTempFileRepeat and TestMockFileMultiSection expected the extra section; they now expect one section per Next. Co-Authored-By: Claude Opus 5.5 --- regression_test.go | 48 +++++++++++++++++++++------- tempfile/mockfile.go | 33 ++++++++++++++------ tempfile/regression_test.go | 62 +++++++++++++++++++++++++++++++++++++ tempfile/tempfile.go | 36 +++++++++++++++------ tempfile/tempfile_test.go | 4 +-- 5 files changed, 151 insertions(+), 32 deletions(-) diff --git a/regression_test.go b/regression_test.go index 607b0cf..b933ff7 100644 --- a/regression_test.go +++ b/regression_test.go @@ -90,9 +90,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) { @@ -136,6 +137,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 } @@ -244,16 +248,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) @@ -386,16 +390,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) } }) } @@ -768,3 +772,23 @@ func TestLegacyFromBytesIsNotCalledConcurrently(t *testing.T) { 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) + } + }) + } +} diff --git a/tempfile/mockfile.go b/tempfile/mockfile.go index 645d2cf..fce031b 100644 --- a/tempfile/mockfile.go +++ b/tempfile/mockfile.go @@ -12,6 +12,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 +34,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 +56,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 +78,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()) } diff --git a/tempfile/regression_test.go b/tempfile/regression_test.go index d1df033..74108a9 100644 --- a/tempfile/regression_test.go +++ b/tempfile/regression_test.go @@ -3,6 +3,7 @@ package tempfile // Regression tests for temp directory selection and cleanup. import ( + "io" "os" "path/filepath" "runtime" @@ -148,3 +149,64 @@ func BenchmarkWriteAndSave(b *testing.B) { } } } + +// 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) { + writers := 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) }, + } + for name, newWriter := range writers { + 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) + } + } + }) + } + }) + } +} diff --git a/tempfile/tempfile.go b/tempfile/tempfile.go index 6ce2aac..994edb9 100644 --- a/tempfile/tempfile.go +++ b/tempfile/tempfile.go @@ -49,6 +49,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) } @@ -110,11 +111,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 +156,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,19 +187,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) { // 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. - _, err := w.Next() - if err != nil { - return nil, err + 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 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) } From 9ade0cf0870ca0ff0ae6f0e7acd28a48b4236c47 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:09:13 -0700 Subject: [PATCH 09/11] tempfile: create read buffers lazily and size them to the section Save allocated a 64 KiB bufio.Reader for every section up front, so merge memory grew by 64 KiB per chunk: 10k chunks took 640 MB before reading a byte. Readers are now created on the first Read of each section, with a buffer that shrinks to the section length for small sections and to an equal share of a 64 MiB budget (at least 4 KiB) when there are many sections. Sorting 200k records in 2,000 chunks allocates 6.9 MB instead of 136.7 MB (BenchmarkSortManyChunks). Multi-pass merging for very high chunk counts is still out of scope. Co-Authored-By: Claude Opus 5.5 --- regression_test.go | 23 ++++++++++ tempfile/mockfile.go | 18 ++++---- tempfile/regression_test.go | 91 ++++++++++++++++++++++++++++++++----- tempfile/tempfile.go | 59 ++++++++++++++++-------- 4 files changed, 153 insertions(+), 38 deletions(-) diff --git a/regression_test.go b/regression_test.go index b933ff7..73c3fbd 100644 --- a/regression_test.go +++ b/regression_test.go @@ -792,3 +792,26 @@ func TestTempFileHasOneSectionPerChunk(t *testing.T) { }) } } + +// 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) + } + } +} diff --git a/tempfile/mockfile.go b/tempfile/mockfile.go index fce031b..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. @@ -104,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 } @@ -127,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 74108a9..666cd59 100644 --- a/tempfile/regression_test.go +++ b/tempfile/regression_test.go @@ -10,6 +10,18 @@ import ( "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 @@ -153,17 +165,7 @@ func BenchmarkWriteAndSave(b *testing.B) { // 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) { - writers := 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) }, - } - for name, newWriter := range writers { + for name, newWriter := range testWriters { t.Run(name, func(t *testing.T) { for _, tc := range []struct { name string @@ -210,3 +212,70 @@ func TestSaveAddsNoEmptySection(t *testing.T) { }) } } + +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/tempfile.go b/tempfile/tempfile.go index 994edb9..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()) @@ -57,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. @@ -246,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 } @@ -267,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 } @@ -305,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] } From 9ac54e9b7f60def986b7f5eae1262bdd8d38c855 Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:13:59 -0700 Subject: [PATCH 10/11] Replace Ordered's gob codec with a binary codec Ordered created a new gob encoder and decoder for every record, which made it about 2x slower than Generic and allocated 1.74 GB per 1M ints. Encode cmp.Ordered values directly, choosing the codec once from the underlying kind of T, so named types like `type ID int64` work too: integers as varints, floats as their exact IEEE 754 bits (NaN and -0 round-trip), and strings as their bytes. Decoding rejects truncated, oversized and out-of-range data. The codecs access T through a pointer to its underlying type, which the kind check makes safe. Sorting 1M ints with Ordered drops from 748 ms and 1.74 GB to 367 ms and 46.7 MB (BenchmarkOrderedInts), the same as Generic with a varint codec. Add round-trip tests for every kind, named types, NaN and -0, invalid input, and sorts through disk. Co-Authored-By: Claude Opus 5.5 --- ordered_codec_test.go | 150 +++++++++++++++++++++++++++++++++++ regression_test.go | 54 +++++++++++++ sort_ordered.go | 180 ++++++++++++++++++++++++++++++------------ 3 files changed, 332 insertions(+), 52 deletions(-) create mode 100644 ordered_codec_test.go 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 73c3fbd..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" @@ -815,3 +817,55 @@ func BenchmarkSortManyChunks(b *testing.B) { } } } + +// 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_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 } From 753530ff145951dc09de8e21d16e61bc8abdf86f Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 19:14:21 -0700 Subject: [PATCH 11/11] Make three error-scenario tests reach the disk path TestLargeDataElements, TestMixedTypeComparison and TestOrderedSerializationError sorted 2 to 4 records with the default 1M ChunkSize, so everything stayed in one in-memory chunk and the codecs were never called. Use one record per chunk so the 1 MB elements, both mixed types and the Ordered codec go through disk. Co-Authored-By: Claude Opus 5.5 --- error_scenarios_test.go | 8 ++++++-- generic_error_test.go | 9 +++++---- 2 files changed, 11 insertions(+), 6 deletions(-) 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