Luigit
repositories / bugabinga.net

bugabinga.net

personal infrastructure for bugabinga!

owned by admin

services/luci/internal/cache/cache.go

Raw
package cache

import (
	"crypto/sha256"
	"encoding/hex"
	"errors"
	"fmt"
	"io"
	"os"
	"path/filepath"
	"sort"
	"strings"
	"syscall"

	"bugabinga.net/luci/internal/ciconfig"
)

type Store struct {
	DataDir string
	// MaxVersions keeps at most this many cached versions per cache name;
	// zero means one.
	MaxVersions int
	// MaxBytes caps bytes across every cache name; zero means uncapped.
	MaxBytes int64
}

func Key(workspace string, entry ciconfig.Cache) (string, error) {
	hash := sha256.New()
	for _, relative := range entry.KeyFiles {
		path, err := securePath(workspace, relative, true)
		if err != nil {
			return "", err
		}
		_, _ = io.WriteString(hash, relative)
		_, _ = hash.Write([]byte{0})
		data, err := os.ReadFile(path)
		if os.IsNotExist(err) {
			_, _ = hash.Write([]byte{0})
		} else if err != nil {
			return "", err
		} else {
			_, _ = hash.Write([]byte{1})
			_, _ = hash.Write(data)
		}
		_, _ = hash.Write([]byte{0})
	}
	return hex.EncodeToString(hash.Sum(nil)), nil
}

// Restore materializes a cache entry into the workspace and reports the
// outcome: "exact", "fallback" (newest other version), or "miss".
func (s Store) Restore(repo, job string, entry ciconfig.Cache, key, workspace string) (string, error) {
	path, err := s.path(repo, job, entry.Name, key)
	if err != nil {
		return "", err
	}
	outcome := "exact"
	lockErr := withLock(filepath.Join(s.DataDir, "cache", ".lock"), func() error {
		return withLock(path+".lock", func() error {
			source := path
			if _, err := os.Lstat(source); os.IsNotExist(err) {
				fallback, err := s.newestVersion(repo, job, entry.Name, key)
				if err != nil {
					return err
				}
				if fallback == "" {
					outcome = "miss"
					return nil
				}
				source = fallback
				outcome = "fallback"
			} else if err != nil {
				return err
			}
			target, err := securePath(workspace, entry.Path, true)
			if err != nil {
				return err
			}
			return copyTree(source, target)
		})
	})
	return outcome, lockErr
}

// newestVersion returns the newest stored version of a cache other than the
// exact key, enabling warm rebuilds when the exact key was never saved.
func (s Store) newestVersion(repo, job, name, key string) (string, error) {
	dir, err := s.path(repo, job, name, "probe")
	if err != nil {
		return "", err
	}
	entries, err := os.ReadDir(filepath.Dir(dir))
	if err != nil {
		if os.IsNotExist(err) {
			return "", nil
		}
		return "", err
	}
	newest, newestTime := "", int64(-1)
	for _, entry := range entries {
		name := entry.Name()
		if name == key || strings.HasPrefix(name, ".") || strings.HasSuffix(name, ".lock") {
			continue
		}
		info, err := entry.Info()
		if err != nil {
			continue
		}
		if modified := info.ModTime().UnixNano(); modified > newestTime {
			newest, newestTime = name, modified
		}
	}
	if newest == "" {
		return "", nil
	}
	return filepath.Join(filepath.Dir(dir), newest), nil
}

func (s Store) Save(repo, job string, entry ciconfig.Cache, key, workspace string) error {
	path, err := s.path(repo, job, entry.Name, key)
	if err != nil {
		return err
	}
	source, err := securePath(workspace, entry.Path, true)
	if err != nil {
		return err
	}
	if _, err := os.Lstat(source); os.IsNotExist(err) {
		return nil
	} else if err != nil {
		return err
	}
	sourceBytes, err := treeSize(source)
	if err != nil {
		return err
	}
	if s.MaxBytes > 0 && sourceBytes > s.MaxBytes {
		return fmt.Errorf("cache source %s is %d bytes, exceeding byte budget %d", source, sourceBytes, s.MaxBytes)
	}
	return withLock(filepath.Join(s.DataDir, "cache", ".lock"), func() error {
		return withLock(path+".lock", func() error {
			parent := filepath.Dir(path)
			if err := os.MkdirAll(parent, 0o755); err != nil {
				return err
			}
			if err := s.reserve(path, sourceBytes); err != nil {
				return err
			}
			stage, err := os.MkdirTemp(parent, ".cache-stage-")
			if err != nil {
				return err
			}
			cleanupStage := true
			defer func() {
				if cleanupStage {
					_ = os.RemoveAll(stage)
				}
			}()
			staged := filepath.Join(stage, "value")
			remaining := sourceBytes
			if err := copyTreeLimited(source, staged, &remaining); err != nil {
				return err
			}
			backup := filepath.Join(stage, "previous")
			hadPrevious := false
			if _, err := os.Lstat(path); err == nil {
				if err := os.Rename(path, backup); err != nil {
					return err
				}
				hadPrevious = true
			} else if !os.IsNotExist(err) {
				return err
			}
			if err := os.Rename(staged, path); err != nil {
				if hadPrevious {
					rollbackErr := os.Rename(backup, path)
					if rollbackErr != nil {
						cleanupStage = false
						rollbackErr = fmt.Errorf("cache rollback failed; previous data remains at %s: %w", backup, rollbackErr)
					}
					return errors.Join(err, rollbackErr)
				}
				return err
			}
			return s.evict(parent)
		})
	})
}

// reserve clears oldest cache versions across every cache name before staging
// a new version, so staging cannot grow the cache beyond MaxBytes.
func (s Store) reserve(protected string, bytes int64) error {
	if s.MaxBytes <= 0 {
		return nil
	}
	root := filepath.Join(s.DataDir, "cache")
	total, err := treeSize(root)
	if err != nil {
		return err
	}
	if total+bytes <= s.MaxBytes {
		return nil
	}
	versions, err := cacheVersions(root)
	if err != nil {
		return err
	}
	sort.SliceStable(versions, func(i, j int) bool { return versions[i].modified < versions[j].modified })
	for _, version := range versions {
		if version.path == protected {
			// The protected version is replaced after staging, so only its
			// replacement contributes to the final aggregate budget.
			total -= version.bytes
			continue
		}
		if err := os.RemoveAll(version.path); err != nil {
			return err
		}
		total -= version.bytes
		if total+bytes <= s.MaxBytes {
			return nil
		}
	}
	if total+bytes <= s.MaxBytes {
		return nil
	}
	return fmt.Errorf("cache staging of %d bytes exceeds aggregate byte budget %d", bytes, s.MaxBytes)
}

type cacheVersion struct {
	path     string
	bytes    int64
	modified int64
}

func cacheVersions(root string) ([]cacheVersion, error) {
	var versions []cacheVersion
	var walk func(string, int) error
	walk = func(dir string, depth int) error {
		entries, err := os.ReadDir(dir)
		if os.IsNotExist(err) {
			return nil
		}
		if err != nil {
			return err
		}
		for _, entry := range entries {
			if strings.HasPrefix(entry.Name(), ".") || strings.HasSuffix(entry.Name(), ".lock") {
				continue
			}
			path := filepath.Join(dir, entry.Name())
			if depth == 4 {
				info, err := os.Lstat(path)
				if err != nil {
					return err
				}
				bytes, err := treeSize(path)
				if err != nil {
					return err
				}
				versions = append(versions, cacheVersion{path: path, bytes: bytes, modified: info.ModTime().UnixNano()})
				continue
			}
			if entry.IsDir() {
				if err := walk(path, depth+1); err != nil {
					return err
				}
			}
		}
		return nil
	}
	if err := walk(root, 1); err != nil {
		return nil, err
	}
	return versions, nil
}

// evict trims one cache-name directory to MaxVersions newest entries.
func (s Store) evict(dir string) error {
	entries, err := os.ReadDir(dir)
	if err != nil {
		if os.IsNotExist(err) {
			return nil
		}
		return err
	}
	var versions []os.DirEntry
	for _, entry := range entries {
		name := entry.Name()
		if strings.HasPrefix(name, ".") || strings.HasSuffix(name, ".lock") {
			continue
		}
		versions = append(versions, entry)
	}
	keep := s.MaxVersions
	if keep < 1 {
		keep = 1
	}
	if err := sortVersionsByTime(dir, versions); err != nil {
		return err
	}
	if len(versions) > keep {
		for _, entry := range versions[keep:] {
			if err := os.RemoveAll(filepath.Join(dir, entry.Name())); err != nil {
				return err
			}
		}
	}
	return nil
}

func sortVersionsByTime(dir string, versions []os.DirEntry) error {
	times := make(map[string]int64, len(versions))
	for _, entry := range versions {
		info, err := entry.Info()
		if err != nil {
			return err
		}
		times[entry.Name()] = info.ModTime().UnixNano()
	}
	sort.SliceStable(versions, func(i, j int) bool {
		return times[versions[i].Name()] > times[versions[j].Name()]
	})
	return nil
}

func treeSize(root string) (int64, error) {
	var total int64
	err := filepath.WalkDir(root, func(_ string, entry os.DirEntry, err error) error {
		if err != nil {
			return err
		}
		info, err := entry.Info()
		if err != nil {
			return err
		}
		if info.Mode().IsRegular() {
			total += info.Size()
		}
		return nil
	})
	return total, err
}

func (s Store) path(repo, job, name, key string) (string, error) {
	for label, value := range map[string]string{"repo": repo, "job": job, "cache": name, "key": key} {
		if value == "" || value != filepath.Base(value) || value == "." || value == ".." {
			return "", fmt.Errorf("unsafe %s %q", label, value)
		}
	}
	return filepath.Join(s.DataDir, "cache", repo, job, name, key), nil
}

func securePath(root, relative string, missingAllowed bool) (string, error) {
	clean := filepath.Clean(relative)
	if clean == "." || filepath.IsAbs(relative) || clean == ".." || strings.HasPrefix(clean, ".."+string(filepath.Separator)) {
		return "", fmt.Errorf("unsafe workspace path %q", relative)
	}
	current := root
	for _, part := range strings.Split(clean, string(filepath.Separator)) {
		current = filepath.Join(current, part)
		info, err := os.Lstat(current)
		if os.IsNotExist(err) && missingAllowed {
			continue
		}
		if err != nil {
			return "", err
		}
		if info.Mode()&os.ModeSymlink != 0 {
			return "", fmt.Errorf("workspace path %q traverses symlink", relative)
		}
	}
	return filepath.Join(root, clean), nil
}

func withLock(path string, fn func() error) error {
	if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
		return err
	}
	file, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0o600)
	if err != nil {
		return err
	}
	defer file.Close()
	if err := syscall.Flock(int(file.Fd()), syscall.LOCK_EX); err != nil {
		return err
	}
	defer syscall.Flock(int(file.Fd()), syscall.LOCK_UN)
	return fn()
}

func copyTree(source, target string) error {
	return copyTreeLimited(source, target, nil)
}

func copyTreeLimited(source, target string, remaining *int64) error {
	info, err := os.Lstat(source)
	if err != nil {
		return err
	}
	if existing, err := os.Lstat(target); err == nil && existing.Mode().Type() != info.Mode().Type() {
		if err := os.RemoveAll(target); err != nil {
			return err
		}
	} else if err != nil && !os.IsNotExist(err) {
		return err
	}
	switch {
	case info.Mode().IsDir():
		if err := os.MkdirAll(target, info.Mode().Perm()); err != nil {
			return err
		}
		entries, err := os.ReadDir(source)
		if err != nil {
			return err
		}
		for _, entry := range entries {
			if err := copyTreeLimited(filepath.Join(source, entry.Name()), filepath.Join(target, entry.Name()), remaining); err != nil {
				return err
			}
		}
		return nil
	case info.Mode().IsRegular():
		if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil {
			return err
		}
		input, err := os.Open(source)
		if err != nil {
			return err
		}
		defer input.Close()
		output, err := os.OpenFile(target, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, info.Mode().Perm())
		if err != nil {
			return err
		}
		var writer io.Writer = output
		if remaining != nil {
			writer = &cacheLimitWriter{writer: output, remaining: remaining}
		}
		_, copyErr := io.Copy(writer, input)
		closeErr := output.Close()
		if copyErr != nil {
			return copyErr
		}
		return closeErr
	case info.Mode()&os.ModeSymlink != 0:
		link, err := os.Readlink(source)
		if err != nil {
			return err
		}
		if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil {
			return err
		}
		_ = os.Remove(target)
		return os.Symlink(link, target)
	default:
		return fmt.Errorf("unsupported cache file %s", source)
	}
}

type cacheLimitWriter struct {
	writer    io.Writer
	remaining *int64
}

func (w *cacheLimitWriter) Write(data []byte) (int, error) {
	if int64(len(data)) > *w.remaining {
		return 0, errors.New("cache source changed while staging")
	}
	n, err := w.writer.Write(data)
	*w.remaining -= int64(n)
	return n, err
}