Skip to content

Commit 12b116a

Browse files
committed
Test
1 parent 7bc818e commit 12b116a

5 files changed

Lines changed: 55 additions & 19 deletions

File tree

‎runner/internal/shim/api/schemas.go‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,10 @@ type TaskInfoResponse struct {
2929
TerminationReason string `json:"termination_reason"`
3030
TerminationMessage string `json:"termination_message"`
3131
Ports []shim.PortMapping `json:"ports"`
32+
33+
ImagePullCurrentBytes *int64 `json:"image_pull_current_bytes"`
34+
ImagePullTotalBytes *int64 `json:"image_pull_total_bytes"`
35+
3236
// The following fields are for debugging only, server doesn't need them
3337
ContainerName string `json:"container_name"`
3438
ContainerID string `json:"container_id"`

‎runner/internal/shim/docker.go‎

Lines changed: 35 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -231,14 +231,16 @@ func (d *DockerRunner) TaskInfo(taskID string) TaskInfo {
231231
return TaskInfo{}
232232
}
233233
return TaskInfo{
234-
ID: task.ID,
235-
Status: task.Status,
236-
TerminationReason: task.TerminationReason,
237-
TerminationMessage: task.TerminationMessage,
238-
Ports: task.ports,
239-
ContainerName: task.containerName,
240-
ContainerID: task.containerID,
241-
GpuIDs: task.gpuIDs,
234+
ID: task.ID,
235+
Status: task.Status,
236+
TerminationReason: task.TerminationReason,
237+
TerminationMessage: task.TerminationMessage,
238+
Ports: task.ports,
239+
ContainerName: task.containerName,
240+
ContainerID: task.containerID,
241+
GpuIDs: task.gpuIDs,
242+
ImagePullCurrentBytes: task.imagePullCurrentBytes,
243+
ImagePullTotalBytes: task.imagePullTotalBytes,
242244
}
243245
}
244246

@@ -350,7 +352,14 @@ func (d *DockerRunner) Run(ctx context.Context, taskID string) error {
350352
// Although it's called "runner dir", we also use it for shim task-related data.
351353
// Maybe we should rename it to "task dir" (including the `/root/.dstack/runners` dir on the host).
352354
pullLogPath := filepath.Join(runnerDir, "pull.log")
353-
if err = pullImage(pullCtx, d.client, cfg, pullLogPath); err != nil {
355+
onPullProgress := func(currentBytes, totalBytes int64) {
356+
task.imagePullCurrentBytes = &currentBytes
357+
task.imagePullTotalBytes = &totalBytes
358+
if updateErr := d.tasks.Update(task); updateErr != nil {
359+
log.Debug(ctx, "pull progress update skipped", "err", updateErr)
360+
}
361+
}
362+
if err = pullImage(pullCtx, d.client, cfg, pullLogPath, onPullProgress); err != nil {
354363
errMessage := fmt.Sprintf("pullImage error: %s", err.Error())
355364
log.Error(ctx, errMessage)
356365
task.SetStatusTerminated(string(types.TerminationReasonCreatingContainerError), errMessage)
@@ -670,7 +679,7 @@ func mountDisk(ctx context.Context, deviceName, mountPoint string, fsRootPerms o
670679
return nil
671680
}
672681

673-
func pullImage(ctx context.Context, client docker.APIClient, taskConfig TaskConfig, logPath string) error {
682+
func pullImage(ctx context.Context, client docker.APIClient, taskConfig TaskConfig, logPath string, onProgress func(currentBytes, totalBytes int64)) error {
674683
if !strings.Contains(taskConfig.ImageName, ":") {
675684
taskConfig.ImageName += ":latest"
676685
}
@@ -730,6 +739,20 @@ func pullImage(ctx context.Context, client docker.APIClient, taskConfig TaskConf
730739
} `json:"errorDetail"`
731740
}
732741

742+
reportProgress := func() {
743+
if onProgress == nil {
744+
return
745+
}
746+
var cb, tb int64
747+
for _, v := range current {
748+
cb += int64(v)
749+
}
750+
for _, v := range total {
751+
tb += int64(v)
752+
}
753+
onProgress(cb, tb)
754+
}
755+
733756
var pullCompleted bool
734757
pullErrors := make([]string, 0)
735758

@@ -743,9 +766,11 @@ func pullImage(ctx context.Context, client docker.APIClient, taskConfig TaskConf
743766
if pullMessage.Status == "Downloading" {
744767
current[pullMessage.Id] = pullMessage.ProgressDetail.Current
745768
total[pullMessage.Id] = pullMessage.ProgressDetail.Total
769+
reportProgress()
746770
}
747771
if pullMessage.Status == "Download complete" {
748772
current[pullMessage.Id] = total[pullMessage.Id]
773+
reportProgress()
749774
}
750775
if pullMessage.ErrorDetail.Message != "" {
751776
log.Error(ctx, "error pulling image", "name", taskConfig.ImageName, "err", pullMessage.ErrorDetail.Message)

‎runner/internal/shim/models.go‎

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -109,12 +109,14 @@ type TaskListItem struct {
109109
}
110110

111111
type TaskInfo struct {
112-
ID string
113-
Status TaskStatus
114-
TerminationReason string
115-
TerminationMessage string
116-
Ports []PortMapping
117-
ContainerName string
118-
ContainerID string
119-
GpuIDs []string
112+
ID string
113+
Status TaskStatus
114+
TerminationReason string
115+
TerminationMessage string
116+
Ports []PortMapping
117+
ImagePullCurrentBytes *int64
118+
ImagePullTotalBytes *int64
119+
ContainerName string
120+
ContainerID string
121+
GpuIDs []string
120122
}

‎runner/internal/shim/task.go‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,9 @@ type Task struct {
4242
ports []PortMapping
4343
runnerDir string // path on host mapped to consts.RunnerDir in container
4444

45+
imagePullCurrentBytes *int64
46+
imagePullTotalBytes *int64
47+
4548
mu *sync.Mutex
4649
}
4750

@@ -75,7 +78,8 @@ func (t *Task) IsTransitionAllowed(toStatus TaskStatus) bool {
7578
case TaskStatusPreparing:
7679
return t.Status == TaskStatusPending
7780
case TaskStatusPulling:
78-
return t.Status == TaskStatusPreparing
81+
// allow pulling->pulling to update pull progress
82+
return t.Status == TaskStatusPulling || t.Status == TaskStatusPreparing
7983
case TaskStatusCreating:
8084
return t.Status == TaskStatusPulling
8185
case TaskStatusRunning:

‎runner/internal/shim/task_test.go‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,7 @@ func TestTask_IsTransitionAllowed_true(t *testing.T) {
9393
{TaskStatusPending, TaskStatusTerminated},
9494
{TaskStatusPreparing, TaskStatusPulling},
9595
{TaskStatusPreparing, TaskStatusTerminated},
96+
{TaskStatusPulling, TaskStatusPulling},
9697
{TaskStatusPulling, TaskStatusCreating},
9798
{TaskStatusPulling, TaskStatusTerminated},
9899
{TaskStatusCreating, TaskStatusRunning},

0 commit comments

Comments
 (0)