| | |
| | | package model |
| | | |
| | | import ( |
| | | "context" |
| | | "errors" |
| | | "fmt" |
| | | "io" |
| | |
| | | } |
| | | |
| | | func DownloadProfile(profile ModelProfile, urls []string, modelsDir string, progress ProgressCallback) error { |
| | | return DownloadProfileWithContext(context.Background(), profile, urls, modelsDir, progress) |
| | | } |
| | | |
| | | func DownloadProfileWithContext(ctx context.Context, profile ModelProfile, urls []string, modelsDir string, progress ProgressCallback) error { |
| | | if err := os.MkdirAll(modelsDir, 0755); err != nil { |
| | | return fmt.Errorf("failed to create models dir: %w", err) |
| | | } |
| | |
| | | |
| | | var lastErr error |
| | | for i, url := range urls { |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | if url == "" { |
| | | continue |
| | | } |
| | | if err := downloadAndInstallFromURL(profile, url, modelsDir, i+1, progress); err != nil { |
| | | if err := downloadAndInstallFromURL(ctx, profile, url, modelsDir, i+1, progress); err != nil { |
| | | if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { |
| | | return err |
| | | } |
| | | lastErr = err |
| | | logger.Info("Model install from URL %d failed: %v", i+1, err) |
| | | continue |
| | |
| | | } |
| | | } |
| | | |
| | | func downloadAndInstallFromURL(profile ModelProfile, url, modelsDir string, attempt int, progress ProgressCallback) error { |
| | | func downloadAndInstallFromURL(ctx context.Context, profile ModelProfile, url, modelsDir string, attempt int, progress ProgressCallback) error { |
| | | runID := fmt.Sprintf("%d-%d", time.Now().UnixNano(), attempt) |
| | | downloadDir := filepath.Join(modelsDir, ".downloads", profile.ID, runID) |
| | | extractDir := filepath.Join(downloadDir, "extract") |
| | |
| | | defer os.RemoveAll(stagingDir) |
| | | |
| | | logger.Info("Downloading model from URL %d: %s", attempt, url) |
| | | if err := downloadFile(url, archivePath, progress); err != nil { |
| | | if err := downloadFileWithContext(ctx, url, archivePath, progress); err != nil { |
| | | return err |
| | | } |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | |
| | | logger.Info("Extracting model archive...") |
| | | if err := extractArchive(archivePath, extractDir); err != nil { |
| | | if err := extractArchiveWithContext(ctx, archivePath, extractDir); err != nil { |
| | | return fmt.Errorf("extraction failed: %w", err) |
| | | } |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | |
| | | sourceDir, err := findInstallSource(profile, extractDir) |
| | | if err != nil { |
| | | return err |
| | | } |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | |
| | | if err := moveOrCopyDir(sourceDir, stagingDir); err != nil { |
| | | return fmt.Errorf("failed to stage model: %w", err) |
| | | } |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | |
| | | validation := ValidateModelDir(profile, stagingDir) |
| | | if !validation.Valid { |
| | | return fmt.Errorf("downloaded model is incomplete: missing=%v problems=%v", validation.Missing, validation.Problems) |
| | | } |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | |
| | | finalDir := filepath.Join(modelsDir, profile.InstallDirName) |
| | |
| | | } |
| | | |
| | | func downloadFile(url, destPath string, progress ProgressCallback) error { |
| | | return downloadFileWithContext(context.Background(), url, destPath, progress) |
| | | } |
| | | |
| | | func downloadFileWithContext(ctx context.Context, url, destPath string, progress ProgressCallback) error { |
| | | const maxAttempts = 6 |
| | | var lastErr error |
| | | var lastPercent = -1.0 |
| | | |
| | | for attempt := 1; attempt <= maxAttempts; attempt++ { |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | downloaded := existingFileSize(destPath) |
| | | err := downloadFileAttempt(url, destPath, downloaded, &lastPercent, progress) |
| | | err := downloadFileAttempt(ctx, url, destPath, downloaded, &lastPercent, progress) |
| | | if err == nil { |
| | | return nil |
| | | } |
| | | if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { |
| | | return err |
| | | } |
| | | if isNonRetryableDownloadError(err) { |
| | | return err |
| | |
| | | lastErr = err |
| | | logger.Info("Download attempt %d/%d failed: %v", attempt, maxAttempts, err) |
| | | if attempt < maxAttempts { |
| | | time.Sleep(time.Duration(attempt) * time.Second) |
| | | select { |
| | | case <-ctx.Done(): |
| | | return ctx.Err() |
| | | case <-time.After(time.Duration(attempt) * time.Second): |
| | | } |
| | | } |
| | | } |
| | | |
| | | return fmt.Errorf("download failed after %d attempts: %w", maxAttempts, lastErr) |
| | | } |
| | | |
| | | func downloadFileAttempt(url, destPath string, resumeFrom int64, lastPercent *float64, progress ProgressCallback) error { |
| | | req, err := http.NewRequest(http.MethodGet, url, nil) |
| | | func downloadFileAttempt(ctx context.Context, url, destPath string, resumeFrom int64, lastPercent *float64, progress ProgressCallback) error { |
| | | req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) |
| | | if err != nil { |
| | | return nonRetryableDownloadError{err: err} |
| | | } |
| | |
| | | } |
| | | buf := make([]byte, 32*1024) |
| | | for { |
| | | if err := ctx.Err(); err != nil { |
| | | return err |
| | | } |
| | | n, readErr := resp.Body.Read(buf) |
| | | if n > 0 { |
| | | if _, writeErr := out.Write(buf[:n]); writeErr != nil { |
| | |
| | | } |
| | | |
| | | func extractArchive(archivePath, destDir string) error { |
| | | cmd := exec.Command("tar", "-xf", archivePath, "-C", destDir) |
| | | return extractArchiveWithContext(context.Background(), archivePath, destDir) |
| | | } |
| | | |
| | | func extractArchiveWithContext(ctx context.Context, archivePath, destDir string) error { |
| | | cmd := exec.CommandContext(ctx, "tar", "-xf", archivePath, "-C", destDir) |
| | | output, err := cmd.CombinedOutput() |
| | | if err != nil { |
| | | if ctxErr := ctx.Err(); ctxErr != nil { |
| | | return ctxErr |
| | | } |
| | | return fmt.Errorf("tar extraction failed: %v, output: %s", err, string(output)) |
| | | } |
| | | return nil |