mirror of
https://github.com/containerd/containerd.git
synced 2026-08-10 01:48:39 +00:00
multipart fetch stability fixes
Signed-off-by: Adrien Delorme <azr@users.noreply.github.com>
This commit is contained in:
@@ -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...)),
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user