feat(coordinator): upload a dataset and chunk it into shard tasks (CTX-05, part 4)
The coordinator can now ingest a dataset itself, not only accept client-supplied
chunk URIs.
- internal/chunk: a deterministic, generic TSV row splitter — repeats the header
per shard, buffers one shard at a time, rejects header-only input. Unit-tested.
- POST /jobs/upload (multipart): streams the dataset into an input artifact,
splits it into shard artifacts, and creates one shard task per shard, all in
one transaction; blobs are cleaned up if the transaction fails.
- GET /tasks/{id}/input streams a task's input shard back to the worker.
- domain: NewUploadedJob, NewShardTask, Task/Job.InputArtifactID; a shard task's
input is an artifact, not a URI. Claim response nests input:{uri,sha256} per
the contract, with uri = /tasks/{id}/input for shards.
- migration 0005 makes input_uri nullable and adds a has-input check.
- The existing URI-based POST /jobs path is untouched; both coexist.
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
// Package chunk splits a tabular input into deterministic shards. It is generic
|
||||
// row splitting only — no workload semantics (SMILES, top-k) live here.
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
// ErrNoRows is returned when the input has a header but no data rows: a job with
|
||||
// zero tasks could never complete, so it is rejected at the source.
|
||||
var ErrNoRows = fmt.Errorf("input has no data rows")
|
||||
|
||||
// SplitTSV reads a header-plus-rows text stream and cuts it into shards of at
|
||||
// most rowsPerShard data rows. Every shard repeats the header, so a worker can
|
||||
// parse its shard in isolation. emit is called once per shard, in order, with a
|
||||
// reader over that shard's bytes; the reader is valid only for the duration of
|
||||
// the call.
|
||||
//
|
||||
// Splitting is deterministic: the same input and rowsPerShard always produce the
|
||||
// same shards, byte for byte — which is what lets chunk_index refer to a stable
|
||||
// piece and makes a re-run reproducible.
|
||||
//
|
||||
// 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 {
|
||||
if rowsPerShard <= 0 {
|
||||
return fmt.Errorf("rowsPerShard must be positive, got %d", rowsPerShard)
|
||||
}
|
||||
|
||||
sc := bufio.NewScanner(r)
|
||||
// Allow long lines: a SMILES row can be far wider than bufio's 64 KB default.
|
||||
sc.Buffer(make([]byte, 0, 64*1024), 8*1024*1024)
|
||||
|
||||
if !sc.Scan() {
|
||||
if err := sc.Err(); err != nil {
|
||||
return fmt.Errorf("read header: %w", err)
|
||||
}
|
||||
return ErrNoRows // completely empty input
|
||||
}
|
||||
header := append([]byte(nil), sc.Bytes()...)
|
||||
|
||||
var (
|
||||
buf bytes.Buffer
|
||||
rows int
|
||||
index int
|
||||
)
|
||||
|
||||
// flush emits the buffered shard and resets for the next one.
|
||||
flush := func() error {
|
||||
if err := emit(index, bytes.NewReader(buf.Bytes())); err != nil {
|
||||
return err
|
||||
}
|
||||
index++
|
||||
buf.Reset()
|
||||
rows = 0
|
||||
return nil
|
||||
}
|
||||
|
||||
for sc.Scan() {
|
||||
if rows == 0 {
|
||||
buf.Write(header)
|
||||
buf.WriteByte('\n')
|
||||
}
|
||||
buf.Write(sc.Bytes())
|
||||
buf.WriteByte('\n')
|
||||
rows++
|
||||
|
||||
if rows == rowsPerShard {
|
||||
if err := flush(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := sc.Err(); err != nil {
|
||||
return fmt.Errorf("read rows: %w", err)
|
||||
}
|
||||
|
||||
// A partial final shard still has to go out.
|
||||
if rows > 0 {
|
||||
if err := flush(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if index == 0 {
|
||||
return ErrNoRows // header only, no data
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
package chunk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// collect runs SplitTSV and returns every shard as a string.
|
||||
func collect(t *testing.T, input string, rowsPerShard int) []string {
|
||||
t.Helper()
|
||||
var shards []string
|
||||
err := SplitTSV(strings.NewReader(input), rowsPerShard, func(index int, shard io.Reader) error {
|
||||
b, _ := io.ReadAll(shard)
|
||||
if index != len(shards) {
|
||||
t.Fatalf("emit index = %d, want %d (out of order)", index, len(shards))
|
||||
}
|
||||
shards = append(shards, string(b))
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SplitTSV: %v", err)
|
||||
}
|
||||
return shards
|
||||
}
|
||||
|
||||
func TestSplitCountsShardsAndRepeatsHeader(t *testing.T) {
|
||||
input := "id\tsmiles\nA\tCC\nB\tCCC\nC\tCCCC\nD\tCCCCC\nE\tCCCCCC\n"
|
||||
shards := collect(t, input, 2)
|
||||
|
||||
if len(shards) != 3 { // 5 rows / 2 per shard = ceil = 3
|
||||
t.Fatalf("got %d shards, want 3", len(shards))
|
||||
}
|
||||
for i, s := range shards {
|
||||
if !strings.HasPrefix(s, "id\tsmiles\n") {
|
||||
t.Errorf("shard %d missing header: %q", i, s)
|
||||
}
|
||||
}
|
||||
if shards[0] != "id\tsmiles\nA\tCC\nB\tCCC\n" {
|
||||
t.Errorf("shard 0 = %q", shards[0])
|
||||
}
|
||||
if shards[2] != "id\tsmiles\nE\tCCCCCC\n" { // partial final shard
|
||||
t.Errorf("shard 2 = %q", shards[2])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitExactMultipleHasNoEmptyTrailingShard(t *testing.T) {
|
||||
input := "h\nr1\nr2\nr3\nr4\n"
|
||||
shards := collect(t, input, 2)
|
||||
if len(shards) != 2 { // exactly 4/2, no empty third shard
|
||||
t.Fatalf("got %d shards, want 2", len(shards))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitIsDeterministic(t *testing.T) {
|
||||
input := "h\n" + strings.Repeat("row\n", 100)
|
||||
a := collect(t, input, 7)
|
||||
b := collect(t, input, 7)
|
||||
if fmt.Sprint(a) != fmt.Sprint(b) {
|
||||
t.Error("two runs produced different shards")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitRejectsHeaderOnly(t *testing.T) {
|
||||
err := SplitTSV(strings.NewReader("id\tsmiles\n"), 10, func(int, io.Reader) error { return nil })
|
||||
if !errors.Is(err, ErrNoRows) {
|
||||
t.Errorf("err = %v, want ErrNoRows", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitRejectsEmptyInput(t *testing.T) {
|
||||
err := SplitTSV(strings.NewReader(""), 10, func(int, io.Reader) error { return nil })
|
||||
if !errors.Is(err, ErrNoRows) {
|
||||
t.Errorf("err = %v, want ErrNoRows", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitRejectsNonPositiveSize(t *testing.T) {
|
||||
err := SplitTSV(strings.NewReader("h\nr\n"), 0, func(int, io.Reader) error { return nil })
|
||||
if err == nil {
|
||||
t.Error("expected an error for rowsPerShard = 0")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitPropagatesEmitError(t *testing.T) {
|
||||
boom := errors.New("boom")
|
||||
err := SplitTSV(strings.NewReader("h\nr1\nr2\n"), 1, func(int, io.Reader) error { return boom })
|
||||
if !errors.Is(err, boom) {
|
||||
t.Errorf("err = %v, want boom", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitSingleShardWhenSizeExceedsRows(t *testing.T) {
|
||||
shards := collect(t, "h\nr1\nr2\n", 100)
|
||||
if len(shards) != 1 {
|
||||
t.Fatalf("got %d shards, want 1", len(shards))
|
||||
}
|
||||
if shards[0] != "h\nr1\nr2\n" {
|
||||
t.Errorf("shard 0 = %q", shards[0])
|
||||
}
|
||||
}
|
||||
|
||||
// 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) {
|
||||
var got bytes.Buffer
|
||||
_ = SplitTSV(strings.NewReader("h\naaaa\nbbbb\n"), 2, func(_ int, shard io.Reader) error {
|
||||
_, _ = io.Copy(&got, shard)
|
||||
return nil
|
||||
})
|
||||
if want := "h\naaaa\nbbbb\n"; got.String() != want {
|
||||
t.Errorf("got %q, want %q", got.String(), want)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user