diff --git a/pkg/clip/clip.go b/pkg/clip/clip.go index 50b85d3..5250fef 100644 --- a/pkg/clip/clip.go +++ b/pkg/clip/clip.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "strconv" "strings" "time" @@ -211,13 +212,22 @@ func MountArchive(options MountOptions) (func() error, <-chan error, *fuse.Serve } root, _ := clipfs.Root() + maxWrite := 1024 * 1024 + if limit := spliceSafeMaxWrite(); limit > 0 && maxWrite > limit { + // go-fuse sets max_read = MaxWrite and serves fd-backed reads by + // splicing header+payload+one page through a pipe bounded by + // fs.pipe-max-size. If a full read doesn't fit, every read falls back + // to a copy and go-fuse logs "trySplice: splice: want N bytes". + log.Info().Int("max_write", limit).Msg("capping FUSE read/write size to fit fs.pipe-max-size") + maxWrite = limit + } server, err := fuse.NewServer(fs.NewNodeFS(root, immutableFilesystemOptions()), options.MountPoint, &fuse.MountOptions{ MaxBackground: 512, DisableXAttrs: true, EnableSymlinkCaching: true, SyncRead: false, RememberInodes: true, - MaxWrite: 1024 * 1024, + MaxWrite: maxWrite, MaxReadAhead: 1024 * 1024, }) if err != nil { @@ -345,6 +355,7 @@ type CreateFromOCIImageOptions struct { Platform *v1.Platform ContentCache storage.ContentCache // Optional cache to warm with decompressed layer streams ContentCacheDir string // Optional temp directory for layer cache upload spooling + SeedDecompressed bool // Store freshly indexed layers' decompressed bytes in ContentCache LayerIndexCache storage.LayerIndexCache // Optional per-layer index artifact cache (skips pull+index on hit) IndexConcurrency int // Max layers indexed concurrently (default 4) } @@ -380,6 +391,7 @@ func CreateFromOCIImage(ctx context.Context, options CreateFromOCIImageOptions) Platform: options.Platform, ContentCache: options.ContentCache, ContentCacheDir: options.ContentCacheDir, + SeedDecompressed: options.SeedDecompressed, LayerIndexCache: options.LayerIndexCache, IndexConcurrency: options.IndexConcurrency, }, options.OutputPath) @@ -428,3 +440,26 @@ func CreateAndUploadOCIArchive(ctx context.Context, options CreateFromOCIImageOp return nil } + +// spliceSafeMaxWrite returns the largest page-aligned FUSE read/write size for +// which go-fuse's zero-copy splice path (header + payload + one extra page) +// fits in the kernel's maximum pipe size. Returns 0 if the limit is unknown. +func spliceSafeMaxWrite() int { + content, err := os.ReadFile("/proc/sys/fs/pipe-max-size") + if err != nil { + return 0 + } + pipeMax, err := strconv.Atoi(strings.TrimSpace(string(content))) + if err != nil || pipeMax <= 0 { + return 0 + } + return spliceSafeMaxWriteFor(pipeMax, os.Getpagesize()) +} + +func spliceSafeMaxWriteFor(pipeMax, pageSize int) int { + if pageSize <= 0 || pipeMax <= 2*pageSize { + return 0 + } + limit := pipeMax - 2*pageSize + return limit - limit%pageSize +} diff --git a/pkg/clip/clip_test.go b/pkg/clip/clip_test.go index 03fa79d..08afbee 100644 --- a/pkg/clip/clip_test.go +++ b/pkg/clip/clip_test.go @@ -18,3 +18,21 @@ func TestImmutableFilesystemOptionsCacheMetadata(t *testing.T) { } } } + +func TestSpliceSafeMaxWriteFor(t *testing.T) { + const page = 4096 + for _, c := range []struct{ pipeMax, want int }{ + {1 << 20, (1 << 20) - 2*page}, + {4 << 20, (4 << 20) - 2*page}, + {2 * page, 0}, + {0, 0}, + } { + got := spliceSafeMaxWriteFor(c.pipeMax, page) + if got != c.want { + t.Fatalf("spliceSafeMaxWriteFor(%d) = %d, want %d", c.pipeMax, got, c.want) + } + if got > 0 && 16+got+page > c.pipeMax { + t.Fatalf("pipeMax %d: %d leaves no room for header and extra page", c.pipeMax, got) + } + } +} diff --git a/pkg/clip/fsnode.go b/pkg/clip/fsnode.go index b91ef9b..64e2a3c 100644 --- a/pkg/clip/fsnode.go +++ b/pkg/clip/fsnode.go @@ -226,6 +226,9 @@ func (n *FSNode) readData(ctx context.Context, dest []byte, off int64) (fuse.Rea nRead, err = n.filesystem.storage.ReadFile(n.clipNode, dest[:readLen], off) } if err != nil { + // The process sees EIO with no detail; leave the cause where an + // operator can find it. + log.Warn().Err(err).Str("path", n.clipNode.Path).Int64("offset", off).Int64("length", readLen).Msg("read failed, returning EIO to the container") return nil, syscall.EIO } } else { diff --git a/pkg/clip/layer_artifact.go b/pkg/clip/layer_artifact.go index b67d028..fdb1e2d 100644 --- a/pkg/clip/layer_artifact.go +++ b/pkg/clip/layer_artifact.go @@ -112,15 +112,12 @@ func (ca *ClipArchiver) applyLayerArtifact(index *btree.BTree, artifact *LayerAr case LayerEntryOpaqueWhiteout: ca.deleteRange(index, entry.Path+"/") case LayerEntryHardLink: - targetNode := index.Get(&common.ClipNode{Path: entry.Target}) - if targetNode != nil { - tn := targetNode.(*common.ClipNode) - index.Set(&common.ClipNode{ - Path: entry.Path, - NodeType: common.FileNode, - Attr: tn.Attr, - Remote: tn.Remote, - }) + // The same inode under another name: a copy of the target, symlinks included + // (nix's optimised store hard-links symlinks into /nix/store/.links). + if targetNode := index.Get(&common.ClipNode{Path: entry.Target}); targetNode != nil { + linked := *targetNode.(*common.ClipNode) + linked.Path = entry.Path + index.Set(&linked) } } } diff --git a/pkg/clip/layer_artifact_test.go b/pkg/clip/layer_artifact_test.go index c1d08e6..a5b4889 100644 --- a/pkg/clip/layer_artifact_test.go +++ b/pkg/clip/layer_artifact_test.go @@ -127,6 +127,25 @@ func TestLayerArtifactRoundTripDeterminism(t *testing.T) { assert.Equal(t, common.SymLinkNode, nodes["/link"].NodeType) } +func TestLayerArtifactHardLinkToSymlink(t *testing.T) { + archiver := NewClipArchiver() + + layer := buildLayer(t, []tarEntry{ + {name: "dir/", typeflag: tar.TypeDir}, + {name: "dir/a.txt", typeflag: tar.TypeReg, content: "hello"}, + {name: "link", typeflag: tar.TypeSymlink, linkname: "dir/a.txt"}, + {name: "hard-to-link", typeflag: tar.TypeLink, linkname: "link"}, + }) + + index := archiver.newIndex() + archiver.applyLayerArtifact(index, indexLayerHelper(t, archiver, layer, "sha256:layer1")) + + nodes := indexPaths(index) + require.Contains(t, nodes, "/hard-to-link") + assert.Equal(t, common.SymLinkNode, nodes["/hard-to-link"].NodeType, "a hard link to a symlink is a symlink") + assert.Equal(t, nodes["/link"].Target, nodes["/hard-to-link"].Target) +} + func TestLayerArtifactSanitizesUnsetTarTimes(t *testing.T) { archiver := NewClipArchiver() diff --git a/pkg/clip/layer_blob_cache_test.go b/pkg/clip/layer_blob_cache_test.go index a2508f4..3acb25e 100644 --- a/pkg/clip/layer_blob_cache_test.go +++ b/pkg/clip/layer_blob_cache_test.go @@ -3,6 +3,7 @@ package clip import ( "archive/tar" "bytes" + "compress/gzip" "context" "crypto/sha256" "encoding/hex" @@ -190,3 +191,34 @@ func TestCompressedLayerContentCacheReadThrough(t *testing.T) { assert.Equal(t, bytes1, bytes2, "index must be identical regardless of layer source") assert.Equal(t, hashes1, hashes2) } + +func TestSeedDecompressedStoresIndexedLayerBytes(t *testing.T) { + compressed := buildLayer(t, []tarEntry{ + {name: "dir/", typeflag: tar.TypeDir}, + {name: "dir/a.txt", typeflag: tar.TypeReg, content: "hello"}, + }) + sum := sha256.Sum256(compressed) + digest := "sha256:" + hex.EncodeToString(sum[:]) + + for _, seed := range []bool{false, true} { + cache := newFakeBlobContentCache() + artifact, err := NewClipArchiver().indexLayerToArtifact( + context.Background(), + io.NopCloser(bytes.NewReader(compressed)), + digest, + IndexOCIImageOptions{CheckpointMiB: 2, ContentCache: cache, ContentCacheDir: t.TempDir(), SeedDecompressed: seed}, + nil, + ) + require.NoError(t, err) + assert.Equal(t, seed, cache.has(artifact.DecompressedHash), "seed=%v", seed) + if seed { + data, err := cache.GetContent(artifact.DecompressedHash, 0, artifact.UncompressedSize, struct{ RoutingKey string }{}) + require.NoError(t, err) + gz, err := gzip.NewReader(bytes.NewReader(compressed)) + require.NoError(t, err) + want, err := io.ReadAll(gz) + require.NoError(t, err) + assert.Equal(t, want, data) + } + } +} diff --git a/pkg/clip/oci_indexer.go b/pkg/clip/oci_indexer.go index 61b78c8..056f14e 100644 --- a/pkg/clip/oci_indexer.go +++ b/pkg/clip/oci_indexer.go @@ -12,6 +12,7 @@ import ( "path" "runtime" "strings" + "sync" "sync/atomic" "syscall" "time" @@ -71,8 +72,13 @@ type IndexOCIImageOptions struct { Platform *v1.Platform // Target platform (defaults to linux/runtime.GOARCH) ContentCache storage.ContentCache // optional remote cache for fully decompressed layers ContentCacheDir string // optional temp directory for cache upload spooling + SeedDecompressed bool // store each freshly indexed layer's decompressed bytes in ContentCache LayerIndexCache storage.LayerIndexCache // optional cache of per-layer index artifacts (skips pull+index on hit) IndexConcurrency int // max layers indexed concurrently (default 4) + + // seeder runs SeedDecompressed uploads off the indexing goroutines so a + // slow store does not hold an indexing slot; set by IndexOCIImage. + seeder *layerSeeder } const defaultIndexConcurrency = 4 @@ -384,6 +390,14 @@ func (ca *ClipArchiver) IndexOCIImage(ctx context.Context, opts IndexOCIImageOpt g, gctx := errgroup.WithContext(ctx) g.SetLimit(concurrency) + // Content cache seeds run alongside indexing, on the parent context + // (gctx is cancelled once g.Wait returns), and are all awaited before + // the index is returned: the indexing process may exit right after. + if opts.SeedDecompressed && opts.ContentCache != nil { + opts.seeder = newLayerSeeder(ctx, concurrency) + defer opts.seeder.wait() + } + var completedLayers atomic.Int64 for i := range layers { @@ -557,12 +571,20 @@ func (ca *ClipArchiver) indexLayerToArtifact( } defer gzr.Close() - // Streaming hash computation via TeeReader. - // The runtime warms decompressed layers after first access; the build path - // keeps indexing strictly streaming so large layers do not pay extra disk - // writes just to seed an optional warm cache. + // Streaming hash computation via TeeReader. With SeedDecompressed the same + // decompressed stream is spooled to disk once and stored in the content + // cache after indexing, so the first container to use the layer reads it + // page-wise from the cache instead of materializing the whole layer. hasher := sha256.New() hashWriter := io.Writer(hasher) + var cacheSpool *indexedLayerContentCacheSpool + if opts.SeedDecompressed && opts.ContentCache != nil { + cacheSpool = newIndexedLayerContentCacheSpool(opts.ContentCacheDir, layerDigest) + if cacheSpool != nil { + defer cacheSpool.closeAndRemove() + hashWriter = io.MultiWriter(hasher, cacheSpool) + } + } hashingReader := io.TeeReader(gzr, hashWriter) uncompressedCounter := &countingReader{r: hashingReader, onRead: func(total int64) { if onBytes != nil { @@ -645,6 +667,36 @@ func (ca *ClipArchiver) indexLayerToArtifact( // Finalize hash (includes all bytes: file contents + tar headers + padding) decompressedHash := hex.EncodeToString(hasher.Sum(nil)) + // The seed is awaited before indexing returns (the indexing process, a + // build worker, often exits right after); through the seeder it runs + // off this goroutine so the upload overlaps with indexing other layers. + if cacheSpool != nil { + if cacheSpool.err != nil { + log.Warn().Err(cacheSpool.err).Str("layer_digest", layerDigest).Msg("indexed layer not seeded: spool write failed") + } else if path, ok := cacheSpool.detach(); ok { + uncompressedBytes := uncompressedCounter.n + seed := func(ctx context.Context) { + seedStart := time.Now() + err := ca.storeIndexedLayerInContentCache(ctx, opts.ContentCache, path, decompressedHash, layerDigest) + os.Remove(path) + if err != nil { + log.Warn().Err(err).Str("layer_digest", layerDigest).Msg("indexed layer not seeded in content cache") + return + } + log.Info(). + Str("layer_digest", layerDigest). + Int64("bytes", uncompressedBytes). + Dur("duration", time.Since(seedStart)). + Msg("seeded decompressed layer into content cache") + } + if opts.seeder != nil { + opts.seeder.run(seed) + } else { + seed(ctx) + } + } + } + return &LayerArtifact{ Version: LayerArtifactVersion, LayerDigest: layerDigest, @@ -656,6 +708,34 @@ func (ca *ClipArchiver) indexLayerToArtifact( }, nil } +// layerSeeder runs content cache seeds concurrently, up to a limit, and lets +// the indexer wait for all of them. +type layerSeeder struct { + ctx context.Context + wg sync.WaitGroup + sem chan struct{} +} + +func newLayerSeeder(ctx context.Context, limit int) *layerSeeder { + if limit < 1 { + limit = 1 + } + return &layerSeeder{ctx: ctx, sem: make(chan struct{}, limit)} +} + +// run starts fn once a slot is free; it blocks only while every slot is busy. +func (s *layerSeeder) run(fn func(ctx context.Context)) { + s.wg.Add(1) + s.sem <- struct{}{} + go func() { + defer s.wg.Done() + defer func() { <-s.sem }() + fn(s.ctx) + }() +} + +func (s *layerSeeder) wait() { s.wg.Wait() } + func (ca *ClipArchiver) storeIndexedLayerInContentCache(ctx context.Context, contentCache storage.ContentCache, filePath, decompressedHash, layerDigest string) error { return ca.storeLayerBlobInContentCache(ctx, contentCache, filePath, decompressedHash, layerDigest, "indexed layer") } diff --git a/pkg/storage/content_cache_read_ahead.go b/pkg/storage/content_cache_read_ahead.go index 5ccb13b..d8c9a4f 100644 --- a/pkg/storage/content_cache_read_ahead.go +++ b/pkg/storage/content_cache_read_ahead.go @@ -9,8 +9,11 @@ import ( ) const ( - DefaultContentCacheReadAheadBytes = 1024 * 1024 - DefaultContentCacheReadAheadSlots = 32 + // Windows are fetched whole: a bigger window means fewer round trips to + // the cache host for a sequential reader, at the cost of latency and + // waste for a reader that touches one page of it. + DefaultContentCacheReadAheadBytes = 4 * 1024 * 1024 + DefaultContentCacheReadAheadSlots = 24 ) type ContentCacheReadAheadOptions struct { @@ -27,8 +30,33 @@ type ContentCacheReadAhead struct { windows map[contentCacheWindowKey][]byte order []contentCacheWindowKey group singleflight.Group + + inflight sync.Mutex + prefetching map[contentCacheWindowKey]struct{} + prefetchWG sync.WaitGroup + + streamsMu sync.Mutex + streams map[string]*contentCacheStream +} + +// contentCacheStream remembers where the last window read of a layer ended so +// a reader that keeps arriving at the next window is recognised as sequential +// and gets progressively deeper prefetch. +type contentCacheStream struct { + lastEnd int64 + depth int } +const ( + // maxPrefetchInFlight bounds background window fetches per read-ahead cache. + maxPrefetchInFlight = 16 + // maxPrefetchDepth is how many windows ahead a sequential reader is + // fetched. One window ahead bounds throughput at two windows per round + // trip (~270 MiB/s at 4 MiB windows and 30 ms); eight keeps 32 MiB in + // flight, enough to run at the cache host's transfer rate. + maxPrefetchDepth = 8 +) + type contentCacheWindowKey struct { hash string routingKey string @@ -77,11 +105,23 @@ func (r *ContentCacheReadAhead) Read(hash string, offset int64, dest []byte, opt return readContentCacheInto(r.cache, hash, offset, dest, opts) } + // A read that straddles a window boundary is served from each aligned + // window in turn. Windows must stay aligned: an oversized window for the + // straddling read would be a separate fetch of almost the same bytes, and + // its odd end would break the sequential-stream detection below. start := (offset / r.windowBytes) * r.windowBytes - end := start + r.windowBytes - if needEnd := offset + length; end < needEnd { - end = needEnd + if offset+length > start+r.windowBytes { + var done int64 + for done < length { + n, err := r.Read(hash, offset+done, dest[done:min(length, ((offset+done)/r.windowBytes+1)*r.windowBytes-offset)], opts, limit) + if err != nil { + return done, err + } + done += n + } + return done, nil } + end := start + r.windowBytes if end > limit { end = limit } @@ -90,21 +130,82 @@ func (r *ContentCacheReadAhead) Read(hash string, offset int64, dest []byte, opt } key := contentCacheWindowKey{hash: hash, routingKey: opts.RoutingKey, start: start, end: end} - if data, ok := r.get(key); ok { - copy(dest, data[offset-start:offset-start+length]) - return length, nil + // A reader in the second half of a window is likely sequential: fetch the + // next window now so its transfer overlaps with the pages being consumed. + // A reader that has already walked consecutive windows is sequential for + // sure and is kept several windows ahead from the moment it enters a new + // window. + if end < limit { + depth := r.sequentialDepth(key) + if depth > 1 || offset+length > start+r.windowBytes/2 { + for i := 0; i < depth; i++ { + next := end + int64(i)*r.windowBytes + if next >= limit { + break + } + r.prefetch(hash, next, limit, opts) + } + } + } + data, err := r.window(key, opts) + if err != nil { + return 0, err + } + if int64(len(data)) < offset-start+length { + return 0, io.ErrUnexpectedEOF + } + copy(dest, data[offset-start:offset-start+length]) + return length, nil +} + +// sequentialDepth records that a reader is in window key and returns how many +// windows ahead to prefetch: 1 for a new or non-contiguous reader, doubling +// with every consecutive window up to maxPrefetchDepth. +func (r *ContentCacheReadAhead) sequentialDepth(key contentCacheWindowKey) int { + r.streamsMu.Lock() + defer r.streamsMu.Unlock() + if r.streams == nil { + r.streams = map[string]*contentCacheStream{} + } + id := key.hash + "\x00" + key.routingKey + stream := r.streams[id] + if stream == nil { + if len(r.streams) >= 1024 { + r.streams = map[string]*contentCacheStream{} + } + stream = &contentCacheStream{} + r.streams[id] = stream + } + if stream.lastEnd == key.end { + return stream.depth // same window as last time + } + if stream.lastEnd == key.start && stream.depth > 0 { + stream.depth *= 2 + if stream.depth > maxPrefetchDepth { + stream.depth = maxPrefetchDepth + } + } else { + stream.depth = 1 } + stream.lastEnd = key.end + return stream.depth +} +// window returns the bytes of key, fetching them once across concurrent callers. +func (r *ContentCacheReadAhead) window(key contentCacheWindowKey, opts struct{ RoutingKey string }) ([]byte, error) { + if data, ok := r.get(key); ok { + return data, nil + } value, err, _ := r.group.Do(key.String(), func() (any, error) { if data, ok := r.get(key); ok { return data, nil } - size := end - start + size := key.end - key.start if size <= 0 || size > int64(int(size)) { return nil, fmt.Errorf("invalid content cache read-ahead size: %d", size) } data := make([]byte, int(size)) - n, err := readContentCacheInto(r.cache, hash, start, data, opts) + n, err := readContentCacheInto(r.cache, key.hash, key.start, data, opts) if err != nil { return nil, err } @@ -115,14 +216,57 @@ func (r *ContentCacheReadAhead) Read(hash string, offset int64, dest []byte, opt return data, nil }) if err != nil { - return 0, err + return nil, err } data, ok := value.([]byte) - if !ok || int64(len(data)) < offset-start+length { - return 0, io.ErrUnexpectedEOF + if !ok { + return nil, io.ErrUnexpectedEOF + } + return data, nil +} + +// prefetch fetches the window starting at start in the background. The +// singleflight group keeps one fetch in flight per window, and a fetch that +// is already running when the reader arrives is simply joined. +func (r *ContentCacheReadAhead) prefetch(hash string, start, limit int64, opts struct{ RoutingKey string }) { + end := start + r.windowBytes + if end > limit { + end = limit + } + if end <= start { + return + } + key := contentCacheWindowKey{hash: hash, routingKey: opts.RoutingKey, start: start, end: end} + if _, ok := r.get(key); ok { + return + } + r.inflight.Lock() + if r.prefetching == nil { + r.prefetching = map[contentCacheWindowKey]struct{}{} + } + if _, busy := r.prefetching[key]; busy || len(r.prefetching) >= maxPrefetchInFlight { + r.inflight.Unlock() + return + } + r.prefetching[key] = struct{}{} + r.inflight.Unlock() + r.prefetchWG.Add(1) + go func() { + defer func() { + r.inflight.Lock() + delete(r.prefetching, key) + r.inflight.Unlock() + r.prefetchWG.Done() + }() + _, _ = r.window(key, opts) + }() +} + +// WaitPrefetches blocks until no background window fetch is in flight. +func (r *ContentCacheReadAhead) WaitPrefetches() { + if r != nil { + r.prefetchWG.Wait() } - copy(dest, data[offset-start:offset-start+length]) - return length, nil } func (k contentCacheWindowKey) String() string { diff --git a/pkg/storage/content_cache_read_ahead_test.go b/pkg/storage/content_cache_read_ahead_test.go new file mode 100644 index 0000000..a5d1e8c --- /dev/null +++ b/pkg/storage/content_cache_read_ahead_test.go @@ -0,0 +1,215 @@ +package storage + +import ( + "bytes" + "sync" + "testing" + "time" +) + +// countingContentCache records every window fetch offset. +type countingContentCache struct { + mu sync.Mutex + data []byte + offsets []int64 +} + +func (c *countingContentCache) GetContent(hash string, offset int64, length int64, opts struct{ RoutingKey string }) ([]byte, error) { + c.mu.Lock() + c.offsets = append(c.offsets, offset) + c.mu.Unlock() + return c.data[offset : offset+length], nil +} + +func (c *countingContentCache) StoreContent(chunks chan []byte, hash string, opts struct{ RoutingKey string }) (string, error) { + return "", nil +} + +// Small sequential reads (2000 x 100 KiB files laid out back to back in a +// layer) straddle window boundaries constantly. Each window must be fetched +// exactly once and the reader must be recognised as sequential so prefetch +// depth ramps up instead of resetting at every straddling read. +func TestContentCacheReadAheadStraddlingReadsFetchEachWindowOnce(t *testing.T) { + const fileSize = 100 << 10 + size := int64(2000 * fileSize) + cache := &countingContentCache{data: make([]byte, size)} + for i := range cache.data { + cache.data[i] = byte(i % 251) + } + ra := NewContentCacheReadAhead(cache, ContentCacheReadAheadOptions{}) + dest := make([]byte, fileSize) + for i := int64(0); i < 2000; i++ { + off := i * fileSize + n, err := ra.Read("h", off, dest, struct{ RoutingKey string }{}, size) + if err != nil || n != fileSize { + t.Fatalf("read %d: n=%d err=%v", i, n, err) + } + if !bytes.Equal(dest, cache.data[off:off+fileSize]) { + t.Fatalf("read %d returned wrong bytes", i) + } + } + ra.WaitPrefetches() + + counts := map[int64]int{} + for _, o := range cache.offsets { + if o%ra.windowBytes != 0 { + t.Fatalf("unaligned window fetch at offset %d", o) + } + counts[o]++ + } + windows := int((size + ra.windowBytes - 1) / ra.windowBytes) + if len(counts) != windows { + t.Fatalf("fetched %d distinct windows, want %d", len(counts), windows) + } + for off, c := range counts { + if c != 1 { + t.Fatalf("window at %d fetched %d times", off, c) + } + } +} + +func TestWarmPacerOnlyThrottlesWhileForegroundReadsAreRecent(t *testing.T) { + s := &OCIClipStorage{} + pace := s.warmPacer("layer-a") + + // No foreground reads: 64 MiB of chunks pass through with no sleeping. + start := time.Now() + for i := 0; i < 16; i++ { + pace(4 << 20) + } + if el := time.Since(start); el > 50*time.Millisecond { + t.Fatalf("idle warm was paced: %s", el) + } + + // A reader is active: 32 MiB should take about half a second at the + // contended rate. + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + start = time.Now() + for i := 0; i < 8; i++ { + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + pace(4 << 20) + } + el := time.Since(start) + want := time.Duration(float64(32<<20) / float64(warmPacerContendedBytesPerSec) * float64(time.Second)) + if el < want*8/10 || el > want*2 { + t.Fatalf("contended warm ran for %s, want about %s", el, want) + } + + // Reader went quiet: full speed again. + time.Sleep(warmPacerActiveWindow + 20*time.Millisecond) + start = time.Now() + for i := 0; i < 16; i++ { + pace(4 << 20) + } + if el := time.Since(start); el > 50*time.Millisecond { + t.Fatalf("warm stayed paced after reader went idle: %s", el) + } +} + +func TestWarmPacerBudgetIsSharedAcrossConcurrentRestores(t *testing.T) { + s := &OCIClipStorage{} + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + + // Two restores writing at once should together take about as long as one + // writing the same total: 32 MiB across both at the contended rate. + start := time.Now() + var wg sync.WaitGroup + for r := 0; r < 2; r++ { + wg.Add(1) + go func() { + defer wg.Done() + pace := s.warmPacer("layer-a") + for i := 0; i < 4; i++ { + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + pace(4 << 20) + } + }() + } + wg.Wait() + el := time.Since(start) + want := time.Duration(float64(32<<20) / float64(warmPacerContendedBytesPerSec) * float64(time.Second)) + if el < want*8/10 || el > want*2 { + t.Fatalf("two contended restores ran for %s, want about %s (shared budget)", el, want) + } +} + +func TestWarmPacerYieldsToForegroundLayerWaiters(t *testing.T) { + s := &OCIClipStorage{} + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + s.foregroundLayerWaiters.Add(1) + pace := s.warmPacer("layer-a") + start := time.Now() + for i := 0; i < 16; i++ { + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + pace(4 << 20) + } + if el := time.Since(start); el > 50*time.Millisecond { + t.Fatalf("restore was paced while a foreground caller waited on a layer: %s", el) + } +} + +func TestWarmPacerDoesNotThrottleALayerBeingRead(t *testing.T) { + s := &OCIClipStorage{} + now := time.Now().UnixNano() + s.lastForegroundReadNanos.Store(now) + s.lastReadByLayer.Store("hot", now) + + // The layer the container is reading restores at full speed... + pace := s.warmPacer("hot") + start := time.Now() + for i := 0; i < 16; i++ { + s.lastReadByLayer.Store("hot", time.Now().UnixNano()) + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + pace(4 << 20) + } + if el := time.Since(start); el > 50*time.Millisecond { + t.Fatalf("restore of a layer being read was paced: %s", el) + } + + // ...while a layer nobody is reading is paced on the same mount. + cold := s.warmPacer("cold") + start = time.Now() + for i := 0; i < 4; i++ { + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + cold(4 << 20) + } + want := time.Duration(float64(16<<20) / float64(warmPacerContendedBytesPerSec) * float64(time.Second)) + if el := time.Since(start); el < want*8/10 { + t.Fatalf("restore of an unread layer was not paced: %s, want about %s", el, want) + } +} + +func TestWarmPacerExemptsOnlyTheLayerBeingReadNow(t *testing.T) { + s := &OCIClipStorage{} + now := time.Now().UnixNano() + s.lastForegroundReadNanos.Store(now) + // The reader touched "earlier" and then moved on to "current" (an import + // walks many layers within the first second). + s.lastReadByLayer.Store("earlier", now-int64(10*time.Millisecond)) + s.lastReadByLayer.Store("current", now) + + // The layer being read now restores flat out... + current := s.warmPacer("current") + start := time.Now() + for i := 0; i < 16; i++ { + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + s.lastReadByLayer.Store("current", time.Now().UnixNano()) + current(4 << 20) + } + if el := time.Since(start); el > 50*time.Millisecond { + t.Fatalf("restore of the layer being read was paced: %s", el) + } + + // ...but a layer the reader has already left is paced like any other. + earlier := s.warmPacer("earlier") + start = time.Now() + for i := 0; i < 4; i++ { + s.lastForegroundReadNanos.Store(time.Now().UnixNano()) + s.lastReadByLayer.Store("current", time.Now().UnixNano()) + earlier(4 << 20) + } + want := time.Duration(float64(16<<20) / float64(warmPacerContendedBytesPerSec) * float64(time.Second)) + if el := time.Since(start); el < want*8/10 { + t.Fatalf("restore of a layer the reader left was not paced: %s, want about %s", el, want) + } +} diff --git a/pkg/storage/oci.go b/pkg/storage/oci.go index bf0c509..31a9e2d 100644 --- a/pkg/storage/oci.go +++ b/pkg/storage/oci.go @@ -13,6 +13,7 @@ import ( "runtime" "strings" "sync" + "sync/atomic" "time" "github.com/beam-cloud/clip/pkg/common" @@ -28,6 +29,20 @@ import ( // OCIClipStorage implements lazy, range-based reading from OCI registries with disk + remote caching type OCIClipStorage struct { + // lastForegroundReadNanos is when a container last read through this + // mount from the content cache; background layer warms pace themselves + // while it is recent so they do not starve the reads they are meant to + // speed up (see warmPacer). + lastForegroundReadNanos atomic.Int64 + foregroundLayerWaiters atomic.Int32 + // lastReadByLayer maps a layer's decompressed hash to the unix-nanos of + // the last content-cache read a container made from it; a layer being + // read is restored at full speed since that directly shortens the reads. + lastReadByLayer sync.Map + warmPaceMu sync.Mutex + warmPaceSince time.Time + warmPaceBytes int64 + metadata *common.ClipArchiveMetadata storageInfo *common.OCIStorageInfo layerCache map[string]v1.Layer @@ -43,6 +58,7 @@ type OCIClipStorage struct { contentCacheWarmOnce map[string]struct{} layerWarmMu sync.Mutex layerWarmOnce map[string]struct{} + contentCacheServed map[string]int64 // bytes range-read from the content cache, per decompressed hash checkpointLogMu sync.Mutex checkpointSuccessOnce map[string]struct{} checkpointFailureOnce map[string]struct{} @@ -60,6 +76,15 @@ var globalLayerDecompress = newLayerDecompressGroup() const maxBackgroundLayerWarms = 2 +// contentCacheWarmThreshold is how much of a layer a mount has to range-read +// from the content cache before the whole layer is pulled to local disk in +// the background. Range reads go a window at a time and top out well below +// the network; a layer read this much (a big shared library being mapped, a +// model being loaded) is going to be read much more, and one streamed copy +// makes every further read local. Layers touched only lightly are left in +// the cache, so a 5 GiB layer is not copied for one small file. +const contentCacheWarmThreshold = 64 << 20 + var backgroundLayerWarmSlots = make(chan struct{}, maxBackgroundLayerWarms) const maxBackgroundContentCacheWarms = 2 @@ -231,7 +256,9 @@ func (s *OCIClipStorage) Prepare(ctx context.Context, opts PrepareOptions) error opts.Progress(PrepareProgress{Total: len(layers)}) } - group, groupCtx := errgroup.WithContext(ctx) + // Prepare copies whole layers in behind a running container; its writes + // yield to that container's reads like any other background warm. + group, groupCtx := errgroup.WithContext(withBackgroundLayerWarm(ctx)) group.SetLimit(concurrency) for _, layerDigest := range layers { layerDigest := layerDigest @@ -558,6 +585,8 @@ func (s *OCIClipStorage) ReadFileContext(ctx context.Context, node *common.ClipN // Try remote ContentCache range read if s.contentCache != nil && decompressedHash != "" && s.contentCacheAvailable { cacheStart := time.Now() + s.lastForegroundReadNanos.Store(cacheStart.UnixNano()) + s.lastReadByLayer.Store(decompressedHash, cacheStart.UnixNano()) if n, err := s.tryRangeReadFromContentCache(decompressedHash, wantUStart, dest[:readLen], s.contentCacheReadLimit(decompressedHash, remote)); err == nil { metrics.RecordReadHit() metrics.RecordRangeGet(decompressedHash, int64(n)) @@ -591,6 +620,9 @@ func (s *OCIClipStorage) ReadFileContext(ctx context.Context, node *common.ClipN Int64("length", readLen). Int("bytes_read", n). Msg("content cache hit - range read from remote") + if s.noteContentCacheServed(decompressedHash, int64(n)) { + readAttrs["layer_warm"] = s.scheduleLayerDecompressWarm(remote.LayerDigest, "content_cache_read") + } return n, nil } else { metrics.RecordReadMiss() @@ -729,56 +761,69 @@ func (s *OCIClipStorage) ensureLayerCached(ctx context.Context, digest string) ( return decompressedHash, layerPath, nil } + // A foreground caller is about to wait on this layer. Background restores + // on this mount stop pacing while one is waiting, since the restore it + // joins (or competes with) is now on a container's critical path. + if !isBackgroundLayerWarm(ctx) { + s.foregroundLayerWaiters.Add(1) + defer s.foregroundLayerWaiters.Add(-1) + } + waitStart := time.Now() decompressKey := layerDecompressKey(decompressedHash, layerPath) - shared, err := globalLayerDecompress.Do(ctx, decompressKey, func() error { - // Double-check disk cache inside the process-wide singleflight. A - // separate OCIClipStorage instance may have materialized the same layer - // between our fast-path stat and entering this call. - if _, err := os.Stat(layerPath); err == nil { - log.Debug().Str("digest", digest).Str("decompressed_hash", decompressedHash).Msg("disk cache hit (after global lock)") - s.scheduleDecompressedLayerContentCacheWarm(decompressedHash, layerPath) - return nil - } + var shared bool + err := retryLayerMaterialize(ctx, digest, func() error { + var attemptErr error + shared, attemptErr = globalLayerDecompress.Do(ctx, decompressKey, func() error { + // Double-check disk cache inside the process-wide singleflight. A + // separate OCIClipStorage instance may have materialized the same layer + // between our fast-path stat and entering this call. + if _, err := os.Stat(layerPath); err == nil { + log.Debug().Str("digest", digest).Str("decompressed_hash", decompressedHash).Msg("disk cache hit (after global lock)") + s.scheduleDecompressedLayerContentCacheWarm(decompressedHash, layerPath) + return nil + } - fileLock := flock.New(layerPath + ".lock") - locked, err := fileLock.TryLockContext(ctx, 100*time.Millisecond) - if err != nil { - return fmt.Errorf("wait for layer cache lock: %w", err) - } - if !locked { - return fmt.Errorf("failed to acquire layer cache lock: %s", layerPath) - } - defer fileLock.Unlock() + fileLock := flock.New(layerPath + ".lock") + locked, err := fileLock.TryLockContext(ctx, 100*time.Millisecond) + if err != nil { + return fmt.Errorf("wait for layer cache lock: %w", err) + } + if !locked { + return fmt.Errorf("failed to acquire layer cache lock: %s", layerPath) + } + defer fileLock.Unlock() - // Another worker process may have completed the layer while this process - // was waiting on the shared disk lock. - if _, err := os.Stat(layerPath); err == nil { - s.scheduleDecompressedLayerContentCacheWarm(decompressedHash, layerPath) - return nil - } + // Another worker process may have completed the layer while this process + // was waiting on the shared disk lock. + if _, err := os.Stat(layerPath); err == nil { + s.scheduleDecompressedLayerContentCacheWarm(decompressedHash, layerPath) + return nil + } - log.Info(). - Str("layer_digest", digest). - Str("decompressed_hash", decompressedHash). - Msg("oci layer cache miss - materializing layer") + log.Info(). + Str("layer_digest", digest). + Str("decompressed_hash", decompressedHash). + Msg("oci layer cache miss - materializing layer") - decompressStart := time.Now() - source, err := s.decompressAndCacheLayerContext(ctx, digest, layerPath) - if source == "" { - source = "oci_registry" - } - s.observeRead(ctx, common.ReadTraceEvent{ - Operation: "clip.layer_decompress", - Source: source, - LayerDigest: digest, - DecompressedHash: decompressedHash, - StartedAt: decompressStart, - Duration: time.Since(decompressStart), - Success: err == nil, - Error: errorString(err), + decompressStart := time.Now() + source, err := s.decompressAndCacheLayerContext(ctx, digest, layerPath) + if source == "" { + source = "oci_registry" + } + s.observeRead(ctx, common.ReadTraceEvent{ + Operation: "clip.layer_decompress", + Source: source, + LayerDigest: digest, + DecompressedHash: decompressedHash, + StartedAt: decompressStart, + Duration: time.Since(decompressStart), + Success: err == nil, + Error: errorString(err), + }) + return err }) - return err + return attemptErr }) if shared { log.Info().Str("digest", digest).Msg("waited for in-progress layer decompression") @@ -806,6 +851,39 @@ func (s *OCIClipStorage) ensureLayerCached(ctx context.Context, digest string) ( return decompressedHash, layerPath, nil } +// Materializing a layer streams it whole from the registry or the content +// cache, so one reset connection or throttled response fails the stream, and +// with it every process whose read is waiting on that layer in the +// singleflight: each of them gets EIO at the same moment, which in a fresh +// container is an ImportError. Those failures are transient, so the read +// retries a few times before giving up. Waiters that shared a failed attempt +// retry too, and one of them leads the next attempt. +const ( + layerMaterializeAttempts = 4 + layerMaterializeRetryBackoff = 500 * time.Millisecond +) + +func retryLayerMaterialize(ctx context.Context, digest string, attempt func() error) error { + var err error + for i := 1; i <= layerMaterializeAttempts; i++ { + err = attempt() + if err == nil || ctx.Err() != nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return err + } + if i == layerMaterializeAttempts { + break + } + backoff := layerMaterializeRetryBackoff * time.Duration(1<<(i-1)) + log.Warn().Err(err).Str("layer_digest", digest).Int("attempt", i).Dur("retry_in", backoff).Msg("layer materialization failed, retrying") + select { + case <-time.After(backoff): + case <-ctx.Done(): + return ctx.Err() + } + } + return err +} + // getDecompressedCachePath returns the cache path for a decompressed hash func (s *OCIClipStorage) getDecompressedCachePath(decompressedHash string) string { return filepath.Join(s.diskCacheDir, decompressedHash) @@ -973,6 +1051,20 @@ func (s *OCIClipStorage) scheduleLayerDecompressWarm(layerDigest string, reason return "scheduled" } +// noteContentCacheServed adds n to the bytes of the layer served from the +// content cache and reports whether the total just crossed the warm +// threshold. +func (s *OCIClipStorage) noteContentCacheServed(decompressedHash string, n int64) bool { + s.layerWarmMu.Lock() + defer s.layerWarmMu.Unlock() + if s.contentCacheServed == nil { + s.contentCacheServed = make(map[string]int64) + } + before := s.contentCacheServed[decompressedHash] + s.contentCacheServed[decompressedHash] = before + n + return before < contentCacheWarmThreshold && before+n >= contentCacheWarmThreshold +} + func (s *OCIClipStorage) markLayerWarmAttempt(decompressedHash string) bool { s.layerWarmMu.Lock() defer s.layerWarmMu.Unlock() @@ -997,7 +1089,7 @@ func (s *OCIClipStorage) runLayerDecompressWarm(layerDigest string, decompressed backgroundLayerWarmSlots <- struct{}{} defer func() { <-backgroundLayerWarmSlots }() - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Minute) + ctx, cancel := context.WithTimeout(withBackgroundLayerWarm(context.Background()), 30*time.Minute) defer cancel() startedAt := time.Now() @@ -1075,7 +1167,11 @@ func (s *OCIClipStorage) decompressAndCacheLayerContext(ctx context.Context, dig if s.contentCacheAvailable && s.contentCache != nil { if cacheStream, ok := s.contentCache.(ContentCacheStream); ok { - written, err := restoreLayerFromContentCache(ctx, cacheStream, decompressedHash, diskPath) + var pace func(int) + if isBackgroundLayerWarm(ctx) { + pace = s.warmPacer(decompressedHash) + } + written, err := restoreLayerFromContentCache(ctx, cacheStream, decompressedHash, diskPath, pace) if err == nil { log.Info(). Str("layer", digest). @@ -1331,7 +1427,92 @@ func (r *contextReader) Read(p []byte) (int, error) { return n, err } -func restoreLayerFromContentCache(ctx context.Context, cacheStream ContentCacheStream, decompressedHash, diskPath string) (int64, error) { +// Background layer warms write a whole layer to local disk while the +// container that triggered them is still reading through the mount. On a +// host whose disk is the bottleneck the warm's writes starve those reads, +// which is backwards: the warm exists to make later reads fast. While the +// mount has been read in the last warmPacerActiveWindow the warm is held to +// warmPacerContendedBytesPerSec; once the reader goes quiet it runs flat out. +const ( + warmPacerActiveWindow = 250 * time.Millisecond + warmPacerContendedBytesPerSec = 64 << 20 +) + +// layerIsCurrentRead reports whether decompressedHash is the layer the mount's +// reader is in right now: the most recently read layer, read within the +// active window. Only that one layer's restore runs unpaced. Exempting every +// layer the reader had touched let a torch import on a fresh worker (14 +// layers, 5.4 GiB, most of them touched within the first second) run every +// restore flat out, and the import's own page reads then queued behind +// 5 GiB of background copy (first import 5–17 s against 1.5 s once local). +func (s *OCIClipStorage) layerIsCurrentRead(decompressedHash string) bool { + if decompressedHash == "" { + return false + } + var newest int64 + current := "" + s.lastReadByLayer.Range(func(k, v any) bool { + if ts := v.(int64); ts > newest { + newest, current = ts, k.(string) + } + return true + }) + if current != decompressedHash { + return false + } + return time.Since(time.Unix(0, newest)) <= warmPacerActiveWindow +} + +type backgroundLayerWarmKey struct{} + +// withBackgroundLayerWarm marks ctx as belonging to a background warm, whose +// restore is paced; a foreground materialization (a reader waiting on the +// whole layer) never is. +func withBackgroundLayerWarm(ctx context.Context) context.Context { + return context.WithValue(ctx, backgroundLayerWarmKey{}, true) +} + +func isBackgroundLayerWarm(ctx context.Context) bool { + v, _ := ctx.Value(backgroundLayerWarmKey{}).(bool) + return v +} + +// warmPacer returns a function the restore of layer decompressedHash calls +// after writing each chunk; it sleeps as needed to hold the contended write +// rate. The budget is shared by every paced restore on this mount (Prepare +// runs several layers at once), so the cap is on the mount's total background +// write rate. The one layer the container is reading right now is never +// paced: once it is local those reads stop going to the network, so finishing +// it fast is the point, and the bandwidth comes out of the other layers. +func (s *OCIClipStorage) warmPacer(decompressedHash string) func(int) { + if s == nil { + return nil + } + return func(n int) { + last := s.lastForegroundReadNanos.Load() + s.warmPaceMu.Lock() + if last == 0 || time.Since(time.Unix(0, last)) > warmPacerActiveWindow || s.foregroundLayerWaiters.Load() > 0 || s.layerIsCurrentRead(decompressedHash) { + s.warmPaceSince = time.Time{} + s.warmPaceBytes = 0 + s.warmPaceMu.Unlock() + return + } + now := time.Now() + if s.warmPaceSince.IsZero() { + s.warmPaceSince = now + s.warmPaceBytes = 0 + } + s.warmPaceBytes += int64(n) + allowed := time.Duration(float64(s.warmPaceBytes) / float64(warmPacerContendedBytesPerSec) * float64(time.Second)) + sleep := allowed - now.Sub(s.warmPaceSince) + s.warmPaceMu.Unlock() + if sleep > 0 { + time.Sleep(sleep) + } + } +} + +func restoreLayerFromContentCache(ctx context.Context, cacheStream ContentCacheStream, decompressedHash, diskPath string, pace func(int)) (int64, error) { chunks, expectedSize, err := cacheStream.GetContentStream(decompressedHash, struct{ RoutingKey string }{RoutingKey: decompressedHash}) if err != nil { return 0, err @@ -1372,6 +1553,9 @@ func restoreLayerFromContentCache(ctx context.Context, cacheStream ContentCacheS go drainContentChunks(chunks) return written, io.ErrShortWrite } + if pace != nil { + pace(n) + } } } }) diff --git a/pkg/storage/oci_parallel_pull.go b/pkg/storage/oci_parallel_pull.go index 319f52a..25acde6 100644 --- a/pkg/storage/oci_parallel_pull.go +++ b/pkg/storage/oci_parallel_pull.go @@ -6,11 +6,16 @@ import ( "fmt" "io" "math" + "math/rand" + "net" "net/http" + "net/http/httptrace" "net/url" "os" + "sort" "strconv" "strings" + "sync" "sync/atomic" "syscall" "time" @@ -23,15 +28,39 @@ import ( ) const ( - parallelBlobPullThreshold int64 = 1 << 30 // 1 GiB - parallelBlobPullConcurrency = 16 - parallelBlobPullPartSize int64 = 256 << 20 // 256 MiB - parallelBlobPullDiskReserve int64 = 1 << 30 // 1 GiB + // Every layer a container waits on is worth fanning out: a single registry + // stream from a remote worker is bounded by one TCP flow to one blob-store + // front-end, and some front-ends are an order of magnitude slower than + // others from the same site. + parallelBlobPullThreshold int64 = 32 << 20 // 32 MiB + parallelBlobPullConcurrency = 8 // per layer + parallelBlobPullPartSize int64 = 32 << 20 // 32 MiB + parallelBlobPullDiskReserve int64 = 1 << 30 // 1 GiB parallelBlobPullAttempts = 3 + + // Ranges across all layers materializing on this worker share one + // connection budget so eight concurrent layers do not open 64 flows. + parallelBlobPullGlobalConcurrency = 24 + + // A range whose connection is starved is abandoned and its remaining bytes + // re-fetched on a new connection (to a different front-end where possible) + // instead of holding the whole layer hostage. See stragglerRule. + parallelBlobStragglerGrace = 3 * time.Second + parallelBlobStragglerFloorPerSec = 4 << 20 // 4 MiB/s + parallelBlobStragglerMedianDivisor = 6 + parallelBlobStragglerTailDivisor = 4 + parallelBlobStragglerMaxAbortsPerPart = 3 + parallelBlobStragglerRateHistory = 32 + parallelBlobSlowEndpointTTL = 2 * time.Minute ) var errParallelBlobRangeUnsupported = errors.New("registry blob range request unsupported") +// errParallelBlobStraggler marks an attempt the straggler monitor cut short. +var errParallelBlobStraggler = errors.New("registry blob range attempt abandoned as straggler") + +var parallelBlobRangeSlots = make(chan struct{}, parallelBlobPullGlobalConcurrency) + type parallelBlobPullConfig struct { inner http.RoundTripper registryHost string @@ -45,6 +74,16 @@ type parallelBlobPullConfig struct { attempts int retryBackoff time.Duration availableBytes func(string) (int64, error) + + // stragglerGrace <= 0 disables the straggler monitor (tests). + stragglerGrace time.Duration + stragglerFloor int64 + // slots bounds concurrent range requests across transports; nil means the + // package-wide budget. + slots chan struct{} + // markSlowEndpoint records the remote address of an abandoned attempt so + // later dials avoid it. nil disables. + markSlowEndpoint func(string) } // parallelBlobTransport turns one large authenticated registry blob GET into @@ -57,6 +96,71 @@ type parallelBlobTransport struct { parallelBlobPullConfig } +// rangePart is one in-flight range attempt, observed by the straggler monitor. +type rangePart struct { + start, end int64 + startedAt time.Time + // bytes is the progress of the current attempt; done is the total already + // written to the file across attempts, so a resumed attempt only asks for + // the remainder. + bytes atomic.Int64 + done int64 + remote atomic.Pointer[string] + cancel context.CancelFunc + aborts atomic.Int32 +} + +func (p *rangePart) rate(now time.Time) (float64, time.Duration) { + elapsed := now.Sub(p.startedAt) + if elapsed <= 0 { + return 0, 0 + } + return float64(p.bytes.Load()) / elapsed.Seconds(), elapsed +} + +// stragglerRule picks the in-flight parts to abandon. A part is a straggler +// when, past the grace period, it is below the absolute floor and far below +// the median rate of its progressing peers and recently completed parts. The +// relative test keeps a uniformly slow link from churning every connection; +// the floor keeps a lone tail part from crawling when there is nothing to +// compare against, and a part with no bytes at all past grace is always +// abandoned. In the tail (no parts left to hand out) idle workers have nothing +// better to do, so any part well below the reference rate is abandoned even +// above the floor; the remainder is resumed, so an abort costs a reconnect. +func stragglerRule(parts []*rangePart, completed []float64, tail bool, now time.Time, grace time.Duration, floor int64, maxAborts int) []*rangePart { + rates := make([]float64, 0, len(parts)+len(completed)) + for _, part := range parts { + if rate, elapsed := part.rate(now); elapsed >= time.Second && rate > 0 { + rates = append(rates, rate) + } + } + rates = append(rates, completed...) + sort.Float64s(rates) + var reference float64 + if len(rates) > 0 { + reference = rates[len(rates)/2] + } + if tail { + grace /= 2 + } + var stragglers []*rangePart + for _, part := range parts { + rate, elapsed := part.rate(now) + if elapsed < grace || int(part.aborts.Load()) >= maxAborts { + continue + } + switch { + case rate == 0: + case tail && len(rates) >= 2 && rate < reference/parallelBlobStragglerTailDivisor: + case rate < float64(floor) && (len(rates) < 2 || rate < reference/parallelBlobStragglerMedianDivisor): + default: + continue + } + stragglers = append(stragglers, part) + } + return stragglers +} + func (s *OCIClipStorage) parallelBlobTransport(digest string) http.RoundTripper { size := s.compressedLayerSize(digest) if size < parallelBlobPullThreshold { @@ -70,19 +174,129 @@ func (s *OCIClipStorage) parallelBlobTransport(digest string) http.RoundTripper } return newParallelBlobTransport(parallelBlobPullConfig{ - inner: remote.DefaultTransport, - registryHost: normalizedRegistryHost(s.storageInfo.RegistryURL), - digest: digest, - size: size, - tempDir: s.diskCacheDir, - minimumFree: saturatingAdd(size, uncompressedSize, parallelBlobPullDiskReserve), - threshold: parallelBlobPullThreshold, - partSize: parallelBlobPullPartSize, - concurrency: parallelBlobPullConcurrency, - attempts: parallelBlobPullAttempts, - retryBackoff: 100 * time.Millisecond, - availableBytes: filesystemAvailableBytes, + inner: blobRangeTransport(), + registryHost: normalizedRegistryHost(s.storageInfo.RegistryURL), + digest: digest, + size: size, + tempDir: s.diskCacheDir, + minimumFree: saturatingAdd(size, uncompressedSize, parallelBlobPullDiskReserve), + threshold: parallelBlobPullThreshold, + partSize: parallelBlobPullPartSize, + concurrency: parallelBlobPullConcurrency, + attempts: parallelBlobPullAttempts, + retryBackoff: 100 * time.Millisecond, + availableBytes: filesystemAvailableBytes, + stragglerGrace: parallelBlobStragglerGrace, + stragglerFloor: parallelBlobStragglerFloorPerSec, + slots: parallelBlobRangeSlots, + markSlowEndpoint: slowBlobEndpoints.mark, + }) +} + +// slowEndpointSet remembers blob-store front-ends that starved a range so new +// dials prefer the others. Entries expire; a front-end is not slow forever. +type slowEndpointSet struct { + mu sync.Mutex + ttl time.Duration + seen map[string]time.Time + closeIdle func() +} + +var slowBlobEndpoints = &slowEndpointSet{ttl: parallelBlobSlowEndpointTTL, seen: map[string]time.Time{}} + +func (s *slowEndpointSet) mark(remoteAddr string) { + host, _, err := net.SplitHostPort(remoteAddr) + if err != nil { + host = remoteAddr + } + if host == "" { + return + } + s.mu.Lock() + s.seen[host] = time.Now().Add(s.ttl) + s.mu.Unlock() + if s.closeIdle != nil { + // Keep-alive would otherwise hand the next range the same starved + // connection. Re-dialing every pooled connection costs a handshake. + s.closeIdle() + } +} + +func (s *slowEndpointSet) isSlow(ip string, now time.Time) bool { + s.mu.Lock() + defer s.mu.Unlock() + until, ok := s.seen[ip] + if !ok { + return false + } + if now.After(until) { + delete(s.seen, ip) + return false + } + return true +} + +// dial resolves addr and connects to a randomly chosen address that is not +// currently marked slow, so concurrent ranges spread over the blob store's +// front-ends instead of piling onto whichever one the resolver listed first. +func (s *slowEndpointSet) dial(ctx context.Context, network, addr string) (net.Conn, error) { + dialer := &net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second} + host, port, err := net.SplitHostPort(addr) + if err != nil || net.ParseIP(host) != nil { + return dialer.DialContext(ctx, network, addr) + } + ips, err := net.DefaultResolver.LookupIPAddr(ctx, host) + if err != nil || len(ips) == 0 { + return dialer.DialContext(ctx, network, addr) + } + now := time.Now() + candidates := make([]net.IPAddr, 0, len(ips)) + for _, ip := range ips { + if !s.isSlow(ip.String(), now) { + candidates = append(candidates, ip) + } + } + if len(candidates) == 0 { + candidates = ips + } + rand.Shuffle(len(candidates), func(i, j int) { candidates[i], candidates[j] = candidates[j], candidates[i] }) + var lastErr error + for _, ip := range candidates { + conn, err := dialer.DialContext(ctx, network, net.JoinHostPort(ip.String(), port)) + if err == nil { + return conn, nil + } + lastErr = err + if ctx.Err() != nil { + break + } + } + return nil, lastErr +} + +var ( + blobRangeTransportOnce sync.Once + blobRangeTransportInst http.RoundTripper +) + +// blobRangeTransport is the shared HTTP transport for range subrequests. It +// mirrors go-containerregistry's defaults but dials through slowBlobEndpoints +// and keeps enough idle connections for the range fan-out to reuse. +func blobRangeTransport() http.RoundTripper { + blobRangeTransportOnce.Do(func() { + base, ok := remote.DefaultTransport.(*http.Transport) + if !ok { + blobRangeTransportInst = remote.DefaultTransport + return + } + transport := base.Clone() + transport.DialContext = slowBlobEndpoints.dial + transport.MaxIdleConnsPerHost = parallelBlobPullGlobalConcurrency * 2 + transport.MaxIdleConns = parallelBlobPullGlobalConcurrency * 4 + slowBlobEndpoints.closeIdle = transport.CloseIdleConnections + blobRangeTransportInst = transport }) + return blobRangeTransportInst } func (s *OCIClipStorage) fetchLayerByDigestWithTransport(ctx context.Context, digest string, transport http.RoundTripper) (v1.Layer, error) { @@ -128,7 +342,7 @@ func (t *parallelBlobTransport) RoundTrip(req *http.Request) (*http.Response, er available, err := t.availableBytes(t.tempDir) if err != nil || (t.minimumFree > 0 && available < t.minimumFree) { - log.Debug(). + log.Warn(). Err(err). Int64("available_bytes", available). Int64("required_bytes", t.minimumFree). @@ -191,6 +405,7 @@ func (t *parallelBlobTransport) prefetch(req *http.Request) (*http.Response, err return nil, err } + pullStart := time.Now() group, groupCtx := errgroup.WithContext(req.Context()) var nextOffset atomic.Int64 nextOffset.Store(1) @@ -203,6 +418,58 @@ func (t *parallelBlobTransport) prefetch(req *http.Request) (*http.Response, err if int64(workers) > partCount { workers = int(partCount) } + + var ( + activeMu sync.Mutex + active = map[*rangePart]struct{}{} + completed []float64 + aborts atomic.Int64 + retries atomic.Int64 + ) + track := func(part *rangePart, on bool) { + activeMu.Lock() + if on { + active[part] = struct{}{} + } else { + delete(active, part) + if rate, elapsed := part.rate(time.Now()); elapsed > 0 && part.done == part.end-part.start+1 { + completed = append(completed, rate) + if len(completed) > parallelBlobStragglerRateHistory { + completed = completed[1:] + } + } + } + activeMu.Unlock() + } + monitorDone := make(chan struct{}) + if t.stragglerGrace > 0 { + go func() { + ticker := time.NewTicker(t.stragglerGrace / 3) + defer ticker.Stop() + for { + select { + case <-monitorDone: + return + case <-groupCtx.Done(): + return + case now := <-ticker.C: + activeMu.Lock() + parts := make([]*rangePart, 0, len(active)) + for part := range active { + parts = append(parts, part) + } + tail := nextOffset.Load() >= t.size + stragglers := stragglerRule(parts, completed, tail, now, t.stragglerGrace, t.stragglerFloor, parallelBlobStragglerMaxAbortsPerPart) + for _, part := range stragglers { + part.aborts.Add(1) + part.cancel() + } + activeMu.Unlock() + } + } + }() + } + for i := 0; i < workers; i++ { group.Go(func() error { for { @@ -214,18 +481,32 @@ func (t *parallelBlobTransport) prefetch(req *http.Request) (*http.Response, err if end < start || end >= t.size { end = t.size - 1 } - if err := t.downloadRange(groupCtx, req, tempFile, start, end); err != nil { + part := &rangePart{start: start, end: end} + if err := t.downloadPart(groupCtx, req, tempFile, part, track, &aborts, &retries); err != nil { return err } } }) } - if err := group.Wait(); err != nil { + err = group.Wait() + close(monitorDone) + if err != nil { return nil, err } if err := req.Context().Err(); err != nil { return nil, err } + elapsed := time.Since(pullStart) + log.Info(). + Str("layer_digest", t.digest). + Int64("compressed_bytes", t.size). + Int64("parts", partCount). + Int("concurrency", workers). + Int64("straggler_aborts", aborts.Load()). + Int64("retries", retries.Load()). + Dur("duration", elapsed). + Float64("mib_per_s", float64(t.size)/(1<<20)/elapsed.Seconds()). + Msg("parallel registry blob prefetch complete") if _, err := tempFile.Seek(0, io.SeekStart); err != nil { return nil, fmt.Errorf("rewind compressed layer prefetch file: %w", err) } @@ -247,10 +528,19 @@ func (t *parallelBlobTransport) prefetch(req *http.Request) (*http.Response, err }, nil } +// downloadRange fetches one range with the normal retry policy and without +// straggler supervision (used for the probe). func (t *parallelBlobTransport) downloadRange(ctx context.Context, original *http.Request, dest *os.File, start, end int64) error { - want := end - start + 1 + return t.downloadPart(ctx, original, dest, &rangePart{start: start, end: end}, nil, nil, nil) +} + +// downloadPart fetches one range, retrying transient failures with backoff and +// immediately re-issuing attempts the straggler monitor abandons. Straggler +// aborts do not consume the transient-failure budget; they are bounded per part +// by parallelBlobStragglerMaxAbortsPerPart inside stragglerRule. +func (t *parallelBlobTransport) downloadPart(ctx context.Context, original *http.Request, dest *os.File, part *rangePart, track func(*rangePart, bool), aborts, retries *atomic.Int64) error { var lastErr error - for attempt := 0; attempt < t.attempts; attempt++ { + for attempt := 0; attempt < t.attempts; { if err := ctx.Err(); err != nil { return err } @@ -265,55 +555,150 @@ func (t *parallelBlobTransport) downloadRange(ctx context.Context, original *htt } } - req := original.Clone(ctx) - req.Header = original.Header.Clone() - req.Header.Set("Range", fmt.Sprintf("bytes=%d-%d", start, end)) - resp, err := (&http.Client{Transport: t.inner}).Do(req) - if err != nil { - lastErr = err - continue + if t.slots != nil { + select { + case t.slots <- struct{}{}: + case <-ctx.Done(): + return ctx.Err() + } } - - if resp.StatusCode != http.StatusPartialContent { - _ = resp.Body.Close() - if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusRequestedRangeNotSatisfiable { - return fmt.Errorf("%w: status %s", errParallelBlobRangeUnsupported, resp.Status) + partCtx, cancel := context.WithCancel(ctx) + part.startedAt = time.Now() + part.bytes.Store(0) + part.remote.Store(nil) + part.cancel = cancel + if track != nil { + track(part, true) + } + err := t.downloadRangeAttempt(partCtx, original, dest, part) + if track != nil { + track(part, false) + } + // Only the straggler monitor cancels partCtx while ctx is alive; read + // that before our own cancel() below makes partCtx.Err() non-nil. + straggled := partCtx.Err() != nil && ctx.Err() == nil + cancel() + if t.slots != nil { + <-t.slots + } + if err == nil { + return nil + } + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } + if straggled { + if aborts != nil { + aborts.Add(1) + } + remote := "" + if addr := part.remote.Load(); addr != nil { + remote = *addr + } + if t.markSlowEndpoint != nil && remote != "" { + t.markSlowEndpoint(remote) } - lastErr = fmt.Errorf("range %d-%d returned status %s", start, end, resp.Status) + rate, elapsed := part.rate(time.Now()) + log.Debug(). + Str("layer_digest", t.digest). + Int64("range_start", part.start). + Int64("range_end", part.end). + Int64("resume_at", part.start+part.done). + Str("remote", remote). + Float64("mib_per_s", rate/(1<<20)). + Dur("elapsed", elapsed). + Int32("aborts", part.aborts.Load()). + Msg("registry blob range straggler abandoned; resuming on a new connection") + lastErr = errParallelBlobStraggler continue } - - gotStart, gotEnd, gotTotal, err := parseContentRange(resp.Header.Get("Content-Range")) - if err != nil || gotStart != start || gotEnd != end || gotTotal != t.size { - _ = resp.Body.Close() - return fmt.Errorf("%w: requested %d-%d/%d, got %q", errParallelBlobRangeUnsupported, start, end, t.size, resp.Header.Get("Content-Range")) + if errors.Is(err, errParallelBlobRangeUnsupported) { + return err } - if resp.ContentLength >= 0 && resp.ContentLength != want { - _ = resp.Body.Close() - return fmt.Errorf("%w: range %d-%d content length %d, expected %d", errParallelBlobRangeUnsupported, start, end, resp.ContentLength, want) + lastErr = err + attempt++ + if retries != nil && attempt < t.attempts { + retries.Add(1) } + } + return fmt.Errorf("download range %d-%d after %d attempts: %w", part.start, part.end, t.attempts, lastErr) +} - written, copyErr := io.CopyN(io.NewOffsetWriter(dest, start), resp.Body, want) - if copyErr == nil { - var extra [1]byte - n, readErr := resp.Body.Read(extra[:]) - if n != 0 || (readErr != nil && !errors.Is(readErr, io.EOF)) { - copyErr = fmt.Errorf("range response exceeded declared length") +// downloadRangeAttempt performs a single validated range request into dest. +func (t *parallelBlobTransport) downloadRangeAttempt(ctx context.Context, original *http.Request, dest *os.File, part *rangePart) error { + start, end := part.start+part.done, part.end + want := end - start + 1 + if want <= 0 { + return nil + } + + ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{ + GotConn: func(info httptrace.GotConnInfo) { + if info.Conn != nil && info.Conn.RemoteAddr() != nil { + addr := info.Conn.RemoteAddr().String() + part.remote.Store(&addr) } + }, + }) + req := original.Clone(ctx) + req.Header = original.Header.Clone() + req.Header.Set("Range", fmt.Sprintf("bytes=%d-%d", start, end)) + resp, err := (&http.Client{Transport: t.inner}).Do(req) + if err != nil { + return err + } + + if resp.StatusCode != http.StatusPartialContent { + _ = resp.Body.Close() + if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusRequestedRangeNotSatisfiable { + return fmt.Errorf("%w: status %s", errParallelBlobRangeUnsupported, resp.Status) } - closeErr := resp.Body.Close() - if copyErr == nil && closeErr == nil && written == want { - return nil - } - if copyErr != nil { - lastErr = copyErr - } else if closeErr != nil { - lastErr = closeErr - } else { - lastErr = fmt.Errorf("short range write: wrote %d, expected %d", written, want) + return fmt.Errorf("range %d-%d returned status %s", start, end, resp.Status) + } + + gotStart, gotEnd, gotTotal, err := parseContentRange(resp.Header.Get("Content-Range")) + if err != nil || gotStart != start || gotEnd != end || gotTotal != t.size { + _ = resp.Body.Close() + return fmt.Errorf("%w: requested %d-%d/%d, got %q", errParallelBlobRangeUnsupported, start, end, t.size, resp.Header.Get("Content-Range")) + } + if resp.ContentLength >= 0 && resp.ContentLength != want { + _ = resp.Body.Close() + return fmt.Errorf("%w: range %d-%d content length %d, expected %d", errParallelBlobRangeUnsupported, start, end, resp.ContentLength, want) + } + + body := &rangeCountingReader{r: resp.Body, n: &part.bytes} + written, copyErr := io.CopyN(io.NewOffsetWriter(dest, start), body, want) + part.done += written + if copyErr == nil { + var extra [1]byte + n, readErr := resp.Body.Read(extra[:]) + if n != 0 || (readErr != nil && !errors.Is(readErr, io.EOF)) { + copyErr = fmt.Errorf("range response exceeded declared length") } } - return fmt.Errorf("download range %d-%d after %d attempts: %w", start, end, t.attempts, lastErr) + closeErr := resp.Body.Close() + if copyErr == nil && closeErr == nil && written == want { + return nil + } + if copyErr != nil { + return copyErr + } + if closeErr != nil { + return closeErr + } + return fmt.Errorf("short range write: wrote %d, expected %d", written, want) +} + +// rangeCountingReader adds bytes read to a shared counter (progress for stragglerRule). +type rangeCountingReader struct { + r io.Reader + n *atomic.Int64 +} + +func (c *rangeCountingReader) Read(p []byte) (int, error) { + n, err := c.r.Read(p) + c.n.Add(int64(n)) + return n, err } func parseContentRange(value string) (start, end, total int64, err error) { diff --git a/pkg/storage/oci_parallel_pull_test.go b/pkg/storage/oci_parallel_pull_test.go index e2fb1a6..a029b89 100644 --- a/pkg/storage/oci_parallel_pull_test.go +++ b/pkg/storage/oci_parallel_pull_test.go @@ -9,6 +9,7 @@ import ( "fmt" "io" "math" + "net" "net/http" "net/http/httptest" "net/url" @@ -488,3 +489,179 @@ func TestRemoveOnCloseFileIsIdempotentForRemovedPath(t *testing.T) { require.NoError(t, os.Remove(path)) require.NoError(t, (&removeOnCloseFile{File: file, path: path}).Close()) } + +func TestStragglerRule(t *testing.T) { + now := time.Now() + mk := func(age time.Duration, bytes int64, aborts int32) *rangePart { + p := &rangePart{startedAt: now.Add(-age)} + p.bytes.Store(bytes) + p.aborts.Store(aborts) + return p + } + const floor = 4 << 20 + grace := 3 * time.Second + rule := func(parts []*rangePart, completed []float64, tail bool) []*rangePart { + return stragglerRule(parts, completed, tail, now, grace, floor, 3) + } + + fast := mk(4*time.Second, 400<<20, 0) // 100 MiB/s + fast2 := mk(4*time.Second, 320<<20, 0) // 80 MiB/s + stalled := mk(4*time.Second, 0, 0) + slow := mk(4*time.Second, 4<<20, 0) // 1 MiB/s: below floor and far below median + young := mk(time.Second, 0, 0) // inside grace + exhausted := mk(4*time.Second, 0, 3) // hit the abort cap + + got := rule([]*rangePart{fast, fast2, stalled, slow, young, exhausted}, nil, false) + require.ElementsMatch(t, []*rangePart{stalled, slow}, got) + + // Uniformly slow link: everything is below the floor but nothing is far + // below the median, so nothing is abandoned. + uniform := []*rangePart{mk(4*time.Second, 8<<20, 0), mk(4*time.Second, 9<<20, 0), mk(4*time.Second, 12<<20, 0)} + require.Empty(t, rule(uniform, nil, false)) + + // ...unless recently completed parts show the link can do far better. + require.ElementsMatch(t, uniform, rule(uniform, []float64{60 << 20, 70 << 20, 80 << 20}, false)) + + // A lone part has no peers; the floor alone decides. + require.Equal(t, []*rangePart{slow}, rule([]*rangePart{slow}, nil, false)) + require.Empty(t, rule([]*rangePart{fast}, nil, false)) + + // Tail: a part above the floor but well below the reference is abandoned, + // and the grace period is halved. + tailSlow := mk(2*time.Second, 20<<20, 0) // 10 MiB/s + require.Empty(t, rule([]*rangePart{tailSlow}, []float64{60 << 20, 70 << 20}, false)) + require.Equal(t, []*rangePart{tailSlow}, rule([]*rangePart{tailSlow}, []float64{60 << 20, 70 << 20}, true)) + require.Empty(t, rule([]*rangePart{mk(2*time.Second, 60<<20, 0)}, []float64{60 << 20, 70 << 20}, true)) + // Tail with no reference falls back to the floor. + require.Empty(t, rule([]*rangePart{tailSlow}, nil, true)) + require.Equal(t, []*rangePart{slow}, rule([]*rangePart{slow}, nil, true)) +} + +func TestSlowEndpointSetExpires(t *testing.T) { + set := &slowEndpointSet{ttl: time.Minute, seen: map[string]time.Time{}} + set.mark("52.216.1.2:443") + set.mark("[2600::1]:443") + now := time.Now() + require.True(t, set.isSlow("52.216.1.2", now)) + require.True(t, set.isSlow("2600::1", now)) + require.False(t, set.isSlow("52.216.1.3", now)) + require.False(t, set.isSlow("52.216.1.2", now.Add(2*time.Minute))) + require.False(t, set.isSlow("52.216.1.2", now), "expired entries are dropped") +} + +// stallingBlobServer serves ranges normally except that the first attempt at +// each range whose start is in stallStarts sends headers and a few bytes, then +// hangs until the client abandons it. +type stallingBlobServer struct { + data []byte + stallStarts map[int64]bool + mu sync.Mutex + stalled map[int64]int + resumeFrom map[int64]bool + resumed atomic.Int64 + requests atomic.Int64 +} + +func (s *stallingBlobServer) serve(w http.ResponseWriter, r *http.Request) { + start, end, err := parseTestRange(r.Header.Get("Range"), int64(len(s.data))) + if err != nil { + http.Error(w, err.Error(), http.StatusRequestedRangeNotSatisfiable) + return + } + s.requests.Add(1) + s.mu.Lock() + if s.resumeFrom[start] { + s.resumed.Add(1) + } + s.mu.Unlock() + w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(s.data))) + w.Header().Set("Content-Length", strconv.FormatInt(end-start+1, 10)) + w.WriteHeader(http.StatusPartialContent) + s.mu.Lock() + stall := s.stallStarts[start] && s.stalled[start] == 0 + if stall { + s.stalled[start]++ + } + s.mu.Unlock() + if stall { + // Partial progress, then starvation: the resume must pick up at +16. + s.mu.Lock() + s.resumeFrom[start+16] = true + s.mu.Unlock() + _, _ = w.Write(s.data[start : start+16]) + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + <-r.Context().Done() + return + } + _, _ = w.Write(s.data[start : end+1]) +} + +func TestParallelBlobTransportAbandonsStragglersAndRefetches(t *testing.T) { + const partSize = 64 << 10 + data := make([]byte, 16*partSize+1) + for i := range data { + data[i] = byte((i*17 + 3) % 253) + } + // Parts start at 1 + k*partSize (byte 0 is the probe). + blobState := &stallingBlobServer{ + data: data, + stallStarts: map[int64]bool{1 + 2*partSize: true, 1 + 9*partSize: true}, + stalled: map[int64]int{}, + resumeFrom: map[int64]bool{}, + } + blob := httptest.NewServer(http.HandlerFunc(blobState.serve)) + defer blob.Close() + registryState := ®istryRedirectServer{data: data, redirectURL: blob.URL + "/signed"} + registry := httptest.NewServer(http.HandlerFunc(registryState.serve)) + defer registry.Close() + parsed, err := url.Parse(registry.URL) + require.NoError(t, err) + + var marked sync.Map + tempDir := t.TempDir() + client := &http.Client{Transport: newParallelBlobTransport(parallelBlobPullConfig{ + inner: registry.Client().Transport, + registryHost: parsed.Host, + digest: parallelPullTestDigest, + size: int64(len(data)), + tempDir: tempDir, + minimumFree: 1, + threshold: 1, + partSize: partSize, + concurrency: 4, + attempts: 3, + retryBackoff: time.Millisecond, + availableBytes: func(string) (int64, error) { return math.MaxInt64, nil }, + stragglerGrace: 150 * time.Millisecond, + stragglerFloor: 1 << 20, + slots: make(chan struct{}, 3), + markSlowEndpoint: func(addr string) { marked.Store(addr, true) }, + })} + + started := time.Now() + resp, err := parallelTestRequest(t, client, registry.URL, context.Background()) + require.NoError(t, err) + actual, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + require.Equal(t, data, actual, "abandoned ranges must be resumed byte-exact") + require.Less(t, time.Since(started), 5*time.Second) + require.Zero(t, registryState.fullRequests.Load(), "stragglers must not trigger the single-stream fallback") + require.Equal(t, 1, blobState.stalled[1+2*partSize]) + require.Equal(t, 1, blobState.stalled[1+9*partSize]) + // 1 probe + 16 parts + 2 resumes, each resume starting after the 16 bytes + // the starved attempt delivered. + require.Equal(t, int64(19), blobState.requests.Load()) + require.Equal(t, int64(2), blobState.resumed.Load()) + blobHost, _, _ := net.SplitHostPort(strings.TrimPrefix(blob.URL, "http://")) + var markedHosts []string + marked.Range(func(k, _ any) bool { + h, _, _ := net.SplitHostPort(k.(string)) + markedHosts = append(markedHosts, h) + return true + }) + require.Equal(t, []string{blobHost}, markedHosts, "the starving front-end is marked slow") + requireNoParallelTemps(t, tempDir) +} diff --git a/pkg/storage/oci_test.go b/pkg/storage/oci_test.go index cdb5b25..eaeb205 100644 --- a/pkg/storage/oci_test.go +++ b/pkg/storage/oci_test.go @@ -13,6 +13,7 @@ import ( "path/filepath" "strings" "sync" + "sync/atomic" "testing" "time" @@ -806,11 +807,78 @@ func TestOCIStorage_ContentCacheReadAheadCoalescesAdjacentReads(t *testing.T) { require.Equal(t, 4, n) require.Equal(t, []byte("klmn"), third) + // The read at 10 reached the second half of window 0-16 and prefetched + // 16-32, which the read at 20 then joined instead of fetching again. The + // read at 20 entered the window right after the previous one, so the + // reader is treated as sequential and the window after it is prefetched + // too. Every window is fetched exactly once. + storage.contentCacheReadAhead.WaitPrefetches() cache.mu.Lock() defer cache.mu.Unlock() - require.Equal(t, 2, cache.getCalls) - require.Equal(t, []int64{0, 16}, cache.getOffsets) - require.Equal(t, []int64{16, 16}, cache.getLengths) + require.Equal(t, 3, cache.getCalls) + require.ElementsMatch(t, []int64{0, 16, 32}, cache.getOffsets) + require.ElementsMatch(t, []int64{16, 16, int64(len(testData)) - 32}, cache.getLengths) +} + +func TestContentCacheReadAheadPrefetchesNextWindowForSequentialReaders(t *testing.T) { + data := make([]byte, 256) + for i := range data { + data[i] = byte(i) + } + cache := newMockCache() + cache.store["h"] = data + ra := NewContentCacheReadAhead(cache, ContentCacheReadAheadOptions{WindowBytes: 16, MaxWindows: 16}) + opts := struct{ RoutingKey string }{RoutingKey: "h"} + offsets := func() []int64 { + ra.WaitPrefetches() + cache.mu.Lock() + defer cache.mu.Unlock() + return append([]int64(nil), cache.getOffsets...) + } + + // First half of the first window: nothing speculative. + buf := make([]byte, 4) + _, err := ra.Read("h", 0, buf, opts, 256) + require.NoError(t, err) + require.Equal(t, []int64{0}, offsets()) + + // Second half: the next window is fetched in the background, once. + _, err = ra.Read("h", 12, buf, opts, 256) + require.NoError(t, err) + _, err = ra.Read("h", 8, buf, opts, 256) + require.NoError(t, err) + require.ElementsMatch(t, []int64{0, 16}, offsets()) + + // Entering the window right after the previous one marks the reader as + // sequential: the prefetched window is served without another fetch and + // two windows are fetched ahead of it at once. + _, err = ra.Read("h", 16, buf, opts, 256) + require.NoError(t, err) + require.Equal(t, data[16:20], buf) + require.ElementsMatch(t, []int64{0, 16, 32, 48}, offsets()) + + // Each further consecutive window doubles the depth (4, then 8) ... + _, err = ra.Read("h", 32, buf, opts, 256) + require.NoError(t, err) + require.ElementsMatch(t, []int64{0, 16, 32, 48, 64, 80, 96}, offsets()) + _, err = ra.Read("h", 48, buf, opts, 256) + require.NoError(t, err) + require.Len(t, offsets(), 12) // 48+8*16 = 176 -> windows through 160-176 are in flight or cached + + // ... and a jump elsewhere in the layer resets the reader to one window + // ahead, fetched only once it is past the midpoint of its window. + _, err = ra.Read("h", 224, buf, opts, 256) + require.NoError(t, err) + require.Len(t, offsets(), 13) + _, err = ra.Read("h", 236, buf, opts, 256) + require.NoError(t, err) + require.Contains(t, offsets(), int64(240)) + require.Len(t, offsets(), 14) + + // Nothing is ever fetched past the limit. + for _, off := range offsets() { + require.Less(t, off, int64(256)) + } } func TestOCIStorage_CacheMiss(t *testing.T) { @@ -3449,3 +3517,75 @@ func TestCheckpointEmptyList(t *testing.T) { assert.Equal(t, int64(0), cOff, "should return 0 for empty checkpoint list") assert.Equal(t, int64(0), uOff, "should return 0 for empty checkpoint list") } + +func TestNoteContentCacheServedCrossesThresholdOnce(t *testing.T) { + s := &OCIClipStorage{} + half := int64(contentCacheWarmThreshold / 2) + require.False(t, s.noteContentCacheServed("a", half), "below threshold") + require.True(t, s.noteContentCacheServed("a", half), "crossing threshold schedules the warm") + require.False(t, s.noteContentCacheServed("a", half), "only once per layer") + require.False(t, s.noteContentCacheServed("b", half), "layers are counted separately") +} + +// A transient failure of the layer stream must not fail the read: the caller +// retries, and a waiter that shared the failed attempt retries as well. +func TestRetryLayerMaterializeRetriesTransientFailuresForOwnerAndWaiters(t *testing.T) { + group := newLayerDecompressGroup() + var attempts atomic.Int32 + ownerStarted := make(chan struct{}) + waiterJoined := make(chan struct{}) + + materialize := func(ctx context.Context) error { + return retryLayerMaterialize(ctx, "sha256:layer", func() error { + _, err := group.Do(ctx, "layer", func() error { + n := attempts.Add(1) + if n == 1 { + close(ownerStarted) + <-waiterJoined + return errors.New("failed to decompress layer to disk: connection reset by peer") + } + return nil + }) + return err + }) + } + + ownerDone := make(chan error, 1) + go func() { ownerDone <- materialize(context.Background()) }() + <-ownerStarted + + waiterDone := make(chan error, 1) + go func() { + // Join the in-flight (failing) attempt, then release it. + go func() { + time.Sleep(20 * time.Millisecond) + close(waiterJoined) + }() + waiterDone <- materialize(context.Background()) + }() + + require.NoError(t, <-ownerDone) + require.NoError(t, <-waiterDone) + require.GreaterOrEqual(t, attempts.Load(), int32(2), "the failed stream must be attempted again") + require.LessOrEqual(t, attempts.Load(), int32(3), "retries share one attempt through the singleflight") +} + +func TestRetryLayerMaterializeGivesUpAfterTheBudgetAndNotOnCancellation(t *testing.T) { + var attempts int + err := retryLayerMaterialize(context.Background(), "sha256:layer", func() error { + attempts++ + return errors.New("layer not found") + }) + require.EqualError(t, err, "layer not found") + require.Equal(t, layerMaterializeAttempts, attempts) + + ctx, cancel := context.WithCancel(context.Background()) + attempts = 0 + err = retryLayerMaterialize(ctx, "sha256:layer", func() error { + attempts++ + cancel() + return context.Canceled + }) + require.ErrorIs(t, err, context.Canceled) + require.Equal(t, 1, attempts, "an abandoned container's read is not retried") +}