diff --git a/snapshot/pull.go b/snapshot/pull.go index a5626b8..3e39325 100644 --- a/snapshot/pull.go +++ b/snapshot/pull.go @@ -176,9 +176,17 @@ func pickIndexChild(ctx context.Context, m *manifest.OCIManifest) (manifest.Inde return manifest.IndexManifest{}, errors.New("image-index has no usable platform child") } -// writeImportTar assembles the import tar: raw layers (the whole v1 format) -// stream straight through, encoded layers decode via pullencoded.go. func writeImportTar(ctx context.Context, dl Downloader, name, localName string, cfg *manifest.SnapshotConfig, layers []manifest.Descriptor, w io.Writer, progress func(string), prefetch int, budget int64) error { + entries, err := planLayers(cfg, layers) + if err != nil { + return err + } + pipe, err := newChunkPipeline(dl, name, entries, prefetch, budget) + if err != nil { + return err + } + defer pipe.Close() + bw := bufio.NewWriterSize(w, 256<<10) tw := tar.NewWriter(bw) @@ -187,11 +195,37 @@ func writeImportTar(ctx context.Context, dl Downloader, name, localName string, return err } + for _, e := range entries { + if !e.encoded() { + if progress != nil { + progress(fmt.Sprintf(" %s (%d bytes)", e.title, e.layer.Size)) + } + if err := streamLayerToTar(ctx, dl, name, e.layer, e.meta, tw, now); err != nil { + return err + } + continue + } + if progress != nil { + progress(fmt.Sprintf(" %s (%d bytes, %d chunks)", e.title, e.meta.Size, len(e.chunks))) + } + if err := pipe.streamFile(ctx, tw, e, now); err != nil { + return err + } + } + + if err := tw.Close(); err != nil { + return fmt.Errorf("close tar: %w", err) + } + return bw.Flush() +} + +func planLayers(cfg *manifest.SnapshotConfig, layers []manifest.Descriptor) ([]layerPlan, error) { byDigest := make(map[string]manifest.Descriptor, len(layers)) for _, layer := range layers { byDigest[layer.Digest] = layer } + entries := make([]layerPlan, 0, len(layers)) emitted := map[string]bool{} for _, layer := range layers { title := layer.Title() @@ -201,37 +235,25 @@ func writeImportTar(ctx context.Context, dl Downloader, name, localName string, } if len(fileMeta.Chunks) == 0 && !manifest.IsZstdMediaType(layer.MediaType) { - if progress != nil { - progress(fmt.Sprintf(" %s (%d bytes)", title, layer.Size)) - } - if err := streamLayerToTar(ctx, dl, name, layer, fileMeta, tw, now); err != nil { - return err - } + entries = append(entries, layerPlan{title: title, meta: fileMeta, layer: layer}) continue } - if emitted[title] { - continue // continuation chunk of a file already streamed + continue } emitted[title] = true descs, err := resolveEncodedFile(layer, fileMeta, byDigest) if err != nil { - return err + return nil, err } - if progress != nil { - progress(fmt.Sprintf(" %s (%d bytes, %d chunks)", title, fileMeta.Size, len(descs))) - } - // Decode by this file's own layer mediaType: a dedup winner may carry another file's. - compressed := manifest.IsZstdMediaType(layer.MediaType) - if err := streamEncodedFile(ctx, dl, name, title, descs, fileMeta, compressed, tw, now, prefetch, budget); err != nil { - return err - } - } - - if err := tw.Close(); err != nil { - return fmt.Errorf("close tar: %w", err) - } - return bw.Flush() + entries = append(entries, layerPlan{ + title: title, + meta: fileMeta, + layer: layer, + chunks: descs, + }) + } + return entries, nil } func writeSnapshotEnvelope(tw *tar.Writer, cfg *manifest.SnapshotConfig, localName string, now time.Time) error { diff --git a/snapshot/pullencoded.go b/snapshot/pullencoded.go index 1a6aa41..7ae1ebe 100644 --- a/snapshot/pullencoded.go +++ b/snapshot/pullencoded.go @@ -2,7 +2,6 @@ package snapshot import ( "archive/tar" - "bytes" "context" "errors" "fmt" @@ -22,32 +21,92 @@ const ( maxBufferedChunkBytes = 1 << 30 ) -// resolveEncodedFile returns one encoded file's descriptors in Files[].Chunks -// order. Identical chunks dedup across files, so a resolved descriptor may -// carry another file's annotations; digest+size verification is the gate. -func resolveEncodedFile(layer manifest.Descriptor, fileMeta manifest.SnapshotFile, byDigest map[string]manifest.Descriptor) ([]manifest.Descriptor, error) { - title := layer.Title() - if fileMeta.Size <= 0 { - return nil, fmt.Errorf("%s: encoded layer missing files[].size in snapshot config", title) +type layerPlan struct { + title string + meta manifest.SnapshotFile + layer manifest.Descriptor + chunks []manifest.Descriptor +} + +func (e layerPlan) encoded() bool { return e.chunks != nil } + +func (e layerPlan) zstd() bool { return manifest.IsZstdMediaType(e.layer.MediaType) } + +func (e layerPlan) bufferCaps() (input, output int64) { + stored := int64(0) + for _, d := range e.chunks { + stored = max(stored, d.Size) } - if len(fileMeta.Chunks) == 0 { - return []manifest.Descriptor{layer}, nil + if !e.zstd() { + return 0, stored } - descs := make([]manifest.Descriptor, len(fileMeta.Chunks)) - for i, digest := range fileMeta.Chunks { - desc, ok := byDigest[digest] - if !ok { - return nil, fmt.Errorf("%s chunk %d (%s) missing from manifest layers", title, i, digest) + return stored, rawChunkStride(e.meta.Size, len(e.chunks)) +} + +type chunkPipeline struct { + dl Downloader + name string + window int + inputCap int64 + outputCap int64 + in *bufPool + out *bufPool + dec *zstd.Decoder +} + +func newChunkPipeline(dl Downloader, name string, entries []layerPlan, prefetch int, budget int64) (*chunkPipeline, error) { + p := &chunkPipeline{dl: dl, name: name} + var anyZstd bool + for _, e := range entries { + if !e.encoded() { + continue } - descs[i] = desc + inputCap, outputCap := e.bufferCaps() + p.inputCap = max(p.inputCap, inputCap) + p.outputCap = max(p.outputCap, outputCap) + anyZstd = anyZstd || e.zstd() } - return descs, nil + if p.outputCap <= 0 || p.inputCap > maxBufferedChunkBytes || p.outputCap > maxBufferedChunkBytes { + return p, nil + } + unit := p.inputCap + p.outputCap + if budget > p.outputCap { + if window := min(int64(prefetch), (budget-p.outputCap)/unit); window >= 2 { + p.window = int(window) + } + } + if p.window == 0 { + return p, nil + } + p.out = newBufPool(p.window + 1) + if anyZstd { + dec, err := zstd.NewReader(nil, + zstd.WithDecoderConcurrency(p.window), + zstd.WithDecoderMaxMemory(uint64(p.outputCap)), + zstd.WithDecodeAllCapLimit(true)) + if err != nil { + return nil, fmt.Errorf("init zstd decoder: %w", err) + } + p.dec, p.in = dec, newBufPool(p.window) + } + return p, nil +} + +func (p *chunkPipeline) Close() { + if p.dec != nil { + p.dec.Close() + } +} + +func (p *chunkPipeline) fileWindow(e layerPlan) int { + if p.window < 2 || len(e.chunks) < 2 { + return 0 + } + return min(p.window, len(e.chunks)) } -// streamEncodedFile reconstructs one file into a single tar entry; chunks are -// independent zstd frames, so their in-order concatenation is one valid stream. -func streamEncodedFile(ctx context.Context, dl Downloader, name, title string, descs []manifest.Descriptor, fileMeta manifest.SnapshotFile, compressed bool, tw *tar.Writer, modTime time.Time, prefetch int, budget int64) error { - hdr, err := layerHeader(title, fileMeta.Size, fileMeta, modTime) +func (p *chunkPipeline) streamFile(ctx context.Context, tw *tar.Writer, e layerPlan, modTime time.Time) error { + hdr, err := layerHeader(e.title, e.meta.Size, e.meta, modTime) if err != nil { return err } @@ -57,75 +116,92 @@ func streamEncodedFile(ctx context.Context, dl Downloader, name, title string, d ctx, cancel := context.WithCancel(ctx) defer cancel() - var body io.Reader - if window := prefetchWindow(descs, prefetch, budget); window >= 2 { - body = newChunkSource(ctx, dl, name, descs, window) + + var written int64 + if window := p.fileWindow(e); window >= 2 { + written, err = newChunkSource(ctx, p, e, window).WriteTo(tw) } else { - cs := &chunkStream{ctx: ctx, dl: dl, name: name, descs: descs} + cs := &chunkStream{ctx: ctx, dl: p.dl, name: p.name, descs: e.chunks} defer func() { _ = cs.Close() }() - body = cs - } - if compressed { - dec, decErr := zstd.NewReader(body) - if decErr != nil { - return fmt.Errorf("init zstd decoder for %s: %w", title, decErr) + var body io.Reader = cs + if e.zstd() { + dec, decErr := zstd.NewReader(body) + if decErr != nil { + return fmt.Errorf("init zstd decoder for %s: %w", e.title, decErr) + } + defer dec.Close() + body = dec } - defer dec.Close() - body = dec + written, err = io.Copy(tw, body) } - - written, err := io.Copy(tw, body) if err != nil { - return fmt.Errorf("stream %s: %w", title, err) + return fmt.Errorf("stream %s: %w", e.title, err) } - if written != fileMeta.Size { - return fmt.Errorf("%s reconstructed to %d bytes, want %d", title, written, fileMeta.Size) + if written != e.meta.Size { + return fmt.Errorf("%s reconstructed to %d bytes, want %d", e.title, written, e.meta.Size) } return nil } -// prefetchWindow sizes the buffered window by the file's largest chunk (sizes -// are heterogeneous) so window×max ≤ budget holds at every point of the -// stream; 0 sends the caller to the O(1) sequential path. -func prefetchWindow(descs []manifest.Descriptor, prefetch int, budget int64) int { - if len(descs) < 2 { - return 0 - } - var maxSize int64 - for _, d := range descs { - if d.Size > maxBufferedChunkBytes { - return 0 +func (p *chunkPipeline) fetch(ctx context.Context, desc manifest.Descriptor, compressed bool) chunkFetch { + if !compressed { + buf := p.out.take(p.outputCap) + stored, err := p.read(ctx, desc, buf) + if err != nil { + p.out.put(buf) + return chunkFetch{err: err} } - maxSize = max(maxSize, d.Size) + return chunkFetch{data: stored, buf: buf} } - if maxSize <= 0 { - return 0 + comp := p.in.take(p.inputCap) + stored, err := p.read(ctx, desc, comp) + if err != nil { + p.in.put(comp) + return chunkFetch{err: err} } - window := min(int64(prefetch), budget/maxSize, int64(len(descs))) - if window < 2 { - return 0 + dst := p.out.take(p.outputCap) + out, err := p.dec.DecodeAll(stored, dst[:0]) + p.in.put(comp) + if err != nil { + p.out.put(dst) + return chunkFetch{err: fmt.Errorf("decode chunk %s: %w", desc.Digest, err)} } - return int(window) + return chunkFetch{data: out, buf: dst} +} + +func (p *chunkPipeline) read(ctx context.Context, desc manifest.Descriptor, buf []byte) ([]byte, error) { + body, err := p.dl.GetBlob(ctx, p.name, desc.Digest) + if err != nil { + return nil, fmt.Errorf("get blob %s: %w", desc.Digest, err) + } + defer func() { _ = body.Close() }() + stored := buf[:desc.Size] + v := ociutil.NewBlobSizeChecker(body, desc.Digest, desc.Size) + if _, err = io.ReadFull(v, stored); err != nil { + return nil, err + } + if _, err = io.Copy(io.Discard, v); err != nil { + return nil, err + } + return stored, nil } type chunkFetch struct { data []byte + buf []byte err error } -// chunkSource yields verified chunk bodies in order with fetches running -// ahead; futures enter the queue before their fetch spawns, so buffered -// chunks never exceed the window. type chunkSource struct { futures chan chan chunkFetch - cur *bytes.Reader + pipe *chunkPipeline } -func newChunkSource(ctx context.Context, dl Downloader, name string, descs []manifest.Descriptor, window int) *chunkSource { - futures := make(chan chan chunkFetch, max(window-1, 0)) +func newChunkSource(ctx context.Context, p *chunkPipeline, e layerPlan, window int) *chunkSource { + futures := make(chan chan chunkFetch, window-1) go func() { defer close(futures) - for _, desc := range descs { + for _, desc := range e.chunks { fut := make(chan chunkFetch, 1) select { case futures <- fut: @@ -133,57 +209,33 @@ func newChunkSource(ctx context.Context, dl Downloader, name string, descs []man return } go func() { - data, err := fetchChunk(ctx, dl, name, desc) - fut <- chunkFetch{data: data, err: err} + fut <- p.fetch(ctx, desc, e.zstd()) }() } }() - return &chunkSource{futures: futures} + return &chunkSource{futures: futures, pipe: p} } -func (s *chunkSource) Read(p []byte) (int, error) { - for { - if s.cur != nil { - n, err := s.cur.Read(p) - if errors.Is(err, io.EOF) { - s.cur = nil - if n > 0 { - return n, nil - } - continue - } - return n, err - } - fut, ok := <-s.futures - if !ok { - return 0, io.EOF - } +func (s *chunkSource) WriteTo(w io.Writer) (int64, error) { + var written int64 + for fut := range s.futures { res := <-fut if res.err != nil { - return 0, res.err + return written, res.err + } + n, err := w.Write(res.data) + written += int64(n) + s.pipe.out.put(res.buf) + if n != len(res.data) && err == nil { + err = io.ErrShortWrite + } + if err != nil { + return written, err } - s.cur = bytes.NewReader(res.data) - } -} - -func fetchChunk(ctx context.Context, dl Downloader, name string, desc manifest.Descriptor) ([]byte, error) { - if desc.Size < 0 || desc.Size > maxBufferedChunkBytes { - return nil, fmt.Errorf("blob %s size %d outside bufferable range", desc.Digest, desc.Size) - } - body, err := dl.GetBlob(ctx, name, desc.Digest) - if err != nil { - return nil, fmt.Errorf("get blob %s: %w", desc.Digest, err) - } - defer func() { _ = body.Close() }() - buf := bytes.NewBuffer(make([]byte, 0, desc.Size)) - if err := ociutil.CopyBlobSized(buf, body, desc.Digest, desc.Size); err != nil { - return nil, err } - return buf.Bytes(), nil + return written, nil } -// chunkStream reads chunks one at a time straight off the registry stream with -// O(1) memory; BlobVerifier returns io.EOF only after the blob checks pass. type chunkStream struct { ctx context.Context dl Downloader @@ -230,3 +282,29 @@ func (s *chunkStream) Close() error { } return nil } + +func resolveEncodedFile(layer manifest.Descriptor, fileMeta manifest.SnapshotFile, byDigest map[string]manifest.Descriptor) ([]manifest.Descriptor, error) { + title := layer.Title() + if fileMeta.Size <= 0 { + return nil, fmt.Errorf("%s: encoded layer missing files[].size in snapshot config", title) + } + if len(fileMeta.Chunks) == 0 { + return []manifest.Descriptor{layer}, nil + } + descs := make([]manifest.Descriptor, len(fileMeta.Chunks)) + for i, digest := range fileMeta.Chunks { + desc, ok := byDigest[digest] + if !ok { + return nil, fmt.Errorf("%s chunk %d (%s) missing from manifest layers", title, i, digest) + } + descs[i] = desc + } + return descs, nil +} + +func rawChunkStride(size int64, n int) int64 { + if n <= 1 || size <= 0 { + return size + } + return (size-1)/int64(n-1) + 1 +} diff --git a/snapshot/roundtrip_test.go b/snapshot/roundtrip_test.go index 0c69785..2481d78 100644 --- a/snapshot/roundtrip_test.go +++ b/snapshot/roundtrip_test.go @@ -12,9 +12,12 @@ import ( "strings" "sync/atomic" "testing" + "testing/iotest" "testing/synctest" "time" + "github.com/klauspost/compress/zstd" + "github.com/cocoonstack/cocoon-common/manifest" "github.com/cocoonstack/cocoon-common/ociutil" ) @@ -395,40 +398,139 @@ func TestPullTinyBudgetStreamsSequentially(t *testing.T) { } } -// The window must be sized by the file's largest chunk: a compressible prefix -// must never authorize a window the dense middle can blow past the budget with. -func TestPrefetchWindowBoundedByLargestChunk(t *testing.T) { - sized := func(sizes ...int64) []manifest.Descriptor { - descs := make([]manifest.Descriptor, len(sizes)) - for i, s := range sizes { - descs[i] = manifest.Descriptor{Digest: fmt.Sprintf("sha256:%064d", i), Size: s} - } - return descs - } +func TestChunkPipelineWindowBoundedByWidestFile(t *testing.T) { for _, tc := range []struct { name string - descs []manifest.Descriptor + plans []layerPlan prefetch int budget int64 want int }{ - {"tiny prefix, dense tail, budget holds none", sized(2, 2, 100, 100, 100, 100), 8, 10, 0}, - {"prefix sum fits but window x max would not", sized(2, 100, 100), 8, 202, 2}, - {"budget holds two largest", sized(2, 2, 100, 100, 100, 100), 8, 200, 2}, - {"homogeneous small capped by chunk count", sized(2, 2, 2, 2), 8, 1 << 20, 4}, - {"homogeneous small capped by prefetch", sized(2, 2, 2, 2, 2, 2, 2, 2, 2, 2), 8, 1 << 20, 8}, - {"single chunk streams", sized(100), 8, 1 << 20, 0}, - {"oversized chunk streams", sized(2, maxBufferedChunkBytes+1), 8, 1 << 40, 0}, - {"all zero-size streams", sized(0, 0, 0), 8, 1 << 20, 0}, + {"tiny prefix, dense tail, budget holds none", []layerPlan{planOf("a", false, 404, 2, 2, 100, 100, 100, 100)}, 8, 10, 0}, + {"current plus two outputs fit exactly", []layerPlan{planOf("a", false, 202, 2, 100, 100)}, 8, 300, 2}, + {"current output does not fit", []layerPlan{planOf("a", false, 202, 2, 100, 100)}, 8, 299, 0}, + {"narrow file must not widen the window", []layerPlan{ + planOf("small", false, 4, 2, 2), + planOf("dense", false, 400, 100, 100, 100, 100), + }, 8, 300, 2}, + {"widest file decides even when it comes first", []layerPlan{ + planOf("dense", false, 400, 100, 100, 100, 100), + planOf("small", false, 4, 2, 2), + }, 8, 300, 2}, + {"compressed pools fit exactly", []layerPlan{planOf("a", true, 300, 50, 50, 50)}, 8, 550, 2}, + {"compressed pools one short stream", []layerPlan{planOf("a", true, 300, 50, 50, 50)}, 8, 549, 0}, + {"cross-file input and output maxima fit exactly", []layerPlan{ + planOf("wide-output", true, 1800, 100, 100), + planOf("wide-input", true, 100, 900, 900), + }, 8, 7200, 2}, + {"cross-file input and output maxima one short stream", []layerPlan{ + planOf("wide-output", true, 1800, 100, 100), + planOf("wide-input", true, 100, 900, 900), + }, 8, 7199, 0}, + {"prefetch caps the window", []layerPlan{planOf("a", false, 20, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2)}, 8, 1 << 20, 8}, + {"oversized chunk streams", []layerPlan{planOf("a", false, 1, 2, maxBufferedChunkBytes+1)}, 8, 1 << 40, 0}, + {"all zero-size streams", []layerPlan{planOf("a", false, 0, 0, 0, 0)}, 8, 1 << 20, 0}, + {"whole-layer entries need no buffers", []layerPlan{{title: "a", layer: manifest.Descriptor{Size: 100}}}, 8, 1 << 20, 0}, } { - if got := prefetchWindow(tc.descs, tc.prefetch, tc.budget); got != tc.want { - t.Errorf("%s: window = %d, want %d", tc.name, got, tc.want) + p, err := newChunkPipeline(nil, "myvm", tc.plans, tc.prefetch, tc.budget) + if err != nil { + t.Fatalf("%s: %v", tc.name, err) + } + p.Close() + if p.window != tc.want { + t.Errorf("%s: window = %d, want %d", tc.name, p.window, tc.want) + } + } +} + +func TestChunkPipelineFileWindowNarrowsPerFile(t *testing.T) { + plans := []layerPlan{planOf("dense", false, 8, 2, 2, 2, 2)} + p, err := newChunkPipeline(nil, "myvm", plans, 8, 1<<20) + if err != nil { + t.Fatal(err) + } + defer p.Close() + if p.window != 8 { + t.Fatalf("shared window = %d, want 8 (prefetch, budget is ample)", p.window) + } + for _, tc := range []struct { + name string + plan layerPlan + want int + }{ + {"fewer chunks than the window", planOf("a", false, 6, 2, 2, 2), 3}, + {"more chunks than the window", planOf("a", false, 20, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2), 8}, + {"single chunk streams", planOf("a", false, 2, 2), 0}, + } { + if got := p.fileWindow(tc.plan); got != tc.want { + t.Errorf("%s: fileWindow = %d, want %d", tc.name, got, tc.want) + } + } + var off chunkPipeline + if got := off.fileWindow(plans[0]); got != 0 { + t.Errorf("window-less pipeline: fileWindow = %d, want 0", got) + } +} + +func TestRawChunkStrideNeverUnderstatesTheCut(t *testing.T) { + for _, tc := range []struct { + stride int64 + n int + }{{256 << 20, 32}, {256 << 20, 2}, {1 << 20, 7}, {3, 5}} { + size := tc.stride*int64(tc.n-1) + 1 + for _, last := range []int64{1, tc.stride / 2, tc.stride} { + if last == 0 { + continue + } + total := tc.stride*int64(tc.n-1) + last + if got := rawChunkStride(total, tc.n); got < tc.stride { + t.Errorf("stride %d n %d last %d: bound %d understates the cut", tc.stride, tc.n, last, got) + } } + if got := rawChunkStride(size, tc.n); got < tc.stride { + t.Errorf("stride %d n %d: bound %d understates the cut", tc.stride, tc.n, got) + } + } + if got := rawChunkStride(100, 1); got != 100 { + t.Errorf("single chunk: got %d, want the whole file", got) + } +} + +func TestChunkPipelineRejectsDecodedChunkOverCap(t *testing.T) { + raw := bytes.Repeat([]byte("x"), 32) + enc, err := zstd.NewWriter(nil) + if err != nil { + t.Fatal(err) + } + defer func() { _ = enc.Close() }() + stored := enc.EncodeAll(raw, nil) + digest := "sha256:" + ociutil.SHA256Hex(stored) + uploader := newFakeUploader() + uploader.blobs[digest] = stored + desc := manifest.Descriptor{Digest: digest, Size: int64(len(stored))} + entry := layerPlan{ + title: "memory", + meta: manifest.SnapshotFile{Size: 8}, + layer: manifest.Descriptor{MediaType: manifest.ZstdMediaType(manifest.MediaTypeVMMemory)}, + chunks: []manifest.Descriptor{desc, desc}, + } + budget := int64(8 + 2*(len(stored)+8)) + p, err := newChunkPipeline(uploader, "myvm", []layerPlan{entry}, 2, budget) + if err != nil { + t.Fatal(err) + } + defer p.Close() + + res := p.fetch(t.Context(), desc, true) + if res.err == nil { + p.out.put(res.buf) + t.Fatalf("decoded %d bytes into an 8-byte slot", len(res.data)) + } + if !errors.Is(res.err, zstd.ErrDecoderSizeExceeded) && !errors.Is(res.err, zstd.ErrWindowSizeExceeded) { + t.Fatalf("decode error = %v, want a zstd size limit", res.err) } } -// Futures must enter the queue before their fetch spawns: with window 2 and no -// consumer yet, exactly one fetch may start (queue cap 1, second send blocks). func TestChunkSourceSpawnsWithinWindow(t *testing.T) { synctest.Test(t, func(t *testing.T) { g := &gatedDownloader{blobs: map[string][]byte{}, release: make(chan struct{})} @@ -441,7 +543,8 @@ func TestChunkSourceSpawnsWithinWindow(t *testing.T) { descs = append(descs, manifest.Descriptor{Digest: d, Size: 2}) want = append(want, b...) } - src := newChunkSource(t.Context(), g, "myvm", descs, 2) + pipe := &chunkPipeline{dl: g, name: "myvm", window: 2, outputCap: 2, out: newBufPool(3)} + src := newChunkSource(t.Context(), pipe, layerPlan{title: "myvm", chunks: descs}, 2) // After Wait every goroutine is durably blocked, so started is final. synctest.Wait() @@ -451,7 +554,7 @@ func TestChunkSourceSpawnsWithinWindow(t *testing.T) { } close(g.release) var out bytes.Buffer - if _, err := io.Copy(&out, src); err != nil { + if _, err := src.WriteTo(&out); err != nil { t.Fatalf("drain: %v", err) } if !bytes.Equal(out.Bytes(), want) { @@ -460,6 +563,23 @@ func TestChunkSourceSpawnsWithinWindow(t *testing.T) { }) } +func TestChunkSourceRejectsShortWrite(t *testing.T) { + pool := newBufPool(1) + buf := pool.take(4) + copy(buf, "data") + fut := make(chan chunkFetch, 1) + fut <- chunkFetch{data: buf[:4], buf: buf} + futures := make(chan chan chunkFetch, 1) + futures <- fut + close(futures) + + src := &chunkSource{futures: futures, pipe: &chunkPipeline{out: pool}} + written, err := src.WriteTo(shortWriter{}) + if written != 3 || !errors.Is(err, io.ErrShortWrite) { + t.Fatalf("WriteTo = (%d, %v), want (3, %v)", written, err, io.ErrShortWrite) + } +} + func TestPullRejectsNegativeLayerSize(t *testing.T) { pinClock(t) uploader := newFakeUploader() @@ -506,7 +626,55 @@ func TestChunkStreamEnforcesBlobLength(t *testing.T) { } } -// Re-pushing identical content must be a pure HasBlob no-op (chunk-level dedup). +func TestChunkPipelineEnforcesBlobLength(t *testing.T) { + body := []byte("buffered chunk body") + digest := "sha256:" + ociutil.SHA256Hex(body) + uploader := newFakeUploader() + uploader.blobs[digest] = body + + p := &chunkPipeline{dl: uploader, name: "myvm", outputCap: int64(len(body) + 5), out: newBufPool(1)} + fetch := func(size int64) error { + res := p.fetch(t.Context(), manifest.Descriptor{Digest: digest, Size: size}, false) + if res.err == nil { + p.out.put(res.buf) + } + return res.err + } + + if err := fetch(int64(len(body))); err != nil { + t.Fatalf("clean blob: %v", err) + } + if err := fetch(int64(len(body)) + 5); err == nil { + t.Fatal("short blob accepted") + } + if err := fetch(int64(len(body)) - 5); err == nil || !strings.Contains(err.Error(), "longer than") { + t.Fatalf("err = %v, want longer-than-size", err) + } + if err := fetch(int64(len(body))); err != nil { + t.Fatalf("pool starved after the error paths: %v", err) + } +} + +func TestChunkPipelinePropagatesTransportError(t *testing.T) { + body := []byte("verified body") + verifyErr := errors.New("transport digest mismatch") + p := &chunkPipeline{ + dl: terminalErrorDownloader{ + body: body, + err: verifyErr, + }, + name: "myvm", + } + + _, err := p.read(t.Context(), manifest.Descriptor{ + Digest: "sha256:ignored", + Size: int64(len(body)), + }, make([]byte, len(body))) + if !errors.Is(err, verifyErr) { + t.Fatalf("read error = %v, want %v", err, verifyErr) + } +} + func TestV2SecondPushSkipsAllBlobs(t *testing.T) { pinClock(t) corpus := v2Corpus(t) @@ -519,6 +687,18 @@ func TestV2SecondPushSkipsAllBlobs(t *testing.T) { } } +func planOf(title string, zstd bool, fileSize int64, sizes ...int64) layerPlan { + descs := make([]manifest.Descriptor, len(sizes)) + for i, s := range sizes { + descs[i] = manifest.Descriptor{Digest: fmt.Sprintf("sha256:%s%063d", title, i), Size: s} + } + mt := manifest.MediaTypeGeneric + if zstd { + mt = manifest.ZstdMediaType(mt) + } + return layerPlan{title: title, meta: manifest.SnapshotFile{Size: fileSize}, layer: manifest.Descriptor{MediaType: mt}, chunks: descs} +} + // v2Corpus exercises every codec branch: empty, small-raw, exactly-at-chunk-size, // one-byte-over, multi-chunk sparse, and an unknown-name generic layer. func v2Corpus(t *testing.T) []byte { @@ -676,6 +856,25 @@ func readTarEntries(t *testing.T, r io.Reader) []tarEntry { } } +type terminalErrorDownloader struct { + body []byte + err error +} + +func (d terminalErrorDownloader) GetManifest(context.Context, string, string) ([]byte, string, error) { + return nil, "", errors.ErrUnsupported +} + +func (d terminalErrorDownloader) GetBlob(context.Context, string, string) (io.ReadCloser, error) { + return io.NopCloser(io.MultiReader(bytes.NewReader(d.body), iotest.ErrReader(d.err))), nil +} + +type shortWriter struct{} + +func (shortWriter) Write(p []byte) (int, error) { + return max(len(p)-1, 0), nil +} + // gatedDownloader blocks every GetBlob until released and counts starts, so // tests can observe how many fetches the prefetcher launches. type gatedDownloader struct {