diff --git a/coordinator/api/requests.http b/coordinator/api/requests.http index 303c3a3..6c50a59 100644 --- a/coordinator/api/requests.http +++ b/coordinator/api/requests.http @@ -34,6 +34,41 @@ Content-Type: application/json @workerId = {{register.response.body.worker_id}} +### 0b. Upload a dataset — the coordinator splits it into shard tasks (201) +# Text fields first, the file part last (it is streamed, not buffered). +# @name uploadJob +POST {{host}}/jobs/upload +Authorization: Bearer {{token}} +Content-Type: multipart/form-data; boundary=----scimesh + +------scimesh +Content-Disposition: form-data; name="workload" + +similarity_search +------scimesh +Content-Disposition: form-data; name="parameters" + +{"top_k":10} +------scimesh +Content-Disposition: form-data; name="chunk_rows" + +2 +------scimesh +Content-Disposition: form-data; name="file"; filename="chembl.tsv" +Content-Type: text/tab-separated-values + +id smiles +A CC +B CCC +C CCCC +D CCCCC +------scimesh-- + +### Download a task's input shard (200) — taskId must be a shard task from an +### uploaded job (claim one first; its input.uri is /tasks/{id}/input). +GET {{host}}/tasks/{{taskId}}/input +Authorization: Bearer {{token}} + ### 1. Create a job and its chunks (201) # The coordinator splits the submission into one task per chunk, transactionally. # @name createJob diff --git a/coordinator/cmd/coordinator/main.go b/coordinator/cmd/coordinator/main.go index 95481fc..767ebeb 100644 --- a/coordinator/cmd/coordinator/main.go +++ b/coordinator/cmd/coordinator/main.go @@ -70,6 +70,7 @@ func run() error { useCases := httptransport.UseCases{ RegisterWorker: usecase.NewRegisterWorker(workerRepo, clk), CreateJob: usecase.NewCreateJob(jobRepo, taskRepo, tx, clk), + SubmitDataset: usecase.NewSubmitDataset(blobStore, artifactRepo, jobRepo, taskRepo, tx, clk), ClaimTask: usecase.NewClaimTask(taskRepo, clk, cfg.LeaseDuration), RenewLease: usecase.NewRenewLease(taskRepo, tx, clk, cfg.LeaseDuration), CompleteTask: usecase.NewCompleteTask(taskRepo, jobRepo, artifactRepo, tx, clk), @@ -77,6 +78,7 @@ func run() error { GetJobStatus: usecase.NewGetJobStatus(jobRepo, taskRepo), UploadArtifact: usecase.NewUploadArtifact(taskRepo, artifactRepo, blobStore, clk), DownloadArtifact: usecase.NewDownloadArtifact(artifactRepo, blobStore), + GetTaskInput: usecase.NewGetTaskInput(taskRepo, artifactRepo, blobStore), } // Background workers are tracked so shutdown can wait for them. Without diff --git a/coordinator/internal/chunk/tsv.go b/coordinator/internal/chunk/tsv.go new file mode 100644 index 0000000..6b968eb --- /dev/null +++ b/coordinator/internal/chunk/tsv.go @@ -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 +} diff --git a/coordinator/internal/chunk/tsv_test.go b/coordinator/internal/chunk/tsv_test.go new file mode 100644 index 0000000..ba12917 --- /dev/null +++ b/coordinator/internal/chunk/tsv_test.go @@ -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) + } +} diff --git a/coordinator/internal/domain/job.go b/coordinator/internal/domain/job.go index b6e9f4a..0aa4827 100644 --- a/coordinator/internal/domain/job.go +++ b/coordinator/internal/domain/job.go @@ -18,13 +18,31 @@ const ( // Job is one user submission that fans out into one or more tasks. type Job struct { - ID uuid.UUID - Workload string - InputURI string - Parameters map[string]any - Status JobStatus - CreatedAt time.Time - CompletedAt *time.Time + ID uuid.UUID + Workload string + InputURI string // external input URI; empty for uploaded datasets + InputArtifactID *uuid.UUID // uploaded input artifact; nil for URI submissions + Parameters map[string]any + Status JobStatus + CreatedAt time.Time + CompletedAt *time.Time +} + +// NewUploadedJob builds a job whose input was uploaded to the coordinator. The +// job's id is generated here so the input artifact can reference it; the reverse +// link (jobs.input_artifact_id) is left unset — the input is found via the +// artifact's job_id — which also sidesteps the circular job↔artifact FK. +func NewUploadedJob(workload string, params map[string]any, now time.Time) (*Job, error) { + if workload == "" { + return nil, ErrInvalidInput + } + return &Job{ + ID: uuid.New(), + Workload: workload, + Parameters: params, + Status: JobPending, + CreatedAt: now, + }, nil } // ChunkSpec describes one piece a job is split into. Callers build these from diff --git a/coordinator/internal/domain/task.go b/coordinator/internal/domain/task.go index 87fc555..8908e44 100644 --- a/coordinator/internal/domain/task.go +++ b/coordinator/internal/domain/task.go @@ -32,7 +32,8 @@ type Task struct { JobID uuid.UUID ChunkIndex int Workload string - InputURI string + InputURI string // external input URI; empty for uploaded shards + InputArtifactID *uuid.UUID // coordinator-stored shard; nil for URI inputs InputSHA256 string Parameters map[string]any Status TaskStatus @@ -81,6 +82,33 @@ func NewTask(jobID uuid.UUID, chunkIndex int, workload, inputURI, inputSHA256 st }, nil } +// NewShardTask builds a pending task whose input is a coordinator-stored shard +// artifact rather than an external URI. The worker fetches it from the +// coordinator, so no InputURI is set — inputSHA256 is the shard's checksum. +func NewShardTask(jobID uuid.UUID, chunkIndex int, workload string, inputArtifactID uuid.UUID, + inputSHA256 string, params map[string]any, maxAttempts int, now time.Time) (*Task, error) { + + if inputArtifactID == uuid.Nil || inputSHA256 == "" || chunkIndex < 0 { + return nil, ErrInvalidInput + } + if maxAttempts <= 0 { + maxAttempts = DefaultMaxAttempts + } + return &Task{ + ID: uuid.New(), + JobID: jobID, + ChunkIndex: chunkIndex, + Workload: workload, + InputArtifactID: &inputArtifactID, + InputSHA256: inputSHA256, + Parameters: params, + Status: TaskPending, + Attempt: 0, + MaxAttempts: maxAttempts, + CreatedAt: now, + }, nil +} + // DefaultMaxAttempts applies when a task does not specify its own ceiling. const DefaultMaxAttempts = 3 @@ -96,14 +124,15 @@ func (t *Task) IsLeaseHeldBy(worker string, attempt int) bool { // everything needed to execute, nothing it has no business seeing. func (t *Task) AsClaimed() ClaimedTask { ct := ClaimedTask{ - TaskID: t.ID, - JobID: t.JobID, - ChunkIndex: t.ChunkIndex, - Workload: t.Workload, - InputURI: t.InputURI, - InputSHA256: t.InputSHA256, - Parameters: t.Parameters, - Attempt: t.Attempt, + TaskID: t.ID, + JobID: t.JobID, + ChunkIndex: t.ChunkIndex, + Workload: t.Workload, + InputURI: t.InputURI, + InputArtifactID: t.InputArtifactID, + InputSHA256: t.InputSHA256, + Parameters: t.Parameters, + Attempt: t.Attempt, } if t.LeaseOwner != nil { ct.LeaseOwner = *t.LeaseOwner @@ -217,18 +246,21 @@ func (t *Task) ExpireLease(now time.Time) { t.CompletedAt = &now } -// ClaimedTask is the worker-facing projection of a leased task. +// ClaimedTask is the worker-facing projection of a leased task. Input is either +// an external URI or a coordinator-stored shard (InputArtifactID set); the +// transport turns the latter into a coordinator download URL. type ClaimedTask struct { - TaskID uuid.UUID - JobID uuid.UUID - ChunkIndex int - Workload string - InputURI string - InputSHA256 string - Parameters map[string]any - Attempt int - LeaseOwner string - LeaseExpiresAt time.Time + TaskID uuid.UUID + JobID uuid.UUID + ChunkIndex int + Workload string + InputURI string + InputArtifactID *uuid.UUID + InputSHA256 string + Parameters map[string]any + Attempt int + LeaseOwner string + LeaseExpiresAt time.Time } // ResultManifest is a completed task's output, ordered for the stitcher. It diff --git a/coordinator/internal/infra/logging.go b/coordinator/internal/infra/logging.go index 4d12e61..b5f49e3 100644 --- a/coordinator/internal/infra/logging.go +++ b/coordinator/internal/infra/logging.go @@ -29,7 +29,7 @@ func NewLogger(cfg Config) (*slog.Logger, io.Closer, error) { ) if cfg.LogFile != "" { - if err := os.MkdirAll(filepath.Dir(cfg.LogFile), 0o755); err != nil { + if err := os.MkdirAll(filepath.Dir(cfg.LogFile), 0o750); err != nil { return nil, nil, fmt.Errorf("create log directory: %w", err) } rotator := &lumberjack.Logger{ diff --git a/coordinator/internal/storage/blob/store.go b/coordinator/internal/storage/blob/store.go index 7c20eac..db98600 100644 --- a/coordinator/internal/storage/blob/store.go +++ b/coordinator/internal/storage/blob/store.go @@ -28,7 +28,7 @@ var _ usecase.BlobStore = (*FSStore)(nil) // rename is only atomic within one filesystem. func NewFSStore(dir string) (*FSStore, error) { staging := filepath.Join(dir, ".staging") - if err := os.MkdirAll(staging, 0o755); err != nil { + if err := os.MkdirAll(staging, 0o750); err != nil { return nil, fmt.Errorf("create blob dirs: %w", err) } return &FSStore{dir: dir, staging: staging}, nil @@ -85,7 +85,9 @@ func (s *FSStore) Open(ctx context.Context, key string) (io.ReadCloser, error) { if err := checkKey(key); err != nil { return nil, err } - f, err := os.Open(filepath.Join(s.dir, key)) + // checkKey has rejected any traversal, so the joined path stays under s.dir. + f, err := os.Open(filepath.Join(s.dir, key)) //nolint:gosec // key validated by checkKey + if err != nil { return nil, err } diff --git a/coordinator/internal/storage/postgres/task_repo.go b/coordinator/internal/storage/postgres/task_repo.go index d34f122..827ab77 100644 --- a/coordinator/internal/storage/postgres/task_repo.go +++ b/coordinator/internal/storage/postgres/task_repo.go @@ -30,7 +30,7 @@ var _ usecase.TaskRepository = (*TaskRepo)(nil) // Every query that returns a task selects exactly this list, in this order — // three hand-written column lists would drift apart within a week. var taskColumns = []string{ - "id", "job_id", "chunk_index", "workload", "input_uri", "input_sha256", + "id", "job_id", "chunk_index", "workload", "input_uri", "input_artifact_id", "input_sha256", "parameters", "status", "attempt", "max_attempts", "lease_owner", "lease_expires_at", "result_artifact_id", "metrics", "error_code", "error_message", "created_at", "started_at", "completed_at", "version", @@ -49,9 +49,12 @@ func scanTask(row pgx.Row) (*domain.Task, error) { var ( t domain.Task status string + // input_uri is nullable now (uploaded shards have none), so it cannot + // scan straight into a string; NULL becomes the empty InputURI. + inputURI *string ) err := row.Scan( - &t.ID, &t.JobID, &t.ChunkIndex, &t.Workload, &t.InputURI, &t.InputSHA256, + &t.ID, &t.JobID, &t.ChunkIndex, &t.Workload, &inputURI, &t.InputArtifactID, &t.InputSHA256, &t.Parameters, &status, &t.Attempt, &t.MaxAttempts, &t.LeaseOwner, &t.LeaseExpiresAt, &t.ResultArtifactID, &t.Metrics, &t.ErrorCode, &t.ErrorMessage, &t.CreatedAt, &t.StartedAt, &t.CompletedAt, &t.Version, @@ -59,6 +62,9 @@ func scanTask(row pgx.Row) (*domain.Task, error) { if err != nil { return nil, err } + if inputURI != nil { + t.InputURI = *inputURI + } t.Status = domain.TaskStatus(status) return &t, nil } @@ -206,10 +212,12 @@ func (r *TaskRepo) InsertBatch(ctx context.Context, tasks []*domain.Task) error batch := &pgx.Batch{} for _, t := range tasks { sql, args, err := psql.Insert("tasks"). - Columns("id", "job_id", "chunk_index", "workload", "input_uri", "input_sha256", - "parameters", "status", "attempt", "max_attempts", "created_at", "version"). - Values(t.ID, t.JobID, t.ChunkIndex, t.Workload, t.InputURI, t.InputSHA256, - jsonbOrEmpty(t.Parameters), string(t.Status), t.Attempt, t.MaxAttempts, t.CreatedAt, t.Version). + Columns("id", "job_id", "chunk_index", "workload", "input_uri", "input_artifact_id", + "input_sha256", "parameters", "status", "attempt", "max_attempts", "created_at", "version"). + // input_uri is stored NULL (not "") when empty, so the ck_tasks_has_input + // check actually bites: a task with neither a URI nor an artifact fails. + Values(t.ID, t.JobID, t.ChunkIndex, t.Workload, nullIfEmpty(t.InputURI), t.InputArtifactID, + t.InputSHA256, jsonbOrEmpty(t.Parameters), string(t.Status), t.Attempt, t.MaxAttempts, t.CreatedAt, t.Version). ToSql() if err != nil { return err diff --git a/coordinator/internal/storage/postgres/tx.go b/coordinator/internal/storage/postgres/tx.go index 4b70e17..38adaa7 100644 --- a/coordinator/internal/storage/postgres/tx.go +++ b/coordinator/internal/storage/postgres/tx.go @@ -84,6 +84,15 @@ func jsonbOrEmpty(m map[string]any) map[string]any { return m } +// nullIfEmpty maps "" to a SQL NULL, so an absent optional string is stored as +// NULL rather than an empty string that would defeat a NOT-NULL-or check. +func nullIfEmpty(s string) any { + if s == "" { + return nil + } + return s +} + // conn returns the transaction bound to ctx, or the pool when there is none. func conn(ctx context.Context, pool *pgxpool.Pool) querier { if tx, ok := ctx.Value(txKey{}).(pgx.Tx); ok { diff --git a/coordinator/internal/transport/http/dto.go b/coordinator/internal/transport/http/dto.go index 6e014a3..c652b25 100644 --- a/coordinator/internal/transport/http/dto.go +++ b/coordinator/internal/transport/http/dto.go @@ -88,18 +88,28 @@ type taskResponse struct { Status string `json:"status"` } +type inputRef struct { + URI string `json:"uri"` + SHA256 string `json:"sha256"` +} + type claimedTaskResponse struct { TaskID uuid.UUID `json:"task_id"` JobID uuid.UUID `json:"job_id"` ChunkIndex int `json:"chunk_index"` Workload string `json:"workload"` - InputURI string `json:"input_uri"` - InputSHA256 string `json:"input_sha256"` + Input inputRef `json:"input"` Parameters map[string]any `json:"parameters"` Attempt int `json:"attempt"` LeaseExpiresAt time.Time `json:"lease_expires_at"` } +type uploadJobResponse struct { + JobID uuid.UUID `json:"job_id"` + TaskCount int `json:"task_count"` + InputArtifactID uuid.UUID `json:"input_artifact_id"` +} + type jobProgressResponse struct { ID uuid.UUID `json:"id"` Status string `json:"status"` @@ -123,13 +133,18 @@ type errorResponse struct { } func toClaimedTaskResponse(c domain.ClaimedTask) claimedTaskResponse { + // A shard's input lives in the coordinator; hand the worker a URL to fetch + // it from. A URI-based task keeps its external URI. + uri := c.InputURI + if c.InputArtifactID != nil { + uri = "/tasks/" + c.TaskID.String() + "/input" + } return claimedTaskResponse{ TaskID: c.TaskID, JobID: c.JobID, ChunkIndex: c.ChunkIndex, Workload: c.Workload, - InputURI: c.InputURI, - InputSHA256: c.InputSHA256, + Input: inputRef{URI: uri, SHA256: c.InputSHA256}, Parameters: c.Parameters, Attempt: c.Attempt, LeaseExpiresAt: c.LeaseExpiresAt, diff --git a/coordinator/internal/transport/http/handlers.go b/coordinator/internal/transport/http/handlers.go index d622e0f..1165105 100644 --- a/coordinator/internal/transport/http/handlers.go +++ b/coordinator/internal/transport/http/handlers.go @@ -2,10 +2,13 @@ package http import ( "context" + "encoding/json" + "errors" "fmt" "io" "net/http" "strconv" + "strings" "github.com/google/uuid" @@ -175,6 +178,107 @@ func (s *Server) handleFailure(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusOK, taskResponse{ID: task.ID, JobID: task.JobID, Status: string(task.Status)}) } +// defaultChunkRows is the shard size used when a request omits chunk_rows. +const defaultChunkRows = 1000 + +// handleUploadDataset accepts a multipart submission — the dataset file plus the +// workload/parameters/chunk_rows fields — and hands the file, streamed, to the +// chunker. The text fields MUST precede the file part: the file is streamed, not +// buffered, so by the time it arrives the other fields are already parsed. +func (s *Server) handleUploadDataset(w http.ResponseWriter, r *http.Request) { + mr, err := r.MultipartReader() + if err != nil { + s.writeError(w, r, domain.ErrInvalidInput) + return + } + + var ( + workload string + params map[string]any + rows = defaultChunkRows + result usecase.SubmitDatasetResult + gotDataset bool + ) + + for { + part, err := mr.NextPart() + if errors.Is(err, io.EOF) { + break + } + if err != nil { + s.writeError(w, r, domain.ErrInvalidInput) + return + } + + switch part.FormName() { + case "workload": + b, _ := io.ReadAll(io.LimitReader(part, 1<<10)) + workload = strings.TrimSpace(string(b)) + case "parameters": + b, _ := io.ReadAll(io.LimitReader(part, 1<<16)) + if len(b) > 0 { + if err := json.Unmarshal(b, ¶ms); err != nil { + s.writeError(w, r, domain.ErrInvalidInput) + return + } + } + case "chunk_rows": + b, _ := io.ReadAll(io.LimitReader(part, 32)) + if n, err := strconv.Atoi(strings.TrimSpace(string(b))); err == nil { + rows = n + } + case "file", "dataset": + filename := part.FileName() + if filename == "" { + filename = "dataset" + } + result, err = s.uc.SubmitDataset.Execute(r.Context(), usecase.SubmitDatasetInput{ + Workload: workload, + Parameters: params, + RowsPerShard: rows, + Filename: filename, + ContentType: part.Header.Get("Content-Type"), + Body: part, + }) + if err != nil { + s.writeError(w, r, err) + return + } + gotDataset = true + } + _ = part.Close() + } + + if !gotDataset { + s.writeError(w, r, domain.ErrInvalidInput) // no file part + return + } + writeJSON(w, http.StatusCreated, uploadJobResponse{ + JobID: result.JobID, + TaskCount: result.TaskCount, + InputArtifactID: result.InputArtifactID, + }) +} + +// handleGetTaskInput streams a task's input shard back to the worker. +func (s *Server) handleGetTaskInput(w http.ResponseWriter, r *http.Request) { + taskID, ok := s.pathUUID(w, r, "task_id") + if !ok { + return + } + art, body, err := s.uc.GetTaskInput.Execute(r.Context(), taskID) + if err != nil { + s.writeError(w, r, err) + return + } + defer func() { _ = body.Close() }() + + w.Header().Set("Content-Type", art.ContentType) + w.Header().Set("Content-Length", strconv.FormatInt(art.SizeBytes, 10)) + w.Header().Set("X-Checksum-SHA256", art.SHA256) + _, _ = io.Copy(w, body) +} + // handleUploadArtifact streams a worker's partial result into blob storage. It // deliberately does not use the short request timeout — a large shard upload // would trip it — and reads identity from headers per the contract (§5.5). @@ -220,7 +324,7 @@ func (s *Server) handleDownloadArtifact(w http.ResponseWriter, r *http.Request) s.writeError(w, r, err) return } - defer body.Close() + defer func() { _ = body.Close() }() w.Header().Set("Content-Type", art.ContentType) w.Header().Set("Content-Length", strconv.FormatInt(art.SizeBytes, 10)) diff --git a/coordinator/internal/transport/http/server.go b/coordinator/internal/transport/http/server.go index 6a65d60..109392c 100644 --- a/coordinator/internal/transport/http/server.go +++ b/coordinator/internal/transport/http/server.go @@ -18,6 +18,7 @@ import ( type UseCases struct { RegisterWorker *usecase.RegisterWorker CreateJob *usecase.CreateJob + SubmitDataset *usecase.SubmitDataset ClaimTask *usecase.ClaimTask RenewLease *usecase.RenewLease CompleteTask *usecase.CompleteTask @@ -25,6 +26,7 @@ type UseCases struct { GetJobStatus *usecase.GetJobStatus UploadArtifact *usecase.UploadArtifact DownloadArtifact *usecase.DownloadArtifact + GetTaskInput *usecase.GetTaskInput } type Server struct { @@ -54,8 +56,10 @@ func (s *Server) Handler(token string) http.Handler { protected := http.NewServeMux() protected.HandleFunc("POST /workers/register", s.handleRegister) protected.HandleFunc("POST /jobs", s.handleCreateJob) + protected.HandleFunc("POST /jobs/upload", s.handleUploadDataset) protected.HandleFunc("GET /jobs/{job_id}", s.handleGetJob) protected.HandleFunc("POST /tasks/claim", s.handleClaim) + protected.HandleFunc("GET /tasks/{task_id}/input", s.handleGetTaskInput) protected.HandleFunc("POST /tasks/{task_id}/heartbeat", s.handleHeartbeat) protected.HandleFunc("POST /tasks/{task_id}/result", s.handleResult) protected.HandleFunc("POST /tasks/{task_id}/failure", s.handleFailure) diff --git a/coordinator/internal/usecase/dto.go b/coordinator/internal/usecase/dto.go index d784b71..6491cbf 100644 --- a/coordinator/internal/usecase/dto.go +++ b/coordinator/internal/usecase/dto.go @@ -49,6 +49,21 @@ type CompleteTaskInput struct { Metrics map[string]any } +type SubmitDatasetInput struct { + Workload string + Parameters map[string]any + RowsPerShard int + Filename string + ContentType string + Body io.Reader +} + +type SubmitDatasetResult struct { + JobID uuid.UUID + TaskCount int + InputArtifactID uuid.UUID +} + type UploadArtifactInput struct { TaskID uuid.UUID WorkerID string diff --git a/coordinator/internal/usecase/upload.go b/coordinator/internal/usecase/upload.go new file mode 100644 index 0000000..f04f582 --- /dev/null +++ b/coordinator/internal/usecase/upload.go @@ -0,0 +1,151 @@ +package usecase + +import ( + "context" + "fmt" + "io" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/chunk" + "github.com/emil28092005/SciMesh/coordinator/internal/domain" +) + +// SubmitDataset accepts an uploaded dataset, splits it into shard artifacts, and +// creates the job with one task per shard — the coordinator-side counterpart of +// a client submitting pre-chunked URIs. +type SubmitDataset struct { + blobs BlobStore + artifacts ArtifactRepository + jobs JobRepository + tasks TaskRepository + tx TxManager + clk Clock +} + +func NewSubmitDataset(blobs BlobStore, artifacts ArtifactRepository, jobs JobRepository, + tasks TaskRepository, tx TxManager, clk Clock) *SubmitDataset { + return &SubmitDataset{blobs: blobs, artifacts: artifacts, jobs: jobs, tasks: tasks, tx: tx, clk: clk} +} + +func (uc *SubmitDataset) Execute(ctx context.Context, in SubmitDatasetInput) (SubmitDatasetResult, error) { + now := uc.clk.Now() + + job, err := domain.NewUploadedJob(in.Workload, in.Parameters, now) + if err != nil { + return SubmitDatasetResult{}, err + } + + // Everything written to blob storage, so a failed transaction can undo it. + var putKeys []string + cleanup := func() { + for _, k := range putKeys { + _ = uc.blobs.Delete(ctx, k) + } + } + + // 1. Stream the upload into the input artifact; we measure size and sha256. + input, err := domain.NewArtifact(job.ID, nil, domain.ArtifactInput, in.Filename, in.ContentType, now) + if err != nil { + return SubmitDatasetResult{}, err + } + sum, size, err := uc.blobs.Put(ctx, input.StorageKey, in.Body) + if err != nil { + return SubmitDatasetResult{}, err + } + putKeys = append(putKeys, input.StorageKey) + input.SetContent(sum, size) + + // 2. Re-open the stored input and split it into shard artifacts + tasks. + shards := []*domain.Artifact{} + tasks := []*domain.Task{} + rc, err := uc.blobs.Open(ctx, input.StorageKey) + if err != nil { + cleanup() + return SubmitDatasetResult{}, err + } + splitErr := chunk.SplitTSV(rc, in.RowsPerShard, func(index int, shard io.Reader) error { + art, err := domain.NewArtifact(job.ID, nil, domain.ArtifactShard, + fmt.Sprintf("shard-%d.tsv", index), in.ContentType, now) + if err != nil { + return err + } + ssum, ssize, err := uc.blobs.Put(ctx, art.StorageKey, shard) + if err != nil { + return err + } + putKeys = append(putKeys, art.StorageKey) + art.SetContent(ssum, ssize) + + task, err := domain.NewShardTask(job.ID, index, in.Workload, art.ID, ssum, in.Parameters, 0, now) + if err != nil { + return err + } + shards = append(shards, art) + tasks = append(tasks, task) + return nil + }) + _ = rc.Close() + if splitErr != nil { + cleanup() + return SubmitDatasetResult{}, splitErr + } + + // 3. Persist job + all artifacts + all tasks atomically. + err = uc.tx.WithinTx(ctx, func(ctx context.Context) error { + if err := uc.jobs.Insert(ctx, job); err != nil { + return err + } + if err := uc.artifacts.Insert(ctx, input); err != nil { + return err + } + for _, a := range shards { + if err := uc.artifacts.Insert(ctx, a); err != nil { + return err + } + } + return uc.tasks.InsertBatch(ctx, tasks) + }) + if err != nil { + cleanup() + return SubmitDatasetResult{}, err + } + + return SubmitDatasetResult{ + JobID: job.ID, + TaskCount: len(tasks), + InputArtifactID: input.ID, + }, nil +} + +// GetTaskInput resolves a task's input shard and opens it for streaming. The +// caller closes the reader. +type GetTaskInput struct { + tasks TaskRepository + artifacts ArtifactRepository + blobs BlobStore +} + +func NewGetTaskInput(tasks TaskRepository, artifacts ArtifactRepository, blobs BlobStore) *GetTaskInput { + return &GetTaskInput{tasks: tasks, artifacts: artifacts, blobs: blobs} +} + +func (uc *GetTaskInput) Execute(ctx context.Context, taskID uuid.UUID) (*domain.Artifact, io.ReadCloser, error) { + task, err := uc.tasks.Get(ctx, taskID) + if err != nil { + return nil, nil, err + } + if task.InputArtifactID == nil { + // A URI-based task keeps its input outside the coordinator. + return nil, nil, domain.ErrArtifactNotFound + } + art, err := uc.artifacts.Get(ctx, *task.InputArtifactID) + if err != nil { + return nil, nil, err + } + rc, err := uc.blobs.Open(ctx, art.StorageKey) + if err != nil { + return nil, nil, err + } + return art, rc, nil +} diff --git a/coordinator/migrations/0005_uploaded_input.down.sql b/coordinator/migrations/0005_uploaded_input.down.sql new file mode 100644 index 0000000..87731e2 --- /dev/null +++ b/coordinator/migrations/0005_uploaded_input.down.sql @@ -0,0 +1,9 @@ +BEGIN; + +ALTER TABLE tasks DROP CONSTRAINT IF EXISTS ck_tasks_has_input; + +-- Restoring NOT NULL requires the columns to be populated; safe on a fresh DB. +ALTER TABLE tasks ALTER COLUMN input_uri SET NOT NULL; +ALTER TABLE jobs ALTER COLUMN input_uri SET NOT NULL; + +COMMIT; diff --git a/coordinator/migrations/0005_uploaded_input.up.sql b/coordinator/migrations/0005_uploaded_input.up.sql new file mode 100644 index 0000000..04d79fb --- /dev/null +++ b/coordinator/migrations/0005_uploaded_input.up.sql @@ -0,0 +1,13 @@ +BEGIN; + +-- Inputs can now arrive as uploaded artifacts (POST /jobs/upload), not only as +-- external URIs. Relax the URI requirement and require every task to have an +-- input one way or the other. +ALTER TABLE jobs ALTER COLUMN input_uri DROP NOT NULL; +ALTER TABLE tasks ALTER COLUMN input_uri DROP NOT NULL; + +ALTER TABLE tasks ADD CONSTRAINT ck_tasks_has_input CHECK ( + input_uri IS NOT NULL OR input_artifact_id IS NOT NULL +); + +COMMIT; diff --git a/coordinator/scripts/smoke.sh b/coordinator/scripts/smoke.sh index 762b784..2d9a1f2 100755 --- a/coordinator/scripts/smoke.sh +++ b/coordinator/scripts/smoke.sh @@ -147,6 +147,56 @@ check "unknown json field → 400" 400 -X POST "${HOST}/tasks/claim" " -d '{"worker_id":"w1","totally_unknown":1}' check "unknown job → 404" 404 "${HOST}/jobs/00000000-0000-0000-0000-000000000000" "${auth[@]}" +echo +echo "dataset upload → chunking" +# Upload a 5-row TSV split at 2 rows/shard → expect 3 shard tasks. The text +# fields precede the file part, which the coordinator streams. +up=$(curl -sS "${bearer[@]}" -X POST "${HOST}/jobs/upload" \ + -F 'workload=similarity_search' \ + -F 'parameters={"top_k":10}' \ + -F 'chunk_rows=2' \ + -F 'file=@-;filename=chembl.tsv;type=text/tab-separated-values' <<'TSV' +id smiles +A CC +B CCC +C CCCC +D CCCCC +E CCCCCC +TSV +) +up_job=$(printf '%s' "$up" | python3 -c 'import json,sys;print(json.load(sys.stdin)["job_id"])' 2>/dev/null) +up_count=$(printf '%s' "$up" | python3 -c 'import json,sys;print(json.load(sys.stdin)["task_count"])' 2>/dev/null) + +if [[ "$up_count" == "3" ]]; then + printf ' \033[32m✓\033[0m %-46s task_count=3\n' "POST /jobs/upload (5 rows / 2)" + pass=$((pass + 1)) +else + printf ' \033[31m✗\033[0m %-46s got task_count=%s, want 3\n' "POST /jobs/upload" "${up_count:-?}" + printf ' %s\n' "$up" + fail=$((fail + 1)) +fi + +# Claim one of this job's shard tasks and pull its input shard from the coordinator. +up_input="" +for _ in $(seq 1 30); do + c=$(curl -sS "${bearer[@]}" -H 'Content-Type: application/json' -X POST "${HOST}/tasks/claim" \ + -d '{"worker_id":"up-w","capabilities":["similarity_search"]}') + [[ -z "$c" ]] && break + cj=$(printf '%s' "$c" | python3 -c 'import json,sys;print(json.load(sys.stdin)["job_id"])' 2>/dev/null) + [[ "$cj" != "$up_job" ]] && continue + up_input=$(printf '%s' "$c" | python3 -c 'import json,sys;print(json.load(sys.stdin)["input"]["uri"])' 2>/dev/null) + break +done + +if [[ "$up_input" == /tasks/*/input ]]; then + printf ' \033[32m✓\033[0m %-46s %s\n' "claim → input.uri points at coordinator" "$up_input" + pass=$((pass + 1)) +else + printf ' \033[31m✗\033[0m %-46s got %q\n' "claim shard input.uri" "$up_input" + fail=$((fail + 1)) +fi +check "download shard input" 200 "${HOST}${up_input}" "${bearer[@]}" + echo curl -sS "${HOST}/jobs/${job_id}" "${auth[@]}" echo diff --git a/docs/api-contract.md b/docs/api-contract.md index e67346c..d025e7c 100644 --- a/docs/api-contract.md +++ b/docs/api-contract.md @@ -24,7 +24,8 @@ must be updated in the same change as any behaviour it describes. | `GET /jobs/{id}` | progress | ✅ done | | `PUT /tasks/{id}/artifacts/{name}` | upload partial | ✅ done | | `GET /artifacts/{id}/download` | download by id | ✅ done | -| `GET /tasks/{id}/input` | download shard | ❌ needs input artifacts (upload+chunking) | +| `POST /jobs/upload` | upload dataset, coordinator chunks it | ✅ done | +| `GET /tasks/{id}/input` | download shard | ✅ done | --- @@ -37,6 +38,28 @@ GET /health `200 {"status":"ok"}` when the database is reachable; `503 {"status":"unavailable"}` otherwise. Unauthenticated. +## Submit a dataset (submitter-side) + +```http +POST /jobs/upload +Authorization: Bearer +Content-Type: multipart/form-data +``` + +Fields, in order (text fields first, file last — the file is streamed): +`workload`, `parameters` (JSON), `chunk_rows` (int, default 1000), and the file +part `file`. The coordinator stores the input, splits the TSV into shard +artifacts (header repeated per shard), and creates one task per shard. + +`201`: + +```json +{ "job_id": "uuid", "task_count": 3, "input_artifact_id": "uuid" } +``` + +Each resulting task's claim response carries `input.uri = /tasks/{id}/input`, +served by §5.4. + ## Register worker ```http