This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user