Add job cancellation and dataset row limit
coordinator / test (push) Canceled after 0s

This commit is contained in:
Emil
2026-07-23 23:14:25 +03:00
parent 6bac7dad3c
commit 7547a30bde
30 changed files with 438 additions and 54 deletions
+6 -3
View File
@@ -53,9 +53,12 @@ type SubmitDatasetInput struct {
Workload string
Parameters map[string]any
RowsPerShard int
Filename string
ContentType string
Body io.Reader
// MaxRows limits how many data rows are turned into shards. Zero means the
// whole uploaded dataset; the input artifact itself remains stored intact.
MaxRows int
Filename string
ContentType string
Body io.Reader
}
type SubmitDatasetResult struct {
+43 -3
View File
@@ -62,6 +62,45 @@ type GetJobStatus struct {
tasks TaskRepository
}
// --- CancelJob -----------------------------------------------------------
type CancelJob struct {
jobs JobRepository
tasks TaskRepository
tx TxManager
clock Clock
}
func NewCancelJob(jobs JobRepository, tasks TaskRepository, tx TxManager, clock Clock) *CancelJob {
return &CancelJob{jobs: jobs, tasks: tasks, tx: tx, clock: clock}
}
// Execute stops a job atomically. Completed and finally failed tasks are kept
// as historical evidence; all other tasks are cancelled, including leased and
// running ones. A repeated cancel of an already cancelled job is idempotent.
func (uc *CancelJob) Execute(ctx context.Context, jobID uuid.UUID) (int64, error) {
now := uc.clock.Now()
var cancelled int64
err := uc.tx.WithinTx(ctx, func(ctx context.Context) error {
job, err := uc.jobs.Get(ctx, jobID)
if err != nil {
return err
}
if job.Status == domain.JobCancelled {
return nil
}
if job.Status == domain.JobCompleted || job.Status == domain.JobFailed {
return domain.ErrJobNotCancellable
}
cancelled, err = uc.tasks.CancelByJob(ctx, jobID, now)
if err != nil {
return err
}
return uc.jobs.UpdateStatus(ctx, jobID, domain.JobCancelled, &now)
})
return cancelled, err
}
func NewGetJobStatus(jobs JobRepository, tasks TaskRepository) *GetJobStatus {
return &GetJobStatus{jobs: jobs, tasks: tasks}
}
@@ -144,9 +183,10 @@ func progressFrom(job domain.Job, counts map[domain.TaskStatus]int) domain.JobPr
Job: job,
Pending: counts[domain.TaskPending],
// Leased and running are both "in flight" for progress purposes.
Leased: counts[domain.TaskLeased] + counts[domain.TaskRunning],
Done: counts[domain.TaskCompleted],
Failed: counts[domain.TaskFailed],
Leased: counts[domain.TaskLeased] + counts[domain.TaskRunning],
Done: counts[domain.TaskCompleted],
Failed: counts[domain.TaskFailed],
Cancelled: counts[domain.TaskCancelled],
}
for _, n := range counts {
p.Total += n
+4
View File
@@ -56,6 +56,10 @@ type TaskRepository interface {
// CountByStatus aggregates a job's tasks for progress reporting.
CountByStatus(ctx context.Context, jobID uuid.UUID) (map[domain.TaskStatus]int, error)
// CancelByJob marks every non-terminal task as cancelled and invalidates its
// lease. It returns how many tasks changed.
CancelByJob(ctx context.Context, jobID uuid.UUID, now time.Time) (int64, error)
// ExpireLeases applies the lease-expiry rule to every elapsed task and
// reports how many were affected.
ExpireLeases(ctx context.Context, now time.Time) (int64, error)
+4 -1
View File
@@ -30,6 +30,7 @@ type JobCard struct {
Running int `json:"running"`
Completed int `json:"completed"`
Failed int `json:"failed"`
Cancelled int `json:"cancelled"`
}
type TaskCard struct {
@@ -166,9 +167,11 @@ func jobCard(job domain.Job, tasks []domain.Task) JobCard {
c.Completed++
case domain.TaskFailed:
c.Failed++
case domain.TaskCancelled:
c.Cancelled++
}
}
p := domain.JobProgress{Job: job, Total: c.Total, Pending: c.Pending, Leased: c.Leased + c.Running, Done: c.Completed, Failed: c.Failed}
p := domain.JobProgress{Job: job, Total: c.Total, Pending: c.Pending, Leased: c.Leased + c.Running, Done: c.Completed, Failed: c.Failed, Cancelled: c.Cancelled}
c.Status = string(p.DeriveStatus())
return c
}
+1 -1
View File
@@ -64,7 +64,7 @@ func (uc *SubmitDataset) Execute(ctx context.Context, in SubmitDatasetInput) (Su
cleanup()
return SubmitDatasetResult{}, err
}
splitErr := chunk.SplitTSV(rc, in.RowsPerShard, func(index int, shard io.Reader) error {
splitErr := chunk.SplitTSVLimit(rc, in.RowsPerShard, in.MaxRows, 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 {
@@ -54,6 +54,7 @@ type harness struct {
downloadArt *usecase.DownloadArtifact
getInput *usecase.GetTaskInput
expire *usecase.ExpireLeases
cancel *usecase.CancelJob
}
func newHarness() *harness {
@@ -79,6 +80,7 @@ func newHarness() *harness {
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)
h.cancel = usecase.NewCancelJob(h.jobs, h.tasks, tx, h.clk)
return h
}
@@ -468,6 +470,44 @@ func TestSubmitDatasetChunksAndServesInput(t *testing.T) {
}
}
func TestSubmitDatasetLimitsRowsBeforeCreatingShards(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, MaxRows: 3, Filename: "chembl.tsv",
ContentType: "text/tab-separated-values", Body: strings.NewReader(tsv),
})
if err != nil {
t.Fatal(err)
}
if res.TaskCount != 2 {
t.Fatalf("task_count = %d, want 2", res.TaskCount)
}
}
func TestCancelJobInvalidatesClaimedAndPendingTasks(t *testing.T) {
h := newHarness()
jobID := h.seedJob(t, "w", 3)
claimed, err := h.claim.Execute(ctx, usecase.ClaimTaskInput{WorkerID: "w1", Workloads: []string{"w"}})
if err != nil || claimed == nil {
t.Fatalf("claim: %v", err)
}
cancelled, err := h.cancel.Execute(ctx, jobID)
if err != nil || cancelled != 3 {
t.Fatalf("cancel = (%d, %v), want (3, nil)", cancelled, err)
}
if _, err := h.renew.Execute(ctx, usecase.RenewLeaseInput{TaskID: claimed.TaskID, WorkerID: "w1", Attempt: claimed.Attempt}); !errors.Is(err, domain.ErrTaskNotLeased) {
t.Errorf("cancelled lease heartbeat = %v, want ErrTaskNotLeased", err)
}
progress, err := h.status.Execute(ctx, jobID)
if err != nil || progress.DeriveStatus() != domain.JobCancelled || progress.Cancelled != 3 {
t.Errorf("cancelled progress = %+v, err = %v", progress, err)
}
if cancelled, err := h.cancel.Execute(ctx, jobID); err != nil || cancelled != 0 {
t.Errorf("second cancel = (%d, %v), want (0, nil)", cancelled, err)
}
}
func TestGetTaskInputMissingForURITask(t *testing.T) {
h := newHarness()
h.seedJob(t, "w", 1) // URI-based task, no coordinator-stored input