This commit is contained in:
@@ -26,9 +26,19 @@ var ErrNoRows = fmt.Errorf("input has no data rows")
|
||||
// Only one shard is buffered at a time, so memory is bounded by shard size (a
|
||||
// worker-sized slice of the data), not by the size of the whole dataset.
|
||||
func SplitTSV(r io.Reader, rowsPerShard int, emit func(index int, shard io.Reader) error) error {
|
||||
return SplitTSVLimit(r, rowsPerShard, 0, emit)
|
||||
}
|
||||
|
||||
// SplitTSVLimit behaves like SplitTSV but emits no more than maxRows data rows.
|
||||
// A maxRows value of zero means unlimited. This lets an operator make a small,
|
||||
// representative pipeline check without materialising a second dataset file.
|
||||
func SplitTSVLimit(r io.Reader, rowsPerShard, maxRows int, emit func(index int, shard io.Reader) error) error {
|
||||
if rowsPerShard <= 0 {
|
||||
return fmt.Errorf("rowsPerShard must be positive, got %d", rowsPerShard)
|
||||
}
|
||||
if maxRows < 0 {
|
||||
return fmt.Errorf("maxRows must be non-negative, got %d", maxRows)
|
||||
}
|
||||
|
||||
sc := bufio.NewScanner(r)
|
||||
// Allow long lines: a SMILES row can be far wider than bufio's 64 KB default.
|
||||
@@ -73,6 +83,9 @@ func SplitTSV(r io.Reader, rowsPerShard int, emit func(index int, shard io.Reade
|
||||
return err
|
||||
}
|
||||
}
|
||||
if maxRows > 0 && index*rowsPerShard+rows == maxRows {
|
||||
break
|
||||
}
|
||||
}
|
||||
if err := sc.Err(); err != nil {
|
||||
return fmt.Errorf("read rows: %w", err)
|
||||
|
||||
@@ -103,6 +103,22 @@ func TestSplitSingleShardWhenSizeExceedsRows(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitLimitUsesOnlyLeadingDataRows(t *testing.T) {
|
||||
input := "h\nr1\nr2\nr3\nr4\nr5\n"
|
||||
var shards []string
|
||||
err := SplitTSVLimit(strings.NewReader(input), 2, 3, func(_ int, shard io.Reader) error {
|
||||
b, _ := io.ReadAll(shard)
|
||||
shards = append(shards, string(b))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got, want := strings.Join(shards, ""), "h\nr1\nr2\nh\nr3\n"; got != want {
|
||||
t.Errorf("limited shards = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// The scanned bytes are reused by bufio; the shard buffer must copy them, or a
|
||||
// later row would corrupt an earlier one. This guards that copy.
|
||||
func TestSplitDoesNotAliasScannerBuffer(t *testing.T) {
|
||||
|
||||
Reference in New Issue
Block a user