From 76ff0c929dd53450a2c263c212bb6ff202569fa7 Mon Sep 17 00:00:00 2001 From: dawn Date: Fri, 17 Jul 2026 02:08:18 +0300 Subject: [PATCH] spindle: add generic CI cache Signed-off-by: dawn --- docker-compose.yml | 1 + docs/DOCS.md | 76 ++++++ nix/microvm/base.nix | 1 + nix/pkgs/spindle-alpine-image.nix | 3 +- spindle/config/config.go | 10 + spindle/db/cache.go | 233 ++++++++++++++++++ spindle/db/cache_test.go | 302 +++++++++++++++++++++++ spindle/db/db.go | 26 +- spindle/engine/cache.go | 333 +++++++++++++++++++++++++ spindle/engine/cache_prune.go | 85 +++++++ spindle/engine/cache_prune_test.go | 152 ++++++++++++ spindle/engine/cache_test.go | 376 +++++++++++++++++++++++++++++ spindle/engine/engine.go | 78 +++++- spindle/engines/microvm/cache.go | 125 ++++++++++ spindle/engines/microvm/engine.go | 24 +- spindle/engines/microvm/models.go | 17 +- spindle/engines/nixery/cache.go | 167 +++++++++++++ spindle/engines/nixery/engine.go | 13 +- spindle/models/cache.go | 63 +++++ spindle/models/cache_test.go | 49 ++++ spindle/models/clone.go | 5 +- spindle/models/pipeline.go | 10 +- spindle/server.go | 23 +- spindle/storage/disk.go | 87 +++++++ spindle/storage/s3.go | 102 ++++++++ spindle/storage/storage.go | 52 ++++ spindle/storage/storage_test.go | 101 ++++++++ 27 files changed, 2485 insertions(+), 29 deletions(-) create mode 100644 spindle/db/cache.go create mode 100644 spindle/db/cache_test.go create mode 100644 spindle/engine/cache.go create mode 100644 spindle/engine/cache_prune.go create mode 100644 spindle/engine/cache_prune_test.go create mode 100644 spindle/engine/cache_test.go create mode 100644 spindle/engines/microvm/cache.go create mode 100644 spindle/engines/nixery/cache.go create mode 100644 spindle/models/cache.go create mode 100644 spindle/models/cache_test.go create mode 100644 spindle/storage/disk.go create mode 100644 spindle/storage/s3.go create mode 100644 spindle/storage/storage.go create mode 100644 spindle/storage/storage_test.go diff --git a/docker-compose.yml b/docker-compose.yml index 6cf374f1..95ae154f 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -183,6 +183,7 @@ services: SPINDLE_MICROVM_PIPELINES_AGENT_PORT: "11240" SPINDLE_S3_LOG_BUCKET: "" SPINDLE_MICROVM_PIPELINES_ENABLE_CGROUPS: "false" + SPINDLE_CACHE_BACKEND: disk # route guest nix substitution + uploads through the local ncps cache. # ncps re-signs on serve with cache.local's key, so the guest trusts the # matching public key below (no signing happens in spindle itself). diff --git a/docs/DOCS.md b/docs/DOCS.md index b72602bc..d185b846 100644 --- a/docs/DOCS.md +++ b/docs/DOCS.md @@ -965,6 +965,58 @@ triggered by a pull request: - `TANGLED_PR_SOURCE_SHA` - The commit SHA of the source branch +### Cache + +The `cache` field lets a workflow persist directories across +pipeline runs. Before the first step, the engine looks up +each entry's key and extracts the matching archive into the +workspace; after all steps succeed, the paths are archived +again and stored back under the key. + +- `key`: name this cache is saved under. Keys are scoped to + the repository and the engine specified. +- `hash`: **optional** list of repo files (lockfiles, + manifests) whose content is folded into the key. The entry + is stored as `-`, so editing + `go.sum` automatically rotates the cache without bumping + the key by hand. Paths are relative to the repository + root and are read from git at the commit being built. +- `paths`: paths to archive. Relative paths are anchored at + the repository checkout, the directory steps start in + (`/workspace/repo` on microvm, `/tangled/workspace` on + nixery). Absolute paths work too, for caching directories + outside the checkout, but note they name engine-specific + locations. All paths must be writable by the CI user and + contain no spaces. +- `compression-level`: **optional** zstd level, `1` + (fastest) to `19` (smallest). defaults to (`5`). +- `when`: **optional** save policy: `on-success` (the + default) or `always`. + +When the exact key (or generation, with `hash`) misses, the +newest older generation under the same key is restored. + +```yaml +cache: + - key: go-mod + hash: + - go.sum + - go.mod + paths: + - .gocache +``` + +Caches are only restored and saved for trusted pipelines +(pushes and same-repository pull requests). Pipelines +building untrusted code, like pull requests from forks, +skip the cache entirely. Saving follows each entry's `when` +policy, except on timeout, when nothing is saved. A cache +miss or failure never fails the workflow. + +The spindle operator chooses the storage backend. See +[Running spindle](#running-spindle). If no backend is +configured, `cache` entries are ignored. + ### Steps The `steps` field allows you to define what steps should run @@ -1510,6 +1562,30 @@ cache (and read from it), configure the cache (prefix - `SPINDLE_NIX_CACHE_UPLOAD_URL`: Cache URL that paths built in the guest are uploaded to. +The generic CI cache (the workflow-level +[`cache`](#cache) field) is configured via prefix +`SPINDLE_CACHE_`. + +- `SPINDLE_CACHE_BACKEND`: Storage backend, `disk` or `s3` + (default: `""`, caching disabled). +- `SPINDLE_CACHE_DISK_DIR`: Directory for the `disk` backend + (default: a `cache` directory next to the spindle + database). +- `SPINDLE_CACHE_S3_BUCKET`: Unversioned bucket for the `s3` + backend. Credentials come from the standard AWS chain and + need `s3:GetBucketVersioning` in addition to object access. +- `SPINDLE_CACHE_S3_PREFIX`: Key prefix inside the bucket + (default: `"spindle/cache"`). +- `SPINDLE_CACHE_RETENTION`: Time since the last restore or + save before an entry is deleted (default: `720h`, or 30 + days). Set to `0` to keep entries indefinitely. +- `SPINDLE_CACHE_PRUNE_INTERVAL`: How often expired entries + are deleted (default: `1h`). + +Cache metadata and usage are tracked in spindle's SQLite +database. Storage backends contain opaque objects and are +never listed during lookup or cleanup. + ### Running spindle 1. **Set the environment variables.** For example: diff --git a/nix/microvm/base.nix b/nix/microvm/base.nix index 5ae6b9bd..e549232e 100644 --- a/nix/microvm/base.nix +++ b/nix/microvm/base.nix @@ -258,6 +258,7 @@ in { gz-utils bzip2 lz4 + zstd p7zip ]; # disable default nixos packages ([perl rsync strace]) diff --git a/nix/pkgs/spindle-alpine-image.nix b/nix/pkgs/spindle-alpine-image.nix index dea9bb75..9852fedb 100644 --- a/nix/pkgs/spindle-alpine-image.nix +++ b/nix/pkgs/spindle-alpine-image.nix @@ -30,7 +30,8 @@ }); # we don't include gnused, xxd etc. here because busybox has them # we want to keep the image this image small! - guestTools = [nix bash git curl jq]; + # zstd is not a busybox applet, and the spindle cache saves tar|zstd + guestTools = [nix bash git curl jq pkgsStatic.zstd]; # run by busybox at sysinit setupScript = writeText "spindle-setup" '' diff --git a/spindle/config/config.go b/spindle/config/config.go index 661e35f7..ff78451b 100644 --- a/spindle/config/config.go +++ b/spindle/config/config.go @@ -61,6 +61,15 @@ type S3 struct { LogBucket string `env:"LOG_BUCKET"` } +type Cache struct { + Backend string `env:"BACKEND"` // "disk" or "s3" + DiskDir string `env:"DISK_DIR"` + S3Bucket string `env:"S3_BUCKET"` + S3Prefix string `env:"S3_PREFIX, default=spindle/cache"` + Retention time.Duration `env:"RETENTION, default=720h"` + PruneInterval time.Duration `env:"PRUNE_INTERVAL, default=1h"` +} + type MicroVMPipelines struct { ImageDir string `env:"IMAGE_DIR"` OverlayDir string `env:"OVERLAY_DIR, default="` // where microVM temporary disks will live @@ -99,6 +108,7 @@ type Config struct { MicroVMPipelines MicroVMPipelines `env:",prefix=SPINDLE_MICROVM_PIPELINES_"` NixCache NixCache `env:",prefix=SPINDLE_NIX_CACHE_"` S3 S3 `env:",prefix=SPINDLE_S3_"` + Cache Cache `env:",prefix=SPINDLE_CACHE_"` } func Load(ctx context.Context) (*Config, error) { diff --git a/spindle/db/cache.go b/spindle/db/cache.go new file mode 100644 index 00000000..33448653 --- /dev/null +++ b/spindle/db/cache.go @@ -0,0 +1,233 @@ +package db + +import ( + "context" + "sort" + "time" +) + +type CacheEntry struct { + ID string + StorageKey string + OwnerDID string + RepoDID string + Engine string + CacheKey string + CacheHash string + SizeBytes int64 + State string + CreatedAt time.Time + LastUsedAt time.Time +} + +const cacheEntryColumns = ` + id, storage_key, owner_did, repo_did, engine, cache_key, cache_hash, + size_bytes, state, created_at, last_used_at` + +func (d *DB) InsertCacheEntry(ctx context.Context, entry CacheEntry) error { + _, err := d.ExecContext(ctx, ` + insert into cache_entries ( + id, storage_key, owner_did, repo_did, engine, cache_key, cache_hash, + size_bytes, state, created_at, last_used_at + ) values (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + entry.ID, + entry.StorageKey, + entry.OwnerDID, + entry.RepoDID, + entry.Engine, + entry.CacheKey, + entry.CacheHash, + entry.SizeBytes, + entry.State, + entry.CreatedAt.UnixNano(), + entry.LastUsedAt.UnixNano(), + ) + return err +} + +func (d *DB) MarkCacheEntryReady(ctx context.Context, id string, sizeBytes int64, now time.Time) ([]CacheEntry, error) { + tx, err := d.BeginTx(ctx, nil) + if err != nil { + return nil, err + } + defer tx.Rollback() + + var repoDID, engine, key, hash string + if err := tx.QueryRowContext(ctx, ` + update cache_entries + set state = 'ready', size_bytes = ?, last_used_at = ? + where id = ? and state = 'pending' + returning repo_did, engine, cache_key, cache_hash`, + sizeBytes, now.UnixNano(), id).Scan(&repoDID, &engine, &key, &hash); err != nil { + return nil, err + } + + rows, err := tx.QueryContext(ctx, ` + update cache_entries + set state = 'deleting' + where repo_did = ? and engine = ? and cache_key = ? and cache_hash = ? + and state = 'ready' and id <> ? + returning `+cacheEntryColumns, repoDID, engine, key, hash, id) + if err != nil { + return nil, err + } + var superseded []CacheEntry + for rows.Next() { + entry, err := scanCacheEntry(rows) + if err != nil { + rows.Close() + return nil, err + } + superseded = append(superseded, *entry) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + sort.Slice(superseded, func(i, j int) bool { + return superseded[i].CreatedAt.After(superseded[j].CreatedAt) + }) + if err := tx.Commit(); err != nil { + return nil, err + } + return superseded, nil +} + +func (d *DB) FindCacheEntry(ctx context.Context, repoDID, engine, key, hash string) (*CacheEntry, error) { + return scanCacheEntry(d.QueryRowContext(ctx, ` + select `+cacheEntryColumns+` + from cache_entries + where repo_did = ? and engine = ? and cache_key = ? and cache_hash = ? and state = 'ready' + order by created_at desc + limit 1`, repoDID, engine, key, hash)) +} + +func (d *DB) FindFallbackCacheEntry(ctx context.Context, repoDID, engine, key, excludeHash string) (*CacheEntry, error) { + return scanCacheEntry(d.QueryRowContext(ctx, ` + select `+cacheEntryColumns+` + from cache_entries + where repo_did = ? and engine = ? and cache_key = ? + and cache_hash <> ? and cache_hash <> '' and state = 'ready' + order by created_at desc + limit 1`, repoDID, engine, key, excludeHash)) +} + +func (d *DB) TouchCacheEntry(ctx context.Context, id string, now time.Time) error { + _, err := d.ExecContext(ctx, ` + update cache_entries set last_used_at = ? where id = ? and state = 'ready'`, now.UnixNano(), id) + return err +} + +func (d *DB) ClaimCacheEntry(ctx context.Context, id, expectedState string, expectedLastUsed time.Time) (bool, error) { + result, err := d.ExecContext(ctx, ` + update cache_entries + set state = 'deleting' + where id = ? and state = ? and last_used_at = ?`, + id, expectedState, expectedLastUsed.UnixNano()) + if err != nil { + return false, err + } + changed, err := result.RowsAffected() + return changed == 1, err +} + +func (d *DB) RestoreCacheEntryState(ctx context.Context, id, state string) error { + _, err := d.ExecContext(ctx, ` + update cache_entries set state = ? where id = ? and state = 'deleting'`, state, id) + return err +} + +func (d *DB) ExpiredCacheEntries(ctx context.Context, readyBefore, pendingBefore time.Time, limit int) ([]CacheEntry, error) { + ready, err := d.queryCacheEntries(ctx, ` + select `+cacheEntryColumns+` + from cache_entries + where state = 'ready' and last_used_at < ? + order by last_used_at + limit ?`, readyBefore.UnixNano(), limit) + if err != nil { + return nil, err + } + recovery, err := d.queryCacheEntries(ctx, ` + select `+cacheEntryColumns+` + from cache_entries + where state in ('pending', 'deleting') and created_at < ? + order by created_at + limit ?`, pendingBefore.UnixNano(), limit) + if err != nil { + return nil, err + } + entries := append(ready, recovery...) + sort.Slice(entries, func(i, j int) bool { + left, right := entries[i].CreatedAt, entries[j].CreatedAt + if entries[i].State == "ready" { + left = entries[i].LastUsedAt + } + if entries[j].State == "ready" { + right = entries[j].LastUsedAt + } + return left.Before(right) + }) + if len(entries) > limit { + entries = entries[:limit] + } + return entries, nil +} + +func (d *DB) queryCacheEntries(ctx context.Context, query string, args ...any) ([]CacheEntry, error) { + rows, err := d.QueryContext(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + var entries []CacheEntry + for rows.Next() { + entry, err := scanCacheEntry(rows) + if err != nil { + return nil, err + } + entries = append(entries, *entry) + } + return entries, rows.Err() +} + +func (d *DB) DeleteCacheEntry(ctx context.Context, id string) error { + _, err := d.ExecContext(ctx, `delete from cache_entries where id = ?`, id) + return err +} + +func (d *DB) CacheUsageByOwner(ctx context.Context, ownerDID string) (bytes, count int64, err error) { + err = d.QueryRowContext(ctx, ` + select coalesce(sum(size_bytes), 0), count(*) + from cache_entries + where owner_did = ? and state in ('ready', 'deleting')`, ownerDID).Scan(&bytes, &count) + return bytes, count, err +} + +type cacheEntryScanner interface { + Scan(dest ...any) error +} + +func scanCacheEntry(row cacheEntryScanner) (*CacheEntry, error) { + var entry CacheEntry + var createdAt, lastUsedAt int64 + if err := row.Scan( + &entry.ID, + &entry.StorageKey, + &entry.OwnerDID, + &entry.RepoDID, + &entry.Engine, + &entry.CacheKey, + &entry.CacheHash, + &entry.SizeBytes, + &entry.State, + &createdAt, + &lastUsedAt, + ); err != nil { + return nil, err + } + entry.CreatedAt = time.Unix(0, createdAt) + entry.LastUsedAt = time.Unix(0, lastUsedAt) + return &entry, nil +} diff --git a/spindle/db/cache_test.go b/spindle/db/cache_test.go new file mode 100644 index 00000000..f538f5a7 --- /dev/null +++ b/spindle/db/cache_test.go @@ -0,0 +1,302 @@ +package db + +import ( + "context" + "database/sql" + "errors" + "testing" + "time" +) + +func testCacheEntry(id, hash, state string, createdAt time.Time) CacheEntry { + return CacheEntry{ + ID: id, + StorageKey: "objects/" + id, + OwnerDID: "did:plc:owner", + RepoDID: "did:plc:repo", + Engine: "microvm", + CacheKey: "dependencies", + CacheHash: hash, + SizeBytes: 10, + State: state, + CreatedAt: createdAt, + LastUsedAt: createdAt, + } +} + +func insertTestCacheEntry(t *testing.T, d *DB, entry CacheEntry) { + t.Helper() + if err := d.InsertCacheEntry(context.Background(), entry); err != nil { + t.Fatalf("InsertCacheEntry(%s): %v", entry.ID, err) + } +} + +func TestCacheEntryLookup(t *testing.T) { + ctx := context.Background() + d := newTestDB(t) + base := time.Date(2026, 1, 2, 3, 4, 5, 6, time.UTC) + + insertTestCacheEntry(t, d, testCacheEntry("exact-old", "requested", "ready", base)) + insertTestCacheEntry(t, d, testCacheEntry("exact-new", "requested", "ready", base.Add(time.Second))) + insertTestCacheEntry(t, d, testCacheEntry("exact-pending", "requested", "pending", base.Add(2*time.Second))) + insertTestCacheEntry(t, d, testCacheEntry("fallback-old", "old-a", "ready", base.Add(3*time.Second))) + insertTestCacheEntry(t, d, testCacheEntry("fallback-new", "old-b", "ready", base.Add(4*time.Second))) + insertTestCacheEntry(t, d, testCacheEntry("fallback-pending", "old-c", "pending", base.Add(5*time.Second))) + + exact, err := d.FindCacheEntry(ctx, "did:plc:repo", "microvm", "dependencies", "requested") + if err != nil { + t.Fatalf("FindCacheEntry: %v", err) + } + if exact.ID != "exact-new" { + t.Fatalf("FindCacheEntry returned %q, want exact-new", exact.ID) + } + if !exact.CreatedAt.Equal(base.Add(time.Second)) || !exact.LastUsedAt.Equal(base.Add(time.Second)) { + t.Fatalf("timestamps = (%v, %v), want %v", exact.CreatedAt, exact.LastUsedAt, base.Add(time.Second)) + } + + fallback, err := d.FindFallbackCacheEntry(ctx, "did:plc:repo", "microvm", "dependencies", "requested") + if err != nil { + t.Fatalf("FindFallbackCacheEntry: %v", err) + } + if fallback.ID != "fallback-new" { + t.Fatalf("FindFallbackCacheEntry returned %q, want fallback-new", fallback.ID) + } + + if _, err := d.FindCacheEntry(ctx, "did:plc:repo", "microvm", "missing", "requested"); !errors.Is(err, sql.ErrNoRows) { + t.Fatalf("missing FindCacheEntry error = %v, want sql.ErrNoRows", err) + } +} + +func TestMarkCacheEntryReadySupersedesMatchingReady(t *testing.T) { + ctx := context.Background() + d := newTestDB(t) + base := time.Date(2026, 1, 3, 4, 5, 6, 7, time.UTC) + first := testCacheEntry("first-completion", "same-hash", "pending", base) + second := testCacheEntry("second-completion", "same-hash", "pending", base.Add(time.Second)) + insertTestCacheEntry(t, d, first) + insertTestCacheEntry(t, d, second) + + superseded, err := d.MarkCacheEntryReady(ctx, first.ID, 100, base.Add(2*time.Second)) + if err != nil { + t.Fatalf("mark first ready: %v", err) + } + if len(superseded) != 0 { + t.Fatalf("first completion superseded %d entries, want none", len(superseded)) + } + superseded, err = d.MarkCacheEntryReady(ctx, second.ID, 200, base.Add(3*time.Second)) + if err != nil { + t.Fatalf("mark second ready: %v", err) + } + if len(superseded) != 1 || superseded[0].ID != first.ID || superseded[0].State != "deleting" { + t.Fatalf("second completion superseded %+v, want deleting %s", superseded, first.ID) + } + + ready, err := d.FindCacheEntry(ctx, second.RepoDID, second.Engine, second.CacheKey, second.CacheHash) + if err != nil { + t.Fatalf("find surviving ready entry: %v", err) + } + if ready.ID != second.ID || ready.SizeBytes != 200 || !ready.LastUsedAt.Equal(base.Add(3*time.Second)) { + t.Fatalf("surviving entry = %+v, want %s (200 bytes)", ready, second.ID) + } +} + +func TestMarkCacheEntryReadyConcurrentCompletions(t *testing.T) { + ctx := context.Background() + d := newTestDB(t) + base := time.Date(2026, 1, 4, 5, 6, 7, 8, time.UTC) + first := testCacheEntry("concurrent-a", "same-hash", "pending", base) + second := testCacheEntry("concurrent-b", "same-hash", "pending", base.Add(time.Second)) + insertTestCacheEntry(t, d, first) + insertTestCacheEntry(t, d, second) + + type result struct { + superseded []CacheEntry + err error + } + start := make(chan struct{}) + results := make(chan result, 2) + for _, entry := range []CacheEntry{first, second} { + entry := entry + go func() { + <-start + superseded, err := d.MarkCacheEntryReady(ctx, entry.ID, 100, base.Add(2*time.Second)) + results <- result{superseded: superseded, err: err} + }() + } + close(start) + + var superseded []CacheEntry + for range 2 { + result := <-results + if result.err != nil { + t.Fatalf("concurrent MarkCacheEntryReady: %v", result.err) + } + superseded = append(superseded, result.superseded...) + } + if len(superseded) != 1 || superseded[0].State != "deleting" { + t.Fatalf("concurrent completions superseded %+v, want one deleting entry", superseded) + } + + var readyCount, deletingCount int + if err := d.QueryRowContext(ctx, ` + select sum(state = 'ready'), sum(state = 'deleting') + from cache_entries + where repo_did = ? and engine = ? and cache_key = ? and cache_hash = ?`, + first.RepoDID, first.Engine, first.CacheKey, first.CacheHash).Scan(&readyCount, &deletingCount); err != nil { + t.Fatalf("count completion states: %v", err) + } + if readyCount != 1 || deletingCount != 1 { + t.Fatalf("completion states = %d ready, %d deleting; want one each", readyCount, deletingCount) + } +} + +func TestCacheEntryReadyTouchAndExpiry(t *testing.T) { + ctx := context.Background() + d := newTestDB(t) + base := time.Date(2026, 2, 3, 4, 5, 6, 7, time.UTC) + + readyOld := testCacheEntry("ready-old", "a", "ready", base) + readyFresh := testCacheEntry("ready-fresh", "b", "ready", base) + pendingOld := testCacheEntry("pending-old", "c", "pending", base) + pendingFresh := testCacheEntry("pending-fresh", "d", "pending", base.Add(20*time.Minute)) + deletingOld := testCacheEntry("deleting-old", "e", "deleting", base) + deletingFresh := testCacheEntry("deleting-fresh", "f", "deleting", base.Add(20*time.Minute)) + for _, entry := range []CacheEntry{readyOld, readyFresh, pendingOld, pendingFresh, deletingOld, deletingFresh} { + insertTestCacheEntry(t, d, entry) + } + + touchedAt := base.Add(30 * time.Minute) + if err := d.TouchCacheEntry(ctx, readyFresh.ID, touchedAt); err != nil { + t.Fatalf("TouchCacheEntry: %v", err) + } + + expired, err := d.ExpiredCacheEntries(ctx, base.Add(10*time.Minute), base.Add(10*time.Minute), 10) + if err != nil { + t.Fatalf("ExpiredCacheEntries: %v", err) + } + got := make(map[string]bool, len(expired)) + for _, entry := range expired { + got[entry.ID] = true + } + if len(got) != 3 || !got[readyOld.ID] || !got[pendingOld.ID] || !got[deletingOld.ID] { + t.Fatalf("expired IDs = %v, want ready-old, pending-old, and deleting-old", got) + } + + limited, err := d.ExpiredCacheEntries(ctx, base.Add(10*time.Minute), base.Add(10*time.Minute), 1) + if err != nil { + t.Fatalf("limited ExpiredCacheEntries: %v", err) + } + if len(limited) != 1 { + t.Fatalf("limited expiry count = %d, want 1", len(limited)) + } +} + +func TestCacheEntryClaim(t *testing.T) { + ctx := context.Background() + d := newTestDB(t) + base := time.Date(2026, 2, 4, 5, 6, 7, 8, time.UTC) + entry := testCacheEntry("claim-me", "hash", "ready", base) + insertTestCacheEntry(t, d, entry) + + touchedAt := base.Add(time.Minute) + if err := d.TouchCacheEntry(ctx, entry.ID, touchedAt); err != nil { + t.Fatalf("TouchCacheEntry: %v", err) + } + claimed, err := d.ClaimCacheEntry(ctx, entry.ID, "ready", base) + if err != nil { + t.Fatalf("stale ClaimCacheEntry: %v", err) + } + if claimed { + t.Fatal("stale ClaimCacheEntry claimed a touched entry") + } + claimed, err = d.ClaimCacheEntry(ctx, entry.ID, "ready", touchedAt) + if err != nil { + t.Fatalf("ClaimCacheEntry: %v", err) + } + if !claimed { + t.Fatal("ClaimCacheEntry did not claim unchanged entry") + } + + if err := d.TouchCacheEntry(ctx, entry.ID, base.Add(2*time.Minute)); err != nil { + t.Fatalf("TouchCacheEntry while deleting: %v", err) + } + if _, err := d.MarkCacheEntryReady(ctx, entry.ID, 999, base.Add(3*time.Minute)); !errors.Is(err, sql.ErrNoRows) { + t.Fatalf("MarkCacheEntryReady while deleting error = %v, want sql.ErrNoRows", err) + } + var state string + var lastUsedAt int64 + var sizeBytes int64 + if err := d.QueryRowContext(ctx, ` + select state, last_used_at, size_bytes from cache_entries where id = ?`, entry.ID).Scan(&state, &lastUsedAt, &sizeBytes); err != nil { + t.Fatalf("query claimed entry: %v", err) + } + if state != "deleting" || lastUsedAt != touchedAt.UnixNano() || sizeBytes != entry.SizeBytes { + t.Fatalf("claimed entry = state %q, last used %d, size %d; want deleting, %d, %d", + state, lastUsedAt, sizeBytes, touchedAt.UnixNano(), entry.SizeBytes) + } + + if err := d.RestoreCacheEntryState(ctx, entry.ID, "ready"); err != nil { + t.Fatalf("RestoreCacheEntryState: %v", err) + } + restored, err := d.FindCacheEntry(ctx, entry.RepoDID, entry.Engine, entry.CacheKey, entry.CacheHash) + if err != nil { + t.Fatalf("find restored entry: %v", err) + } + if restored.State != "ready" || !restored.LastUsedAt.Equal(touchedAt) { + t.Fatalf("restored entry = state %q, last used %v", restored.State, restored.LastUsedAt) + } +} + +func TestCacheEntryDelete(t *testing.T) { + ctx := context.Background() + d := newTestDB(t) + entry := testCacheEntry("delete-me", "hash", "ready", time.Now()) + insertTestCacheEntry(t, d, entry) + + if err := d.DeleteCacheEntry(ctx, entry.ID); err != nil { + t.Fatalf("DeleteCacheEntry: %v", err) + } + if err := d.DeleteCacheEntry(ctx, entry.ID); err != nil { + t.Fatalf("second DeleteCacheEntry: %v", err) + } + if _, err := d.FindCacheEntry(ctx, entry.RepoDID, entry.Engine, entry.CacheKey, entry.CacheHash); !errors.Is(err, sql.ErrNoRows) { + t.Fatalf("find deleted error = %v, want sql.ErrNoRows", err) + } +} + +func TestCacheUsageByOwner(t *testing.T) { + ctx := context.Background() + d := newTestDB(t) + base := time.Date(2026, 3, 4, 5, 6, 7, 8, time.UTC) + + first := testCacheEntry("owner-ready-a", "a", "ready", base) + first.SizeBytes = 40 + second := testCacheEntry("owner-ready-b", "b", "ready", base) + second.SizeBytes = 2 + deleting := testCacheEntry("owner-deleting", "old", "deleting", base) + deleting.SizeBytes = 3 + pending := testCacheEntry("owner-pending", "c", "pending", base) + pending.SizeBytes = 1000 + other := testCacheEntry("other-ready", "d", "ready", base) + other.OwnerDID = "did:plc:other" + other.SizeBytes = 500 + for _, entry := range []CacheEntry{first, second, deleting, pending, other} { + insertTestCacheEntry(t, d, entry) + } + + bytes, count, err := d.CacheUsageByOwner(ctx, "did:plc:owner") + if err != nil { + t.Fatalf("CacheUsageByOwner: %v", err) + } + if bytes != 45 || count != 3 { + t.Fatalf("usage = (%d bytes, %d entries), want (45, 3)", bytes, count) + } + + bytes, count, err = d.CacheUsageByOwner(ctx, "did:plc:missing") + if err != nil { + t.Fatalf("empty CacheUsageByOwner: %v", err) + } + if bytes != 0 || count != 0 { + t.Fatalf("empty usage = (%d bytes, %d entries), want zero", bytes, count) + } +} diff --git a/spindle/db/db.go b/spindle/db/db.go index 8a74d6d2..741ba0d1 100644 --- a/spindle/db/db.go +++ b/spindle/db/db.go @@ -107,6 +107,30 @@ func Make(ctx context.Context, dbPath string) (*DB, error) { updated_at text not null ); + create table if not exists cache_entries ( + id text primary key, + storage_key text unique not null, + owner_did text not null, + repo_did text not null, + engine text not null, + cache_key text not null, + cache_hash text not null, + size_bytes integer not null default 0, + state text not null check (state in ('pending', 'ready', 'deleting')), + created_at integer not null, + last_used_at integer not null + ); + + create index if not exists cache_entries_lookup + on cache_entries (repo_did, engine, cache_key, cache_hash, created_at desc) + where state = 'ready'; + create index if not exists cache_entries_ready_expiry + on cache_entries (last_used_at) where state = 'ready'; + create index if not exists cache_entries_pending_expiry + on cache_entries (created_at) where state in ('pending', 'deleting'); + create index if not exists cache_entries_owner_usage + on cache_entries (owner_did) where state in ('ready', 'deleting'); + create table if not exists pipelines ( id text primary key, repo_did text not null, @@ -136,7 +160,7 @@ func Make(ctx context.Context, dbPath string) (*DB, error) { return nil, err } - return &DB{db}, nil + return &DB{DB: db}, nil } func runMigrations(_ context.Context, conn *sql.Conn, logger *slog.Logger) error { diff --git a/spindle/engine/cache.go b/spindle/engine/cache.go new file mode 100644 index 00000000..55137b2d --- /dev/null +++ b/spindle/engine/cache.go @@ -0,0 +1,333 @@ +package engine + +import ( + "bufio" + "bytes" + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "io" + "log/slog" + "os/exec" + "strings" + "sync" + "time" + + "github.com/google/uuid" + "tangled.org/core/spindle/db" + + "tangled.org/core/spindle/models" + "tangled.org/core/spindle/storage" +) + +// cache log steps live below the setup step (-1) +const ( + CacheRestoreStepIdx = -2 + CacheSaveStepIdx = -3 +) + +type cacheStep struct { + name string + command string +} + +func (s cacheStep) Name() string { return s.name } +func (s cacheStep) Command() string { return s.command } +func (s cacheStep) Kind() models.StepKind { return models.StepKindSystem } + +var ( + CacheRestoreStep models.Step = cacheStep{name: "restore cache", command: "restore cached paths"} + CacheSaveStep models.Step = cacheStep{name: "save cache", command: "persist changed paths"} +) + +type CacheRunner interface { + RestoreCache(ctx context.Context, wid models.WorkflowId, wf *models.Workflow, store storage.Storage, caches []ResolvedCache, wfLogger models.WorkflowLogger) error + SaveCache(ctx context.Context, wid models.WorkflowId, wf *models.Workflow, store storage.Storage, caches []ResolvedCache, wfLogger models.WorkflowLogger) error +} + +type ResolvedCache struct { + Paths []string + Key string + Hash string + SaveKey string + RestoreID string + RestoreKey string + RestoreName string + CompressionLevel int + When string +} + +func (rc ResolvedCache) saveOn(failed bool) bool { + return rc.When == "always" || !failed +} + +const CacheExitNoPaths = 42 + +// avoids storing an empty archive when zstd is missing +const CacheExitNoCompressor = 43 + +func AbsolutizePaths(paths []string, workspaceRoot string) []string { + abs := make([]string, 0, len(paths)) + for _, p := range paths { + if !strings.HasPrefix(p, "/") { + p = workspaceRoot + "/" + p + } + abs = append(abs, p) + } + return abs +} + +// run tar from / so absolute paths survive extraction +// tar exits nonzero on missing paths, so only existing ones reach it +func CacheSaveScript(paths []string, workspaceRoot string, compressionLevel int) string { + trimmed := make([]string, 0, len(paths)) + for _, p := range AbsolutizePaths(paths, workspaceRoot) { + trimmed = append(trimmed, strings.TrimPrefix(p, "/")) + } + tail := fmt.Sprintf(`tar -cf - -C / "$@" | %s`, CacheCompressCmd(compressionLevel)) + return fmt.Sprintf(`set -o pipefail +command -v zstd >/dev/null 2>&1 || { echo "zstd not found in image; cannot save cache" >&2; exit %d; } +set -- +for p in %s; do [ -e "/$p" ] && set -- "$@" "$p"; done +if [ $# -eq 0 ]; then echo "no cache paths exist; skipping" >&2; exit %d; fi +%s`, CacheExitNoCompressor, strings.Join(trimmed, " "), CacheExitNoPaths, tail) +} + +func CacheCompressCmd(level int) string { + if level == 0 { + return "zstd -T0 -5" + } + return fmt.Sprintf("zstd -T0 -%d", level) +} + +// old entries might still be gzip, so detect them instead of trusting config +func CacheDecompressCmd(br *bufio.Reader) string { + head, _ := br.Peek(4) + if bytes.HasPrefix(head, []byte{0x1f, 0x8b}) { + return "gzip -dc" + } + return "zstd -dc" +} + +// Put keeps draining the pipe after it returns, so the guest writer never +// blocks on a full pipe. +type CacheUpload struct { + Writer *io.PipeWriter + done chan error +} + +func NewCacheUpload(ctx context.Context, store storage.Storage, key string) *CacheUpload { + pr, pw := io.Pipe() + u := &CacheUpload{Writer: pw, done: make(chan error, 1)} + go func() { + err := store.Put(ctx, key, pr) + _, _ = io.Copy(io.Discard, pr) + u.done <- err + }() + return u +} + +func (u *CacheUpload) Abort(err error) { + u.Writer.CloseWithError(err) + <-u.done +} + +// storage treats EOF as a complete archive, so only a clean exec may Finish +func (u *CacheUpload) Finish() error { + u.Writer.Close() + return <-u.done +} + +type indexedCacheStore struct { + storage.Storage + index *db.DB + logger *slog.Logger + mu sync.Mutex + entries map[string]string // storage key -> cache entry id +} + +func (s *indexedCacheStore) Get(ctx context.Context, key string) (io.ReadCloser, error) { + s.mu.Lock() + id, ok := s.entries[key] + s.mu.Unlock() + r, err := s.Storage.Get(ctx, key) + if err != nil { + if ok && errors.Is(err, storage.ErrNotExist) { + _ = s.index.DeleteCacheEntry(context.WithoutCancel(ctx), id) + } + return nil, err + } + if ok { + if err := s.index.TouchCacheEntry(ctx, id, time.Now()); err != nil { + s.logger.Warn("cache usage update failed", "id", id, "err", err) + } + } + return r, nil +} + +func (s *indexedCacheStore) Put(ctx context.Context, key string, r io.Reader) error { + s.mu.Lock() + id, ok := s.entries[key] + s.mu.Unlock() + if !ok { + return fmt.Errorf("cache metadata missing for %q", key) + } + counted := &countingReader{r: r} + if err := s.Storage.Put(ctx, key, counted); err != nil { + return err + } + superseded, err := s.index.MarkCacheEntryReady(ctx, id, counted.n, time.Now()) + if err != nil { + _ = s.Storage.Delete(context.WithoutCancel(ctx), key) + return fmt.Errorf("mark cache ready: %w", err) + } + // completed saves leave the pending map so cleanup keeps their object + s.mu.Lock() + delete(s.entries, key) + s.mu.Unlock() + cleanupCtx := context.WithoutCancel(ctx) + for _, old := range superseded { + s.deleteEntry(cleanupCtx, old.StorageKey, old.ID, "replaced") + } + return nil +} + +func (s *indexedCacheStore) cleanup(ctx context.Context) { + s.mu.Lock() + defer s.mu.Unlock() + for key, id := range s.entries { + s.deleteEntry(ctx, key, id, "incomplete") + } +} + +func (s *indexedCacheStore) deleteEntry(ctx context.Context, key, id, reason string) { + if err := s.Storage.Delete(ctx, key); err != nil { + s.logger.Warn("delete "+reason+" cache failed", "key", key, "err", err) + return + } + if err := s.index.DeleteCacheEntry(ctx, id); err != nil { + s.logger.Warn("delete "+reason+" cache metadata failed", "id", id, "err", err) + } +} + +type countingReader struct { + r io.Reader + n int64 +} + +func (r *countingReader) Read(p []byte) (int, error) { + n, err := r.r.Read(p) + r.n += int64(n) + return n, err +} + +func cacheStoreForRestore(base storage.Storage, index *db.DB, l *slog.Logger, entries []ResolvedCache) storage.Storage { + mapped := make(map[string]string, len(entries)) + for _, entry := range entries { + if entry.RestoreKey != "" { + mapped[entry.RestoreKey] = entry.RestoreID + } + } + return &indexedCacheStore{Storage: base, index: index, logger: l, entries: mapped} +} + +func prepareCacheSaves(ctx context.Context, base storage.Storage, index *db.DB, l *slog.Logger, ownerDID, repoDID, engineName string, entries []ResolvedCache) (*indexedCacheStore, error) { + saveStore := &indexedCacheStore{ + Storage: base, + index: index, + logger: l, + entries: make(map[string]string, len(entries)), + } + now := time.Now() + for i := range entries { + id := uuid.NewString() + entries[i].SaveKey = "objects/" + id + if err := index.InsertCacheEntry(ctx, db.CacheEntry{ + ID: id, + StorageKey: entries[i].SaveKey, + OwnerDID: ownerDID, + RepoDID: repoDID, + Engine: engineName, + CacheKey: entries[i].Key, + CacheHash: entries[i].Hash, + State: "pending", + CreatedAt: now, + LastUsedAt: now, + }); err != nil { + saveStore.cleanup(context.WithoutCancel(ctx)) + return nil, err + } + saveStore.entries[entries[i].SaveKey] = id + } + return saveStore, nil +} + +// on a hash miss, the newest older generation still warms the build +// unusable entries degrade to a plain miss +func ResolveCaches(ctx context.Context, l *slog.Logger, index *db.DB, repoDID, engine, repoPath, rev string, entries []models.CacheEntry) []ResolvedCache { + resolved := make([]ResolvedCache, 0, len(entries)) + for _, entry := range entries { + hash := "" + if len(entry.Hash) > 0 && repoPath != "" { + sum, missing := hashKeyFiles(ctx, repoPath, rev, entry.Hash) + for _, m := range missing { + l.Warn("cache hash file not in repo", "key", entry.Key, "path", m) + } + hash = sum + } + rc := ResolvedCache{ + Paths: entry.Paths, + Key: entry.Key, + Hash: hash, + CompressionLevel: entry.CompressionLevel, + When: entry.When, + } + + found, err := index.FindCacheEntry(ctx, repoDID, engine, entry.Key, hash) + if errors.Is(err, sql.ErrNoRows) && hash != "" { + found, err = index.FindFallbackCacheEntry(ctx, repoDID, engine, entry.Key, hash) + if err == nil { + rc.RestoreName = entry.Key + "-" + found.CacheHash + } + } + if err != nil && !errors.Is(err, sql.ErrNoRows) { + l.Warn("cache lookup failed; entry will save but not restore", "key", entry.Key, "err", err) + } else if err == nil { + rc.RestoreID = found.ID + rc.RestoreKey = found.StorageKey + } + resolved = append(resolved, rc) + } + return resolved +} + +func hashKeyFiles(ctx context.Context, repoPath, rev string, paths []string) (string, []string) { + h := sha256.New() + var missing []string + hashed := 0 + for _, p := range paths { + blob, err := gitBlobId(ctx, repoPath, rev, p) + if err != nil { + missing = append(missing, p) + continue + } + fmt.Fprintf(h, "%s=%s\n", p, blob) + hashed++ + } + if hashed == 0 { + return "", missing + } + return hex.EncodeToString(h.Sum(nil))[:12], missing +} + +// sparse checkouts might not have the file, but the object database does +func gitBlobId(ctx context.Context, repoPath, rev, path string) (string, error) { + out, err := exec.CommandContext(ctx, "git", "-C", repoPath, "rev-parse", rev+":"+path).Output() + if err != nil { + return "", err + } + return strings.TrimSpace(string(out)), nil +} diff --git a/spindle/engine/cache_prune.go b/spindle/engine/cache_prune.go new file mode 100644 index 00000000..f2b08c2c --- /dev/null +++ b/spindle/engine/cache_prune.go @@ -0,0 +1,85 @@ +package engine + +import ( + "context" + "log/slog" + "time" + + "tangled.org/core/spindle/db" + "tangled.org/core/spindle/storage" +) + +const ( + cachePruneBatch = 100 + cachePendingMaxAge = time.Hour +) + +func StartCachePruner(ctx context.Context, l *slog.Logger, index *db.DB, store storage.Storage, retention, interval time.Duration) { + if store == nil || interval <= 0 { + return + } + go func() { + prune := func() { + total := 0 + for { + n, err := PruneCaches(ctx, index, store, time.Now(), retention, cachePendingMaxAge, cachePruneBatch) + total += n + if err != nil { + l.Warn("cache prune failed", "count", total, "err", err) + return + } + if n < cachePruneBatch { + if total > 0 { + l.Info("pruned cache entries", "count", total) + } + return + } + } + } + prune() + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + prune() + } + } + }() +} + +func PruneCaches(ctx context.Context, index *db.DB, store storage.Storage, now time.Time, retention, pendingMaxAge time.Duration, limit int) (int, error) { + readyBefore := time.Unix(0, 0) + if retention > 0 { + readyBefore = now.Add(-retention) + } + entries, err := index.ExpiredCacheEntries(ctx, readyBefore, now.Add(-pendingMaxAge), limit) + if err != nil { + return 0, err + } + pruned := 0 + for _, entry := range entries { + if entry.State != "deleting" { + claimed, err := index.ClaimCacheEntry(ctx, entry.ID, entry.State, entry.LastUsedAt) + if err != nil { + return pruned, err + } + if !claimed { + continue + } + } + if err := store.Delete(ctx, entry.StorageKey); err != nil { + if entry.State != "deleting" { + _ = index.RestoreCacheEntryState(context.WithoutCancel(ctx), entry.ID, entry.State) + } + return pruned, err + } + if err := index.DeleteCacheEntry(ctx, entry.ID); err != nil { + return pruned, err + } + pruned++ + } + return pruned, nil +} diff --git a/spindle/engine/cache_prune_test.go b/spindle/engine/cache_prune_test.go new file mode 100644 index 00000000..1b623d34 --- /dev/null +++ b/spindle/engine/cache_prune_test.go @@ -0,0 +1,152 @@ +package engine + +import ( + "context" + "errors" + "slices" + "testing" + "time" +) + +func TestPruneCachesExpiryPolicies(t *testing.T) { + ctx := context.Background() + now := time.Date(2026, 5, 6, 7, 8, 9, 0, time.UTC) + type spec struct { + id string + state string + age time.Duration + } + tests := []struct { + name string + retention time.Duration + pendingMax time.Duration + entries []spec + wantPruned int + wantSurvive []string + }{ + { + name: "ready entries expire by last use", + retention: time.Hour, + pendingMax: 15 * time.Minute, + entries: []spec{{"old", "ready", 2 * time.Hour}, {"fresh", "ready", 30 * time.Minute}}, + wantPruned: 1, + wantSurvive: []string{"fresh"}, + }, + { + name: "zero retention keeps ready entries", + retention: 0, + pendingMax: time.Hour, + entries: []spec{{"ready", "ready", 24 * time.Hour}, {"pending", "pending", 2 * time.Hour}}, + wantPruned: 1, + wantSurvive: []string{"ready"}, + }, + { + name: "pending entries expire by age", + retention: time.Hour, + pendingMax: time.Hour, + entries: []spec{{"old", "pending", 2 * time.Hour}, {"fresh", "pending", 10 * time.Minute}}, + wantPruned: 1, + wantSurvive: []string{"fresh"}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + d := newCacheTestDB(t) + store := &fakeStorage{objects: map[string][]byte{}} + keys := map[string]string{} + for _, e := range tt.entries { + entry := cacheTestEntry(e.id, "did:plc:repo", "microvm", "deps", "hash", e.state, now.Add(-e.age)) + insertCacheTestEntry(t, d, entry) + store.objects[entry.StorageKey] = []byte(e.id) + keys[e.id] = entry.StorageKey + } + + pruned, err := PruneCaches(ctx, d, store, now, tt.retention, tt.pendingMax, 10) + if err != nil { + t.Fatalf("PruneCaches: %v", err) + } + if pruned != tt.wantPruned { + t.Fatalf("pruned %d entries, want %d", pruned, tt.wantPruned) + } + for _, e := range tt.entries { + survives := slices.Contains(tt.wantSurvive, e.id) + if store.has(keys[e.id]) != survives || cacheTestEntryExists(t, d, e.id) != survives { + t.Fatalf("entry %q survived = %v, want %v", e.id, !survives, survives) + } + } + }) + } +} + +func TestPruneCachesSkipsEntryRefreshedAfterScan(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + now := time.Date(2026, 5, 6, 7, 8, 9, 0, time.UTC) + first := cacheTestEntry("first", "did:plc:repo", "microvm", "deps", "first", "ready", now.Add(-3*time.Hour)) + refreshed := cacheTestEntry("refreshed", "did:plc:repo", "microvm", "deps", "refreshed", "ready", now.Add(-2*time.Hour)) + insertCacheTestEntry(t, d, first) + insertCacheTestEntry(t, d, refreshed) + store := &fakeStorage{objects: map[string][]byte{ + first.StorageKey: []byte("first"), + refreshed.StorageKey: []byte("refreshed"), + }} + store.onDelete = func(key string) { + if key != first.StorageKey { + return + } + if err := d.TouchCacheEntry(ctx, refreshed.ID, now); err != nil { + t.Fatalf("TouchCacheEntry: %v", err) + } + } + + pruned, err := PruneCaches(ctx, d, store, now, time.Hour, time.Hour, 10) + if err != nil { + t.Fatalf("PruneCaches: %v", err) + } + if pruned != 1 { + t.Fatalf("pruned %d entries, want 1", pruned) + } + if !store.has(refreshed.StorageKey) || !cacheTestEntryExists(t, d, refreshed.ID) { + t.Fatal("entry refreshed after expiry scan was pruned") + } +} + +func TestPruneCachesDeleteFailureRemainsRetryable(t *testing.T) { + for _, initialState := range []string{"ready", "deleting"} { + t.Run(initialState, func(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + now := time.Date(2026, 5, 6, 7, 8, 9, 0, time.UTC) + entry := cacheTestEntry("retry", "did:plc:repo", "microvm", "deps", "hash", initialState, now.Add(-2*time.Hour)) + insertCacheTestEntry(t, d, entry) + deleteErr := errors.New("delete failed") + store := &fakeStorage{ + objects: map[string][]byte{entry.StorageKey: []byte("archive")}, + deleteErr: deleteErr, + } + + pruned, err := PruneCaches(ctx, d, store, now, time.Hour, time.Hour, 10) + if !errors.Is(err, deleteErr) { + t.Fatalf("first PruneCaches error = %v, want %v", err, deleteErr) + } + if pruned != 0 { + t.Fatalf("first prune count = %d, want 0", pruned) + } + state, _, _ := cacheTestEntryState(t, d, entry.ID) + if state != initialState { + t.Fatalf("state after delete failure = %q, want %q", state, initialState) + } + if !store.has(entry.StorageKey) { + t.Fatal("failed delete removed object") + } + + pruned, err = PruneCaches(ctx, d, store, now, time.Hour, time.Hour, 10) + if err != nil { + t.Fatalf("retry PruneCaches: %v", err) + } + if pruned != 1 || store.has(entry.StorageKey) || cacheTestEntryExists(t, d, entry.ID) { + t.Fatalf("retry result = pruned %d, object %v, metadata %v", pruned, store.has(entry.StorageKey), cacheTestEntryExists(t, d, entry.ID)) + } + }) + } +} diff --git a/spindle/engine/cache_test.go b/spindle/engine/cache_test.go new file mode 100644 index 00000000..944dfdcc --- /dev/null +++ b/spindle/engine/cache_test.go @@ -0,0 +1,376 @@ +package engine + +import ( + "bufio" + "bytes" + "context" + "errors" + "io" + "log/slog" + "os" + "os/exec" + "path/filepath" + "slices" + "strings" + "testing" + "time" + + "tangled.org/core/spindle/db" + "tangled.org/core/spindle/models" + "tangled.org/core/spindle/storage" +) + +type fakeStorage struct { + objects map[string][]byte + putErr error + deleteErr error // fails once, then clears + onDelete func(string) +} + +func (f *fakeStorage) Get(_ context.Context, key string) (io.ReadCloser, error) { + data, ok := f.objects[key] + if !ok { + return nil, storage.ErrNotExist + } + return io.NopCloser(bytes.NewReader(data)), nil +} + +func (f *fakeStorage) Put(_ context.Context, key string, r io.Reader) error { + data, err := io.ReadAll(r) + if err != nil { + return err + } + f.objects[key] = data + return f.putErr +} + +func (f *fakeStorage) Delete(_ context.Context, key string) error { + if f.onDelete != nil { + f.onDelete(key) + } + if f.deleteErr != nil { + err := f.deleteErr + f.deleteErr = nil + return err + } + delete(f.objects, key) + return nil +} + +func (f *fakeStorage) has(key string) bool { + _, ok := f.objects[key] + return ok +} + +func gitRepo(t *testing.T, files map[string]string) (string, string) { + t.Helper() + if _, err := exec.LookPath("git"); err != nil { + t.Skip("git not available") + } + dir := t.TempDir() + run := func(args ...string) { + t.Helper() + cmd := exec.Command("git", append([]string{"-C", dir, "-c", "user.email=t@t", "-c", "user.name=t", "-c", "init.defaultBranch=main"}, args...)...) + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git %v: %v\n%s", args, err, out) + } + } + run("init") + for name, content := range files { + p := filepath.Join(dir, name) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, []byte(content), 0o644); err != nil { + t.Fatal(err) + } + } + run("add", ".") + run("commit", "-m", "init") + return dir, "HEAD" +} + +func newCacheTestDB(t *testing.T) *db.DB { + t.Helper() + d, err := db.Make(context.Background(), filepath.Join(t.TempDir(), "spindle.db")) + if err != nil { + t.Fatalf("db.Make: %v", err) + } + t.Cleanup(func() { _ = d.Close() }) + return d +} + +func cacheTestEntry(id, repoDID, engine, key, hash, state string, at time.Time) db.CacheEntry { + return db.CacheEntry{ + ID: id, + StorageKey: "objects/" + id, + OwnerDID: "did:plc:owner", + RepoDID: repoDID, + Engine: engine, + CacheKey: key, + CacheHash: hash, + State: state, + CreatedAt: at, + LastUsedAt: at, + } +} + +func insertCacheTestEntry(t *testing.T, d *db.DB, entry db.CacheEntry) { + t.Helper() + if err := d.InsertCacheEntry(context.Background(), entry); err != nil { + t.Fatalf("InsertCacheEntry(%s): %v", entry.ID, err) + } +} + +func cacheTestEntryState(t *testing.T, d *db.DB, id string) (string, int64, time.Time) { + t.Helper() + var state string + var size, lastUsed int64 + if err := d.QueryRow(`select state, size_bytes, last_used_at from cache_entries where id = ?`, id).Scan(&state, &size, &lastUsed); err != nil { + t.Fatalf("query cache entry %s: %v", id, err) + } + return state, size, time.Unix(0, lastUsed) +} + +func cacheTestEntryExists(t *testing.T, d *db.DB, id string) bool { + t.Helper() + var exists bool + if err := d.QueryRow(`select exists(select 1 from cache_entries where id = ?)`, id).Scan(&exists); err != nil { + t.Fatalf("query cache entry existence %s: %v", id, err) + } + return exists +} + +var discardLogger = slog.New(slog.NewTextHandler(io.Discard, nil)) + +func TestResolveCachesExactUnhashed(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + base := time.Date(2026, 4, 5, 6, 7, 8, 0, time.UTC) + + insertCacheTestEntry(t, d, cacheTestEntry("exact", "did:plc:repo", "microvm", "deps", "", "ready", base)) + insertCacheTestEntry(t, d, cacheTestEntry("other-repo", "did:plc:other", "microvm", "deps", "", "ready", base.Add(time.Hour))) + insertCacheTestEntry(t, d, cacheTestEntry("other-engine", "did:plc:repo", "nixery", "deps", "", "ready", base.Add(time.Hour))) + insertCacheTestEntry(t, d, cacheTestEntry("pending", "did:plc:repo", "microvm", "deps", "", "pending", base.Add(time.Hour))) + insertCacheTestEntry(t, d, cacheTestEntry("hashed-only", "did:plc:repo", "microvm", "tools", "old", "ready", base)) + + resolved := ResolveCaches(ctx, discardLogger, d, "did:plc:repo", "microvm", "", "", []models.CacheEntry{ + {Key: "deps", Paths: []string{"/x"}, CompressionLevel: 7, When: "always"}, + {Key: "tools", Paths: []string{"/y"}}, + }) + if len(resolved) != 2 { + t.Fatalf("got %d resolved entries, want 2", len(resolved)) + } + got := resolved[0] + if got.RestoreID != "exact" || got.RestoreKey != "objects/exact" || got.RestoreName != "" { + t.Fatalf("exact restore = (%q, %q, %q)", got.RestoreID, got.RestoreKey, got.RestoreName) + } + if got.SaveKey != "" { + t.Fatalf("resolve allocated a save key %q", got.SaveKey) + } + if got.CompressionLevel != 7 || got.When != "always" { + t.Fatalf("resolved metadata = %+v", got) + } + if resolved[1].RestoreKey != "" { + t.Fatalf("unhashed entry used hashed fallback %q", resolved[1].RestoreKey) + } +} + +func TestResolveCachesHashExactAndFallback(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + repoPath, rev := gitRepo(t, map[string]string{"go.sum": "v1 contents", "go.mod": "module x"}) + request := []models.CacheEntry{{Key: "go-mod", Hash: []string{"go.sum", "go.mod"}, Paths: []string{"/x"}}} + + first := ResolveCaches(ctx, discardLogger, d, "did:plc:repo", "microvm", repoPath, rev, request)[0] + if first.Hash == "" { + t.Fatal("hash is empty") + } + + base := time.Date(2026, 4, 5, 6, 7, 8, 0, time.UTC) + insertCacheTestEntry(t, d, cacheTestEntry("exact", "did:plc:repo", "microvm", "go-mod", first.Hash, "ready", base.Add(time.Minute))) + insertCacheTestEntry(t, d, cacheTestEntry("sibling", "did:plc:repo", "microvm", "go-modules", "newer", "ready", base.Add(time.Hour))) + + exact := ResolveCaches(ctx, discardLogger, d, "did:plc:repo", "microvm", repoPath, rev, request)[0] + if exact.RestoreID != "exact" || exact.RestoreKey != "objects/exact" || exact.RestoreName != "" { + t.Fatalf("exact generation restore = %+v", exact) + } + + if err := os.WriteFile(filepath.Join(repoPath, "go.sum"), []byte("v2 contents"), 0o644); err != nil { + t.Fatal(err) + } + commit := exec.Command("git", "-C", repoPath, "-c", "user.email=t@t", "-c", "user.name=t", "commit", "-qam", "bump") + if out, err := commit.CombinedOutput(); err != nil { + t.Fatalf("commit: %v\n%s", err, out) + } + rotated := ResolveCaches(ctx, discardLogger, d, "did:plc:repo", "microvm", repoPath, rev, request)[0] + if rotated.Hash == first.Hash { + t.Fatal("lockfile change did not rotate the cache hash") + } + if rotated.RestoreID != "exact" || rotated.RestoreKey != "objects/exact" || rotated.RestoreName != "go-mod-"+first.Hash { + t.Fatalf("fallback restore = %+v", rotated) + } +} + +func TestResolveCachesMissingHashFiles(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + repoPath, rev := gitRepo(t, map[string]string{"go.mod": "module x"}) + + partial := ResolveCaches(ctx, discardLogger, d, "did:plc:repo", "microvm", repoPath, rev, []models.CacheEntry{ + {Key: "go-mod", Hash: []string{"go.sum", "go.mod"}, Paths: []string{"/x"}}, + })[0] + if len(partial.Hash) != 12 { + t.Fatalf("partial hash = %q, want 12 characters", partial.Hash) + } + + none := ResolveCaches(ctx, discardLogger, d, "did:plc:repo", "microvm", repoPath, rev, []models.CacheEntry{ + {Key: "go-mod", Hash: []string{"nope.lock"}, Paths: []string{"/x"}}, + })[0] + if none.Hash != "" || none.RestoreKey != "" { + t.Fatalf("all-missing result = hash %q, restore %q", none.Hash, none.RestoreKey) + } +} + +func TestIndexedCacheStoreSaveLifecycle(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + old := cacheTestEntry("superseded", "did:plc:repo", "microvm", "deps", "hash", "ready", time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC)) + insertCacheTestEntry(t, d, old) + base := &fakeStorage{objects: map[string][]byte{old.StorageKey: []byte("old")}} + entries := []ResolvedCache{{Key: "deps", Hash: "hash"}} + + store, err := prepareCacheSaves(ctx, base, d, discardLogger, "did:plc:owner", "did:plc:repo", "microvm", entries) + if err != nil { + t.Fatalf("prepareCacheSaves: %v", err) + } + saveKey := entries[0].SaveKey + id := store.entries[saveKey] + + payload := []byte("cache archive bytes") + if err := store.Put(ctx, saveKey, bytes.NewReader(payload)); err != nil { + t.Fatalf("indexed Put: %v", err) + } + state, size, _ := cacheTestEntryState(t, d, id) + if state != "ready" || size != int64(len(payload)) { + t.Fatalf("saved metadata = state %q, size %d", state, size) + } + if got := base.objects[saveKey]; !bytes.Equal(got, payload) { + t.Fatalf("stored payload = %q, want %q", got, payload) + } + if base.has(old.StorageKey) || cacheTestEntryExists(t, d, old.ID) { + t.Fatal("completed save left the superseded generation") + } + store.cleanup(ctx) + if !base.has(saveKey) || !cacheTestEntryExists(t, d, id) { + t.Fatal("cleanup removed a completed cache") + } +} + +func TestIndexedCacheStoreRestoreTouchesEntry(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + old := time.Date(2025, 1, 2, 3, 4, 5, 0, time.UTC) + entry := cacheTestEntry("restore", "did:plc:repo", "microvm", "deps", "hash", "ready", old) + insertCacheTestEntry(t, d, entry) + base := &fakeStorage{objects: map[string][]byte{entry.StorageKey: []byte("archive")}} + store := cacheStoreForRestore(base, d, discardLogger, []ResolvedCache{{RestoreID: entry.ID, RestoreKey: entry.StorageKey}}) + + r, err := store.Get(ctx, entry.StorageKey) + if err != nil { + t.Fatalf("indexed Get: %v", err) + } + data, err := io.ReadAll(r) + if err != nil { + t.Fatalf("read restored object: %v", err) + } + if err := r.Close(); err != nil { + t.Fatalf("close restored object: %v", err) + } + if string(data) != "archive" { + t.Fatalf("restored payload = %q", data) + } + _, _, touched := cacheTestEntryState(t, d, entry.ID) + if !touched.After(old) { + t.Fatalf("last used = %v, want after %v", touched, old) + } +} + +func TestIndexedCacheStoreFailedPutCleanup(t *testing.T) { + ctx := context.Background() + d := newCacheTestDB(t) + putErr := errors.New("put failed") + base := &fakeStorage{objects: make(map[string][]byte), putErr: putErr} + entries := []ResolvedCache{{Key: "deps"}} + store, err := prepareCacheSaves(ctx, base, d, discardLogger, "did:plc:owner", "did:plc:repo", "microvm", entries) + if err != nil { + t.Fatalf("prepareCacheSaves: %v", err) + } + saveKey := entries[0].SaveKey + + if err := store.Put(ctx, saveKey, strings.NewReader("partial")); !errors.Is(err, putErr) { + t.Fatalf("indexed Put error = %v, want %v", err, putErr) + } + store.cleanup(ctx) + if base.has(saveKey) { + t.Fatal("cleanup left partial object") + } + if cacheTestEntryExists(t, d, store.entries[saveKey]) { + t.Fatal("cleanup left pending metadata") + } +} + +func TestIndexedCacheStoreRejectsUnpreparedKey(t *testing.T) { + store := &indexedCacheStore{} + if err := store.Put(context.Background(), "objects/missing", strings.NewReader("data")); err == nil { + t.Fatal("Put accepted a key without pending metadata") + } +} + +func TestAbsolutizePaths(t *testing.T) { + got := AbsolutizePaths([]string{"node_modules", "/root/.cache", ".gocache"}, "/workspace/repo") + want := []string{"/workspace/repo/node_modules", "/root/.cache", "/workspace/repo/.gocache"} + if !slices.Equal(got, want) { + t.Fatalf("got %v, want %v", got, want) + } +} + +func TestCacheDecompressCmd(t *testing.T) { + cases := []struct { + name string + head []byte + want string + }{ + {"zstd", []byte{0x28, 0xb5, 0x2f, 0xfd, 0x00}, "zstd -dc"}, + {"gzip", []byte{0x1f, 0x8b, 0x08, 0x00}, "gzip -dc"}, + {"empty", nil, "zstd -dc"}, + } + for _, tc := range cases { + got := CacheDecompressCmd(bufio.NewReader(bytes.NewReader(tc.head))) + if got != tc.want { + t.Errorf("%s: got %q, want %q", tc.name, got, tc.want) + } + } +} + +func TestResolvedCacheSaveOn(t *testing.T) { + cases := []struct { + when string + onFail, onPass bool + }{ + {"", false, true}, + {"on-success", false, true}, + {"always", true, true}, + } + for _, tc := range cases { + rc := ResolvedCache{When: tc.when} + if got := rc.saveOn(true); got != tc.onFail { + t.Errorf("when=%q failed run: got %v, want %v", tc.when, got, tc.onFail) + } + if got := rc.saveOn(false); got != tc.onPass { + t.Errorf("when=%q passing run: got %v, want %v", tc.when, got, tc.onPass) + } + } +} + +var _ storage.Storage = (*fakeStorage)(nil) diff --git a/spindle/engine/engine.go b/spindle/engine/engine.go index 6589f9d1..faa1789f 100644 --- a/spindle/engine/engine.go +++ b/spindle/engine/engine.go @@ -13,6 +13,7 @@ import ( "tangled.org/core/spindle/db" "tangled.org/core/spindle/models" "tangled.org/core/spindle/secrets" + "tangled.org/core/spindle/storage" ) var ( @@ -24,18 +25,47 @@ type workflowFinalizer interface { FinalizeWorkflow(ctx context.Context, wid models.WorkflowId, wf *models.Workflow, wfLogger models.WorkflowLogger) error } -func StartWorkflows(l *slog.Logger, vault secrets.Manager, cfg *config.Config, db *db.DB, n *notifier.Notifier, ctx context.Context, pipeline *models.Pipeline, pipelineId models.PipelineId) { +func StartWorkflows(l *slog.Logger, vault secrets.Manager, cfg *config.Config, db *db.DB, n *notifier.Notifier, cacheStore storage.Storage, ctx context.Context, pipeline *models.Pipeline, pipelineId models.PipelineId) { l.Info("starting all workflows in parallel", "pipeline", pipelineId) + isTrustedRepo := pipeline.TrustedSource && pipeline.RepoDid != "" var allSecrets []secrets.UnlockedSecret // never pass secrets to pipelines that run untrusted (e.g. fork) code - if pipeline.TrustedSource && pipeline.RepoDid != "" { + if isTrustedRepo { if res, err := vault.GetSecretsUnlocked(ctx, secrets.RepoIdentifier(pipeline.RepoDid.String())); err == nil { allSecrets = res } } else if !pipeline.TrustedSource { l.Info("skipping secrets for untrusted pipeline source", "pipeline", pipelineId) } + // untrusted runs cant read or write shared caches + cacheOwnerDID := "" + cacheEnabled := cacheStore != nil && isTrustedRepo + if cacheEnabled { + repo, err := db.GetRepoByDid(pipeline.RepoDid) + if err != nil { + l.Warn("cache owner lookup failed; caching disabled", "repo", pipeline.RepoDid, "err", err) + cacheEnabled = false + } else { + cacheOwnerDID = repo.Owner.String() + } + } else if cacheStore != nil && !pipeline.TrustedSource { + l.Info("skipping caches for untrusted pipeline source", "pipeline", pipelineId) + } + + // hash the checked commit, not the working tree + cacheRepoPath, cacheRev := "", "" + if tm := pipeline.TriggerMetadata; tm != nil { + if rev, err := models.ExtractCommitSHA(*tm); err == nil { + did := pipeline.RepoDid.String() + if tm.SourceRepo != nil && *tm.SourceRepo != "" { + did = *tm.SourceRepo + } + cacheRepoPath, cacheRev = filepath.Join(cfg.Server.RepoDir, did), rev + } else { + l.Warn("cannot resolve pipeline commit; cache hashing disabled", "err", err) + } + } secretValues := make([]string, len(allSecrets)) for i, s := range allSecrets { @@ -50,6 +80,7 @@ func StartWorkflows(l *slog.Logger, vault secrets.Manager, cfg *config.Config, d var wg sync.WaitGroup for eng, wfs := range pipeline.Workflows { workflowTimeout := eng.WorkflowTimeout() + cacheRunner, cachesSupported := eng.(CacheRunner) l.Info("using workflow timeout", "timeout", workflowTimeout) for _, w := range wfs { @@ -121,6 +152,46 @@ func StartWorkflows(l *slog.Logger, vault secrets.Manager, cfg *config.Config, d ctx, cancel := context.WithTimeout(ctx, workflowTimeout) defer cancel() + var resolvedCaches []ResolvedCache + if cacheEnabled && len(w.Caches) > 0 { + if cachesSupported { + resolvedCaches = ResolveCaches(ctx, l, db, pipeline.RepoDid.String(), w.Engine, cacheRepoPath, cacheRev, w.Caches) + wfLogger.ControlWriter(CacheRestoreStepIdx, CacheRestoreStep, models.StepStatusStart).Write([]byte{0}) + // caches are an optimization, never a reason to fail the workflow + restoreStore := cacheStoreForRestore(cacheStore, db, l, resolvedCaches) + if err := cacheRunner.RestoreCache(ctx, wid, &w, restoreStore, resolvedCaches, wfLogger); err != nil { + l.Warn("cache restore failed", "wid", wid, "err", err) + } + wfLogger.ControlWriter(CacheRestoreStepIdx, CacheRestoreStep, models.StepStatusEnd).Write([]byte{0}) + } else { + l.Warn("engine does not support caches, skipping restore", "wid", wid) + } + } + + // dont save on timeouts, their context is already dead + saveCaches := func(failed bool) { + toSave := resolvedCaches[:0] + for _, rc := range resolvedCaches { + if rc.saveOn(failed) { + toSave = append(toSave, rc) + } + } + if len(toSave) == 0 { + return + } + wfLogger.ControlWriter(CacheSaveStepIdx, CacheSaveStep, models.StepStatusStart).Write([]byte{0}) + saveStore, err := prepareCacheSaves(ctx, cacheStore, db, l, cacheOwnerDID, pipeline.RepoDid.String(), w.Engine, toSave) + if err != nil { + l.Warn("cache metadata setup failed", "wid", wid, "err", err) + } else { + if err := cacheRunner.SaveCache(ctx, wid, &w, saveStore, toSave, wfLogger); err != nil { + l.Warn("cache save failed", "wid", wid, "err", err) + } + saveStore.cleanup(context.WithoutCancel(ctx)) + } + wfLogger.ControlWriter(CacheSaveStepIdx, CacheSaveStep, models.StepStatusEnd).Write([]byte{0}) + } + for stepIdx, step := range w.Steps { // log start of step if wfLogger != nil { @@ -145,6 +216,7 @@ func StartWorkflows(l *slog.Logger, vault secrets.Manager, cfg *config.Config, d l.Error("failed to set workflow status to timeout", "wid", wid, "err", dbErr) } } else { + saveCaches(true) dbErr := db.StatusFailed(wid, err.Error(), -1, n) if dbErr != nil { l.Error("failed to set workflow status to failed", "wid", wid, "err", dbErr) @@ -154,6 +226,8 @@ func StartWorkflows(l *slog.Logger, vault secrets.Manager, cfg *config.Config, d } } + saveCaches(false) + if finalizer, ok := eng.(workflowFinalizer); ok { if err := finalizer.FinalizeWorkflow(ctx, wid, &w, wfLogger); err != nil { dbErr := db.StatusFailed(wid, err.Error(), -1, n) diff --git a/spindle/engines/microvm/cache.go b/spindle/engines/microvm/cache.go new file mode 100644 index 00000000..c4f7b0da --- /dev/null +++ b/spindle/engines/microvm/cache.go @@ -0,0 +1,125 @@ +package microvm + +import ( + "bufio" + "context" + "fmt" + "io" + + agentv1 "tangled.org/core/spindle/agentproto/gen" + "tangled.org/core/spindle/engine" + "tangled.org/core/spindle/models" + "tangled.org/core/spindle/storage" +) + +func (e *Engine) RestoreCache(ctx context.Context, wid models.WorkflowId, wf *models.Workflow, store storage.Storage, caches []engine.ResolvedCache, wfLogger models.WorkflowLogger) error { + state, ok := wf.Data.(*workflowState) + if !ok || state == nil || state.Agent == nil { + return fmt.Errorf("microVM workflow is not connected to agent") + } + + out := wfLogger.DataWriter(engine.CacheRestoreStepIdx, "stdout") + for _, entry := range caches { + if err := ctx.Err(); err != nil { + return err + } + if entry.RestoreKey == "" { + fmt.Fprintf(out, "cache %q: miss\n", entry.Key) + continue + } + if entry.RestoreName != "" { + fmt.Fprintf(out, "cache %q: restoring from %q\n", entry.Key, entry.RestoreName) + } + + rc, err := store.Get(ctx, entry.RestoreKey) + if err != nil { + fmt.Fprintf(out, "cache %q: fetch failed: %v\n", entry.Key, err) + continue + } + br := bufio.NewReader(rc) + decompress := engine.CacheDecompressCmd(br) + + var restored int64 + exit, err := state.Agent.Exec(ctx, AgentExec{ + ID: fmt.Sprintf("%s-cache-restore", wid.String()), + ExecStart: cacheExecStart(state, fmt.Sprintf("set -o pipefail\n%s | tar -x -C /", decompress)), + Stdin: &countingReader{r: br, n: &restored}, + Stderr: out, + }) + rc.Close() + if err != nil { + fmt.Fprintf(out, "cache %q: restore failed: %v\n", entry.Key, err) + continue + } + if exit != 0 { + fmt.Fprintf(out, "cache %q: restore failed: guest exited %d\n", entry.Key, exit) + continue + } + fmt.Fprintf(out, "cache %q: restored %d bytes\n", entry.Key, restored) + } + return nil +} + +func (e *Engine) SaveCache(ctx context.Context, wid models.WorkflowId, wf *models.Workflow, store storage.Storage, caches []engine.ResolvedCache, wfLogger models.WorkflowLogger) error { + state, ok := wf.Data.(*workflowState) + if !ok || state == nil || state.Agent == nil { + return fmt.Errorf("microVM workflow is not connected to agent") + } + + out := wfLogger.DataWriter(engine.CacheSaveStepIdx, "stdout") + for _, entry := range caches { + if err := ctx.Err(); err != nil { + return err + } + + script := engine.CacheSaveScript(entry.Paths, guestWorkDir, entry.CompressionLevel) + + up := engine.NewCacheUpload(ctx, store, entry.SaveKey) + exit, execErr := state.Agent.Exec(ctx, AgentExec{ + ID: fmt.Sprintf("%s-cache-save", wid.String()), + ExecStart: cacheExecStart(state, script), + Stdout: up.Writer, + Stderr: out, + }) + switch { + case exit == engine.CacheExitNoPaths: + up.Abort(fmt.Errorf("guest exited %d", exit)) + fmt.Fprintf(out, "cache %q: nothing to save\n", entry.Key) + continue + case exit == engine.CacheExitNoCompressor: + up.Abort(fmt.Errorf("guest exited %d", exit)) + fmt.Fprintf(out, "cache %q: zstd not available in image; skipping\n", entry.Key) + continue + case execErr != nil: + up.Abort(execErr) + return fmt.Errorf("save cache %q: %w", entry.Key, execErr) + case exit != 0: + up.Abort(fmt.Errorf("guest exited %d", exit)) + return fmt.Errorf("save cache %q: save script exited %d", entry.Key, exit) + } + if err := up.Finish(); err != nil { + return fmt.Errorf("save cache %q: %w", entry.Key, err) + } + fmt.Fprintf(out, "cache %q: saved\n", entry.Key) + } + return nil +} + +func cacheExecStart(state *workflowState, script string) *agentv1.ExecStart { + return &agentv1.ExecStart{ + Argv: []string{state.ImageSpec.Shell, "-c", script}, + Env: guestBaseEnv(), + User: guestWorkflowUser, + } +} + +type countingReader struct { + r io.Reader + n *int64 +} + +func (c *countingReader) Read(p []byte) (int, error) { + n, err := c.r.Read(p) + *c.n += int64(n) + return n, err +} diff --git a/spindle/engines/microvm/engine.go b/spindle/engines/microvm/engine.go index 2ce3a3ba..909fb08e 100644 --- a/spindle/engines/microvm/engine.go +++ b/spindle/engines/microvm/engine.go @@ -42,6 +42,16 @@ const ( type cleanupFunc func(context.Context) error +// return a fresh slice since callers append to it +func guestBaseEnv() []string { + return []string{ + "HOME=/workspace", + "LOGNAME=" + guestWorkflowUser, + "PATH=" + guestBasePATH, + "USER=" + guestWorkflowUser, + } +} + type Engine struct { l *slog.Logger cfg *config.Config @@ -135,6 +145,13 @@ func (e *Engine) InitWorkflow(twf tangled.Pipeline_Workflow, tpl tangled.Pipelin swf.Name = twf.Name swf.Environment = dwf.Environment + for _, entry := range dwf.Cache { + if err := entry.Validate(); err != nil { + return nil, err + } + } + swf.Caches = dwf.Cache + if tpl.TriggerMetadata != nil { if clone := models.BuildCloneStep(twf, *tpl.TriggerMetadata, e.cfg.Server.Dev); clone.Command() != "" { swf.Steps = append([]models.Step{clone}, swf.Steps...) @@ -360,12 +377,7 @@ func (e *Engine) RunStep(ctx context.Context, wid models.WorkflowId, w *models.W err := e.activateConfig(execCtx, wid, state, s, wfLogger.DataWriter(idx, "stdout")) return e.classifyStepError(ctx, wid, step, state, stderr, vmExited, "Failed to activate config", err) } - env := []string{ - "HOME=/workspace", - "LOGNAME=" + guestWorkflowUser, - "PATH=" + guestBasePATH, - "USER=" + guestWorkflowUser, - } + env := guestBaseEnv() for k, v := range w.Environment { env = append(env, k+"="+v) } diff --git a/spindle/engines/microvm/models.go b/spindle/engines/microvm/models.go index 02f3c096..adc94157 100644 --- a/spindle/engines/microvm/models.go +++ b/spindle/engines/microvm/models.go @@ -3,16 +3,19 @@ package microvm import ( "fmt" "slices" + + "tangled.org/core/spindle/models" ) type manifestWorkflow struct { - Image string `yaml:"image"` - Services map[string]any `yaml:"services"` - Virtualisation map[string]any `yaml:"virtualisation"` - Dependencies []string `yaml:"dependencies"` - Registry map[string]any `yaml:"registry"` - Environment map[string]string `yaml:"environment"` - Caches map[string]string `yaml:"caches"` + Image string `yaml:"image"` + Services map[string]any `yaml:"services"` + Virtualisation map[string]any `yaml:"virtualisation"` + Dependencies []string `yaml:"dependencies"` + Registry map[string]any `yaml:"registry"` + Environment map[string]string `yaml:"environment"` + Caches map[string]string `yaml:"caches"` + Cache []models.CacheEntry `yaml:"cache"` Steps []struct { Name string `yaml:"name"` Command string `yaml:"command"` diff --git a/spindle/engines/nixery/cache.go b/spindle/engines/nixery/cache.go new file mode 100644 index 00000000..0fd8ebed --- /dev/null +++ b/spindle/engines/nixery/cache.go @@ -0,0 +1,167 @@ +package nixery + +import ( + "bufio" + "context" + "fmt" + "io" + + "github.com/docker/docker/api/types" + "github.com/docker/docker/api/types/container" + "github.com/docker/docker/pkg/stdcopy" + + "tangled.org/core/spindle/engine" + "tangled.org/core/spindle/models" + "tangled.org/core/spindle/storage" +) + +func (e *Engine) baseEnv() EnvVars { + envs := EnvVars{} + envs.AddEnv("HOME", homeDir) + envs.AddEnv("PATH", fmt.Sprintf("%s/.nix-profile/bin:/nix/var/nix/profiles/default/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin", homeDir)) + return envs +} + +func (e *Engine) execAttached(ctx context.Context, containerID string, opts container.ExecOptions) (string, types.HijackedResponse, error) { + execResp, err := e.docker.ContainerExecCreate(ctx, containerID, opts) + if err != nil { + return "", types.HijackedResponse{}, fmt.Errorf("create exec: %w", err) + } + attach, err := e.docker.ContainerExecAttach(ctx, execResp.ID, container.ExecAttachOptions{}) + if err != nil { + return "", types.HijackedResponse{}, fmt.Errorf("attach exec: %w", err) + } + return execResp.ID, attach, nil +} + +func (e *Engine) containerID(wf *models.Workflow) (string, error) { + addl, ok := wf.Data.(addlFields) + if !ok || addl.container == "" { + return "", fmt.Errorf("nixery workflow has no container") + } + return addl.container, nil +} + +func (e *Engine) RestoreCache(ctx context.Context, wid models.WorkflowId, wf *models.Workflow, store storage.Storage, caches []engine.ResolvedCache, wfLogger models.WorkflowLogger) error { + containerID, err := e.containerID(wf) + if err != nil { + return err + } + + out := wfLogger.DataWriter(engine.CacheRestoreStepIdx, "stdout") + for _, entry := range caches { + if err := ctx.Err(); err != nil { + return err + } + if entry.RestoreKey == "" { + fmt.Fprintf(out, "cache %q: miss\n", entry.Key) + continue + } + if entry.RestoreName != "" { + fmt.Fprintf(out, "cache %q: restoring from %q\n", entry.Key, entry.RestoreName) + } + + rc, err := store.Get(ctx, entry.RestoreKey) + if err != nil { + fmt.Fprintf(out, "cache %q: fetch failed: %v\n", entry.Key, err) + continue + } + + br := bufio.NewReader(rc) + execID, attach, err := e.execAttached(ctx, containerID, container.ExecOptions{ + Cmd: []string{"bash", "-c", engine.CacheDecompressCmd(br) + " | tar -x -C /"}, + Env: e.baseEnv(), + AttachStdin: true, + AttachStdout: true, + AttachStderr: true, + }) + if err != nil { + rc.Close() + return fmt.Errorf("restore cache %q: %w", entry.Key, err) + } + + // drain this now or tar can block on stderr before reading stdin + copyDone := make(chan error, 1) + go func() { + _, err := io.Copy(attach.Conn, br) + _ = attach.CloseWrite() + copyDone <- err + }() + _, _ = stdcopy.StdCopy(out, out, attach.Reader) + copyErr := <-copyDone + rc.Close() + attach.Close() + if copyErr != nil { + return fmt.Errorf("restore cache %q: stream archive: %w", entry.Key, copyErr) + } + + inspect, err := e.docker.ContainerExecInspect(ctx, execID) + if err != nil { + return fmt.Errorf("restore cache %q: %w", entry.Key, err) + } + if inspect.ExitCode != 0 { + fmt.Fprintf(out, "cache %q: extract failed (exit %d)\n", entry.Key, inspect.ExitCode) + continue + } + fmt.Fprintf(out, "cache %q: restored\n", entry.Key) + } + return nil +} + +func (e *Engine) SaveCache(ctx context.Context, wid models.WorkflowId, wf *models.Workflow, store storage.Storage, caches []engine.ResolvedCache, wfLogger models.WorkflowLogger) error { + containerID, err := e.containerID(wf) + if err != nil { + return err + } + + out := wfLogger.DataWriter(engine.CacheSaveStepIdx, "stdout") + for _, entry := range caches { + if err := ctx.Err(); err != nil { + return err + } + + script := engine.CacheSaveScript(entry.Paths, workspaceDir, entry.CompressionLevel) + + execID, attach, err := e.execAttached(ctx, containerID, container.ExecOptions{ + Cmd: []string{"bash", "-c", script}, + Env: e.baseEnv(), + AttachStdout: true, + AttachStderr: true, + }) + if err != nil { + return fmt.Errorf("save cache %q: %w", entry.Key, err) + } + + up := engine.NewCacheUpload(ctx, store, entry.SaveKey) + + // StdCopy only returns once the archive is fully written + _, copyErr := stdcopy.StdCopy(up.Writer, out, attach.Reader) + attach.Close() + inspect, inspectErr := e.docker.ContainerExecInspect(ctx, execID) + + switch { + case inspectErr == nil && inspect.ExitCode == engine.CacheExitNoPaths: + up.Abort(fmt.Errorf("no cache paths")) + fmt.Fprintf(out, "cache %q: nothing to save\n", entry.Key) + continue + case inspectErr == nil && inspect.ExitCode == engine.CacheExitNoCompressor: + up.Abort(fmt.Errorf("zstd not available")) + fmt.Fprintf(out, "cache %q: zstd not available in image; skipping\n", entry.Key) + continue + case copyErr != nil: + up.Abort(copyErr) + return fmt.Errorf("save cache %q: stream archive: %w", entry.Key, copyErr) + case inspectErr != nil: + up.Abort(inspectErr) + return fmt.Errorf("save cache %q: %w", entry.Key, inspectErr) + case inspect.ExitCode != 0: + up.Abort(fmt.Errorf("exited %d", inspect.ExitCode)) + return fmt.Errorf("save cache %q: tar exited %d", entry.Key, inspect.ExitCode) + } + if err := up.Finish(); err != nil { + return fmt.Errorf("save cache %q: %w", entry.Key, err) + } + fmt.Fprintf(out, "cache %q: saved\n", entry.Key) + } + return nil +} diff --git a/spindle/engines/nixery/engine.go b/spindle/engines/nixery/engine.go index a06d34bf..710ea439 100644 --- a/spindle/engines/nixery/engine.go +++ b/spindle/engines/nixery/engine.go @@ -91,6 +91,7 @@ func (e *Engine) InitWorkflow(twf tangled.Pipeline_Workflow, tpl tangled.Pipelin } `yaml:"steps"` Dependencies map[string][]string `yaml:"dependencies"` Environment map[string]string `yaml:"environment"` + Cache []models.CacheEntry `yaml:"cache"` }{} if err := engine.DescribeManifestError(twf.Raw, dwf); err != nil { return nil, err @@ -109,6 +110,12 @@ func (e *Engine) InitWorkflow(twf tangled.Pipeline_Workflow, tpl tangled.Pipelin } swf.Name = twf.Name swf.Environment = dwf.Environment + for _, entry := range dwf.Cache { + if err := entry.Validate(); err != nil { + return nil, err + } + } + swf.Caches = dwf.Cache addl.image = workflowImage(dwf.Dependencies, e.cfg.NixeryPipelines.Nixery) if sock := e.cfg.Server.DockerSocket; sock != "" { @@ -155,7 +162,7 @@ func workflowImage(deps map[string][]string, nixery string) string { } // load defaults from somewhere else - dependencies = path.Join(dependencies, "bash", "git", "coreutils", "nix") + dependencies = path.Join(dependencies, "bash", "git", "coreutils", "gnutar", "zstd", "nix") if runtime.GOARCH == "arm64" { dependencies = path.Join("arm64", dependencies) @@ -409,9 +416,7 @@ func (e *Engine) RunStep(ctx context.Context, wid models.WorkflowId, w *models.W } } - envs.AddEnv("HOME", homeDir) - existingPath := "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin" - envs.AddEnv("PATH", fmt.Sprintf("%s/.nix-profile/bin:/nix/var/nix/profiles/default/bin:%s", homeDir, existingPath)) + envs = append(envs, e.baseEnv()...) if sock := e.cfg.Server.DockerSocket; sock != "" { envs.AddEnv("DOCKER_HOST", fmt.Sprintf("unix://%s", sock)) } diff --git a/spindle/models/cache.go b/spindle/models/cache.go new file mode 100644 index 00000000..2b022f6b --- /dev/null +++ b/spindle/models/cache.go @@ -0,0 +1,63 @@ +package models + +import ( + "fmt" + "regexp" + "strings" +) + +type CacheEntry struct { + Key string `yaml:"key"` + Hash []string `yaml:"hash"` + Paths []string `yaml:"paths"` + // 1 is fastest, 19 is smallest, 0 is the zstd default + CompressionLevel int `yaml:"compression-level"` + // on-success is the default, always also saves on failed runs + When string `yaml:"when"` +} + +// keys become storage paths, so no slashes +var cacheKeyRe = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$`) + +func (c CacheEntry) Validate() error { + if !cacheKeyRe.MatchString(c.Key) { + return fmt.Errorf("cache: invalid key %q (allowed: letters, digits, '.', '_', '-')", c.Key) + } + for _, f := range c.Hash { + // rev-parse needs plain repo-relative paths, not pathspecs + if f == "" || strings.HasPrefix(f, "/") || strings.HasPrefix(f, "..") { + return fmt.Errorf("cache %q: hash path %q is not repo-relative", c.Key, f) + } + if strings.ContainsAny(f, ": \t\n\"'`$\\*?[") || strings.Contains(f, "/../") || strings.HasSuffix(f, "/..") { + return fmt.Errorf("cache %q: hash path %q contains unsupported characters", c.Key, f) + } + } + if len(c.Paths) == 0 { + return fmt.Errorf("cache %q: no paths", c.Key) + } + if c.CompressionLevel < 0 || c.CompressionLevel > 19 { + return fmt.Errorf("cache %q: compression-level %d out of range (1-19)", c.Key, c.CompressionLevel) + } + switch c.When { + case "", "on-success", "always": + default: + return fmt.Errorf("cache %q: when %q is not one of on-success, always", c.Key, c.When) + } + seen := make(map[string]bool, len(c.Paths)) + for _, p := range c.Paths { + if p == "" { + return fmt.Errorf("cache %q: empty path", c.Key) + } + if strings.ContainsAny(p, " \t\n\"'`$\\") { + return fmt.Errorf("cache %q: path %q contains unsupported characters", c.Key, p) + } + if strings.Contains(p, "..") { + return fmt.Errorf("cache %q: path %q must not contain '..'", c.Key, p) + } + if seen[p] { + return fmt.Errorf("cache %q: duplicate path %q", c.Key, p) + } + seen[p] = true + } + return nil +} diff --git a/spindle/models/cache_test.go b/spindle/models/cache_test.go new file mode 100644 index 00000000..9d112792 --- /dev/null +++ b/spindle/models/cache_test.go @@ -0,0 +1,49 @@ +package models + +import ( + "strings" + "testing" +) + +func TestCacheEntryValidate(t *testing.T) { + valid := CacheEntry{Key: "go-mod-v1", Hash: []string{"go.sum", "sub/dir/package-lock.json"}, Paths: []string{"/workspace/go/pkg/mod", "/root/.cache"}} + if err := valid.Validate(); err != nil { + t.Fatalf("valid entry: %v", err) + } + + cases := []struct { + name string + entry CacheEntry + want string + }{ + {"empty key", CacheEntry{Key: "", Paths: []string{"/x"}}, "invalid key"}, + {"key with slash", CacheEntry{Key: "a/b", Paths: []string{"/x"}}, "invalid key"}, + {"no paths", CacheEntry{Key: "ok"}, "no paths"}, + {"empty path", CacheEntry{Key: "ok", Paths: []string{""}}, "empty path"}, + {"path with space", CacheEntry{Key: "ok", Paths: []string{"/my dir"}}, "unsupported characters"}, + {"path with quote", CacheEntry{Key: "ok", Paths: []string{"/x'$(rm)"}}, "unsupported characters"}, + {"path traversal", CacheEntry{Key: "ok", Paths: []string{"/x/../y"}}, ".."}, + {"duplicate path", CacheEntry{Key: "ok", Paths: []string{"/x", "/x"}}, "duplicate"}, + {"level too high", CacheEntry{Key: "ok", Paths: []string{"/x"}, CompressionLevel: 20}, "out of range"}, + {"negative level", CacheEntry{Key: "ok", Paths: []string{"/x"}, CompressionLevel: -1}, "out of range"}, + {"bad when", CacheEntry{Key: "ok", Paths: []string{"/x"}, When: "sometimes"}, "not one of"}, + {"absolute hash path", CacheEntry{Key: "ok", Hash: []string{"/go.sum"}, Paths: []string{"/x"}}, "not repo-relative"}, + {"empty hash path", CacheEntry{Key: "ok", Hash: []string{""}, Paths: []string{"/x"}}, "not repo-relative"}, + {"hash path traversal", CacheEntry{Key: "ok", Hash: []string{"../secret"}, Paths: []string{"/x"}}, "not repo-relative"}, + {"hash path inner traversal", CacheEntry{Key: "ok", Hash: []string{"a/../../b"}, Paths: []string{"/x"}}, "unsupported characters"}, + {"hash path with colon", CacheEntry{Key: "ok", Hash: []string{"rev:go.sum"}, Paths: []string{"/x"}}, "unsupported characters"}, + {"hash path with space", CacheEntry{Key: "ok", Hash: []string{"my lock"}, Paths: []string{"/x"}}, "unsupported characters"}, + {"hash path with glob", CacheEntry{Key: "ok", Hash: []string{"*.sum"}, Paths: []string{"/x"}}, "unsupported characters"}, + } + for _, tc := range cases { + err := tc.entry.Validate() + if err == nil || !strings.Contains(err.Error(), tc.want) { + t.Errorf("%s: got %v, want error containing %q", tc.name, err, tc.want) + } + } + + ok := CacheEntry{Key: "go-mod", Hash: []string{"go.sum"}, Paths: []string{"node_modules", "/root/.cache"}, CompressionLevel: 19, When: "always"} + if err := ok.Validate(); err != nil { + t.Errorf("relative and absolute paths should validate, got %v", err) + } +} diff --git a/spindle/models/clone.go b/spindle/models/clone.go index d9d54389..1f6be1cb 100644 --- a/spindle/models/clone.go +++ b/spindle/models/clone.go @@ -47,7 +47,7 @@ func BuildCloneStep(twf tangled.Pipeline_Workflow, tr tangled.Pipeline_TriggerMe return CloneStep{} } - commitSHA, err := extractCommitSHA(tr) + commitSHA, err := ExtractCommitSHA(tr) if err != nil { return CloneStep{ kind: StepKindSystem, @@ -83,8 +83,7 @@ func BuildCloneStep(twf tangled.Pipeline_Workflow, tr tangled.Pipeline_TriggerMe } } -// extractCommitSHA extracts the commit SHA from trigger metadata based on trigger type -func extractCommitSHA(tr tangled.Pipeline_TriggerMetadata) (string, error) { +func ExtractCommitSHA(tr tangled.Pipeline_TriggerMetadata) (string, error) { switch workflow.TriggerKind(tr.Kind) { case workflow.TriggerKindPush: if tr.Push == nil { diff --git a/spindle/models/pipeline.go b/spindle/models/pipeline.go index a794644b..fabcce84 100644 --- a/spindle/models/pipeline.go +++ b/spindle/models/pipeline.go @@ -1,12 +1,18 @@ package models -import "github.com/bluesky-social/indigo/atproto/syntax" +import ( + "github.com/bluesky-social/indigo/atproto/syntax" + + "tangled.org/core/api/tangled" +) type Pipeline struct { RepoDid syntax.DID Workflows map[Engine][]Workflow // whether the code being ran was checked out from RepoDid itself TrustedSource bool + // used to resolve cache hash files against the checkout the workflow builds + TriggerMetadata *tangled.Pipeline_TriggerMetadata } type Step interface { @@ -29,4 +35,6 @@ type Workflow struct { Name string Data any Environment map[string]string + Caches []CacheEntry + Engine string } diff --git a/spindle/server.go b/spindle/server.go index dd51d281..b50c0eda 100644 --- a/spindle/server.go +++ b/spindle/server.go @@ -42,6 +42,7 @@ import ( "tangled.org/core/spindle/models" "tangled.org/core/spindle/queue" "tangled.org/core/spindle/secrets" + "tangled.org/core/spindle/storage" "tangled.org/core/spindle/xrpc" "tangled.org/core/tid" "tangled.org/core/workflow" @@ -70,6 +71,7 @@ type Spindle struct { res *idresolver.Resolver verify repoverify.Verifier vault secrets.Manager + cache storage.Storage motd []byte motdMu sync.RWMutex rootCtx context.Context @@ -154,6 +156,14 @@ func New(ctx context.Context, cfg *config.Config, d *db.DB, engines map[string]m resolver := idresolver.DefaultResolver(cfg.Server.PlcUrl) + cacheStore, err := storage.New(ctx, cfg) + if err != nil { + return nil, fmt.Errorf("failed to setup cache storage: %w", err) + } + if cacheStore != nil { + logger.Info("cache storage enabled", "backend", cfg.Cache.Backend) + } + spindle := &Spindle{ jc: jc, e: e, @@ -166,6 +176,7 @@ func New(ctx context.Context, cfg *config.Config, d *db.DB, engines map[string]m res: resolver, verify: repoverify.New(resolver, cfg.Server.Dev), vault: vault, + cache: cacheStore, motd: defaultMotd, rootCtx: ctx, } @@ -224,6 +235,8 @@ func New(ctx context.Context, cfg *config.Config, d *db.DB, engines map[string]m cfg.Server.Tap.AdminPassword = pw logger.Info("embedded tap: using random admin password") } + engine.StartCachePruner(ctx, logger, d, cacheStore, cfg.Cache.Retention, cfg.Cache.PruneInterval) + spindle.tap = NewTapClient(spindle) return spindle, nil @@ -820,16 +833,18 @@ func (s *Spindle) processPipeline(repoDid syntax.DID, tpl tangled.Pipeline, pipe } maps.Copy(ewf.Environment, pipelineEnv) + ewf.Engine = w.Engine workflows[eng] = append(workflows[eng], *ewf) } // enqueue pipeline ok := s.jq.Enqueue(repoDid, queue.Job{ Run: func() error { - engine.StartWorkflows(log.SubLogger(s.l, "engine"), s.vault, s.cfg, s.db, s.n, s.rootCtx, &models.Pipeline{ - RepoDid: repoDid, - Workflows: workflows, - TrustedSource: trustedSource, + engine.StartWorkflows(log.SubLogger(s.l, "engine"), s.vault, s.cfg, s.db, s.n, s.cache, s.rootCtx, &models.Pipeline{ + RepoDid: repoDid, + Workflows: workflows, + TrustedSource: trustedSource, + TriggerMetadata: tpl.TriggerMetadata, }, pipelineId) return nil }, diff --git a/spindle/storage/disk.go b/spindle/storage/disk.go new file mode 100644 index 00000000..158681c8 --- /dev/null +++ b/spindle/storage/disk.go @@ -0,0 +1,87 @@ +package storage + +import ( + "context" + "fmt" + "io" + "os" + "path/filepath" +) + +type Disk struct { + root string +} + +func NewDisk(root string) (*Disk, error) { + if root == "" { + return nil, fmt.Errorf("storage: disk backend requires a directory") + } + if err := os.MkdirAll(root, 0o755); err != nil { + return nil, fmt.Errorf("storage: create disk root: %w", err) + } + return &Disk{root: filepath.Clean(root)}, nil +} + +func (d *Disk) path(key string) (string, error) { + if err := ValidateKey(key); err != nil { + return "", err + } + return filepath.Join(d.root, filepath.FromSlash(key)), nil +} + +func (d *Disk) Get(_ context.Context, key string) (io.ReadCloser, error) { + p, err := d.path(key) + if err != nil { + return nil, err + } + f, err := os.Open(p) + if err != nil { + if os.IsNotExist(err) { + return nil, ErrNotExist + } + return nil, fmt.Errorf("storage: get %q: %w", key, err) + } + return f, nil +} + +func (d *Disk) Put(_ context.Context, key string, r io.Reader) error { + p, err := d.path(key) + if err != nil { + return err + } + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + return fmt.Errorf("storage: put %q: %w", key, err) + } + // dont let readers see a partial object + tmp, err := os.CreateTemp(filepath.Dir(p), ".tmp-*") + if err != nil { + return fmt.Errorf("storage: put %q: %w", key, err) + } + tmpName := tmp.Name() + defer os.Remove(tmpName) + if _, err := io.Copy(tmp, r); err != nil { + _ = tmp.Close() + return fmt.Errorf("storage: put %q: %w", key, err) + } + if err := tmp.Close(); err != nil { + return fmt.Errorf("storage: put %q: %w", key, err) + } + if err := os.Rename(tmpName, p); err != nil { + return fmt.Errorf("storage: put %q: %w", key, err) + } + return nil +} + +func (d *Disk) Delete(_ context.Context, key string) error { + p, err := d.path(key) + if err != nil { + return err + } + if err := os.Remove(p); err != nil { + if os.IsNotExist(err) { + return nil + } + return fmt.Errorf("storage: delete %q: %w", key, err) + } + return nil +} diff --git a/spindle/storage/s3.go b/spindle/storage/s3.go new file mode 100644 index 00000000..438e201b --- /dev/null +++ b/spindle/storage/s3.go @@ -0,0 +1,102 @@ +package storage + +import ( + "context" + "errors" + "fmt" + "io" + "strings" + + awsconfig "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3/types" +) + +type S3 struct { + bucket string + prefix string + client *s3.Client +} + +func NewS3(ctx context.Context, bucket, prefix string) (*S3, error) { + if bucket == "" { + return nil, fmt.Errorf("storage: s3 backend requires a bucket") + } + sdkConfig, err := awsconfig.LoadDefaultConfig(ctx) + if err != nil { + return nil, fmt.Errorf("storage: load s3 config: %w", err) + } + client := s3.NewFromConfig(sdkConfig) + versioning, err := client.GetBucketVersioning(ctx, &s3.GetBucketVersioningInput{ + Bucket: &bucket, + }) + if err != nil { + return nil, fmt.Errorf("storage: check s3 bucket versioning: %w", err) + } + if versioning.Status != "" { + return nil, fmt.Errorf("storage: s3 cache bucket must not use versioning") + } + return &S3{ + bucket: bucket, + prefix: strings.Trim(prefix, "/"), + client: client, + }, nil +} + +func (s *S3) fullKey(key string) (string, error) { + if err := ValidateKey(key); err != nil { + return "", err + } + if s.prefix == "" { + return key, nil + } + return s.prefix + "/" + key, nil +} + +func (s *S3) Get(ctx context.Context, key string) (io.ReadCloser, error) { + full, err := s.fullKey(key) + if err != nil { + return nil, err + } + res, err := s.client.GetObject(ctx, &s3.GetObjectInput{ + Bucket: &s.bucket, + Key: &full, + }) + if err != nil { + var nsk *types.NoSuchKey + if errors.As(err, &nsk) { + return nil, ErrNotExist + } + return nil, fmt.Errorf("storage: get %q: %w", key, err) + } + return res.Body, nil +} + +func (s *S3) Put(ctx context.Context, key string, r io.Reader) error { + full, err := s.fullKey(key) + if err != nil { + return err + } + if _, err := s.client.PutObject(ctx, &s3.PutObjectInput{ + Bucket: &s.bucket, + Key: &full, + Body: r, + }); err != nil { + return fmt.Errorf("storage: put %q: %w", key, err) + } + return nil +} + +func (s *S3) Delete(ctx context.Context, key string) error { + full, err := s.fullKey(key) + if err != nil { + return err + } + if _, err := s.client.DeleteObject(ctx, &s3.DeleteObjectInput{ + Bucket: &s.bucket, + Key: &full, + }); err != nil { + return fmt.Errorf("storage: delete %q: %w", key, err) + } + return nil +} diff --git a/spindle/storage/storage.go b/spindle/storage/storage.go new file mode 100644 index 00000000..6f013f40 --- /dev/null +++ b/spindle/storage/storage.go @@ -0,0 +1,52 @@ +package storage + +import ( + "context" + "errors" + "fmt" + "io" + "path/filepath" + "regexp" + "strings" + + "tangled.org/core/spindle/config" +) + +var ErrNotExist = errors.New("storage: object does not exist") + +type Storage interface { + Get(ctx context.Context, key string) (io.ReadCloser, error) + Put(ctx context.Context, key string, r io.Reader) error + Delete(ctx context.Context, key string) error +} + +var keyRe = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._:/@%+=-]{0,511}$`) + +func ValidateKey(key string) error { + if !keyRe.MatchString(key) { + return fmt.Errorf("storage: invalid key %q", key) + } + for _, seg := range strings.Split(key, "/") { + if seg == "" || seg == "." || seg == ".." { + return fmt.Errorf("storage: invalid key %q", key) + } + } + return nil +} + +func New(ctx context.Context, cfg *config.Config) (Storage, error) { + switch cfg.Cache.Backend { + case "": + return nil, nil + case "disk": + dir := cfg.Cache.DiskDir + if dir == "" { + dir = filepath.Join(filepath.Dir(cfg.Server.DBPath), "cache") + } + return NewDisk(dir) + case "s3": + return NewS3(ctx, cfg.Cache.S3Bucket, cfg.Cache.S3Prefix) + default: + return nil, fmt.Errorf("storage: unknown backend %q", cfg.Cache.Backend) + } +} diff --git a/spindle/storage/storage_test.go b/spindle/storage/storage_test.go new file mode 100644 index 00000000..14ee5b9f --- /dev/null +++ b/spindle/storage/storage_test.go @@ -0,0 +1,101 @@ +package storage + +import ( + "context" + "errors" + "io" + "strings" + "testing" +) + +func TestValidateKey(t *testing.T) { + valid := []string{ + "did:plc:xyz123/go-mod-v1", + "did:web:spindle.example.com/cache.tar", + "abc", + "a/b/c/d", + } + for _, k := range valid { + if err := ValidateKey(k); err != nil { + t.Errorf("ValidateKey(%q) = %v, want nil", k, err) + } + } + + invalid := []string{ + "", + "../escape", + "a/../../b", + "/leading", + "trailing/", + "double//slash", + "with space", + "with\\backslash", + "-leading-dash", + } + for _, k := range invalid { + if err := ValidateKey(k); err == nil { + t.Errorf("ValidateKey(%q) = nil, want error", k) + } + } +} + +func TestDiskRoundTrip(t *testing.T) { + ctx := context.Background() + d, err := NewDisk(t.TempDir()) + if err != nil { + t.Fatal(err) + } + + key := "did:plc:xyz/go-mod-v1" + if err := d.Put(ctx, key, strings.NewReader("archive-bytes")); err != nil { + t.Fatal(err) + } + + rc, err := d.Get(ctx, key) + if err != nil { + t.Fatal(err) + } + got, err := io.ReadAll(rc) + rc.Close() + if err != nil { + t.Fatal(err) + } + if string(got) != "archive-bytes" { + t.Fatalf("got %q, want %q", got, "archive-bytes") + } + + // overwrite + if err := d.Put(ctx, key, strings.NewReader("new-bytes")); err != nil { + t.Fatal(err) + } + rc, _ = d.Get(ctx, key) + got, _ = io.ReadAll(rc) + rc.Close() + if string(got) != "new-bytes" { + t.Fatalf("overwrite: got %q, want %q", got, "new-bytes") + } + + if err := d.Delete(ctx, key); err != nil { + t.Fatalf("delete: %v", err) + } + if _, err := d.Get(ctx, key); !errors.Is(err, ErrNotExist) { + t.Fatalf("get deleted: got %v, want ErrNotExist", err) + } + if err := d.Delete(ctx, key); err != nil { + t.Fatalf("delete missing: %v", err) + } +} + +func TestDiskRejectsTraversal(t *testing.T) { + ctx := context.Background() + d, err := NewDisk(t.TempDir()) + if err != nil { + t.Fatal(err) + } + if err := d.Put(ctx, "../evil", strings.NewReader("x")); err == nil { + t.Fatal("put with traversal key succeeded") + } + if _, err := d.Get(ctx, "../evil"); err == nil { + t.Fatal("get with traversal key succeeded") + } +} -- 2.51.2