@@ -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 )
0 commit comments