185 lines
4.3 KiB
Go
185 lines
4.3 KiB
Go
package provision
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"context"
|
|
"crypto/sha512"
|
|
"encoding/hex"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"net/url"
|
|
"os"
|
|
"path"
|
|
"path/filepath"
|
|
"regexp"
|
|
"strings"
|
|
|
|
"git.g3e.fr/syonad/two/internal/lab/topology"
|
|
)
|
|
|
|
const maxSumsSize = 1 << 20
|
|
|
|
var sha512Pattern = regexp.MustCompile(`^[0-9a-f]{128}$`)
|
|
|
|
type Fetcher struct {
|
|
Client *http.Client
|
|
CacheDir string
|
|
}
|
|
|
|
func (f Fetcher) Image(ctx context.Context, img topology.Image) (string, error) {
|
|
if !filepath.IsAbs(f.CacheDir) {
|
|
return "", fmt.Errorf("image cache dir %q must be an absolute path", f.CacheDir)
|
|
}
|
|
name, err := fileName(img.URL)
|
|
if err != nil {
|
|
return "", fmt.Errorf("image %s: %w", img.Name, err)
|
|
}
|
|
|
|
sums, err := f.get(ctx, img.Sums, maxSumsSize)
|
|
if err != nil {
|
|
return "", fmt.Errorf("image %s: sums: %w", img.Name, err)
|
|
}
|
|
want, err := expectedSum(sums, name)
|
|
if err != nil {
|
|
return "", fmt.Errorf("image %s: %s: %w", img.Name, img.Sums, err)
|
|
}
|
|
|
|
dir := filepath.Join(f.CacheDir, img.Name)
|
|
target := filepath.Join(dir, name)
|
|
if got, err := fileSum(target); err == nil && got == want {
|
|
return target, nil
|
|
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
|
|
return "", fmt.Errorf("image %s: %w", img.Name, err)
|
|
}
|
|
|
|
if err := os.MkdirAll(dir, 0o755); err != nil {
|
|
return "", err
|
|
}
|
|
if err := f.download(ctx, img.URL, dir, target, want); err != nil {
|
|
return "", fmt.Errorf("image %s: %w", img.Name, err)
|
|
}
|
|
return target, nil
|
|
}
|
|
|
|
func (f Fetcher) client() *http.Client {
|
|
if f.Client != nil {
|
|
return f.Client
|
|
}
|
|
return http.DefaultClient
|
|
}
|
|
|
|
func (f Fetcher) open(ctx context.Context, rawURL string) (*http.Response, error) {
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
resp, err := f.client().Do(req)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if resp.StatusCode != http.StatusOK {
|
|
resp.Body.Close()
|
|
return nil, fmt.Errorf("GET %s: %s", rawURL, resp.Status)
|
|
}
|
|
return resp, nil
|
|
}
|
|
|
|
func (f Fetcher) get(ctx context.Context, rawURL string, limit int64) ([]byte, error) {
|
|
resp, err := f.open(ctx, rawURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer resp.Body.Close()
|
|
data, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if int64(len(data)) > limit {
|
|
return nil, fmt.Errorf("GET %s: larger than %d bytes", rawURL, limit)
|
|
}
|
|
return data, nil
|
|
}
|
|
|
|
func (f Fetcher) download(ctx context.Context, rawURL, dir, target, want string) error {
|
|
resp, err := f.open(ctx, rawURL)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer resp.Body.Close()
|
|
|
|
tmp, err := os.CreateTemp(dir, filepath.Base(target)+".part-*")
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer os.Remove(tmp.Name())
|
|
|
|
h := sha512.New()
|
|
if _, err := io.Copy(io.MultiWriter(tmp, h), resp.Body); err != nil {
|
|
tmp.Close()
|
|
return fmt.Errorf("GET %s: %w", rawURL, err)
|
|
}
|
|
if err := tmp.Close(); err != nil {
|
|
return err
|
|
}
|
|
if got := hex.EncodeToString(h.Sum(nil)); got != want {
|
|
return fmt.Errorf("GET %s: sha512 %s, want %s", rawURL, got, want)
|
|
}
|
|
if err := os.Chmod(tmp.Name(), 0o644); err != nil {
|
|
return err
|
|
}
|
|
return os.Rename(tmp.Name(), target)
|
|
}
|
|
|
|
func fileName(rawURL string) (string, error) {
|
|
u, err := url.Parse(rawURL)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
name := path.Base(u.Path)
|
|
if name == "" || name == "." || name == "/" {
|
|
return "", fmt.Errorf("url %q does not name a file", rawURL)
|
|
}
|
|
return name, nil
|
|
}
|
|
|
|
func expectedSum(sums []byte, name string) (string, error) {
|
|
var found []string
|
|
sc := bufio.NewScanner(bytes.NewReader(sums))
|
|
for sc.Scan() {
|
|
fields := strings.Fields(sc.Text())
|
|
if len(fields) != 2 || strings.TrimPrefix(fields[1], "*") != name {
|
|
continue
|
|
}
|
|
sum := strings.ToLower(fields[0])
|
|
if !sha512Pattern.MatchString(sum) {
|
|
return "", fmt.Errorf("entry for %s is not a sha512 sum", name)
|
|
}
|
|
found = append(found, sum)
|
|
}
|
|
if err := sc.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
switch {
|
|
case len(found) == 0:
|
|
return "", fmt.Errorf("no sum for %s", name)
|
|
case len(found) > 1:
|
|
return "", fmt.Errorf("%d sums for %s", len(found), name)
|
|
}
|
|
return found[0], nil
|
|
}
|
|
|
|
func fileSum(p string) (string, error) {
|
|
file, err := os.Open(p)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
defer file.Close()
|
|
h := sha512.New()
|
|
if _, err := io.Copy(h, file); err != nil {
|
|
return "", err
|
|
}
|
|
return hex.EncodeToString(h.Sum(nil)), nil
|
|
}
|