From e5ba27951a86473de8f36a188aadabb8f187ec2f Mon Sep 17 00:00:00 2001 From: Efremenko Arhip Date: Thu, 23 Jul 2026 16:47:23 +0300 Subject: [PATCH] test(coordinator): unit + integration coverage across every layer MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - domain: NewJobWithTasks, DeriveStatus, NewUploadedJob, NewShardTask, NewWorker, NewArtifact/SetContent (domain 47% -> 88%). - internal/memstore: in-memory implementations of every usecase port, so orchestration can be tested without Postgres or a filesystem. - usecase: claim/renew/complete/fail/create/register/upload/submit-dataset flows over memstore, including rule-10 cross-task rejection, idempotent replay, lease sweep-on-claim, and dataset chunking (usecase 0% -> 73%). - transport: httptest end-to-end over real use cases + memstore — auth, readiness, full lifecycle, multipart upload + shard input, error mappings (0% -> 70%). - postgres integration: fix the tests broken by the artifact_id switch and add worker-repo, artifact-repo, and shard-task (nullable input_uri) round-trips. go test -race ./... is clean; golangci-lint (incl. integration tag) reports 0. --- coordinator/internal/domain/artifact_test.go | 55 +++ coordinator/internal/domain/job_test.go | 140 +++++++ coordinator/internal/domain/worker_test.go | 31 ++ coordinator/internal/memstore/memstore.go | 321 +++++++++++++++ .../storage/postgres/integration_test.go | 137 +++++- .../internal/transport/http/server_test.go | 310 ++++++++++++++ coordinator/internal/usecase/usecase_test.go | 389 ++++++++++++++++++ 7 files changed, 1372 insertions(+), 11 deletions(-) create mode 100644 coordinator/internal/domain/artifact_test.go create mode 100644 coordinator/internal/domain/job_test.go create mode 100644 coordinator/internal/domain/worker_test.go create mode 100644 coordinator/internal/memstore/memstore.go create mode 100644 coordinator/internal/transport/http/server_test.go create mode 100644 coordinator/internal/usecase/usecase_test.go diff --git a/coordinator/internal/domain/artifact_test.go b/coordinator/internal/domain/artifact_test.go new file mode 100644 index 0000000..da4bd22 --- /dev/null +++ b/coordinator/internal/domain/artifact_test.go @@ -0,0 +1,55 @@ +package domain + +import ( + "errors" + "testing" + + "github.com/google/uuid" +) + +func TestNewArtifact(t *testing.T) { + jobID := uuid.New() + taskID := uuid.New() + a, err := NewArtifact(jobID, &taskID, ArtifactPartialResult, "result.csv", "text/csv", testNow) + if err != nil { + t.Fatal(err) + } + if a.JobID != jobID || a.TaskID == nil || *a.TaskID != taskID { + t.Error("ownership not recorded") + } + // Storage key is derived from the artifact id, never the filename — no path + // traversal from a hostile "../.." name. + if a.StorageKey != a.ID.String() { + t.Errorf("storage key = %q, want the artifact id", a.StorageKey) + } + if a.SizeBytes != 0 || a.SHA256 != "" { + t.Error("size and checksum are unknown until SetContent") + } +} + +func TestNewArtifactDefaultsContentType(t *testing.T) { + a, err := NewArtifact(uuid.New(), nil, ArtifactInput, "data", "", testNow) + if err != nil { + t.Fatal(err) + } + if a.ContentType != "application/octet-stream" { + t.Errorf("content type = %q, want the default", a.ContentType) + } +} + +func TestNewArtifactRejectsBadInput(t *testing.T) { + if _, err := NewArtifact(uuid.New(), nil, ArtifactInput, "", "text/csv", testNow); !errors.Is(err, ErrInvalidInput) { + t.Errorf("empty filename: err = %v, want ErrInvalidInput", err) + } + if _, err := NewArtifact(uuid.New(), nil, "", "f", "text/csv", testNow); !errors.Is(err, ErrInvalidInput) { + t.Errorf("empty kind: err = %v, want ErrInvalidInput", err) + } +} + +func TestArtifactSetContent(t *testing.T) { + a, _ := NewArtifact(uuid.New(), nil, ArtifactShard, "shard-0.tsv", "text/csv", testNow) + a.SetContent("deadbeef", 42) + if a.SHA256 != "deadbeef" || a.SizeBytes != 42 { + t.Error("SetContent must record checksum and size") + } +} diff --git a/coordinator/internal/domain/job_test.go b/coordinator/internal/domain/job_test.go new file mode 100644 index 0000000..0b58902 --- /dev/null +++ b/coordinator/internal/domain/job_test.go @@ -0,0 +1,140 @@ +package domain + +import ( + "errors" + "testing" + + "github.com/google/uuid" +) + +func TestNewJobWithTasksBuildsBoth(t *testing.T) { + job, tasks, err := NewJobWithTasks("similarity_search", "s3://in", nil, []ChunkSpec{ + {ChunkIndex: 0, InputURI: "s3://c0", InputSHA256: "a"}, + {ChunkIndex: 1, InputURI: "s3://c1", InputSHA256: "b"}, + }, testNow) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(tasks) != 2 { + t.Fatalf("got %d tasks, want 2", len(tasks)) + } + for _, tk := range tasks { + if tk.JobID != job.ID { + t.Error("task not linked to job") + } + if tk.Workload != "similarity_search" { + t.Error("task should inherit the job workload") + } + } + if job.Status != JobPending { + t.Errorf("status = %q, want pending", job.Status) + } +} + +func TestNewJobWithTasksRejectsBadInput(t *testing.T) { + good := []ChunkSpec{{ChunkIndex: 0, InputURI: "s3://c0", InputSHA256: "a"}} + cases := map[string]struct { + workload string + inputURI string + chunks []ChunkSpec + }{ + "empty workload": {"", "s3://in", good}, + "empty input": {"w", "", good}, + "no chunks": {"w", "s3://in", nil}, + "duplicate index": {"w", "s3://in", []ChunkSpec{ + {ChunkIndex: 0, InputURI: "a", InputSHA256: "x"}, + {ChunkIndex: 0, InputURI: "b", InputSHA256: "y"}, + }}, + } + for name, c := range cases { + t.Run(name, func(t *testing.T) { + if _, _, err := NewJobWithTasks(c.workload, c.inputURI, nil, c.chunks, testNow); !errors.Is(err, ErrInvalidInput) { + t.Errorf("err = %v, want ErrInvalidInput", err) + } + }) + } +} + +func TestNewJobWithTasksInheritsAndOverridesWorkload(t *testing.T) { + _, tasks, err := NewJobWithTasks("base", "s3://in", nil, []ChunkSpec{ + {ChunkIndex: 0, InputURI: "a", InputSHA256: "x"}, + {ChunkIndex: 1, InputURI: "b", InputSHA256: "y", Workload: "special"}, + }, testNow) + if err != nil { + t.Fatal(err) + } + if tasks[0].Workload != "base" || tasks[1].Workload != "special" { + t.Errorf("workloads = %q, %q", tasks[0].Workload, tasks[1].Workload) + } +} + +func TestDeriveStatus(t *testing.T) { + cases := []struct { + name string + p JobProgress + want JobStatus + }{ + {"empty", JobProgress{Total: 0}, JobPending}, + {"all pending", JobProgress{Total: 3, Pending: 3}, JobPending}, + {"one leased", JobProgress{Total: 3, Pending: 2, Leased: 1}, JobRunning}, + {"partly done", JobProgress{Total: 3, Pending: 1, Done: 2}, JobRunning}, + {"all done", JobProgress{Total: 3, Done: 3}, JobCompleted}, + {"done and failed", JobProgress{Total: 3, Done: 2, Failed: 1}, JobFailed}, + {"failed but work remains", JobProgress{Total: 3, Pending: 1, Failed: 2}, JobRunning}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + if got := c.p.DeriveStatus(); got != c.want { + t.Errorf("DeriveStatus() = %q, want %q", got, c.want) + } + }) + } +} + +func TestNewUploadedJob(t *testing.T) { + job, err := NewUploadedJob("w", map[string]any{"k": 1}, testNow) + if err != nil { + t.Fatal(err) + } + if job.Status != JobPending || job.InputURI != "" { + t.Error("uploaded job should be pending with no input URI") + } + if _, err := NewUploadedJob("", nil, testNow); !errors.Is(err, ErrInvalidInput) { + t.Errorf("empty workload: err = %v, want ErrInvalidInput", err) + } +} + +func TestNewShardTask(t *testing.T) { + art := uuid.New() + task, err := NewShardTask(uuid.New(), 2, "w", art, "sha", nil, 0, testNow) + if err != nil { + t.Fatal(err) + } + if task.InputArtifactID == nil || *task.InputArtifactID != art { + t.Error("shard task must reference its input artifact") + } + if task.InputURI != "" { + t.Error("shard task must not carry a URI") + } + if task.MaxAttempts != DefaultMaxAttempts { + t.Errorf("maxAttempts = %d, want default %d", task.MaxAttempts, DefaultMaxAttempts) + } + + bad := []struct { + name string + art uuid.UUID + sha string + idx int + }{ + {"nil artifact", uuid.Nil, "sha", 0}, + {"empty sha", art, "", 0}, + {"negative index", art, "sha", -1}, + } + for _, c := range bad { + t.Run(c.name, func(t *testing.T) { + if _, err := NewShardTask(uuid.New(), c.idx, "w", c.art, c.sha, nil, 0, testNow); !errors.Is(err, ErrInvalidInput) { + t.Errorf("err = %v, want ErrInvalidInput", err) + } + }) + } +} diff --git a/coordinator/internal/domain/worker_test.go b/coordinator/internal/domain/worker_test.go new file mode 100644 index 0000000..d39716e --- /dev/null +++ b/coordinator/internal/domain/worker_test.go @@ -0,0 +1,31 @@ +package domain + +import ( + "errors" + "testing" +) + +func TestNewWorker(t *testing.T) { + w, err := NewWorker("lab-01", []string{"similarity_search"}, testNow) + if err != nil { + t.Fatal(err) + } + if w.Status != WorkerOnline { + t.Errorf("status = %q, want online", w.Status) + } + if w.ID.String() == "" { + t.Error("worker must get an id") + } + if !w.LastHeartbeatAt.Equal(testNow) || !w.CreatedAt.Equal(testNow) { + t.Error("timestamps must be stamped") + } +} + +func TestNewWorkerRejectsNoCapabilities(t *testing.T) { + if _, err := NewWorker("lab-01", nil, testNow); !errors.Is(err, ErrInvalidInput) { + t.Errorf("err = %v, want ErrInvalidInput", err) + } + if _, err := NewWorker("lab-01", []string{}, testNow); !errors.Is(err, ErrInvalidInput) { + t.Errorf("empty slice: err = %v, want ErrInvalidInput", err) + } +} diff --git a/coordinator/internal/memstore/memstore.go b/coordinator/internal/memstore/memstore.go new file mode 100644 index 0000000..71243ed --- /dev/null +++ b/coordinator/internal/memstore/memstore.go @@ -0,0 +1,321 @@ +// Package memstore holds in-memory implementations of the usecase ports for +// tests: they exercise use-case orchestration without a database or filesystem. +// The real invariants that depend on Postgres (SKIP LOCKED, row locking) are +// covered separately by the integration tests. +package memstore + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "io" + "sort" + "sync" + "time" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" + "github.com/emil28092005/SciMesh/coordinator/internal/usecase" +) + +// Clock returns a fixed, advanceable time. +type Clock struct{ t time.Time } + +func NewClock(t time.Time) *Clock { return &Clock{t: t} } +func (c *Clock) Now() time.Time { return c.t } +func (c *Clock) Advance(d time.Duration) { c.t = c.t.Add(d) } + +// Tx is a no-op transaction manager: the in-memory stores need no atomicity to +// be observed, so it simply runs the function. +type Tx struct{} + +func (Tx) WithinTx(ctx context.Context, fn func(ctx context.Context) error) error { return fn(ctx) } + +// --- TaskRepo ------------------------------------------------------------ + +type TaskRepo struct { + mu sync.Mutex + tasks map[uuid.UUID]*domain.Task +} + +func NewTaskRepo() *TaskRepo { return &TaskRepo{tasks: map[uuid.UUID]*domain.Task{}} } + +var _ usecase.TaskRepository = (*TaskRepo)(nil) + +// clone returns a copy so a caller's mutations do not touch stored state until +// Update — mirroring how a repository hands back detached entities. +func clone(t *domain.Task) *domain.Task { cp := *t; return &cp } + +func (r *TaskRepo) put(t *domain.Task) { + r.mu.Lock() + defer r.mu.Unlock() + r.tasks[t.ID] = clone(t) +} + +func (r *TaskRepo) ClaimNext(ctx context.Context, f usecase.ClaimFilter) (*domain.Task, error) { + r.mu.Lock() + defer r.mu.Unlock() + + var cands []*domain.Task + for _, t := range r.tasks { + if t.Status != domain.TaskPending || t.Attempt >= t.MaxAttempts { + continue + } + if len(f.Workloads) > 0 && !contains(f.Workloads, t.Workload) { + continue + } + cands = append(cands, t) + } + if len(cands) == 0 { + return nil, nil + } + sort.Slice(cands, func(i, j int) bool { + if cands[i].CreatedAt.Equal(cands[j].CreatedAt) { + return cands[i].ChunkIndex < cands[j].ChunkIndex + } + return cands[i].CreatedAt.Before(cands[j].CreatedAt) + }) + + t := cands[0] + t.Status = domain.TaskLeased + t.Attempt++ + owner := f.Owner + t.LeaseOwner = &owner + t.LeaseExpiresAt = &f.LeaseUntil + if t.StartedAt == nil { + t.StartedAt = &f.Now + } + t.Version++ + return clone(t), nil +} + +func (r *TaskRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Task, error) { + r.mu.Lock() + defer r.mu.Unlock() + t, ok := r.tasks[id] + if !ok { + return nil, domain.ErrTaskNotFound + } + return clone(t), nil +} + +func (r *TaskRepo) GetForUpdate(ctx context.Context, id uuid.UUID) (*domain.Task, error) { + return r.Get(ctx, id) +} + +func (r *TaskRepo) Update(ctx context.Context, t *domain.Task) error { + r.mu.Lock() + defer r.mu.Unlock() + stored, ok := r.tasks[t.ID] + if !ok || stored.Version != t.Version-1 { + return domain.ErrLeaseConflict // vanished or advanced under us + } + r.tasks[t.ID] = clone(t) + return nil +} + +func (r *TaskRepo) InsertBatch(ctx context.Context, tasks []*domain.Task) error { + for _, t := range tasks { + r.put(t) + } + return nil +} + +func (r *TaskRepo) ListCompleted(ctx context.Context, jobID uuid.UUID) ([]*domain.Task, error) { + r.mu.Lock() + defer r.mu.Unlock() + var out []*domain.Task + for _, t := range r.tasks { + if t.JobID == jobID && t.Status == domain.TaskCompleted { + out = append(out, clone(t)) + } + } + sort.Slice(out, func(i, j int) bool { return out[i].ChunkIndex < out[j].ChunkIndex }) + return out, nil +} + +func (r *TaskRepo) CountByStatus(ctx context.Context, jobID uuid.UUID) (map[domain.TaskStatus]int, error) { + r.mu.Lock() + defer r.mu.Unlock() + counts := map[domain.TaskStatus]int{} + for _, t := range r.tasks { + if t.JobID == jobID { + counts[t.Status]++ + } + } + return counts, nil +} + +func (r *TaskRepo) ExpireLeases(ctx context.Context, now time.Time) (int64, error) { + r.mu.Lock() + defer r.mu.Unlock() + var n int64 + for _, t := range r.tasks { + if t.Status == domain.TaskLeased && t.LeaseExpiresAt != nil && t.LeaseExpiresAt.Before(now) { + t.ExpireLease(now) + n++ + } + } + return n, nil +} + +// --- JobRepo ------------------------------------------------------------- + +type JobRepo struct { + mu sync.Mutex + jobs map[uuid.UUID]*domain.Job +} + +func NewJobRepo() *JobRepo { return &JobRepo{jobs: map[uuid.UUID]*domain.Job{}} } + +var _ usecase.JobRepository = (*JobRepo)(nil) + +func (r *JobRepo) Insert(ctx context.Context, j *domain.Job) error { + r.mu.Lock() + defer r.mu.Unlock() + cp := *j + r.jobs[j.ID] = &cp + return nil +} + +func (r *JobRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Job, error) { + r.mu.Lock() + defer r.mu.Unlock() + j, ok := r.jobs[id] + if !ok { + return nil, domain.ErrJobNotFound + } + cp := *j + return &cp, nil +} + +func (r *JobRepo) UpdateStatus(ctx context.Context, id uuid.UUID, status domain.JobStatus, completedAt *time.Time) error { + r.mu.Lock() + defer r.mu.Unlock() + j, ok := r.jobs[id] + if !ok { + return domain.ErrJobNotFound + } + j.Status = status + j.CompletedAt = completedAt + return nil +} + +// --- WorkerRepo ---------------------------------------------------------- + +type WorkerRepo struct { + mu sync.Mutex + workers map[uuid.UUID]*domain.Worker +} + +func NewWorkerRepo() *WorkerRepo { return &WorkerRepo{workers: map[uuid.UUID]*domain.Worker{}} } + +var _ usecase.WorkerRepository = (*WorkerRepo)(nil) + +func (r *WorkerRepo) Insert(ctx context.Context, w *domain.Worker) error { + r.mu.Lock() + defer r.mu.Unlock() + cp := *w + r.workers[w.ID] = &cp + return nil +} + +func (r *WorkerRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Worker, error) { + r.mu.Lock() + defer r.mu.Unlock() + w, ok := r.workers[id] + if !ok { + return nil, domain.ErrWorkerNotFound + } + cp := *w + return &cp, nil +} + +// --- ArtifactRepo -------------------------------------------------------- + +type ArtifactRepo struct { + mu sync.Mutex + arts map[uuid.UUID]*domain.Artifact +} + +func NewArtifactRepo() *ArtifactRepo { return &ArtifactRepo{arts: map[uuid.UUID]*domain.Artifact{}} } + +var _ usecase.ArtifactRepository = (*ArtifactRepo)(nil) + +func (r *ArtifactRepo) Insert(ctx context.Context, a *domain.Artifact) error { + r.mu.Lock() + defer r.mu.Unlock() + cp := *a + r.arts[a.ID] = &cp + return nil +} + +func (r *ArtifactRepo) Get(ctx context.Context, id uuid.UUID) (*domain.Artifact, error) { + r.mu.Lock() + defer r.mu.Unlock() + a, ok := r.arts[id] + if !ok { + return nil, domain.ErrArtifactNotFound + } + cp := *a + return &cp, nil +} + +// --- BlobStore ----------------------------------------------------------- + +type BlobStore struct { + mu sync.Mutex + blobs map[string][]byte +} + +func NewBlobStore() *BlobStore { return &BlobStore{blobs: map[string][]byte{}} } + +var _ usecase.BlobStore = (*BlobStore)(nil) + +func (b *BlobStore) Put(ctx context.Context, key string, r io.Reader) (string, int64, error) { + data, err := io.ReadAll(r) + if err != nil { + return "", 0, err + } + sum := sha256.Sum256(data) + b.mu.Lock() + b.blobs[key] = data + b.mu.Unlock() + return hex.EncodeToString(sum[:]), int64(len(data)), nil +} + +func (b *BlobStore) Open(ctx context.Context, key string) (io.ReadCloser, error) { + b.mu.Lock() + defer b.mu.Unlock() + data, ok := b.blobs[key] + if !ok { + return nil, domain.ErrArtifactNotFound + } + return io.NopCloser(bytes.NewReader(data)), nil +} + +func (b *BlobStore) Delete(ctx context.Context, key string) error { + b.mu.Lock() + defer b.mu.Unlock() + delete(b.blobs, key) + return nil +} + +// Has reports whether a blob exists — handy for asserting cleanup in tests. +func (b *BlobStore) Has(key string) bool { + b.mu.Lock() + defer b.mu.Unlock() + _, ok := b.blobs[key] + return ok +} + +func contains(ss []string, s string) bool { + for _, x := range ss { + if x == s { + return true + } + } + return false +} diff --git a/coordinator/internal/storage/postgres/integration_test.go b/coordinator/internal/storage/postgres/integration_test.go index d6ec8a7..6ee522e 100644 --- a/coordinator/internal/storage/postgres/integration_test.go +++ b/coordinator/internal/storage/postgres/integration_test.go @@ -244,13 +244,25 @@ func TestListCompletedIsOrderedByChunkIndex(t *testing.T) { ctx := context.Background() job, tasks := seedJob(t, pool, 4) - repo, tx := NewTaskRepo(pool), NewTxManager(pool) + repo, artifacts, tx := NewTaskRepo(pool), NewArtifactRepo(pool), NewTxManager(pool) now := time.Now().UTC() // Complete them out of order to prove the ordering comes from SQL. for _, i := range []int{2, 0, 3, 1} { task := tasks[i] err := tx.WithinTx(ctx, func(ctx context.Context) error { + // A completed task must reference a real result artifact (FK + check). + taskID := task.ID + art, err := domain.NewArtifact(job.ID, &taskID, domain.ArtifactPartialResult, + fmt.Sprintf("result-%d.csv", task.ChunkIndex), "text/csv", now) + if err != nil { + return err + } + art.SetContent(fmt.Sprintf("rsha-%d", task.ChunkIndex), 1) + if err := artifacts.Insert(ctx, art); err != nil { + return err + } + fresh, err := repo.GetForUpdate(ctx, task.ID) if err != nil { return err @@ -260,11 +272,7 @@ func TestListCompletedIsOrderedByChunkIndex(t *testing.T) { fresh.LeaseOwner = &owner expires := now.Add(time.Minute) fresh.LeaseExpiresAt = &expires - if err := fresh.CompleteWith( - fmt.Sprintf("s3://result-%d", fresh.ChunkIndex), - fmt.Sprintf("rsha-%d", fresh.ChunkIndex), - nil, owner, fresh.Attempt, now, - ); err != nil { + if err := fresh.CompleteWith(art.ID, nil, owner, fresh.Attempt, now); err != nil { return err } return repo.Update(ctx, fresh) @@ -296,9 +304,9 @@ func TestCompleteTaskReplayIsIdempotent(t *testing.T) { ctx := context.Background() job, _ := seedJob(t, pool, 1) - tasks, jobs, tx := NewTaskRepo(pool), NewJobRepo(pool), NewTxManager(pool) + tasks, jobs, artifacts, tx := NewTaskRepo(pool), NewJobRepo(pool), NewArtifactRepo(pool), NewTxManager(pool) clk := fixedClock{now: time.Now().UTC()} - uc := usecase.NewCompleteTask(tasks, jobs, tx, clk) + uc := usecase.NewCompleteTask(tasks, jobs, artifacts, tx, clk) claimed, err := tasks.ClaimNext(ctx, usecase.ClaimFilter{ Owner: "worker-1", Now: clk.now, LeaseUntil: clk.now.Add(time.Minute), @@ -307,9 +315,12 @@ func TestCompleteTaskReplayIsIdempotent(t *testing.T) { t.Skipf("could not claim this job's task (got %v, %v)", claimed, err) } + // A partial-result artifact the coordinator stored for this task. + art := seedArtifact(t, pool, job.ID, &claimed.ID, domain.ArtifactPartialResult) + in := usecase.CompleteTaskInput{ TaskID: claimed.ID, WorkerID: "worker-1", Attempt: claimed.Attempt, - ResultURI: "s3://r0", ResultSHA256: "rsha", + ResultArtifactID: art.ID, } if _, err := uc.Execute(ctx, in); err != nil { t.Fatalf("first submission: %v", err) @@ -318,9 +329,10 @@ func TestCompleteTaskReplayIsIdempotent(t *testing.T) { t.Errorf("replay must be idempotent, got %v", err) } - // A different manifest for the same task is a genuine conflict. + // A different result artifact for the same task is a genuine conflict. + art2 := seedArtifact(t, pool, job.ID, &claimed.ID, domain.ArtifactPartialResult) other := in - other.ResultURI = "s3://different" + other.ResultArtifactID = art2.ID if _, err := uc.Execute(ctx, other); !errors.Is(err, domain.ErrResultConflict) { t.Errorf("err = %v, want ErrResultConflict", err) } @@ -330,6 +342,109 @@ type fixedClock struct{ now time.Time } func (c fixedClock) Now() time.Time { return c.now } +// seedArtifact inserts an artifact and returns it, cleaned up with its job. +func seedArtifact(t *testing.T, pool *pgxpool.Pool, jobID uuid.UUID, taskID *uuid.UUID, kind domain.ArtifactKind) *domain.Artifact { + t.Helper() + art, err := domain.NewArtifact(jobID, taskID, kind, "f.csv", "text/csv", time.Now().UTC()) + if err != nil { + t.Fatalf("build artifact: %v", err) + } + art.SetContent(fmt.Sprintf("sha-%s", art.ID), 3) + if err := NewArtifactRepo(pool).Insert(context.Background(), art); err != nil { + t.Fatalf("insert artifact: %v", err) + } + return art +} + +func TestWorkerRepoRoundTrip(t *testing.T) { + pool := testPool(t) + ctx := context.Background() + repo := NewWorkerRepo(pool) + + w, err := domain.NewWorker("lab-int", []string{"similarity_search", "similarity_graph"}, time.Now().UTC()) + if err != nil { + t.Fatal(err) + } + if err := repo.Insert(ctx, w); err != nil { + t.Fatalf("insert: %v", err) + } + t.Cleanup(func() { _, _ = pool.Exec(context.Background(), `DELETE FROM workers WHERE id = $1`, w.ID) }) + + got, err := repo.Get(ctx, w.ID) + if err != nil { + t.Fatalf("get: %v", err) + } + if got.Status != domain.WorkerOnline || len(got.Capabilities) != 2 { + t.Errorf("round-trip mismatch: %+v", got) + } + // capabilities must survive the jsonb round-trip. + if got.Capabilities[0] != "similarity_search" { + t.Errorf("capabilities = %v", got.Capabilities) + } + + if _, err := repo.Get(ctx, uuid.New()); !errors.Is(err, domain.ErrWorkerNotFound) { + t.Errorf("missing worker err = %v, want ErrWorkerNotFound", err) + } +} + +func TestArtifactRepoRoundTrip(t *testing.T) { + pool := testPool(t) + ctx := context.Background() + job, _ := seedJob(t, pool, 1) + + art := seedArtifact(t, pool, job.ID, nil, domain.ArtifactInput) + got, err := NewArtifactRepo(pool).Get(ctx, art.ID) + if err != nil { + t.Fatalf("get: %v", err) + } + if got.Kind != domain.ArtifactInput || got.StorageKey != art.StorageKey || got.SizeBytes != 3 { + t.Errorf("round-trip mismatch: %+v", got) + } + if _, err := NewArtifactRepo(pool).Get(ctx, uuid.New()); !errors.Is(err, domain.ErrArtifactNotFound) { + t.Errorf("missing artifact err = %v, want ErrArtifactNotFound", err) + } +} + +// A shard task stores its input as an artifact and no URI: this exercises the +// nullable input_uri column, the input_artifact_id round-trip, and the +// ck_tasks_has_input check that requires one or the other. +func TestShardTaskRoundTrip(t *testing.T) { + pool := testPool(t) + ctx := context.Background() + + job, err := domain.NewUploadedJob("similarity_search", nil, time.Now().UTC()) + if err != nil { + t.Fatal(err) + } + jobs, taskRepo, tx := NewJobRepo(pool), NewTaskRepo(pool), NewTxManager(pool) + if err := jobs.Insert(ctx, job); err != nil { + t.Fatalf("insert job: %v", err) + } + t.Cleanup(func() { _, _ = pool.Exec(context.Background(), `DELETE FROM jobs WHERE id = $1`, job.ID) }) + + shard := seedArtifact(t, pool, job.ID, nil, domain.ArtifactShard) + task, err := domain.NewShardTask(job.ID, 0, "similarity_search", shard.ID, shard.SHA256, nil, 0, time.Now().UTC()) + if err != nil { + t.Fatal(err) + } + if err := tx.WithinTx(ctx, func(ctx context.Context) error { + return taskRepo.InsertBatch(ctx, []*domain.Task{task}) + }); err != nil { + t.Fatalf("insert shard task: %v", err) + } + + got, err := taskRepo.Get(ctx, task.ID) + if err != nil { + t.Fatalf("get: %v", err) + } + if got.InputArtifactID == nil || *got.InputArtifactID != shard.ID { + t.Errorf("input_artifact_id did not round-trip: %v", got.InputArtifactID) + } + if got.InputURI != "" { + t.Errorf("shard task input_uri = %q, want empty (NULL)", got.InputURI) + } +} + func TestExpireLeasesRequeuesElapsedTasks(t *testing.T) { pool := testPool(t) ctx := context.Background() diff --git a/coordinator/internal/transport/http/server_test.go b/coordinator/internal/transport/http/server_test.go new file mode 100644 index 0000000..89aef8a --- /dev/null +++ b/coordinator/internal/transport/http/server_test.go @@ -0,0 +1,310 @@ +package http_test + +import ( + "bytes" + "context" + "encoding/json" + "io" + "log/slog" + "mime/multipart" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + "time" + + "github.com/emil28092005/SciMesh/coordinator/internal/memstore" + coordhttp "github.com/emil28092005/SciMesh/coordinator/internal/transport/http" + "github.com/emil28092005/SciMesh/coordinator/internal/usecase" +) + +const token = "secret" + +type env struct { + ts *httptest.Server + blobs *memstore.BlobStore +} + +func newEnv(t *testing.T, ready func(context.Context) error) *env { + t.Helper() + tasks := memstore.NewTaskRepo() + jobs := memstore.NewJobRepo() + work := memstore.NewWorkerRepo() + arts := memstore.NewArtifactRepo() + blobs := memstore.NewBlobStore() + clk := memstore.NewClock(time.Date(2026, 7, 21, 12, 0, 0, 0, time.UTC)) + tx := memstore.Tx{} + lease := 2 * time.Minute + + uc := coordhttp.UseCases{ + RegisterWorker: usecase.NewRegisterWorker(work, clk), + CreateJob: usecase.NewCreateJob(jobs, tasks, tx, clk), + SubmitDataset: usecase.NewSubmitDataset(blobs, arts, jobs, tasks, tx, clk), + ClaimTask: usecase.NewClaimTask(tasks, clk, lease), + RenewLease: usecase.NewRenewLease(tasks, tx, clk, lease), + CompleteTask: usecase.NewCompleteTask(tasks, jobs, arts, tx, clk), + FailTask: usecase.NewFailTask(tasks, jobs, tx, clk), + GetJobStatus: usecase.NewGetJobStatus(jobs, tasks), + UploadArtifact: usecase.NewUploadArtifact(tasks, arts, blobs, clk), + DownloadArtifact: usecase.NewDownloadArtifact(arts, blobs), + GetTaskInput: usecase.NewGetTaskInput(tasks, arts, blobs), + } + srv := coordhttp.NewServer(uc, slog.New(slog.NewTextHandler(io.Discard, nil)), 5*time.Second, 15*time.Second, ready) + ts := httptest.NewServer(srv.Handler(token)) + t.Cleanup(ts.Close) + return &env{ts: ts, blobs: blobs} +} + +func healthy(context.Context) error { return nil } + +// do sends an authenticated JSON request and returns status + decoded body. +func (e *env) do(t *testing.T, method, path, body string) (int, map[string]any) { + t.Helper() + req, _ := http.NewRequestWithContext(context.Background(), method, e.ts.URL+path, strings.NewReader(body)) + req.Header.Set("Authorization", "Bearer "+token) + if body != "" { + req.Header.Set("Content-Type", "application/json") + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("%s %s: %v", method, path, err) + } + defer resp.Body.Close() + var m map[string]any + b, _ := io.ReadAll(resp.Body) + _ = json.Unmarshal(b, &m) + return resp.StatusCode, m +} + +// get issues an unauthenticated GET and returns the response, failing on error. +func (e *env) get(t *testing.T, path string) *http.Response { + t.Helper() + req, _ := http.NewRequestWithContext(context.Background(), "GET", e.ts.URL+path, nil) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("GET %s: %v", path, err) + } + return resp +} + +func TestHealthOK(t *testing.T) { + e := newEnv(t, healthy) + resp := e.get(t, "/health") // unauthenticated + defer resp.Body.Close() + if resp.StatusCode != 200 { + t.Errorf("status = %d, want 200", resp.StatusCode) + } +} + +func TestHealthUnavailableWhenDBDown(t *testing.T) { + e := newEnv(t, func(context.Context) error { return context.DeadlineExceeded }) + resp := e.get(t, "/health") + defer resp.Body.Close() + if resp.StatusCode != http.StatusServiceUnavailable { + t.Errorf("status = %d, want 503", resp.StatusCode) + } +} + +func TestAuthRequired(t *testing.T) { + e := newEnv(t, healthy) + send := func(authz string) int { + req, _ := http.NewRequestWithContext(context.Background(), "POST", e.ts.URL+"/tasks/claim", + strings.NewReader(`{"worker_id":"w1"}`)) + req.Header.Set("Content-Type", "application/json") + if authz != "" { + req.Header.Set("Authorization", authz) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("claim: %v", err) + } + defer resp.Body.Close() + return resp.StatusCode + } + if code := send(""); code != 401 { + t.Errorf("no token: status = %d, want 401", code) + } + if code := send("Bearer nope"); code != 401 { + t.Errorf("wrong token: status = %d, want 401", code) + } +} + +func TestRegisterWorker(t *testing.T) { + e := newEnv(t, healthy) + code, body := e.do(t, "POST", "/workers/register", `{"name":"lab","capabilities":["w"]}`) + if code != 201 { + t.Fatalf("status = %d, want 201", code) + } + if body["worker_id"] == nil || body["heartbeat_interval_seconds"] == nil { + t.Errorf("missing fields in %v", body) + } +} + +func TestRegisterRejectsNoCapabilities(t *testing.T) { + e := newEnv(t, healthy) + if code, _ := e.do(t, "POST", "/workers/register", `{"name":"lab"}`); code != 400 { + t.Errorf("status = %d, want 400", code) + } +} + +func TestFullLifecycle(t *testing.T) { + e := newEnv(t, healthy) + + // Create a one-chunk job. + code, job := e.do(t, "POST", "/jobs", `{ + "workload":"w","input_uri":"s3://in", + "chunks":[{"chunk_index":0,"input_uri":"s3://c0","input_sha256":"sha"}]}`) + if code != 201 { + t.Fatalf("create job: %d", code) + } + jobID := job["id"].(string) + + // Claim it. + code, claim := e.do(t, "POST", "/tasks/claim", `{"worker_id":"w1","capabilities":["w"]}`) + if code != 200 { + t.Fatalf("claim: %d", code) + } + taskID := claim["task_id"].(string) + attempt := int(claim["attempt"].(float64)) + + // Heartbeat. + if code, _ := e.do(t, "POST", "/tasks/"+taskID+"/heartbeat", + `{"worker_id":"w1","attempt":`+itoa(attempt)+`}`); code != 200 { + t.Fatalf("heartbeat: %d", code) + } + + // Upload a result artifact (PUT, headers carry identity). + artID := e.putArtifact(t, taskID, "w1", attempt, "q,m\nA,B\n") + + // Submit the result by artifact id. + if code, _ := e.do(t, "POST", "/tasks/"+taskID+"/result", + `{"worker_id":"w1","attempt":`+itoa(attempt)+`,"result":{"artifact_id":"`+artID+`"}}`); code != 200 { + t.Fatalf("result: %d", code) + } + + // Job is now completed. + code, prog := e.do(t, "GET", "/jobs/"+jobID, "") + if code != 200 || prog["status"] != "completed" { + t.Errorf("job status = %v (code %d), want completed", prog["status"], code) + } +} + +func TestForeignArtifactResultConflict(t *testing.T) { + e := newEnv(t, healthy) + e.do(t, "POST", "/jobs", `{"workload":"w","input_uri":"s3://in", + "chunks":[{"chunk_index":0,"input_uri":"s3://c0","input_sha256":"sha"}, + {"chunk_index":1,"input_uri":"s3://c1","input_sha256":"sha"}]}`) + + _, cA := e.do(t, "POST", "/tasks/claim", `{"worker_id":"w1","capabilities":["w"]}`) + _, cB := e.do(t, "POST", "/tasks/claim", `{"worker_id":"w1","capabilities":["w"]}`) + taskA, attA := cA["task_id"].(string), int(cA["attempt"].(float64)) + taskB, attB := cB["task_id"].(string), int(cB["attempt"].(float64)) + artA := e.putArtifact(t, taskA, "w1", attA, "data") + + // Complete taskB with taskA's artifact → 409. + if code, _ := e.do(t, "POST", "/tasks/"+taskB+"/result", + `{"worker_id":"w1","attempt":`+itoa(attB)+`,"result":{"artifact_id":"`+artA+`"}}`); code != 409 { + t.Errorf("cross-task result: status = %d, want 409", code) + } +} + +func TestUploadDatasetChunksAndServesInput(t *testing.T) { + e := newEnv(t, healthy) + tsv := "id\tsmiles\nA\tCC\nB\tCCC\nC\tCCCC\nD\tCCCCC\nE\tCCCCCC\n" + + code, body := e.uploadDataset(t, "w", 2, tsv) + if code != 201 { + t.Fatalf("upload: status = %d", code) + } + if int(body["task_count"].(float64)) != 3 { + t.Fatalf("task_count = %v, want 3", body["task_count"]) + } + + // Claim a shard, follow its input.uri, and pull the shard bytes. + _, claim := e.do(t, "POST", "/tasks/claim", `{"worker_id":"w1","capabilities":["w"]}`) + input := claim["input"].(map[string]any) + uri := input["uri"].(string) + if !strings.HasPrefix(uri, "/tasks/") || !strings.HasSuffix(uri, "/input") { + t.Fatalf("input.uri = %q", uri) + } + req, _ := http.NewRequestWithContext(context.Background(), "GET", e.ts.URL+uri, nil) + req.Header.Set("Authorization", "Bearer "+token) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("get input: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != 200 { + t.Fatalf("get input: status = %d", resp.StatusCode) + } + shard, _ := io.ReadAll(resp.Body) + if !strings.HasPrefix(string(shard), "id\tsmiles\n") { + t.Errorf("shard missing header: %q", shard) + } +} + +func TestErrorMappings(t *testing.T) { + e := newEnv(t, healthy) + zero := "00000000-0000-0000-0000-000000000000" + + if code, _ := e.do(t, "GET", "/jobs/"+zero, ""); code != 404 { + t.Errorf("unknown job: %d, want 404", code) + } + if code, _ := e.do(t, "POST", "/tasks/not-a-uuid/heartbeat", `{"worker_id":"w1","attempt":1}`); code != 400 { + t.Errorf("malformed uuid: %d, want 400", code) + } + if code, _ := e.do(t, "POST", "/tasks/claim", `{"worker_id":"w1","totally_unknown":1}`); code != 400 { + t.Errorf("unknown field: %d, want 400", code) + } +} + +// --- helpers ------------------------------------------------------------- + +func (e *env) putArtifact(t *testing.T, taskID, worker string, attempt int, data string) string { + t.Helper() + req, _ := http.NewRequestWithContext(context.Background(), "PUT", + e.ts.URL+"/tasks/"+taskID+"/artifacts/r.csv", strings.NewReader(data)) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", "text/csv") + req.Header.Set("X-Worker-ID", worker) + req.Header.Set("X-Task-Attempt", itoa(attempt)) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if resp.StatusCode != 200 { + t.Fatalf("put artifact: status = %d", resp.StatusCode) + } + var m map[string]any + b, _ := io.ReadAll(resp.Body) + _ = json.Unmarshal(b, &m) + return m["artifact_id"].(string) +} + +func (e *env) uploadDataset(t *testing.T, workload string, rows int, tsv string) (int, map[string]any) { + t.Helper() + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + _ = mw.WriteField("workload", workload) + _ = mw.WriteField("chunk_rows", itoa(rows)) + fw, _ := mw.CreateFormFile("file", "chembl.tsv") + _, _ = io.Copy(fw, strings.NewReader(tsv)) + _ = mw.Close() + + req, _ := http.NewRequestWithContext(context.Background(), "POST", e.ts.URL+"/jobs/upload", &buf) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", mw.FormDataContentType()) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + var m map[string]any + b, _ := io.ReadAll(resp.Body) + _ = json.Unmarshal(b, &m) + return resp.StatusCode, m +} + +func itoa(n int) string { return strconv.Itoa(n) } diff --git a/coordinator/internal/usecase/usecase_test.go b/coordinator/internal/usecase/usecase_test.go new file mode 100644 index 0000000..9c3b458 --- /dev/null +++ b/coordinator/internal/usecase/usecase_test.go @@ -0,0 +1,389 @@ +package usecase_test + +import ( + "context" + "errors" + "fmt" + "strings" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/emil28092005/SciMesh/coordinator/internal/domain" + "github.com/emil28092005/SciMesh/coordinator/internal/memstore" + "github.com/emil28092005/SciMesh/coordinator/internal/usecase" +) + +var ctx = context.Background() + +const lease = 2 * time.Minute + +// harness wires every use case to in-memory stores so orchestration can be +// tested without a database. +type harness struct { + tasks *memstore.TaskRepo + jobs *memstore.JobRepo + work *memstore.WorkerRepo + arts *memstore.ArtifactRepo + blobs *memstore.BlobStore + clk *memstore.Clock + + createJob *usecase.CreateJob + submit *usecase.SubmitDataset + claim *usecase.ClaimTask + renew *usecase.RenewLease + complete *usecase.CompleteTask + fail *usecase.FailTask + status *usecase.GetJobStatus + results *usecase.ListResults + register *usecase.RegisterWorker + uploadArt *usecase.UploadArtifact + downloadArt *usecase.DownloadArtifact + getInput *usecase.GetTaskInput + expire *usecase.ExpireLeases +} + +func newHarness() *harness { + h := &harness{ + tasks: memstore.NewTaskRepo(), + jobs: memstore.NewJobRepo(), + work: memstore.NewWorkerRepo(), + arts: memstore.NewArtifactRepo(), + blobs: memstore.NewBlobStore(), + clk: memstore.NewClock(time.Date(2026, 7, 21, 12, 0, 0, 0, time.UTC)), + } + tx := memstore.Tx{} + h.createJob = usecase.NewCreateJob(h.jobs, h.tasks, tx, h.clk) + h.submit = usecase.NewSubmitDataset(h.blobs, h.arts, h.jobs, h.tasks, tx, h.clk) + h.claim = usecase.NewClaimTask(h.tasks, h.clk, lease) + h.renew = usecase.NewRenewLease(h.tasks, tx, h.clk, lease) + h.complete = usecase.NewCompleteTask(h.tasks, h.jobs, h.arts, tx, h.clk) + h.fail = usecase.NewFailTask(h.tasks, h.jobs, tx, h.clk) + h.status = usecase.NewGetJobStatus(h.jobs, h.tasks) + h.results = usecase.NewListResults(h.tasks) + h.register = usecase.NewRegisterWorker(h.work, h.clk) + h.uploadArt = usecase.NewUploadArtifact(h.tasks, h.arts, h.blobs, h.clk) + h.downloadArt = usecase.NewDownloadArtifact(h.arts, h.blobs) + h.getInput = usecase.NewGetTaskInput(h.tasks, h.arts, h.blobs) + h.expire = usecase.NewExpireLeases(h.tasks, h.clk) + return h +} + +// seedJob creates a URI-chunked job with n chunks and returns its id. +func (h *harness) seedJob(t *testing.T, workload string, n int) uuid.UUID { + t.Helper() + in := usecase.CreateJobInput{Workload: workload, InputURI: "s3://in"} + for i := 0; i < n; i++ { + in.Chunks = append(in.Chunks, usecase.ChunkInput{ + ChunkIndex: i, InputURI: fmt.Sprintf("s3://c%d", i), InputSHA256: "sha", + }) + } + job, err := h.createJob.Execute(ctx, in) + if err != nil { + t.Fatalf("seedJob: %v", err) + } + return job.ID +} + +// leaseOne claims a single task for worker and returns its id and attempt. +func (h *harness) leaseOne(t *testing.T, worker, workload string) (uuid.UUID, int) { + t.Helper() + c, err := h.claim.Execute(ctx, usecase.ClaimTaskInput{WorkerID: worker, Workloads: []string{workload}}) + if err != nil || c == nil { + t.Fatalf("leaseOne: claim returned (%v, %v)", c, err) + } + return c.TaskID, c.Attempt +} + +// uploadResult stores a partial-result artifact for a leased task. +func (h *harness) uploadResult(t *testing.T, taskID uuid.UUID, worker string, attempt int) uuid.UUID { + t.Helper() + art, err := h.uploadArt.Execute(ctx, usecase.UploadArtifactInput{ + TaskID: taskID, WorkerID: worker, Attempt: attempt, + Filename: "r.csv", ContentType: "text/csv", Body: strings.NewReader("q,m\nA,B\n"), + }) + if err != nil { + t.Fatalf("uploadResult: %v", err) + } + return art.ID +} + +// --- ClaimTask ----------------------------------------------------------- + +func TestClaimLeasesAndAdvancesAttempt(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + + c, err := h.claim.Execute(ctx, usecase.ClaimTaskInput{WorkerID: "w1", Workloads: []string{"w"}}) + if err != nil || c == nil { + t.Fatalf("claim = (%v, %v)", c, err) + } + if c.Attempt != 1 || c.LeaseOwner != "w1" { + t.Errorf("attempt=%d owner=%q, want 1/w1", c.Attempt, c.LeaseOwner) + } +} + +func TestClaimEmptyQueueReturnsNil(t *testing.T) { + h := newHarness() + c, err := h.claim.Execute(ctx, usecase.ClaimTaskInput{WorkerID: "w1", Workloads: []string{"w"}}) + if err != nil || c != nil { + t.Errorf("claim on empty queue = (%v, %v), want (nil, nil)", c, err) + } +} + +func TestClaimRequiresWorkerID(t *testing.T) { + h := newHarness() + if _, err := h.claim.Execute(ctx, usecase.ClaimTaskInput{}); !errors.Is(err, domain.ErrInvalidInput) { + t.Errorf("err = %v, want ErrInvalidInput", err) + } +} + +func TestClaimSweepsExpiredLeaseFirst(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + // w1 leases it, then goes silent past the lease. + taskID, _ := h.leaseOne(t, "w1", "w") + h.clk.Advance(lease + time.Minute) + + // w2 claims: the sweep requeues the dead lease, so w2 gets the same task at attempt 2. + c, err := h.claim.Execute(ctx, usecase.ClaimTaskInput{WorkerID: "w2", Workloads: []string{"w"}}) + if err != nil || c == nil { + t.Fatalf("claim = (%v, %v)", c, err) + } + if c.TaskID != taskID || c.Attempt != 2 || c.LeaseOwner != "w2" { + t.Errorf("got task=%v attempt=%d owner=%q", c.TaskID, c.Attempt, c.LeaseOwner) + } +} + +// --- RenewLease ---------------------------------------------------------- + +func TestRenewExtendsForHolder(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + taskID, attempt := h.leaseOne(t, "w1", "w") + + c, err := h.renew.Execute(ctx, usecase.RenewLeaseInput{TaskID: taskID, WorkerID: "w1", Attempt: attempt}) + if err != nil { + t.Fatalf("renew: %v", err) + } + if !c.LeaseExpiresAt.Equal(h.clk.Now().Add(lease)) { + t.Error("lease not extended to now+lease") + } +} + +func TestRenewRejectsForeignWorker(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + taskID, attempt := h.leaseOne(t, "w1", "w") + + _, err := h.renew.Execute(ctx, usecase.RenewLeaseInput{TaskID: taskID, WorkerID: "intruder", Attempt: attempt}) + if !errors.Is(err, domain.ErrLeaseConflict) { + t.Errorf("err = %v, want ErrLeaseConflict", err) + } +} + +// --- CompleteTask -------------------------------------------------------- + +func TestCompleteHappyPathClosesJob(t *testing.T) { + h := newHarness() + jobID := h.seedJob(t, "w", 1) + taskID, attempt := h.leaseOne(t, "w1", "w") + artID := h.uploadResult(t, taskID, "w1", attempt) + + if _, err := h.complete.Execute(ctx, usecase.CompleteTaskInput{ + TaskID: taskID, WorkerID: "w1", Attempt: attempt, ResultArtifactID: artID, + }); err != nil { + t.Fatalf("complete: %v", err) + } + + prog, _ := h.status.Execute(ctx, jobID) + if prog.DeriveStatus() != domain.JobCompleted { + t.Errorf("job status = %q, want completed", prog.DeriveStatus()) + } +} + +func TestCompleteRejectsForeignArtifact(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 2) + // Lease two tasks; upload an artifact for taskA, try to complete taskB with it. + taskA, attA := h.leaseOne(t, "w1", "w") + taskB, attB := h.leaseOne(t, "w1", "w") + artA := h.uploadResult(t, taskA, "w1", attA) + + _, err := h.complete.Execute(ctx, usecase.CompleteTaskInput{ + TaskID: taskB, WorkerID: "w1", Attempt: attB, ResultArtifactID: artA, + }) + if !errors.Is(err, domain.ErrResultConflict) { + t.Errorf("cross-task artifact: err = %v, want ErrResultConflict", err) + } +} + +func TestCompleteIsIdempotentOnReplay(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + taskID, attempt := h.leaseOne(t, "w1", "w") + artID := h.uploadResult(t, taskID, "w1", attempt) + in := usecase.CompleteTaskInput{TaskID: taskID, WorkerID: "w1", Attempt: attempt, ResultArtifactID: artID} + + if _, err := h.complete.Execute(ctx, in); err != nil { + t.Fatalf("first complete: %v", err) + } + if _, err := h.complete.Execute(ctx, in); err != nil { + t.Errorf("replay must be idempotent, got %v", err) + } +} + +// --- FailTask ------------------------------------------------------------ + +func TestFailRequeuesWhileAttemptsRemain(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + taskID, attempt := h.leaseOne(t, "w1", "w") + + task, err := h.fail.Execute(ctx, usecase.FailTaskInput{ + TaskID: taskID, WorkerID: "w1", Attempt: attempt, + ErrorCode: "boom", ErrorMessage: "exploded", Retryable: true, + }) + if err != nil { + t.Fatalf("fail: %v", err) + } + if task.Status != domain.TaskPending { + t.Errorf("status = %q, want pending (requeued)", task.Status) + } + // It should be claimable again. + if c, _ := h.claim.Execute(ctx, usecase.ClaimTaskInput{WorkerID: "w2", Workloads: []string{"w"}}); c == nil { + t.Error("requeued task should be claimable") + } +} + +// --- CreateJob / status -------------------------------------------------- + +func TestCreateJobFansOutIntoTasks(t *testing.T) { + h := newHarness() + jobID := h.seedJob(t, "w", 3) + prog, err := h.status.Execute(ctx, jobID) + if err != nil { + t.Fatalf("status: %v", err) + } + if prog.Total != 3 || prog.Pending != 3 { + t.Errorf("progress total=%d pending=%d, want 3/3", prog.Total, prog.Pending) + } +} + +// --- RegisterWorker ------------------------------------------------------ + +func TestRegisterWorkerPersists(t *testing.T) { + h := newHarness() + w, err := h.register.Execute(ctx, usecase.RegisterWorkerInput{Name: "lab", Capabilities: []string{"w"}}) + if err != nil { + t.Fatalf("register: %v", err) + } + got, err := h.work.Get(ctx, w.ID) + if err != nil || got.Status != domain.WorkerOnline { + t.Errorf("worker not stored online: %v %v", got, err) + } +} + +func TestRegisterWorkerRejectsNoCapabilities(t *testing.T) { + h := newHarness() + if _, err := h.register.Execute(ctx, usecase.RegisterWorkerInput{Name: "lab"}); !errors.Is(err, domain.ErrInvalidInput) { + t.Errorf("err = %v, want ErrInvalidInput", err) + } +} + +// --- UploadArtifact ------------------------------------------------------ + +func TestUploadArtifactRejectsForeignWorker(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + taskID, attempt := h.leaseOne(t, "w1", "w") + + _, err := h.uploadArt.Execute(ctx, usecase.UploadArtifactInput{ + TaskID: taskID, WorkerID: "intruder", Attempt: attempt, + Filename: "r.csv", ContentType: "text/csv", Body: strings.NewReader("x"), + }) + if !errors.Is(err, domain.ErrLeaseConflict) { + t.Errorf("err = %v, want ErrLeaseConflict", err) + } +} + +func TestDownloadArtifactRoundTrips(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + taskID, attempt := h.leaseOne(t, "w1", "w") + artID := h.uploadResult(t, taskID, "w1", attempt) + + art, rc, err := h.downloadArt.Execute(ctx, artID) + if err != nil { + t.Fatalf("download: %v", err) + } + defer rc.Close() + if art.Kind != domain.ArtifactPartialResult { + t.Errorf("kind = %q", art.Kind) + } +} + +// --- SubmitDataset / GetTaskInput --------------------------------------- + +func TestSubmitDatasetChunksAndServesInput(t *testing.T) { + h := newHarness() + tsv := "id\tsmiles\nA\tCC\nB\tCCC\nC\tCCCC\nD\tCCCCC\nE\tCCCCCC\n" + + res, err := h.submit.Execute(ctx, usecase.SubmitDatasetInput{ + Workload: "w", RowsPerShard: 2, Filename: "chembl.tsv", + ContentType: "text/tab-separated-values", Body: strings.NewReader(tsv), + }) + if err != nil { + t.Fatalf("submit: %v", err) + } + if res.TaskCount != 3 { // 5 rows / 2 + t.Fatalf("task_count = %d, want 3", res.TaskCount) + } + + // The job now has three claimable shard tasks; each serves its own input. + prog, _ := h.status.Execute(ctx, res.JobID) + if prog.Total != 3 { + t.Errorf("job total = %d, want 3", prog.Total) + } + + c, err := h.claim.Execute(ctx, usecase.ClaimTaskInput{WorkerID: "w1", Workloads: []string{"w"}}) + if err != nil || c == nil { + t.Fatalf("claim shard: %v", err) + } + if c.InputArtifactID == nil { + t.Fatal("shard task must reference an input artifact") + } + art, rc, err := h.getInput.Execute(ctx, c.TaskID) + if err != nil { + t.Fatalf("get input: %v", err) + } + defer rc.Close() + if art.Kind != domain.ArtifactShard { + t.Errorf("input kind = %q, want shard", art.Kind) + } +} + +func TestGetTaskInputMissingForURITask(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) // URI-based task, no coordinator-stored input + taskID, _ := h.leaseOne(t, "w1", "w") + + if _, _, err := h.getInput.Execute(ctx, taskID); !errors.Is(err, domain.ErrArtifactNotFound) { + t.Errorf("err = %v, want ErrArtifactNotFound", err) + } +} + +// --- ExpireLeases -------------------------------------------------------- + +func TestExpireLeasesReclaims(t *testing.T) { + h := newHarness() + h.seedJob(t, "w", 1) + h.leaseOne(t, "w1", "w") + h.clk.Advance(lease + time.Minute) + + n, err := h.expire.Execute(ctx) + if err != nil || n != 1 { + t.Errorf("expire = (%d, %v), want (1, nil)", n, err) + } +}