From 9719edbb8cfbe9bd7d697beb1ad844ab7ef55adb Mon Sep 17 00:00:00 2001 From: Ian Foster Date: Tue, 29 Sep 2026 18:22:33 -0700 Subject: [PATCH] Fix Tier 1 correctness bugs in the sort pipeline - Run the build, sort and save stages in one errgroup, so a save error (toBytes failure, full disk) no longer deadlocks Sort() and a sort error no longer leaks the save goroutine and its temp file. - Report a read error on a chunk's first record instead of silently dropping the chunk, and treat a length header without its payload as io.ErrUnexpectedEOF instead of a clean end of chunk. - Return the recovered error from the legacy FromBytes wrapper (named results) instead of (nil, nil). - Convert panics in compareFunc, fromBytes and toBytes during the save and merge stages into ComparisonError, DeserializationError and SerializationError instead of crashing the process. - Close the temp reader before closing the result channels, fixing the "send on closed channel" panic when Close fails. - Create the temp file lazily, once a second chunk exists, and close it on every path. Creation errors are reported on the error channel, so the constructors never return a nil sorter. - Release the .extsort_ directory reference when the reader closes. - Temp dir selection: return an explicit TempFilesDir unchanged so New reports it when unusable, and skip default candidates that are missing or not writable (no /var/tmp, read-only root) instead of failing. Add regression tests for each bug, and make TestDeserializationError, TestNilInputs and TestComparisonFunctionPanic reach FromBytes and the merge and require an error. Co-Authored-By: Claude Opus 5.5 --- error_scenarios_test.go | 154 ++++----- regression_test.go | 630 ++++++++++++++++++++++++++++++++++++ sort_generic.go | 300 +++++++++++------ sort_ordered.go | 6 - sort_sorttype_legacy.go | 13 +- sort_strings.go | 6 - tempfile/regression_test.go | 121 +++++++ tempfile/tempdir.go | 52 ++- tempfile/tempfile.go | 21 +- tempfile_error_test.go | 18 +- 10 files changed, 1092 insertions(+), 229 deletions(-) create mode 100644 regression_test.go create mode 100644 tempfile/regression_test.go diff --git a/error_scenarios_test.go b/error_scenarios_test.go index cc5460d..4072f62 100644 --- a/error_scenarios_test.go +++ b/error_scenarios_test.go @@ -76,14 +76,9 @@ func TestDeserializationError(t *testing.T) { panic("deserialization failed") // Simulate critical failure } - sort, outChan, errChan := extsort.New(inputChan, failingFromBytes, KeyLessThan, nil) - - // This should panic or fail gracefully during merge phase - defer func() { - if r := recover(); r != nil { - t.Logf("Expected panic during deserialization: %v", r) - } - }() + // One record per chunk, so the records are written to disk and read back with FromBytes + config := &extsort.Config{ChunkSize: 1} + sort, outChan, errChan := extsort.New(inputChan, failingFromBytes, KeyLessThan, config) sort.Sort(context.Background()) @@ -92,51 +87,54 @@ func TestDeserializationError(t *testing.T) { // Consume output } - // If we get here, check for error - if err := <-errChan; err != nil { - t.Logf("Got expected deserialization error: %v", err) - // Verify it's our specific error type - var deserErr *extsort.DeserializationError - if !errors.As(err, &deserErr) { - t.Errorf("Expected DeserializationError, got: %T", err) - } + // The panic must be reported as an error + err := <-errChan + if err == nil { + t.Fatal("Expected deserialization error, got nil") + } + t.Logf("Got expected deserialization error: %v", err) + // Verify it's our specific error type + var deserErr *extsort.DeserializationError + if !errors.As(err, &deserErr) { + t.Errorf("Expected DeserializationError, got: %T", err) } } // TestNilInputs tests behavior with nil function parameters func TestNilInputs(t *testing.T) { - inputChan := make(chan extsort.SortType, 1) - inputChan <- val{Key: 1, Order: 1} - close(inputChan) - - // Test with nil fromBytes function - func() { - defer func() { - if r := recover(); r != nil { - t.Logf("Expected panic with nil fromBytes: %v", r) - } - }() - - sort, _, _ := extsort.New(inputChan, nil, KeyLessThan, nil) + newInput := func() chan extsort.SortType { + inputChan := make(chan extsort.SortType, 4) + for i := 4; i > 0; i-- { + inputChan <- val{Key: i, Order: i} + } + close(inputChan) + return inputChan + } + // Two records per chunk, so chunks are sorted, written to disk and merged + config := func() *extsort.Config { return &extsort.Config{ChunkSize: 2} } + sortErr := func(sort *extsort.SortTypeSorter, outChan <-chan extsort.SortType, errChan <-chan error) error { sort.Sort(context.Background()) - }() - - // Recreate input for next test - inputChan2 := make(chan extsort.SortType, 1) - inputChan2 <- val{Key: 1, Order: 1} - close(inputChan2) - - // Test with nil comparison function - func() { - defer func() { - if r := recover(); r != nil { - t.Logf("Expected panic with nil lessFunc: %v", r) - } - }() + for range outChan { + // Consume output + } + return <-errChan + } - sort, _, _ := extsort.New(inputChan2, fromBytesForTest, nil, nil) - sort.Sort(context.Background()) - }() + // Test with nil fromBytes function: calling it during the merge must be reported + sort, outChan, errChan := extsort.New(newInput(), nil, KeyLessThan, config()) + err := sortErr(sort, outChan, errChan) + var deserErr *extsort.DeserializationError + if !errors.As(err, &deserErr) { + t.Errorf("nil fromBytes: expected DeserializationError, got: %v", err) + } + + // Test with nil comparison function: calling it while sorting chunks must be reported + sort, outChan, errChan = extsort.New(newInput(), fromBytesForTest, nil, config()) + err = sortErr(sort, outChan, errChan) + var compErr *extsort.ComparisonError + if !errors.As(err, &compErr) { + t.Errorf("nil lessFunc: expected ComparisonError, got: %v", err) + } } // TestLargeDataElements tests with unusually large individual elements @@ -227,41 +225,47 @@ func largeLessThan(a, b extsort.SortType) bool { // TestComparisonFunctionPanic tests handling of panics in comparison function func TestComparisonFunctionPanic(t *testing.T) { - inputChan := make(chan extsort.SortType, 3) - inputChan <- val{Key: 1, Order: 1} - inputChan <- val{Key: 2, Order: 2} - inputChan <- val{Key: 3, Order: 3} - close(inputChan) - // Comparison function that panics panicLessFunc := func(a, b extsort.SortType) bool { panic("comparison function panic") } - sort, outChan, errChan := extsort.New(inputChan, fromBytesForTest, panicLessFunc, nil) - - // Should handle the panic gracefully - defer func() { - if r := recover(); r != nil { - t.Logf("Caught expected panic from comparison function: %v", r) - } - }() - - sort.Sort(context.Background()) - - // Drain channels - for range outChan { - // Consume output - } + for _, tc := range []struct { + name string + chunkSize int + }{ + {"while sorting a chunk", 3}, // all records in one chunk + {"while merging chunks", 1}, // a one-record chunk needs no comparison until the merge + } { + t.Run(tc.name, func(t *testing.T) { + inputChan := make(chan extsort.SortType, 3) + inputChan <- val{Key: 1, Order: 1} + inputChan <- val{Key: 2, Order: 2} + inputChan <- val{Key: 3, Order: 3} + close(inputChan) + + config := &extsort.Config{ChunkSize: tc.chunkSize} + sort, outChan, errChan := extsort.New(inputChan, fromBytesForTest, panicLessFunc, config) + + sort.Sort(context.Background()) + + // Drain channels + for range outChan { + // Consume output + } - // Check for error - if err := <-errChan; err != nil { - t.Logf("Got expected error from panicking comparison: %v", err) - // Verify it's our specific error type - var compErr *extsort.ComparisonError - if !errors.As(err, &compErr) { - t.Errorf("Expected ComparisonError, got: %T", err) - } + // The panic must be reported as an error + err := <-errChan + if err == nil { + t.Fatal("Expected comparison error, got nil") + } + t.Logf("Got expected error from panicking comparison: %v", err) + // Verify it's our specific error type + var compErr *extsort.ComparisonError + if !errors.As(err, &compErr) { + t.Errorf("Expected ComparisonError, got: %T", err) + } + }) } } diff --git a/regression_test.go b/regression_test.go new file mode 100644 index 0000000..d791b23 --- /dev/null +++ b/regression_test.go @@ -0,0 +1,630 @@ +package extsort + +// Regression tests for correctness bugs in the sort pipeline. Several of them inject +// failing temp files through package internals, so they live in package extsort. + +import ( + "bufio" + "bytes" + "cmp" + "context" + "errors" + "io" + "os" + "path/filepath" + "runtime" + "slices" + "strconv" + "strings" + "sync" + "testing" + "testing/iotest" + "time" + + "github.com/lanrat/extsort/tempfile" +) + +func itoaBytes(i int) ([]byte, error) { return []byte(strconv.Itoa(i)), nil } +func atoiBytes(b []byte) (int, error) { return strconv.Atoi(string(b)) } + +// descendingInts sends n..1 on an unbuffered channel, then closes it. +func descendingInts(n int) chan int { + ch := make(chan int) + go func() { + defer close(ch) + for i := n; i > 0; i-- { + ch <- i + } + }() + return ch +} + +// runSort calls Sort, drains the output, then reads the error channel. +// It fails the test if that does not finish in time. +func runSort[E any](t *testing.T, ctx context.Context, s *GenericSorter[E]) ([]E, error) { + t.Helper() + type result struct { + got []E + err error + } + done := make(chan result, 1) + go func() { + s.Sort(ctx) + var got []E + for v := range s.mergeChunkChan { + got = append(got, v) + } + done <- result{got, <-s.mergeErrChan} + }() + select { + case r := <-done: + return r.got, r.err + case <-time.After(10 * time.Second): + t.Fatal("sort did not finish within 10s") + return nil, nil + } +} + +// drainWithTimeout reads out until it closes, then reads errc. +func drainWithTimeout[E any](t *testing.T, out <-chan E, errc <-chan error) ([]E, error) { + t.Helper() + var got []E + timeout := time.After(10 * time.Second) + for { + select { + case v, ok := <-out: + if !ok { + return got, <-errc + } + got = append(got, v) + case <-timeout: + t.Fatal("sort did not finish within 10s") + } + } +} + +// trackedTemp hands out in-memory temp files and counts how many the sorter +// leaves open. readErr and closeErr inject failures. +type trackedTemp struct { + readErr map[int]error // section index -> error returned when reading it + closeErr error // returned by the reader's Close + + mu sync.Mutex + created int + closed int +} + +func (tt *trackedTemp) newWriter() (tempfile.TempWriter, error) { + tt.mu.Lock() + defer tt.mu.Unlock() + tt.created++ + return &trackedWriter{TempWriter: tempfile.Mock(0), tt: tt}, nil +} + +func (tt *trackedTemp) markClosed() { + tt.mu.Lock() + defer tt.mu.Unlock() + tt.closed++ +} + +func (tt *trackedTemp) createdCount() int { + tt.mu.Lock() + defer tt.mu.Unlock() + return tt.created +} + +// open returns how many temp files were created but not closed exactly once. +func (tt *trackedTemp) open() int { + tt.mu.Lock() + defer tt.mu.Unlock() + return tt.created - tt.closed +} + +type trackedWriter struct { + tempfile.TempWriter + tt *trackedTemp +} + +func (w *trackedWriter) Close() error { + w.tt.markClosed() + return w.TempWriter.Close() +} + +func (w *trackedWriter) Save() (tempfile.TempReader, error) { + r, err := w.TempWriter.Save() + if err != nil { + return nil, err + } + return &trackedReader{TempReader: r, tt: w.tt}, nil +} + +type trackedReader struct { + tempfile.TempReader + tt *trackedTemp +} + +func (r *trackedReader) Read(i int) *bufio.Reader { + if err, ok := r.tt.readErr[i]; ok { + return bufio.NewReader(iotest.ErrReader(err)) + } + return r.TempReader.Read(i) +} + +func (r *trackedReader) Close() error { + r.tt.markClosed() + if err := r.TempReader.Close(); err != nil { + return err + } + return r.tt.closeErr +} + +// newTrackedSorter returns an int sorter whose temp files are tracked by tt. +func newTrackedSorter(input <-chan int, toBytes ToBytesGeneric[int], compare CompareGeneric[int], config *Config, tt *trackedTemp) *GenericSorter[int] { + s := newSorter(input, atoiBytes, toBytes, compare, config) + s.newTempWriter = tt.newWriter + return s +} + +// A failing save stage used to leave the sort workers blocked on saveChunkChan, +// so Sort never returned once the input spanned more than 2+2*NumWorkers chunks. +func TestSaveErrorDoesNotDeadlock(t *testing.T) { + errEncode := errors.New("encode failed") + failing := func(int) ([]byte, error) { return nil, errEncode } + for _, n := range []int{6, 7, 1000} { + t.Run(strconv.Itoa(n)+" chunks", func(t *testing.T) { + tt := &trackedTemp{} + s := newTrackedSorter(descendingInts(n), failing, cmp.Compare[int], &Config{ChunkSize: 1, NumWorkers: 2, ChanBuffSize: 16}, tt) + _, err := runSort(t, context.Background(), s) + var serErr *SerializationError + if !errors.As(err, &serErr) || !errors.Is(err, errEncode) { + t.Fatalf("got error %v, want a SerializationError wrapping %q", err, errEncode) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + }) + } +} + +// An error in the sort stage used to make Sort return without closing saveChunkChan, +// leaking the save goroutine and its temp file. +func TestSortErrorStopsSaveStage(t *testing.T) { + tt := &trackedTemp{} + created := make(chan struct{}) + var once sync.Once + // Chunks are [6 5] [4 3] [2 1]. Sorting [2 1] fails, but only after the + // first two chunks reached the save stage and it created the temp file. + compare := func(a, b int) int { + if a == 1 || b == 1 { + select { + case <-created: + case <-time.After(5 * time.Second): + } + panic("comparison failed") + } + return cmp.Compare(a, b) + } + s := newTrackedSorter(descendingInts(6), itoaBytes, compare, &Config{ChunkSize: 2, NumWorkers: 2}, tt) + s.newTempWriter = func() (tempfile.TempWriter, error) { + defer once.Do(func() { close(created) }) + return tt.newWriter() + } + + _, err := runSort(t, context.Background(), s) + var cmpErr *ComparisonError + if !errors.As(err, &cmpErr) { + t.Fatalf("got error %v, want a ComparisonError", err) + } + if tt.createdCount() != 1 { + t.Fatalf("temp file created %d times, want 1", tt.createdCount()) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + buf := make([]byte, 1<<20) + for deadline := time.Now().Add(2 * time.Second); ; { + stacks := string(buf[:runtime.Stack(buf, true)]) + if !strings.Contains(stacks, "saveChunksOptimized") { + break + } + if time.Now().After(deadline) { + t.Fatal("save goroutine still running after Sort returned") + } + time.Sleep(10 * time.Millisecond) + } +} + +// A read error on a chunk's first record used to be treated as an empty chunk, +// silently dropping that chunk's records. +func TestMergeReportsChunkReadErrors(t *testing.T) { + errDisk := errors.New("disk read error") + for _, tc := range []struct { + name string + 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}, + {"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) + 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) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + }) + } +} + +// getNext used to treat a length header without its payload as a clean end of chunk. +func TestGetNextTruncatedRecord(t *testing.T) { + for _, tc := range []struct { + name string + data []byte + wantOK bool + want error + }{ + {"end of chunk", nil, false, nil}, + {"complete record", []byte{1, '7'}, true, nil}, + {"header without payload", []byte{5}, false, io.ErrUnexpectedEOF}, + {"partial payload", []byte{5, '1', '2'}, false, io.ErrUnexpectedEOF}, + {"partial header", []byte{0x80}, false, io.ErrUnexpectedEOF}, + } { + t.Run(tc.name, func(t *testing.T) { + m := &mergeFile[int]{fromBytes: atoiBytes, reader: bufio.NewReader(bytes.NewReader(tc.data))} + _, ok, err := m.getNext() + if ok != tc.wantOK || !errors.Is(err, tc.want) { + t.Fatalf("getNext() = ok %v, err %v; want ok %v, err %v", ok, err, tc.wantOK, tc.want) + } + if ok && m.nextRec != 7 { + t.Errorf("getNext() read %d, want 7", m.nextRec) + } + }) + } +} + +type legacyInt int + +func (v legacyInt) ToBytes() []byte { return []byte(strconv.Itoa(int(v))) } + +// The legacy FromBytes wrapper recovered panics into a local variable that was +// never returned, so a panicking FromBytes produced (nil, nil) and nil records. +func TestLegacyFromBytesPanicIsReported(t *testing.T) { + fromBytes := makeSortTypeFromBytes(func([]byte) SortType { panic("corrupt record") }) + var deserErr *DeserializationError + if v, err := fromBytes([]byte("x")); v != nil || !errors.As(err, &deserErr) { + t.Fatalf("got (%v, %v), want (nil, DeserializationError)", v, err) + } + + in := make(chan SortType, 10) + for i := 10; i > 0; i-- { + in <- legacyInt(i) + } + close(in) + panicOn5 := func(b []byte) SortType { + n, _ := strconv.Atoi(string(b)) + if n == 5 { + panic("corrupt record") + } + return legacyInt(n) + } + less := func(a, b SortType) bool { // tolerates nil so the old behavior shows up as output + ai, aok := a.(legacyInt) + bi, bok := b.(legacyInt) + if !aok || !bok { + return !aok && bok + } + return ai < bi + } + sorter, out, errc := NewMock(in, panicOn5, less, &Config{ChunkSize: 2}, 0) + sorter.Sort(context.Background()) + got, err := drainWithTimeout(t, out, errc) + if !errors.As(err, &deserErr) { + t.Errorf("got error %v, want a DeserializationError", err) + } + if slices.Contains(got, nil) { + t.Errorf("output contains nil records: %v", got) + } +} + +// Panics in user callbacks during the save and merge stages used to crash the +// process. Each must be reported on the error channel instead. +func TestCallbackPanicsBecomeErrors(t *testing.T) { + panicCompare := func(a, b int) int { panic("compare failed") } + panicFromBytes := func([]byte) (int, error) { panic("decode failed") } + panicToBytes := func(int) ([]byte, error) { panic("encode failed") } + var cmpErr *ComparisonError + var deserErr *DeserializationError + var serErr *SerializationError + for _, tc := range []struct { + name string + n int + workers int + fromBytes FromBytesGeneric[int] + toBytes ToBytesGeneric[int] + compare CompareGeneric[int] + target any + }{ + // With ChunkSize 1, sorting a chunk never calls compare: the first call is in the merge. + {"compare in single-threaded merge", 3, 4, atoiBytes, itoaBytes, panicCompare, &cmpErr}, + {"compare in parallel merge", 20, 2, atoiBytes, itoaBytes, panicCompare, &cmpErr}, + {"fromBytes in merge", 5, 2, panicFromBytes, itoaBytes, cmp.Compare[int], &deserErr}, + {"toBytes in save", 50, 2, atoiBytes, panicToBytes, cmp.Compare[int], &serErr}, + } { + t.Run(tc.name, func(t *testing.T) { + sorter, out, errc := MockGeneric(descendingInts(tc.n), tc.fromBytes, tc.toBytes, tc.compare, &Config{ChunkSize: 1, NumWorkers: tc.workers}, 0) + sorter.Sort(context.Background()) + if _, err := drainWithTimeout(t, out, errc); !errors.As(err, tc.target) { + t.Fatalf("got error %v, want %T", err, tc.target) + } + }) + } + + t.Run("compare in final merge", func(t *testing.T) { + s := newSorter(nil, atoiBytes, itoaBytes, panicCompare, nil) + a, b := make(chan int, 1), make(chan int, 1) + a <- 1 + b <- 2 + close(a) + close(b) + if err := s.finalMergeSimple(context.Background(), []chan int{a, b}); !errors.As(err, &cmpErr) { + t.Fatalf("got error %v, want a ComparisonError", err) + } + }) +} + +// A failing tempReader.Close used to panic with "send on closed channel", because +// 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 + 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) + 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) + } + }) + } +} + +// The temp file used to be created in the constructor and never closed on the empty, +// single-chunk and error paths. It is now only created for multi-chunk sorts, and +// released on every path. +func TestTempFileLifecycle(t *testing.T) { + errEncode := errors.New("encode failed") + failing := func(int) ([]byte, error) { return nil, errEncode } + for _, tc := range []struct { + name string + n int + toBytes ToBytesGeneric[int] + wantCreated int + wantErr error + }{ + {"empty input", 0, itoaBytes, 0, nil}, + {"single chunk", 10, itoaBytes, 0, nil}, + {"multiple chunks", 100, itoaBytes, 1, nil}, + {"save error", 100, failing, 1, errEncode}, + } { + t.Run(tc.name, func(t *testing.T) { + tt := &trackedTemp{} + s := newTrackedSorter(descendingInts(tc.n), tc.toBytes, cmp.Compare[int], &Config{ChunkSize: 10}, tt) + got, err := runSort(t, context.Background(), s) + if !errors.Is(err, tc.wantErr) { + t.Fatalf("got error %v, want %v", err, tc.wantErr) + } + if tc.wantErr == nil && (len(got) != tc.n || !slices.IsSorted(got)) { + t.Errorf("got %d records (sorted: %v), want %d sorted", len(got), slices.IsSorted(got), tc.n) + } + if tt.createdCount() != tc.wantCreated { + t.Errorf("temp file created %d times, want %d", tt.createdCount(), tc.wantCreated) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + }) + } + + t.Run("cancelled while saving", func(t *testing.T) { + tt := &trackedTemp{} + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + in := make(chan int) + go func() { // endless input, stopped by the cancellation + defer close(in) + for i := 0; ; i++ { + select { + case in <- i: + case <-ctx.Done(): + return + } + } + }() + s := newTrackedSorter(in, itoaBytes, cmp.Compare[int], &Config{ChunkSize: 10}, tt) + var once sync.Once + s.newTempWriter = func() (tempfile.TempWriter, error) { + defer once.Do(cancel) + return tt.newWriter() + } + if _, err := runSort(t, ctx, s); !errors.Is(err, context.Canceled) { + t.Fatalf("got error %v, want %v", err, context.Canceled) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + }) + + t.Run("cancelled while merging", func(t *testing.T) { + tt := &trackedTemp{} + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + s := newTrackedSorter(descendingInts(100), itoaBytes, cmp.Compare[int], &Config{ChunkSize: 10}, tt) + s.Sort(ctx) + for i := 0; i < 5; i++ { + <-s.mergeChunkChan + } + cancel() + if _, err := drainWithTimeout(t, s.mergeChunkChan, s.mergeErrChan); !errors.Is(err, context.Canceled) { + t.Fatalf("got error %v, want %v", err, context.Canceled) + } + if open := tt.open(); open != 0 { + t.Errorf("%d temp file(s) left open", open) + } + }) +} + +// Temp files on disk must be closed, and on Windows removed, on success and on error. +func TestTempFilesReleasedOnDisk(t *testing.T) { + errEncode := errors.New("encode failed") + failOn50 := func(i int) ([]byte, error) { + if i == 50 { + return nil, errEncode + } + return itoaBytes(i) + } + for _, tc := range []struct { + name string + toBytes ToBytesGeneric[int] + wantErr error + }{ + {"success", itoaBytes, nil}, + {"save error", failOn50, errEncode}, + } { + t.Run(tc.name, func(t *testing.T) { + dir := t.TempDir() + s, _, _ := Generic(descendingInts(100), atoiBytes, tc.toBytes, cmp.Compare[int], &Config{ChunkSize: 10, TempFilesDir: dir}) + if _, err := runSort(t, context.Background(), s); !errors.Is(err, tc.wantErr) { + t.Fatalf("got error %v, want %v", err, tc.wantErr) + } + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + for _, e := range entries { + t.Errorf("temp file left on disk: %s", e.Name()) + } + if n := openFilesUnder(dir); n > 0 { + t.Errorf("%d temp file(s) still open", n) + } + }) + } +} + +// openFilesUnder counts this process's open files under dir. It needs /proc/self/fd +// (Linux) and returns 0 where that is unavailable. +func openFilesUnder(dir string) int { + fds, err := os.ReadDir("/proc/self/fd") + if err != nil { + return 0 + } + if resolved, err := filepath.EvalSymlinks(dir); err == nil { + dir = resolved + } + n := 0 + for _, fd := range fds { + target, err := os.Readlink(filepath.Join("/proc/self/fd", fd.Name())) + if err == nil && strings.HasPrefix(target, dir+string(filepath.Separator)) { + n++ + } + } + return n +} + +// When the temp file could not be created, the constructors used to return a nil +// sorter, so the README pattern `go sorter.Sort(ctx)` crashed instead of reporting +// the error on the error channel. +func TestTempFileCreationErrorIsReportedOnErrChan(t *testing.T) { + notADir := filepath.Join(t.TempDir(), "file") + if err := os.WriteFile(notADir, nil, 0o600); err != nil { + t.Fatal(err) + } + config := func() *Config { return &Config{ChunkSize: 2, TempFilesDir: notADir} } + strs := func(n int) chan string { + ch := make(chan string) + go func() { + defer close(ch) + for i := n; i > 0; i-- { + ch <- strconv.Itoa(i) + } + }() + return ch + } + legacy := func(n int) chan SortType { + ch := make(chan SortType) + go func() { + defer close(ch) + for i := n; i > 0; i-- { + ch <- legacyInt(i) + } + }() + return ch + } + legacyLess := func(a, b SortType) bool { return a.(legacyInt) < b.(legacyInt) } + legacyFromBytes := func(b []byte) SortType { + n, _ := strconv.Atoi(string(b)) + return legacyInt(n) + } + wantErr := func(t *testing.T, err error) { + t.Helper() + if err == nil { + t.Fatal("got nil error, want the temp file creation error") + } + } + + t.Run("Generic", func(t *testing.T) { + sorter, out, errc := Generic(descendingInts(10), atoiBytes, itoaBytes, cmp.Compare[int], config()) + if sorter == nil { + t.Fatal("constructor returned a nil sorter") + } + go sorter.Sort(context.Background()) + _, err := drainWithTimeout(t, out, errc) + wantErr(t, err) + }) + t.Run("Ordered", func(t *testing.T) { + sorter, out, errc := Ordered(descendingInts(10), config()) + if sorter == nil { + t.Fatal("constructor returned a nil sorter") + } + go sorter.Sort(context.Background()) + _, err := drainWithTimeout(t, out, errc) + wantErr(t, err) + }) + t.Run("Strings", func(t *testing.T) { + sorter, out, errc := Strings(strs(10), config()) + if sorter == nil { + t.Fatal("constructor returned a nil sorter") + } + go sorter.Sort(context.Background()) + _, err := drainWithTimeout(t, out, errc) + wantErr(t, err) + }) + t.Run("New", func(t *testing.T) { + sorter, out, errc := New(legacy(10), legacyFromBytes, legacyLess, config()) + if sorter == nil { + t.Fatal("constructor returned a nil sorter") + } + go sorter.Sort(context.Background()) + _, err := drainWithTimeout(t, out, errc) + wantErr(t, err) + }) + t.Run("single chunk needs no temp file", func(t *testing.T) { + sorter, out, errc := Ordered(descendingInts(2), config()) + go sorter.Sort(context.Background()) + got, err := drainWithTimeout(t, out, errc) + if err != nil || !slices.Equal(got, []int{1, 2}) { + t.Fatalf("got %v, %v; want [1 2], nil", got, err) + } + }) +} diff --git a/sort_generic.go b/sort_generic.go index 7e670e8..0a895be 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -61,9 +61,9 @@ type memoryPools struct { // and employs memory pools to reduce garbage collection pressure during operation. type GenericSorter[E any] struct { config Config - buildSortCtx context.Context - saveCtx context.Context + sortCtx context.Context // shared by the build, sort and save stages mergeErrChan chan error + newTempWriter func() (tempfile.TempWriter, error) // called once a second chunk exists tempWriter tempfile.TempWriter tempReader tempfile.TempReader input <-chan E @@ -155,19 +155,21 @@ func (s *GenericSorter[E]) initMemoryPools() *memoryPools { // // Call Sort() on the returned sorter to begin the sorting process. // Results are delivered via the output channel, errors via the error channel. -// On error or interruption, temporary files may remain in config.TempFilesDir. +// The temporary file is only created once the input spans more than one chunk; +// failure to create it is reported on the error channel. The file is closed and +// removed when the sort completes, fails or is cancelled. // // IMPORTANT: The input channel must be closed to signal completion. The Sort() method // will block until the input channel is closed. Failure to close it will cause a deadlock. func Generic[E any](input <-chan E, fromBytes FromBytesGeneric[E], toBytes ToBytesGeneric[E], compareFunc CompareGeneric[E], config *Config) (*GenericSorter[E], <-chan E, <-chan error) { - var err error s := newSorter(input, fromBytes, toBytes, compareFunc, config) - s.tempWriter, err = tempfile.New(s.config.TempFilesDir, true) - if err != nil { - s.mergeErrChan <- err - close(s.mergeErrChan) - close(s.mergeChunkChan) - return nil, s.mergeChunkChan, s.mergeErrChan + dir := s.config.TempFilesDir + s.newTempWriter = func() (tempfile.TempWriter, error) { + w, err := tempfile.New(dir, true) + if err != nil { + return nil, err + } + return w, nil } return s, s.mergeChunkChan, s.mergeErrChan } @@ -178,7 +180,9 @@ func Generic[E any](input <-chan E, fromBytes FromBytesGeneric[E], toBytes ToByt // All other behavior is identical to Generic(). 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.tempWriter = tempfile.Mock(n) + s.newTempWriter = func() (tempfile.TempWriter, error) { + return tempfile.Mock(n), nil + } return s, s.mergeChunkChan, s.mergeErrChan } @@ -194,35 +198,38 @@ func MockGeneric[E any](input <-chan E, fromBytes FromBytesGeneric[E], toBytes T // Merge uses the same context and runs in a goroutine after Sort returns(). // for example, if calling sort in an errGroup, you must pass the group's parent context into sort. func (s *GenericSorter[E]) Sort(ctx context.Context) { - var buildSortErrGroup, saveErrGroup *errgroup.Group - buildSortErrGroup, s.buildSortCtx = errgroup.WithContext(ctx) - saveErrGroup, s.saveCtx = errgroup.WithContext(ctx) + // One group for all stages: an error in any stage cancels the others, + // so a failed save cannot leave the sort workers blocked on saveChunkChan. + group, groupCtx := errgroup.WithContext(ctx) + s.sortCtx = groupCtx //start creating chunks - buildSortErrGroup.Go(s.buildChunks) + group.Go(s.buildChunks) // sort chunks + var sorters sync.WaitGroup + sorters.Add(s.config.NumWorkers) for i := 0; i < s.config.NumWorkers; i++ { - buildSortErrGroup.Go(s.sortChunks) - } + group.Go(func() error { + defer sorters.Done() + return s.sortChunks() + }) + } + + // Close saveChunkChan to signal end of chunks once every sort worker has + // exited, successfully or not, so the save worker always returns. + group.Go(func() error { + sorters.Wait() + close(s.saveChunkChan) + return nil + }) // Start the save worker that will handle single-chunk optimization - saveErrGroup.Go(s.saveChunksOptimized) - - err := buildSortErrGroup.Wait() - if err != nil { - s.mergeErrChan <- err - close(s.mergeErrChan) - close(s.mergeChunkChan) - return - } - - // Close saveChunkChan to signal end of chunks - close(s.saveChunkChan) + group.Go(s.saveChunksOptimized) - // Wait for save worker to complete - err = saveErrGroup.Wait() + err := group.Wait() if err != nil { + s.closeTempFiles() s.mergeErrChan <- err close(s.mergeErrChan) close(s.mergeChunkChan) @@ -241,6 +248,19 @@ func (s *GenericSorter[E]) Sort(ctx context.Context) { go s.mergeNChunks(ctx) } +// closeTempFiles releases the temp file on paths that never reach the merge. +// The merge closes the reader itself. +func (s *GenericSorter[E]) closeTempFiles() { + if s.tempReader != nil { + _ = s.tempReader.Close() + s.tempReader = nil + } + if s.tempWriter != nil { + _ = s.tempWriter.Close() + s.tempWriter = nil + } +} + // buildChunks reads data from the input chan to builds chunks and pushes them to chunkChan func (s *GenericSorter[E]) buildChunks() error { defer close(s.chunkChan) // if this is not called on error, causes a deadlock @@ -254,9 +274,9 @@ func (s *GenericSorter[E]) buildChunks() error { break } c.data = append(c.data, rec) - case <-s.buildSortCtx.Done(): + case <-s.sortCtx.Done(): s.putChunk(c) // Return unused chunk to pool - return s.buildSortCtx.Err() + return s.sortCtx.Err() } } if len(c.data) == 0 { @@ -268,9 +288,9 @@ func (s *GenericSorter[E]) buildChunks() error { select { // chunk is now full case s.chunkChan <- c: - case <-s.buildSortCtx.Done(): + case <-s.sortCtx.Done(): s.putChunk(c) // Return unused chunk to pool - return s.buildSortCtx.Err() + return s.sortCtx.Err() } } @@ -310,18 +330,18 @@ func (s *GenericSorter[E]) sortChunks() error { // Sort completed successfully, proceed to save select { case s.saveChunkChan <- b: - case <-s.buildSortCtx.Done(): - return s.buildSortCtx.Err() + case <-s.sortCtx.Done(): + return s.sortCtx.Err() } - case <-s.buildSortCtx.Done(): + case <-s.sortCtx.Done(): // Context cancelled while sorting - abandon this chunk - return s.buildSortCtx.Err() + return s.sortCtx.Err() } } else { return nil } - case <-s.buildSortCtx.Done(): - return s.buildSortCtx.Err() + case <-s.sortCtx.Done(): + return s.sortCtx.Err() } } } @@ -368,8 +388,8 @@ func (s *GenericSorter[E]) saveChunksOptimized() error { // Channel closed, no chunks at all return nil } - case <-s.saveCtx.Done(): - return s.saveCtx.Err() + case <-s.sortCtx.Done(): + return s.sortCtx.Err() } // Try to get a second chunk with context checking @@ -381,12 +401,20 @@ func (s *GenericSorter[E]) saveChunksOptimized() error { s.singleChunk = firstChunk return nil } - case <-s.saveCtx.Done(): + case <-s.sortCtx.Done(): s.putChunk(firstChunk) // Return to pool before exiting - return s.saveCtx.Err() + return s.sortCtx.Err() } - // We have at least 2 chunks - use multi-chunk path + // We have at least 2 chunks - use multi-chunk path, which needs the temp file + tempWriter, err := s.newTempWriter() + if err != nil { + s.putChunk(firstChunk) + s.putChunk(secondChunk) + return err + } + s.tempWriter = tempWriter + // Save the first chunk if err := s.saveChunk(firstChunk); err != nil { s.putChunk(secondChunk) // Return to pool @@ -403,17 +431,25 @@ func (s *GenericSorter[E]) saveChunksOptimized() error { select { case chunk, ok := <-s.saveChunkChan: if !ok { - // Channel closed, we're done + // Channel closed, we're done, unless it closed because another stage failed + if err := s.sortCtx.Err(); err != nil { + return err + } // Finalize the temp writer and save it for reading - var err error - s.tempReader, err = s.tempWriter.Save() - return err + tempReader, err := s.tempWriter.Save() + if err != nil { + return err + } + // The reader now owns the file + s.tempReader = tempReader + s.tempWriter = nil + return nil } if err := s.saveChunk(chunk); err != nil { return err } - case <-s.saveCtx.Done(): - return s.saveCtx.Err() + case <-s.sortCtx.Done(): + return s.sortCtx.Err() } } } @@ -426,10 +462,10 @@ func (s *GenericSorter[E]) saveChunk(b *genericChunk[E]) error { for _, d := range b.data { // binary encoding for size - raw, err := s.toBytes(d) + raw, err := s.encode(d) if err != nil { s.putChunk(b) // Return chunk to pool on error - return NewSerializationError(err, "saveChunk") + return err } n := binary.PutUvarint(scratch, uint64(len(raw))) _, err = s.tempWriter.Write(scratch[:n]) @@ -454,46 +490,62 @@ func (s *GenericSorter[E]) saveChunk(b *genericChunk[E]) error { return nil } +// encode serializes one record with toBytes, converting both a returned error +// and a panic into a SerializationError. +func (s *GenericSorter[E]) encode(d E) (raw []byte, err error) { + defer func() { + if r := recover(); r != nil { + raw = nil + err = NewSerializationError(r, "saveChunk") + } + }() + raw, err = s.toBytes(d) + if err != nil { + return nil, NewSerializationError(err, "saveChunk") + } + return raw, nil +} + // mergeNChunks runs asynchronously in the background feeding data to getNext // sends errors to s.mergeErrorChan. Uses parallel merging for better performance. func (s *GenericSorter[E]) mergeNChunks(ctx context.Context) { + // Deferred calls run last-in first-out: the error channel closes before the output channel. defer close(s.mergeChunkChan) - defer func() { - if s.tempReader != nil { - err := s.tempReader.Close() - if err != nil { - // Try to send error, but don't panic if channel is closed - select { - case s.mergeErrChan <- err: - default: - } - } - } - }() - // Always ensure error channel is closed defer close(s.mergeErrChan) if s.tempReader == nil { return } - numChunks := s.tempReader.Size() - if numChunks == 0 { - return + var err error + if s.tempReader.Size() <= s.config.NumWorkers { + // For small number of chunks, use single-threaded merge + err = s.mergeNChunksSingleThreaded(ctx) + } else { + // Use parallel merging for many chunks + err = s.mergeNChunksParallel(ctx) } - // For small number of chunks, use single-threaded merge - if numChunks <= s.config.NumWorkers { - s.mergeNChunksSingleThreaded(ctx) - return + // Release the temp file before signalling completion + if closeErr := s.tempReader.Close(); closeErr != nil && err == nil { + err = NewDiskError(closeErr, "close temp file", "") } + s.tempReader = nil - // Use parallel merging for many chunks - s.mergeNChunksParallel(ctx) + if err != nil { + s.mergeErrChan <- err + } } // mergeNChunksSingleThreaded is the original single-threaded implementation -func (s *GenericSorter[E]) mergeNChunksSingleThreaded(ctx context.Context) { +func (s *GenericSorter[E]) mergeNChunksSingleThreaded(ctx context.Context) (err error) { + // A panicking compareFunc must not crash the process from this goroutine + defer func() { + if r := recover(); r != nil { + err = NewComparisonError(r, "mergeNChunksSingleThreaded") + } + }() + pq := queue.NewPriorityQueue(func(a, b *mergeFile[E]) int { return s.compareFunc(a.nextRec, b.nextRec) }) @@ -504,12 +556,11 @@ func (s *GenericSorter[E]) mergeNChunksSingleThreaded(ctx context.Context) { reader: s.tempReader.Read(i), } _, ok, err := merge.getNext() // start the merge by preloading the values - if err == io.EOF || !ok { - continue - } if err != nil { - s.mergeErrChan <- err - return + return err + } + if !ok { + continue // empty chunk } pq.Push(merge) } @@ -518,8 +569,7 @@ func (s *GenericSorter[E]) mergeNChunksSingleThreaded(ctx context.Context) { merge := pq.Peek() rec, more, err := merge.getNext() if err != nil { - s.mergeErrChan <- err - return + return err } if more { pq.PeekUpdate() @@ -530,14 +580,14 @@ func (s *GenericSorter[E]) mergeNChunksSingleThreaded(ctx context.Context) { select { case s.mergeChunkChan <- rec: case <-ctx.Done(): - s.mergeErrChan <- ctx.Err() - return + return ctx.Err() } } + return nil } // mergeNChunksParallel implements parallel k-way merging with robust cancellation -func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) { +func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) error { numChunks := s.tempReader.Size() numWorkers := s.config.NumWorkers @@ -605,7 +655,10 @@ func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) { finalMergeWg.Add(1) go func() { defer finalMergeWg.Done() - s.finalMergeSimple(mergeCtx, intermediateChans[:workersStarted]) + if err := s.finalMergeSimple(mergeCtx, intermediateChans[:workersStarted]); err != nil { + errChan <- err + mergeCancel() // Stop the workers, which may be blocked sending to the final merge + } }() // Wait for all workers to complete @@ -619,16 +672,22 @@ func (s *GenericSorter[E]) mergeNChunksParallel(ctx context.Context) { // Wait for error collector to finish processing all errors errorCollectorWg.Wait() - // Send any collected error (now safe to read mergeErr) + // Return any collected error (now safe to read mergeErr) if mergeErr != nil { - s.mergeErrChan <- mergeErr - } else if ctx.Err() != nil { - s.mergeErrChan <- ctx.Err() + return mergeErr } + return ctx.Err() } // mergeWorkerSimple merges a subset of chunks with proper context handling -func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, endChunk int, output chan<- E) error { +func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, endChunk int, output chan<- E) (err error) { + // A panicking compareFunc must not crash the process from this goroutine + defer func() { + if r := recover(); r != nil { + err = NewComparisonError(r, "mergeWorkerSimple") + } + }() + pq := queue.NewPriorityQueue(func(a, b *mergeFile[E]) int { return s.compareFunc(a.nextRec, b.nextRec) }) @@ -640,12 +699,12 @@ func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, en reader: s.tempReader.Read(i), } _, ok, err := merge.getNext() - if err == io.EOF || !ok { - continue - } if err != nil { return err } + if !ok { + continue // empty chunk + } pq.Push(merge) } @@ -679,8 +738,16 @@ func (s *GenericSorter[E]) mergeWorkerSimple(ctx context.Context, startChunk, en return nil } -// finalMergeSimple performs streaming merge with simpler synchronization -func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, intermediateChans []chan E) { +// finalMergeSimple performs streaming merge with simpler synchronization. +// It returns nil when ctx is cancelled; the caller reports the cancellation. +func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, intermediateChans []chan E) (err error) { + // A panicking compareFunc must not crash the process from this goroutine + defer func() { + if r := recover(); r != nil { + err = NewComparisonError(r, "finalMergeSimple") + } + }() + pq := queue.NewPriorityQueue(func(a, b *channelMergeSource[E]) int { return s.compareFunc(a.nextRec, b.nextRec) }) @@ -697,7 +764,7 @@ func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, intermediateCha for pq.Len() > 0 { // Check if context is cancelled before each iteration if ctx.Err() != nil { - return + return nil } source := pq.Peek() @@ -713,9 +780,10 @@ func (s *GenericSorter[E]) finalMergeSimple(ctx context.Context, intermediateCha } case <-ctx.Done(): // Context cancelled, exit immediately - return + return nil } } + return nil } // channelMergeSource represents a source of sorted data from a channel @@ -748,25 +816,43 @@ type mergeFile[E any] struct { // The first call will return nil while the struct is initialized. // It handles deserialization errors by wrapping them in DeserializationError instances. func (m *mergeFile[E]) getNext() (E, bool, error) { - var newRecBytes []byte old := m.nextRec n, err := binary.ReadUvarint(m.reader) - if err == nil { - newRecBytes = make([]byte, int(n)) - _, err = io.ReadFull(m.reader, newRecBytes) + if err == io.EOF { + return old, false, nil // clean end of the chunk } if err != nil { + return old, false, err + } + newRecBytes := make([]byte, int(n)) + if _, err := io.ReadFull(m.reader, newRecBytes); err != nil { if err == io.EOF { - return old, false, nil + // a length header without its payload is a truncated record, not the end of the chunk + err = io.ErrUnexpectedEOF } return old, false, err } - m.nextRec, err = m.fromBytes(newRecBytes) + m.nextRec, err = m.decode(newRecBytes) if err != nil { - return old, true, NewDeserializationError(err, len(newRecBytes), "getNext") + return old, true, err } return old, true, nil } + +// decode deserializes one record with fromBytes, converting both a returned error +// and a panic into a DeserializationError. +func (m *mergeFile[E]) decode(d []byte) (rec E, err error) { + defer func() { + if r := recover(); r != nil { + err = NewDeserializationError(r, len(d), "getNext") + } + }() + rec, err = m.fromBytes(d) + if err != nil { + return rec, NewDeserializationError(err, len(d), "getNext") + } + return rec, nil +} diff --git a/sort_ordered.go b/sort_ordered.go index 72eeb66..cf729cb 100644 --- a/sort_ordered.go +++ b/sort_ordered.go @@ -72,9 +72,6 @@ func (s *OrderedSorter[T]) toBytesOrdered(d T) ([]byte, error) { 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) - if s == nil { - return nil, output, errChan - } orderedSorter.GenericSorter = *s return orderedSorter, output, errChan } @@ -85,9 +82,6 @@ func Ordered[T cmp.Ordered](input <-chan T, config *Config) (*OrderedSorter[T], 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) - if s == nil { - return nil, output, errChan - } orderedSorter.GenericSorter = *s return orderedSorter, output, errChan } diff --git a/sort_sorttype_legacy.go b/sort_sorttype_legacy.go index c2cfc3b..8109e85 100644 --- a/sort_sorttype_legacy.go +++ b/sort_sorttype_legacy.go @@ -45,14 +45,15 @@ func sortTypeToBytes(a SortType) (result []byte, err error) { // 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. func makeSortTypeFromBytes(fromBytes FromBytes) func([]byte) (SortType, error) { - return func(d []byte) (SortType, error) { - var err error + return func(d []byte) (result SortType, err error) { + // named results: the deferred recover must be able to set the returned error defer func() { if r := recover(); r != nil { + result = nil err = NewDeserializationError(r, len(d), "FromBytes") } }() - return fromBytes(d), err + return fromBytes(d), nil } } @@ -80,9 +81,6 @@ func New(input <-chan SortType, fromBytes FromBytes, lessFunc CompareLessFunc, c compareGeneric := makeCompareSortType(lessFunc) genericSorter, output, errChan := Generic(input, fromBytesGeneric, sortTypeToBytes, compareGeneric, config) - if genericSorter == nil { - return nil, output, errChan - } s := &SortTypeSorter{GenericSorter: *genericSorter} return s, output, errChan } @@ -99,9 +97,6 @@ func NewMock(input <-chan SortType, fromBytes FromBytes, lessFunc CompareLessFun compareGeneric := makeCompareSortType(lessFunc) genericSorter, output, errChan := MockGeneric(input, fromBytesGeneric, sortTypeToBytes, compareGeneric, config, n) - if genericSorter == nil { - return nil, output, errChan - } s := &SortTypeSorter{GenericSorter: *genericSorter} return s, output, errChan } diff --git a/sort_strings.go b/sort_strings.go index 66a3b75..b40feca 100644 --- a/sort_strings.go +++ b/sort_strings.go @@ -31,9 +31,6 @@ func toBytesString(s string) ([]byte, error) { // Sort() will continue reading from the input channel until it is closed. func Strings(input <-chan string, config *Config) (*StringSorter, <-chan string, <-chan error) { genericSorter, output, errChan := Generic(input, fromBytesString, toBytesString, cmp.Compare, config) - if genericSorter == nil { - return nil, output, errChan - } s := &StringSorter{GenericSorter: *genericSorter} return s, output, errChan } @@ -43,9 +40,6 @@ func Strings(input <-chan string, config *Config) (*StringSorter, <-chan string, // The parameter n specifies the maximum number of strings to process. func StringsMock(input <-chan string, config *Config, n int) (*StringSorter, <-chan string, <-chan error) { genericSorter, output, errChan := MockGeneric(input, fromBytesString, toBytesString, cmp.Compare, config, n) - if genericSorter == nil { - return nil, output, errChan - } s := &StringSorter{GenericSorter: *genericSorter} return s, output, errChan } diff --git a/tempfile/regression_test.go b/tempfile/regression_test.go new file mode 100644 index 0000000..fcc09a8 --- /dev/null +++ b/tempfile/regression_test.go @@ -0,0 +1,121 @@ +package tempfile + +// Regression tests for temp directory selection and cleanup. + +import ( + "os" + "path/filepath" + "runtime" + "testing" +) + +// canTestPermissions reports whether chmod restrictions apply to this process. +func canTestPermissions() bool { + return runtime.GOOS != "windows" && os.Geteuid() != 0 +} + +// An explicit directory that cannot hold files used to be silently replaced by the +// default directory. It must be reported as an error instead. +func TestExplicitUnusableDirIsAnError(t *testing.T) { + base := t.TempDir() + file := filepath.Join(base, "file") + if err := os.WriteFile(file, nil, 0o600); err != nil { + t.Fatal(err) + } + cases := map[string]string{ + "regular file": file, + "path through a file": filepath.Join(file, "sub"), + } + if canTestPermissions() { + locked := filepath.Join(base, "locked") + if err := os.Mkdir(locked, 0o000); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Chmod(locked, 0o700) }) + cases["permission denied"] = filepath.Join(locked, "sub") + } + for name, dir := range cases { + t.Run(name, func(t *testing.T) { + if got := GetTempDir(dir, true); got != dir { + t.Errorf("GetTempDir(%q) = %q, want it unchanged", dir, got) + } + w, err := New(dir, true) + if err == nil { + t.Errorf("New(%q) created %s, want an error", dir, w.Name()) + _ = w.Close() + } + }) + } +} + +// Default selection used to pick a candidate that did not exist (and then fail to +// create it, as with no /var/tmp in a non-root container) or was read-only (as with a +// read-only root filesystem), instead of falling back to the next candidate. +func TestDefaultSelectionSkipsUnusableCandidates(t *testing.T) { + base := t.TempDir() + missing := filepath.Join(base, "missing") + file := filepath.Join(base, "file") + if err := os.WriteFile(file, nil, 0o600); err != nil { + t.Fatal(err) + } + writable := filepath.Join(base, "writable") + if err := os.Mkdir(writable, 0o700); err != nil { + t.Fatal(err) + } + candidates := []string{missing, file} + if canTestPermissions() { + readOnly := filepath.Join(base, "readonly") + if err := os.Mkdir(readOnly, 0o500); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Chmod(readOnly, 0o700) }) + candidates = append(candidates, readOnly) + } + candidates = append(candidates, writable) + + if got := firstUsableDir(candidates); got != writable { + t.Fatalf("firstUsableDir(%q) = %q, want %q", candidates, got, writable) + } + if _, err := os.Stat(missing); !os.IsNotExist(err) { + t.Errorf("missing candidate %s was created", missing) + } + if entries, err := os.ReadDir(writable); err != nil || len(entries) != 0 { + t.Errorf("probe left files behind in %s: %v %v", writable, entries, err) + } +} + +// Our own .extsort_ fallback directories are still selected before they exist, +// because New creates them on demand. +func TestDefaultSelectionAcceptsMissingExtsortDir(t *testing.T) { + ours := filepath.Join(t.TempDir(), extsortTempDirName) + if got := firstUsableDir([]string{ours}); got != ours { + t.Fatalf("firstUsableDir() = %q, want %q", got, ours) + } + if _, err := os.Stat(ours); !os.IsNotExist(err) { + t.Errorf("selection created %s; New should create it", ours) + } +} + +// Save hands the writer's directory reference to the reader, so closing the reader +// must release it. The .extsort_ directory used to be left behind after every +// successful sort. +func TestReaderCloseRemovesExtsortDir(t *testing.T) { + dir := filepath.Join(t.TempDir(), extsortTempDirName) + w, err := New(dir, true) + if err != nil { + t.Fatal(err) + } + if _, err := w.WriteString("data"); err != nil { + t.Fatal(err) + } + r, err := w.Save() + if err != nil { + t.Fatal(err) + } + if err := r.Close(); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(dir); !os.IsNotExist(err) { + t.Errorf("%s still exists after the reader was closed", dir) + } +} diff --git a/tempfile/tempdir.go b/tempfile/tempdir.go index ea45e51..3536392 100644 --- a/tempfile/tempdir.go +++ b/tempfile/tempdir.go @@ -23,16 +23,14 @@ var ( ) // GetTempDir returns the optimal temporary directory for the given preference. -// If dir is provided and non-empty, it's validated and returned if usable. -// Otherwise, returns a pre-computed optimal directory based on preferDiskBacked. +// If dir is provided and non-empty, it is returned unchanged: an unusable +// caller-provided directory is reported as an error by New rather than being +// silently replaced. Otherwise, returns a pre-computed optimal directory based on preferDiskBacked. // This function is thread-safe and performs O(1) lookups after initialization. func GetTempDir(dir string, preferDiskBacked bool) string { - // If caller provides a specific directory, validate and use it + // If caller provides a specific directory, use it as given if dir != "" { - if isDirectoryUsable(dir) { - return dir - } - // Fall through to use pre-computed directory if provided dir is unusable + return dir } // Ensure directories have been discovered (happens once) @@ -76,16 +74,32 @@ func cacheExpensiveOperations() { // It iterates through candidates in priority order and returns the first usable directory. // Falls back to OS temp directory if no candidates are usable. func findBestDirectory(preferDiskBacked bool) string { - candidates := buildCandidateList(preferDiskBacked) + if dir := firstUsableDir(buildCandidateList(preferDiskBacked)); dir != "" { + return dir + } + + // Final fallback to OS default temp dir + return cachedOSTemp +} +// firstUsableDir returns the first candidate that can hold temp files, or "" if none can. +// Our own process-specific directories are usable if they exist or can be created, since +// New creates them on demand. Any other candidate (such as /var/tmp or os.TempDir()) is +// never created, so it must already be a directory we can write to; a missing or read-only +// one falls through to the next candidate. +func firstUsableDir(candidates []string) string { for _, candidate := range candidates { - if isDirectoryUsable(candidate) { + if isExtsortDirectory(candidate) { + if isDirectoryUsable(candidate) { + return candidate + } + continue + } + if isWritableDirectory(candidate) { return candidate } } - - // Final fallback to OS default temp dir - return cachedOSTemp + return "" } // buildCandidateList returns a prioritized list of temporary directory candidates. @@ -158,6 +172,20 @@ func buildAdditionalFallbacks() []string { return candidates } +// isWritableDirectory reports whether dir is an existing directory in which files can be +// created, by creating and removing a probe file. Unlike isDirectoryUsable, it returns false +// for a directory that does not exist, and detects read-only filesystems and permission errors. +func isWritableDirectory(dir string) bool { + f, err := os.CreateTemp(dir, mergeFilenamePrefix+"probe_") + if err != nil { + return false + } + name := f.Name() + _ = f.Close() + _ = os.Remove(name) + return true +} + // isDirectoryUsable checks if a directory exists and is a directory, or can be created. // It returns true for non-existent directories that could potentially be created. // We don't test writability here to avoid creating unnecessary files - the actual diff --git a/tempfile/tempfile.go b/tempfile/tempfile.go index ca3f91e..78c029e 100644 --- a/tempfile/tempfile.go +++ b/tempfile/tempfile.go @@ -59,6 +59,7 @@ type fileReader struct { 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 } // New creates a new FileWriter for virtual temporary files in the specified directory. @@ -142,6 +143,7 @@ func (w *FileWriter) Close() error { // Clean up directory if we created it and no other writers are using it if w.createdDir != "" { decrementDirRefCount(w.createdDir) + w.createdDir = "" } return err @@ -190,6 +192,7 @@ func (w *FileWriter) Save() (TempReader, error) { return nil, err } + var r *fileReader if w.needsCleanup { // Windows case: close file and reopen for reading filename := w.file.Name() @@ -197,11 +200,19 @@ func (w *FileWriter) Save() (TempReader, error) { if err != nil { return nil, err } - return newTempReader(filename, w.sections, w.needsCleanup) + r, err = newTempReader(filename, w.sections, w.needsCleanup) } else { // Unix case: file is unlinked, reuse the same file handle - return newTempReaderFromFile(w.file, w.sections, w.needsCleanup) + r, err = newTempReaderFromFile(w.file, w.sections, w.needsCleanup) } + if err != nil { + return nil, err + } + + // The reader now owns the directory reference and releases it on Close + r.createdDir = w.createdDir + w.createdDir = "" + return r, nil } // newTempReader creates a TempReader by opening a file by name. @@ -263,6 +274,12 @@ func (r *fileReader) Close() error { } } + // Clean up directory if we created it and no other writers are using it + if r.createdDir != "" { + decrementDirRefCount(r.createdDir) + r.createdDir = "" + } + return err } diff --git a/tempfile_error_test.go b/tempfile_error_test.go index 4ef22a0..b110d27 100644 --- a/tempfile_error_test.go +++ b/tempfile_error_test.go @@ -60,14 +60,11 @@ func TestTempFileCreationFailure(t *testing.T) { // Create the sorter - this should fail when trying to create temp files sort, outChan, errChan := extsort.New(inputChan, testFromBytes, testLess, config) - // Check if the sorter is nil (expected after fix) + // The temp file is created lazily, so the sorter is usable and Sort reports the failure if sort == nil { - t.Log("Sorter is nil as expected when tempfile creation fails") - } else { - // If sorter is not nil, this should still not segfault after our fix - t.Log("Sorter is not nil, attempting to sort (this should not segfault)") - sort.Sort(context.Background()) + t.Fatal("Sorter is nil; Sort() on it would panic instead of reporting the error") } + sort.Sort(context.Background()) // Drain any output that might come through outputCount := 0 @@ -134,14 +131,11 @@ func TestTempFileCreationFailureStrings(t *testing.T) { sort, outChan, errChan := extsort.Strings(inputChan, config) - // Check if the sorter is nil (expected after fix) + // The temp file is created lazily, so the sorter is usable and Sort reports the failure if sort == nil { - t.Log("String sorter is nil as expected when tempfile creation fails") - } else { - // If sorter is not nil, this should still not segfault after our fix - t.Log("String sorter is not nil, attempting to sort (this should not segfault)") - sort.Sort(context.Background()) + t.Fatal("String sorter is nil; Sort() on it would panic instead of reporting the error") } + sort.Sort(context.Background()) // Drain output outputCount := 0