From 276fee072758d9fce08f30e64e499b3f95cd6630 Mon Sep 17 00:00:00 2001 From: Tonis Tiigi Date: Mon, 10 Jun 2019 13:23:31 -0700 Subject: [PATCH] worker: fix gc race on fromremote Signed-off-by: Tonis Tiigi --- worker/base/worker.go | 44 +++++++++++++++++++++++++++++++------------ 1 file changed, 32 insertions(+), 12 deletions(-) diff --git a/worker/base/worker.go b/worker/base/worker.go index c4dc7519e..c12180e56 100644 --- a/worker/base/worker.go +++ b/worker/base/worker.go @@ -355,14 +355,19 @@ func (w *Worker) FromRemote(ctx context.Context, remote *solver.Remote) (cache.I return nil, err } - cs, release := snapshot.NewContainerdSnapshotter(w.Snapshotter) + cd, release := snapshot.NewContainerdSnapshotter(w.Snapshotter) defer release() unpackProgressDone := oneOffProgress(ctx, "unpacking") - chainIDs, err := w.unpack(ctx, remote.Descriptors, cs) + chainIDs, refs, err := w.unpack(ctx, w.CacheManager, remote.Descriptors, cd) if err != nil { return nil, unpackProgressDone(err) } + defer func() { + for _, ref := range refs { + ref.Release(context.TODO()) + } + }() unpackProgressDone(nil) for i, chainID := range chainIDs { @@ -388,31 +393,46 @@ func (w *Worker) FromRemote(ctx context.Context, remote *solver.Remote) (cache.I return nil, errors.Errorf("unreachable") } -func (w *Worker) unpack(ctx context.Context, descs []ocispec.Descriptor, s cdsnapshot.Snapshotter) ([]string, error) { +func (w *Worker) unpack(ctx context.Context, cm cache.Manager, descs []ocispec.Descriptor, s cdsnapshot.Snapshotter) (ids []string, refs []cache.ImmutableRef, err error) { + defer func() { + if err != nil { + for _, r := range refs { + r.Release(context.TODO()) + } + } + }() + layers, err := getLayers(ctx, descs) if err != nil { - return nil, err + return nil, nil, err } var chain []digest.Digest for _, layer := range layers { - if _, err := rootfs.ApplyLayer(ctx, layer, chain, s, w.Applier); err != nil { - return nil, err - } - chain = append(chain, layer.Diff.Digest) + newChain := append(chain, layer.Diff.Digest) + + chainID := ociidentity.ChainID(newChain) + ref, err := cm.Get(ctx, string(chainID)) + if err == nil { + refs = append(refs, ref) + } else { + if _, err := rootfs.ApplyLayer(ctx, layer, chain, s, w.Applier); err != nil { + return nil, nil, err + } + } + chain = newChain - chainID := ociidentity.ChainID(chain) if err := w.Snapshotter.SetBlob(ctx, string(chainID), layer.Diff.Digest, layer.Blob.Digest); err != nil { - return nil, err + return nil, nil, err } } - ids := make([]string, len(chain)) + ids = make([]string, len(chain)) for i := range chain { ids[i] = string(ociidentity.ChainID(chain[:i+1])) } - return ids, nil + return ids, refs, nil } // Labels returns default labels