diff --git a/checksum_merge_test.go b/checksum_merge_test.go new file mode 100644 index 0000000..a4ec562 --- /dev/null +++ b/checksum_merge_test.go @@ -0,0 +1,103 @@ +package extsort + +import ( + "bufio" + "cmp" + "context" + "errors" + "io" + "slices" + "testing" +) + +type checksumFailureReader struct { + data []byte + err error +} + +func (r *checksumFailureReader) Read(p []byte) (int, error) { + if len(r.data) == 0 { + return 0, r.err + } + n := copy(p, r.data) + r.data = r.data[n:] + if len(r.data) == 0 { + return n, r.err + } + return n, nil +} + +type checksumFailureTempReader struct { + readers []*bufio.Reader +} + +func (r *checksumFailureTempReader) Close() error { + return nil +} + +func (r *checksumFailureTempReader) Size() int { + return len(r.readers) +} + +func (r *checksumFailureTempReader) Read(i int) *bufio.Reader { + return r.readers[i] +} + +func TestChecksumFailureStopsMerge(t *testing.T) { + for _, test := range []struct { + name string + numWorkers int + bufferSize int + secondSection []byte + want []int + exact bool + }{ + { + name: "single threaded", + numWorkers: 2, + bufferSize: 10, + secondSection: []byte{1, 2, 1, 4}, + want: []int{1, 2, 3}, + exact: true, + }, + { + name: "parallel", + numWorkers: 1, + secondSection: []byte{1, 2, 1, 4}, + want: []int{1, 2, 3}, + }, + {name: "first block/single threaded", numWorkers: 2, bufferSize: 10, exact: true}, + {name: "first block/parallel", numWorkers: 1, exact: true}, + } { + t.Run(test.name, func(t *testing.T) { + checksumErr := errors.New("temporary section 1 block 0 checksum mismatch") + sorter := newSorter[int]( + nil, + func(data []byte) (int, error) { return int(data[0]), nil }, + func(value int) ([]byte, error) { return []byte{byte(value)}, nil }, + cmp.Compare[int], + &Config{ChunkSize: 2, NumWorkers: test.numWorkers, SortedChanBuffSize: test.bufferSize}, + ) + sorter.tempReader = &checksumFailureTempReader{readers: []*bufio.Reader{ + bufio.NewReader(&checksumFailureReader{data: []byte{1, 1, 1, 3, 1, 5}, err: io.EOF}), + bufio.NewReader(&checksumFailureReader{data: test.secondSection, err: checksumErr}), + }} + + go sorter.mergeNChunks(context.Background()) + + var got []int + for value := range sorter.mergeChunkChan { + got = append(got, value) + } + if len(got) > len(test.want) || !slices.Equal(got, test.want[:len(got)]) { + t.Fatalf("output before checksum failure = %v, want a prefix of %v", got, test.want) + } + if test.exact && !slices.Equal(got, test.want) { + t.Fatalf("output before checksum failure = %v, want %v", got, test.want) + } + if err := <-sorter.mergeErrChan; !errors.Is(err, checksumErr) { + t.Fatalf("merge error = %v, want %v", err, checksumErr) + } + }) + } +} diff --git a/config.go b/config.go index d12d580..0ab3686 100644 --- a/config.go +++ b/config.go @@ -36,6 +36,10 @@ type Config struct { // // Default: "" (intelligent selection). TempFilesDir string + + // Checksum verifies temporary file blocks before records are returned by the + // merge phase. Default: false. + Checksum bool } // DefaultConfig returns a Config with sensible default values optimized for diff --git a/ordered_test.go b/ordered_test.go index 965b3e9..5328c07 100644 --- a/ordered_test.go +++ b/ordered_test.go @@ -3,6 +3,7 @@ package extsort_test import ( "context" "reflect" + "slices" "testing" "github.com/lanrat/extsort" @@ -240,3 +241,32 @@ done5: t.Errorf("Expected empty result, got %v", result) } } + +func TestOrderedWithChecksum(t *testing.T) { + input := make(chan int, 6) + for _, value := range []int{6, 1, 5, 2, 4, 3} { + input <- value + } + close(input) + + config := extsort.DefaultConfig() + config.ChunkSize = 2 + config.TempFilesDir = t.TempDir() + config.Checksum = true + + sorter, output, errChan := extsort.Ordered(input, config) + sorter.Sort(context.Background()) + + var got []int + for value := range output { + got = append(got, value) + } + if err := <-errChan; err != nil { + t.Fatal(err) + } + + want := []int{1, 2, 3, 4, 5, 6} + if !slices.Equal(got, want) { + t.Fatalf("sorted values = %v, want %v", got, want) + } +} diff --git a/sort_generic.go b/sort_generic.go index 7e670e8..a863100 100644 --- a/sort_generic.go +++ b/sort_generic.go @@ -5,6 +5,7 @@ import ( "bufio" "context" "encoding/binary" + "errors" "io" "slices" "sync" @@ -162,7 +163,11 @@ func (s *GenericSorter[E]) initMemoryPools() *memoryPools { 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 s.config.Checksum { + s.tempWriter, err = tempfile.NewChecksummed(s.config.TempFilesDir, true) + } else { + s.tempWriter, err = tempfile.New(s.config.TempFilesDir, true) + } if err != nil { s.mergeErrChan <- err close(s.mergeErrChan) @@ -194,9 +199,12 @@ 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) { + sortCtx, cancel := context.WithCancel(ctx) + defer cancel() + var buildSortErrGroup, saveErrGroup *errgroup.Group - buildSortErrGroup, s.buildSortCtx = errgroup.WithContext(ctx) - saveErrGroup, s.saveCtx = errgroup.WithContext(ctx) + buildSortErrGroup, s.buildSortCtx = errgroup.WithContext(sortCtx) + saveErrGroup, s.saveCtx = errgroup.WithContext(sortCtx) //start creating chunks buildSortErrGroup.Go(s.buildChunks) @@ -207,22 +215,34 @@ func (s *GenericSorter[E]) Sort(ctx context.Context) { } // Start the save worker that will handle single-chunk optimization - saveErrGroup.Go(s.saveChunksOptimized) + saveErrGroup.Go(func() error { + err := s.saveChunksOptimized() + if err != nil { + cancel() + } + return err + }) - err := buildSortErrGroup.Wait() - if err != nil { - s.mergeErrChan <- err - close(s.mergeErrChan) - close(s.mergeChunkChan) - return + buildErr := buildSortErrGroup.Wait() + if buildErr != nil { + cancel() } // Close saveChunkChan to signal end of chunks close(s.saveChunkChan) // Wait for save worker to complete - err = saveErrGroup.Wait() + saveErr := saveErrGroup.Wait() + err := buildErr + if saveErr != nil && (err == nil || errors.Is(err, context.Canceled)) { + err = saveErr + } if err != nil { + if s.tempReader != nil { + _ = s.tempReader.Close() + } else { + _ = s.tempWriter.Close() + } s.mergeErrChan <- err close(s.mergeErrChan) close(s.mergeChunkChan) @@ -241,6 +261,20 @@ func (s *GenericSorter[E]) Sort(ctx context.Context) { go s.mergeNChunks(ctx) } +func (s *GenericSorter[E]) Next(ctx context.Context) (value E, ok bool, err error) { + select { + case value, ok = <-s.mergeChunkChan: + if ok { + return value, true, nil + } + case <-ctx.Done(): + return value, false, ctx.Err() + } + + err, _ = <-s.mergeErrChan + return value, false, err +} + // 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 @@ -457,6 +491,7 @@ func (s *GenericSorter[E]) saveChunk(b *genericChunk[E]) error { // 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) { + defer close(s.mergeErrChan) defer close(s.mergeChunkChan) defer func() { if s.tempReader != nil { @@ -470,8 +505,6 @@ func (s *GenericSorter[E]) mergeNChunks(ctx context.Context) { } } }() - // Always ensure error channel is closed - defer close(s.mergeErrChan) if s.tempReader == nil { return @@ -504,13 +537,13 @@ 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 } + if !ok { + continue + } pq.Push(merge) } @@ -640,12 +673,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 + } pq.Push(merge) } diff --git a/tempfile/checksum.go b/tempfile/checksum.go new file mode 100644 index 0000000..1298b10 --- /dev/null +++ b/tempfile/checksum.go @@ -0,0 +1,132 @@ +package tempfile + +import ( + "bufio" + "fmt" + "hash/crc32" + "io" +) + +const checksumBlockSize = 1 << 16 + +var checksumTable = crc32.MakeTable(crc32.Castagnoli) + +type checksumWriter struct { + blockSize int + blockChecksum uint32 + checksums []uint32 + buf []byte +} + +func (w *checksumWriter) write(dst *bufio.Writer, p []byte) (int, error) { + written := 0 + for len(p) > 0 { + blockRemaining := checksumBlockSize - w.blockSize + if blockRemaining > len(p) { + blockRemaining = len(p) + } + chunk := p[:blockRemaining] + n, err := dst.Write(chunk) + if err != nil { + return 0, err + } + if n != len(chunk) { + return 0, io.ErrShortWrite + } + w.blockChecksum = crc32.Update(w.blockChecksum, checksumTable, chunk) + w.blockSize += n + written += n + if w.blockSize == checksumBlockSize { + w.finishBlock() + } + p = p[n:] + } + return written, nil +} + +func (w *checksumWriter) writeString(dst *bufio.Writer, s string) (int, error) { + bufSize := min(len(s), checksumBlockSize) + if cap(w.buf) < bufSize { + w.buf = make([]byte, bufSize) + } else { + w.buf = w.buf[:bufSize] + } + written := 0 + for len(s) > 0 { + n := copy(w.buf, s) + m, err := w.write(dst, w.buf[:n]) + if err != nil { + return 0, err + } + written += m + s = s[m:] + } + return written, nil +} + +func (w *checksumWriter) finishBlock() { + if w.blockSize == 0 { + return + } + w.checksums = append(w.checksums, w.blockChecksum) + w.blockSize = 0 + w.blockChecksum = 0 +} + +func (w *checksumWriter) finishSection() []uint32 { + w.finishBlock() + checksums := w.checksums + w.checksums = nil + return checksums +} + +func newChecksummedReader(reader io.Reader, section int, checksums []uint32, size int64) *bufio.Reader { + return bufio.NewReaderSize(&checksummedReader{ + reader: reader, + section: section, + checksums: checksums, + remaining: size, + }, 16) +} + +type checksummedReader struct { + reader io.Reader + section int + checksums []uint32 + nextBlock int + remaining int64 + data []byte + offset int +} + +func (r *checksummedReader) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if r.offset == len(r.data) { + if r.nextBlock == len(r.checksums) { + return 0, io.EOF + } + blockSize := checksumBlockSize + if int64(blockSize) > r.remaining { + blockSize = int(r.remaining) + } + if cap(r.data) < blockSize { + r.data = make([]byte, blockSize) + } else { + r.data = r.data[:blockSize] + } + if _, err := io.ReadFull(r.reader, r.data); err != nil { + return 0, fmt.Errorf("read temporary section %d block %d: %w", r.section, r.nextBlock, err) + } + if crc32.Checksum(r.data, checksumTable) != r.checksums[r.nextBlock] { + return 0, fmt.Errorf("temporary section %d block %d checksum mismatch", r.section, r.nextBlock) + } + r.nextBlock++ + r.remaining -= int64(blockSize) + r.offset = 0 + } + n := copy(p, r.data[r.offset:]) + r.offset += n + return n, nil +} diff --git a/tempfile/checksum_test.go b/tempfile/checksum_test.go new file mode 100644 index 0000000..c6b0401 --- /dev/null +++ b/tempfile/checksum_test.go @@ -0,0 +1,74 @@ +package tempfile + +import ( + "bytes" + "io" + "os" + "testing" +) + +func TestChecksummedTempFile(t *testing.T) { + w, err := NewChecksummed(t.TempDir(), true) + if err != nil { + t.Fatal(err) + } + want := bytes.Repeat([]byte("checksum"), checksumBlockSize*2/len("checksum")+17) + if _, err = w.WriteString(string(want)); err != nil { + t.Fatal(err) + } + r, err := w.Save() + if err != nil { + t.Fatal(err) + } + defer r.Close() + + got, err := io.ReadAll(r.Read(0)) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, want) { + t.Fatal("checksummed temporary file data mismatch") + } +} + +func TestChecksummedTempFileDetectsCorruption(t *testing.T) { + w, err := NewChecksummed(t.TempDir(), true) + if err != nil { + t.Fatal(err) + } + want := bytes.Repeat([]byte{0x5a}, checksumBlockSize*2) + if _, err = w.Write(want); err != nil { + t.Fatal(err) + } + tempReader, err := w.Save() + if err != nil { + t.Fatal(err) + } + r := tempReader.(*fileReader) + defer r.Close() + + f := r.file + if w.needsCleanup { + f, err = os.OpenFile(r.filename, os.O_RDWR, 0) + if err != nil { + t.Fatal(err) + } + defer f.Close() + } + if _, err = f.WriteAt([]byte{0xff}, int64(checksumBlockSize+1)); err != nil { + t.Fatal(err) + } + + reader := r.Read(0) + firstBlock := make([]byte, checksumBlockSize) + if _, err = io.ReadFull(reader, firstBlock); err != nil { + t.Fatal(err) + } + if !bytes.Equal(firstBlock, want[:checksumBlockSize]) { + t.Fatal("valid block data mismatch") + } + buf := make([]byte, 1) + if n, err := reader.Read(buf); err == nil || n != 0 { + t.Fatalf("corrupted block read = %d, %v; want 0, error", n, err) + } +} diff --git a/tempfile/tempfile.go b/tempfile/tempfile.go index ca3f91e..d06014b 100644 --- a/tempfile/tempfile.go +++ b/tempfile/tempfile.go @@ -48,14 +48,20 @@ var ( type FileWriter struct { file *os.File bufWriter *bufio.Writer - sections []int64 + sections []sectionMeta needsCleanup bool // true if manual cleanup is needed (Windows) createdDir string // directory we created (for cleanup) + checksum *checksumWriter +} + +type sectionMeta struct { + end int64 + checksums []uint32 } type fileReader struct { file *os.File - sections []int64 + sections []sectionMeta readers []*bufio.Reader needsCleanup bool // true if manual cleanup is needed (Windows) filename string // filename for cleanup @@ -67,8 +73,21 @@ type fileReader struct { // The function attempts automatic cleanup on Unix systems by unlinking the file immediately, // while Windows requires explicit cleanup when the FileWriter is closed. func New(dir string, preferDiskBacked bool) (*FileWriter, error) { + return newFileWriter(dir, preferDiskBacked, false) +} + +// NewChecksummed creates a FileWriter that verifies each temporary file block +// before returning its contents to readers. +func NewChecksummed(dir string, preferDiskBacked bool) (*FileWriter, error) { + return newFileWriter(dir, preferDiskBacked, true) +} + +func newFileWriter(dir string, preferDiskBacked, checksummed bool) (*FileWriter, error) { var w FileWriter var err error + if checksummed { + w.checksum = &checksumWriter{} + } // Use intelligent directory selection if no specific directory provided selectedDir := GetTempDir(dir, preferDiskBacked) @@ -104,7 +123,7 @@ func New(dir string, preferDiskBacked bool) (*FileWriter, error) { } w.bufWriter = bufio.NewWriterSize(w.file, fileBufferSize) - w.sections = make([]int64, 0, 10) + w.sections = make([]sectionMeta, 0, 10) return &w, nil } @@ -131,6 +150,7 @@ func (w *FileWriter) Close() error { err := w.file.Close() w.sections = nil w.bufWriter = nil + w.checksum = nil // Only attempt manual cleanup if needed (Windows case) if w.needsCleanup { @@ -150,12 +170,18 @@ 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) { + if w.checksum != nil { + return w.checksum.write(w.bufWriter, p) + } return w.bufWriter.Write(p) } // 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. +// Without checksumming it avoids a byte slice conversion. func (w *FileWriter) WriteString(s string) (int, error) { + if w.checksum != nil { + return w.checksum.writeString(w.bufWriter, s) + } return w.bufWriter.WriteString(s) } @@ -163,6 +189,10 @@ func (w *FileWriter) WriteString(s string) (int, error) { // It flushes buffered data and records the section boundary for later reading. // Returns the file offset where the next section will begin. func (w *FileWriter) Next() (int64, error) { + var checksums []uint32 + if w.checksum != nil { + checksums = w.checksum.finishSection() + } // save offsets err := w.bufWriter.Flush() if err != nil { @@ -172,7 +202,7 @@ func (w *FileWriter) Next() (int64, error) { if err != nil { return 0, err } - w.sections = append(w.sections, pos) + w.sections = append(w.sections, sectionMeta{end: pos, checksums: checksums}) return pos, nil } @@ -197,16 +227,16 @@ func (w *FileWriter) Save() (TempReader, error) { if err != nil { return nil, err } - return newTempReader(filename, w.sections, w.needsCleanup) + return newTempReader(filename, w.sections, w.checksum != nil, w.needsCleanup) } else { // Unix case: file is unlinked, reuse the same file handle - return newTempReaderFromFile(w.file, w.sections, w.needsCleanup) + return newTempReaderFromFile(w.file, w.sections, w.checksum != nil, w.needsCleanup) } } // newTempReader creates a TempReader by opening a file by name. // This is used on Windows where files need to be closed and reopened for reading. -func newTempReader(filename string, sections []int64, needsCleanup bool) (*fileReader, error) { +func newTempReader(filename string, sections []sectionMeta, checksummed, needsCleanup bool) (*fileReader, error) { // create TempReader by opening file by name var err error var r fileReader @@ -220,10 +250,15 @@ func newTempReader(filename string, sections []int64, needsCleanup bool) (*fileR 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) + for i, meta := range r.sections { + sectionSize := meta.end - offset + section := io.NewSectionReader(r.file, offset, sectionSize) + offset = meta.end + if checksummed { + r.readers[i] = newChecksummedReader(section, i, meta.checksums, sectionSize) + } else { + r.readers[i] = bufio.NewReaderSize(section, fileBufferSize) + } } return &r, nil @@ -231,7 +266,7 @@ func newTempReader(filename string, sections []int64, needsCleanup bool) (*fileR // newTempReaderFromFile creates a TempReader by reusing an existing file handle. // This is used on Unix systems where unlinked files can continue to be accessed. -func newTempReaderFromFile(file *os.File, sections []int64, needsCleanup bool) (*fileReader, error) { +func newTempReaderFromFile(file *os.File, sections []sectionMeta, checksummed, needsCleanup bool) (*fileReader, error) { // create TempReader by reusing existing file handle var r fileReader r.file = file @@ -241,10 +276,15 @@ func newTempReaderFromFile(file *os.File, sections []int64, needsCleanup bool) ( 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) + for i, meta := range r.sections { + sectionSize := meta.end - offset + section := io.NewSectionReader(r.file, offset, sectionSize) + offset = meta.end + if checksummed { + r.readers[i] = newChecksummedReader(section, i, meta.checksums, sectionSize) + } else { + r.readers[i] = bufio.NewReaderSize(section, fileBufferSize) + } } return &r, nil