multipart fetch stability fixes

Signed-off-by: Adrien Delorme <azr@users.noreply.github.com>
This commit is contained in:
Adrien Delorme
2026-01-23 09:45:04 +01:00
parent 10cf3c8bc6
commit e86523ecdb
2 changed files with 133 additions and 19 deletions

View File

@@ -34,6 +34,7 @@ import (
"github.com/klauspost/compress/zstd"
digest "github.com/opencontainers/go-digest"
ocispec "github.com/opencontainers/image-spec/specs-go/v1"
"golang.org/x/sync/errgroup"
"github.com/containerd/containerd/v2/core/images"
"github.com/containerd/containerd/v2/core/remotes"
@@ -60,6 +61,7 @@ func (p *bufferPool) Get() *bytes.Buffer {
}
func (p *bufferPool) Put(buffer *bytes.Buffer) {
buffer.Reset()
p.pool.Put(buffer)
}
@@ -510,9 +512,11 @@ func (r dockerFetcher) open(ctx context.Context, req *request, mediatype string,
if numChunks < parallelism {
parallelism = numChunks
}
// Prepare channels, buffer pool, and readers/writers for parallel fetching.
queue := make(chan int64, parallelism)
ctx, cancelCtx := context.WithCancel(ctx)
done := ctx.Done()
ctx, cancel := context.WithCancel(ctx)
eg, ctx := errgroup.WithContext(ctx)
readers, writers := make([]io.Reader, numChunks), make([]*pipeWriter, numChunks)
bufPool := newbufferPool(chunkSize)
for i := range numChunks {
@@ -520,21 +524,23 @@ func (r dockerFetcher) open(ctx context.Context, req *request, mediatype string,
}
// keep reference of the initial body value to ensure it is closed
ibody := body
go func() {
eg.Go(func() error {
defer close(queue)
for i := range numChunks {
select {
case queue <- i:
case <-done:
case <-ctx.Done():
if i == 0 {
ibody.Close()
}
return // avoid leaking a goroutine if we exit early.
return ctx.Err()
}
}
close(queue)
}()
return nil
})
for range parallelism {
go func() {
eg.Go(func() error {
for i := range queue { // first in first out
copy := func() error {
var body io.ReadCloser
@@ -542,6 +548,7 @@ func (r dockerFetcher) open(ctx context.Context, req *request, mediatype string,
body = ibody
} else {
if err := r.Acquire(ctx, 1); err != nil {
_ = writers[i].CloseWithError(err)
return err
}
defer r.Release(1)
@@ -550,12 +557,6 @@ func (r dockerFetcher) open(ctx context.Context, req *request, mediatype string,
nresp, err := reqClone.doWithRetries(ctx, lastHost, withErrorCheck)
if err != nil {
_ = writers[i].CloseWithError(err)
select {
case <-done:
return ctx.Err()
default:
cancelCtx()
}
return err
}
body = nresp.Body
@@ -564,20 +565,23 @@ func (r dockerFetcher) open(ctx context.Context, req *request, mediatype string,
_ = body.Close()
_ = writers[i].CloseWithError(err)
if err != nil && err != io.EOF {
cancelCtx()
return err
}
return nil
}
if copy() != nil {
return
if err := copy(); err != nil {
return err
}
}
}()
return nil
})
}
body = &fnOnClose{
BeforeClose: func() {
cancelCtx()
cancel()
if err := eg.Wait(); err != nil {
log.G(ctx).WithError(err).Warn("parallel fetch failed")
}
},
ReadCloser: io.NopCloser(io.MultiReader(readers...)),
}

View File

@@ -32,16 +32,23 @@ import (
"net/url"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/klauspost/compress/zstd"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/sync/semaphore"
"github.com/containerd/containerd/v2/core/transfer"
)
type writeFunc func(p []byte) (int, error)
func (f writeFunc) Write(p []byte) (int, error) { return f(p) }
func TestFetcherOpen(t *testing.T) {
content := make([]byte, 128)
rand.New(rand.NewSource(1)).Read(content)
@@ -286,6 +293,109 @@ func TestFetcherOpenParallel(t *testing.T) {
assert.Error(t, err, "this should have failed")
}
func TestFetcherOpenParallel_CloseAfterCopyError(t *testing.T) {
size := int64(10 * 1024 * 1024)
content := make([]byte, size)
rr, err := rand.New(rand.NewSource(1)).Read(content)
require.NoError(t, err)
require.Equal(t, int(size), rr)
type errWriter struct {
max int64
n int64
}
ew := &errWriter{max: 1024}
ewWrite := func(p []byte) (int, error) {
n := len(p)
ew.n += int64(n)
if ew.n >= ew.max {
return 0, errors.New("simulated write failure after limit reached")
}
return n, nil
}
// simulate Close should not wait for download to complete after write error
unblockOnce := sync.Once{}
unblock := make(chan struct{})
unblockAll := func() { unblockOnce.Do(func() { close(unblock) }) }
defer unblockAll()
s := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
rng, err := parseRange(r.Header.Get("Range"), size)
if errors.Is(err, errNoOverlap) {
err = nil
}
assert.NoError(t, err)
if len(rng) == 0 {
rw.Header().Set("content-length", strconv.Itoa(len(content)))
_, _ = rw.Write(content)
return
}
if rng[0].start > 0 {
select {
case <-r.Context().Done():
return
case <-unblock:
}
}
b := content[rng[0].start : rng[0].start+rng[0].length]
rw.Header().Set("content-range", rng[0].contentRange(size))
rw.Header().Set("content-length", strconv.Itoa(len(b)))
_, err = rw.Write(b)
t.Logf("wrote range %s, err=%v", rng[0].contentRange(size), err)
}))
defer s.Close()
u, err := url.Parse(s.URL)
if err != nil {
t.Fatal(err)
}
f := dockerFetcher{
&dockerBase{
repository: "nonempty",
limiter: semaphore.NewWeighted(4),
performances: transfer.ImageResolverPerformanceSettings{
MaxConcurrentDownloads: 4,
ConcurrentLayerFetchBuffer: 1 * 1024 * 1024,
},
},
}
host := RegistryHost{
Client: s.Client(),
Host: u.Host,
Scheme: u.Scheme,
Path: u.Path,
}
req := f.request(host, http.MethodGet)
rc, _, err := f.open(context.Background(), req, "", 0, true)
require.NoError(t, err, "failed to open reader")
_, copyErr := io.Copy(writeFunc(ewWrite), rc)
require.NotNil(t, copyErr, "expected write error during copy")
closeDone := make(chan error, 1)
go func() {
closeDone <- rc.Close()
}()
timer := time.NewTimer(10 * time.Second)
defer timer.Stop()
select {
case err := <-closeDone:
if err != nil {
t.Errorf("close error: %v", err)
}
case <-timer.C:
t.Errorf("close blocked after write error")
unblockAll()
<-closeDone
}
}
func TestContentEncoding(t *testing.T) {
t.Parallel()