Skip to content

Commit 17c7a96

Browse files
authored
feat(backend): add skip support to pull/fetch layer hooks (#536)
Signed-off-by: chlins <chlins.zhang@gmail.com>
1 parent ba60a20 commit 17c7a96

7 files changed

Lines changed: 340 additions & 13 deletions

File tree

pkg/backend/fetch.go

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,12 @@ import (
3838
func (b *backend) Fetch(ctx context.Context, target string, cfg *config.Fetch) error {
3939
logrus.Infof("fetch: fetching from %s", target)
4040

41+
// Apply default hooks when caller leaves it unset to avoid nil deref.
42+
if cfg.Hooks == nil {
43+
defaults := config.NewFetch()
44+
cfg.Hooks = defaults.Hooks
45+
}
46+
4147
// fetchByDragonfly is called if a Dragonfly endpoint is specified in the configuration.
4248
if cfg.DragonflyEndpoint != "" {
4349
logrus.Infof("fetch: using dragonfly for %s", target)
@@ -117,11 +123,19 @@ func (b *backend) Fetch(ctx context.Context, target string, cfg *config.Fetch) e
117123
}
118124

119125
logrus.Debugf("fetch: processing layer %s", layer.Digest)
126+
if cfg.Hooks.BeforePullLayer(layer, manifest) {
127+
logrus.Debugf("fetch: layer %s skipped by hook", layer.Digest)
128+
pb.Complete(layer.Digest.String(), fmt.Sprintf("%s %s", internalpb.NormalizePrompt("Skipped blob"), layer.Digest.String()))
129+
cfg.Hooks.AfterPullLayer(layer, true, nil)
130+
return nil
131+
}
120132
if err := tracker.TrackTransfer(func() error {
121133
return pullAndExtractFromRemote(ctx, pb, internalpb.NormalizePrompt("Fetching blob"), client, cfg.Output, layer, tracker)
122134
}); err != nil {
135+
cfg.Hooks.AfterPullLayer(layer, false, err)
123136
return err
124137
}
138+
cfg.Hooks.AfterPullLayer(layer, false, nil)
125139

126140
logrus.Debugf("fetch: successfully processed layer %s", layer.Digest)
127141
return nil

pkg/backend/fetch_by_d7y.go

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -157,9 +157,14 @@ func (b *backend) fetchByDragonfly(ctx context.Context, target string, cfg *conf
157157
func fetchLayerByDragonfly(ctx context.Context, pb *internalpb.ProgressBar, client dfdaemon.DfdaemonDownloadClient, ref Referencer, manifest ocispec.Manifest, desc ocispec.Descriptor, authToken string, cfg *config.Fetch) error {
158158
err := retry.Do(func() error {
159159
logrus.Debugf("fetch: processing layer %s", desc.Digest)
160-
cfg.Hooks.BeforePullLayer(desc, manifest) // Call before hook
160+
if cfg.Hooks.BeforePullLayer(desc, manifest) {
161+
logrus.Debugf("fetch: layer %s skipped by hook", desc.Digest)
162+
pb.Complete(desc.Digest.String(), fmt.Sprintf("%s %s", internalpb.NormalizePrompt("Skipped blob"), desc.Digest.String()))
163+
cfg.Hooks.AfterPullLayer(desc, true, nil)
164+
return nil
165+
}
161166
err := downloadAndExtractFetchLayer(ctx, pb, client, ref, desc, authToken, cfg)
162-
cfg.Hooks.AfterPullLayer(desc, err) // Call after hook
167+
cfg.Hooks.AfterPullLayer(desc, false, err) // Call after hook
163168
if err != nil {
164169
err = fmt.Errorf("pull: failed to download and extract layer %s: %w", desc.Digest, err)
165170
logrus.Error(err)

pkg/backend/fetch_hooks_test.go

Lines changed: 182 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,182 @@
1+
/*
2+
* Copyright 2025 The ModelPack Authors
3+
*
4+
* Licensed under the Apache License, Version 2.0 (the "License");
5+
* you may not use this file except in compliance with the License.
6+
* You may obtain a copy of the License at
7+
*
8+
* http://www.apache.org/licenses/LICENSE-2.0
9+
*
10+
* Unless required by applicable law or agreed to in writing, software
11+
* distributed under the License is distributed on an "AS IS" BASIS,
12+
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
* See the License for the specific language governing permissions and
14+
* limitations under the License.
15+
*/
16+
17+
package backend
18+
19+
import (
20+
"context"
21+
"encoding/json"
22+
"fmt"
23+
"net/http"
24+
"net/http/httptest"
25+
"os"
26+
"strings"
27+
"sync"
28+
"sync/atomic"
29+
"testing"
30+
31+
modelspec "github.com/modelpack/model-spec/specs-go/v1"
32+
godigest "github.com/opencontainers/go-digest"
33+
ocispec "github.com/opencontainers/image-spec/specs-go/v1"
34+
"github.com/stretchr/testify/assert"
35+
"github.com/stretchr/testify/require"
36+
37+
"github.com/modelpack/modctl/pkg/config"
38+
)
39+
40+
// recordingFetchHook tracks hook invocations and can request specific layers
41+
// to be skipped by digest.
42+
type recordingFetchHook struct {
43+
mu sync.Mutex
44+
skipDigests map[string]bool
45+
beforeCount int32
46+
afterCalls []afterFetchCall
47+
}
48+
49+
type afterFetchCall struct {
50+
digest string
51+
skipped bool
52+
err error
53+
}
54+
55+
func (r *recordingFetchHook) BeforePullLayer(desc ocispec.Descriptor, _ ocispec.Manifest) bool {
56+
atomic.AddInt32(&r.beforeCount, 1)
57+
r.mu.Lock()
58+
defer r.mu.Unlock()
59+
return r.skipDigests[desc.Digest.String()]
60+
}
61+
62+
func (r *recordingFetchHook) AfterPullLayer(desc ocispec.Descriptor, skipped bool, err error) {
63+
r.mu.Lock()
64+
defer r.mu.Unlock()
65+
r.afterCalls = append(r.afterCalls, afterFetchCall{
66+
digest: desc.Digest.String(),
67+
skipped: skipped,
68+
err: err,
69+
})
70+
}
71+
72+
// startFetchTestServer spins up an HTTP server that serves a manifest with
73+
// two layers and tracks how many times each blob is requested.
74+
func startFetchTestServer(t *testing.T) (server *httptest.Server, file1Digest, file2Digest godigest.Digest, blobHits map[string]*int32) {
75+
t.Helper()
76+
77+
const (
78+
file1Content = "file1 content..."
79+
file2Content = "file2 content..."
80+
)
81+
file1Digest = godigest.FromString(file1Content)
82+
file2Digest = godigest.FromString(file2Content)
83+
84+
hits := map[string]*int32{
85+
file1Digest.String(): new(int32),
86+
file2Digest.String(): new(int32),
87+
}
88+
89+
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
90+
switch r.URL.Path {
91+
case "/v2/":
92+
w.WriteHeader(http.StatusOK)
93+
case "/v2/test/model/manifests/latest":
94+
manifest := ocispec.Manifest{
95+
Layers: []ocispec.Descriptor{
96+
{
97+
MediaType: "application/octet-stream.raw",
98+
Digest: file1Digest,
99+
Size: int64(len(file1Content)),
100+
Annotations: map[string]string{
101+
modelspec.AnnotationFilepath: "file1.txt",
102+
},
103+
},
104+
{
105+
MediaType: "application/octet-stream.raw",
106+
Digest: file2Digest,
107+
Size: int64(len(file2Content)),
108+
Annotations: map[string]string{
109+
modelspec.AnnotationFilepath: "file2.txt",
110+
},
111+
},
112+
},
113+
}
114+
w.Header().Set("Content-Type", "application/json")
115+
require.NoError(t, json.NewEncoder(w).Encode(manifest))
116+
case fmt.Sprintf("/v2/test/model/blobs/%s", file1Digest):
117+
atomic.AddInt32(hits[file1Digest.String()], 1)
118+
_, err := w.Write([]byte(file1Content))
119+
require.NoError(t, err)
120+
case fmt.Sprintf("/v2/test/model/blobs/%s", file2Digest):
121+
atomic.AddInt32(hits[file2Digest.String()], 1)
122+
_, err := w.Write([]byte(file2Content))
123+
require.NoError(t, err)
124+
default:
125+
t.Logf("Unexpected request to %s", r.URL.Path)
126+
w.WriteHeader(http.StatusNotFound)
127+
}
128+
}))
129+
130+
return server, file1Digest, file2Digest, hits
131+
}
132+
133+
// TestFetch_HookSkipShortCircuitsLayer verifies that returning skip=true from
134+
// BeforePullLayer prevents the blob from being downloaded and that
135+
// AfterPullLayer is still invoked with skipped=true.
136+
func TestFetch_HookSkipShortCircuitsLayer(t *testing.T) {
137+
tempDir, err := os.MkdirTemp("", "fetch-hook-test")
138+
require.NoError(t, err)
139+
defer os.RemoveAll(tempDir)
140+
141+
server, file1Digest, file2Digest, hits := startFetchTestServer(t)
142+
defer server.Close()
143+
144+
hook := &recordingFetchHook{
145+
skipDigests: map[string]bool{file1Digest.String(): true},
146+
}
147+
148+
b := &backend{}
149+
url := strings.TrimPrefix(server.URL, "http://")
150+
cfg := &config.Fetch{
151+
Output: tempDir,
152+
Patterns: []string{"*.txt"},
153+
PlainHTTP: true,
154+
Concurrency: 2,
155+
Hooks: hook,
156+
}
157+
158+
require.NoError(t, b.Fetch(context.Background(), url+"/test/model:latest", cfg))
159+
160+
// file1 must NOT have been downloaded; file2 must have been.
161+
assert.Equal(t, int32(0), atomic.LoadInt32(hits[file1Digest.String()]),
162+
"skipped layer should not be fetched from remote")
163+
assert.Equal(t, int32(1), atomic.LoadInt32(hits[file2Digest.String()]),
164+
"non-skipped layer should be fetched once")
165+
166+
// BeforePullLayer fires for both layers exactly once (no retries on success).
167+
assert.Equal(t, int32(2), atomic.LoadInt32(&hook.beforeCount))
168+
169+
// AfterPullLayer must be invoked for both layers, with proper skipped flag.
170+
hook.mu.Lock()
171+
defer hook.mu.Unlock()
172+
require.Len(t, hook.afterCalls, 2)
173+
174+
byDigest := map[string]afterFetchCall{}
175+
for _, c := range hook.afterCalls {
176+
byDigest[c.digest] = c
177+
}
178+
assert.True(t, byDigest[file1Digest.String()].skipped, "file1 should be marked skipped")
179+
assert.NoError(t, byDigest[file1Digest.String()].err)
180+
assert.False(t, byDigest[file2Digest.String()].skipped, "file2 should not be marked skipped")
181+
assert.NoError(t, byDigest[file2Digest.String()].err)
182+
}

pkg/backend/pull.go

Lines changed: 14 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,12 @@ import (
4141
func (b *backend) Pull(ctx context.Context, target string, cfg *config.Pull) error {
4242
logrus.Infof("pull: pulling artifact %s", target)
4343

44+
// Apply default hooks when caller leaves it unset to avoid nil deref.
45+
if cfg.Hooks == nil {
46+
defaults := config.NewPull()
47+
cfg.Hooks = defaults.Hooks
48+
}
49+
4450
// pullByDragonfly is called if a Dragonfly endpoint is specified in the configuration.
4551
if cfg.DragonflyEndpoint != "" {
4652
logrus.Infof("pull: using dragonfly for %s", target)
@@ -118,13 +124,18 @@ func (b *backend) Pull(ctx context.Context, target string, cfg *config.Pull) err
118124

119125
return retry.Do(func() error {
120126
logrus.Debugf("pull: processing layer %s", layer.Digest)
121-
// call the before hook.
122-
cfg.Hooks.BeforePullLayer(layer, manifest)
127+
// call the before hook; allow caller to skip this layer.
128+
if cfg.Hooks.BeforePullLayer(layer, manifest) {
129+
logrus.Debugf("pull: layer %s skipped by hook", layer.Digest)
130+
pb.Complete(layer.Digest.String(), fmt.Sprintf("%s %s", internalpb.NormalizePrompt("Skipped blob"), layer.Digest.String()))
131+
cfg.Hooks.AfterPullLayer(layer, true, nil)
132+
return nil
133+
}
123134
err := tracker.TrackTransfer(func() error {
124135
return fn(layer)
125136
})
126137
// call the after hook.
127-
cfg.Hooks.AfterPullLayer(layer, err)
138+
cfg.Hooks.AfterPullLayer(layer, false, err)
128139
if err != nil {
129140
err = fmt.Errorf("pull: failed to process layer %s: %w", layer.Digest, err)
130141
logrus.Error(err)

pkg/backend/pull_by_d7y.go

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -181,9 +181,14 @@ func buildBlobURL(ref Referencer, plainHTTP bool, digest string) string {
181181
func processLayer(ctx context.Context, pb *internalpb.ProgressBar, client dfdaemon.DfdaemonDownloadClient, ref Referencer, manifest ocispec.Manifest, desc ocispec.Descriptor, authToken string, cfg *config.Pull) error {
182182
err := retry.Do(func() error {
183183
logrus.Debugf("pull: processing layer %s", desc.Digest)
184-
cfg.Hooks.BeforePullLayer(desc, manifest) // Call before hook
184+
if cfg.Hooks.BeforePullLayer(desc, manifest) {
185+
logrus.Debugf("pull: layer %s skipped by hook", desc.Digest)
186+
pb.Complete(desc.Digest.String(), fmt.Sprintf("%s %s", internalpb.NormalizePrompt("Skipped blob"), desc.Digest.String()))
187+
cfg.Hooks.AfterPullLayer(desc, true, nil)
188+
return nil
189+
}
185190
err := downloadAndExtractLayer(ctx, pb, client, ref, desc, authToken, cfg)
186-
cfg.Hooks.AfterPullLayer(desc, err) // Call after hook
191+
cfg.Hooks.AfterPullLayer(desc, false, err) // Call after hook
187192
if err != nil {
188193
err = fmt.Errorf("pull: failed to download and extract layer %s: %w", desc.Digest, err)
189194
logrus.Error(err)

pkg/config/pull.go

Lines changed: 19 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -78,16 +78,29 @@ func (p *Pull) Validate() error {
7878
}
7979

8080
// PullHooks is the hook events during the pull operation.
81+
//
82+
// Note: every retry attempt re-invokes BeforePullLayer / AfterPullLayer.
8183
type PullHooks interface {
82-
// BeforePullLayer will execute before pulling the layer described as desc, will carry the manifest as well.
83-
BeforePullLayer(desc ocispec.Descriptor, manifest ocispec.Manifest)
84+
// BeforePullLayer will execute before pulling the layer described as desc,
85+
// will carry the manifest as well.
86+
//
87+
// If the hook returns skip=true, the backend will treat this layer as
88+
// already satisfied and will NOT actually pull/extract it. The caller is
89+
// responsible for ensuring the corresponding content already exists and
90+
// matches the descriptor's digest. AfterPullLayer will still be invoked
91+
// with skipped=true and a nil error.
92+
BeforePullLayer(desc ocispec.Descriptor, manifest ocispec.Manifest) (skip bool)
8493

85-
// AfterPullLayer will execute after pulling the layer described as desc, the error will be nil if pulled successfully.
86-
AfterPullLayer(desc ocispec.Descriptor, err error)
94+
// AfterPullLayer will execute after pulling the layer described as desc.
95+
// skipped indicates whether the layer was skipped by BeforePullLayer's
96+
// decision. err will be nil if pulled (or skipped) successfully.
97+
AfterPullLayer(desc ocispec.Descriptor, skipped bool, err error)
8798
}
8899

89100
// emptyPullHook is the empty pull hook implementation with do nothing.
90101
type emptyPullHook struct{}
91102

92-
func (emptyPullHook) BeforePullLayer(desc ocispec.Descriptor, manifest ocispec.Manifest) {}
93-
func (emptyPullHook) AfterPullLayer(desc ocispec.Descriptor, err error) {}
103+
func (emptyPullHook) BeforePullLayer(desc ocispec.Descriptor, manifest ocispec.Manifest) bool {
104+
return false
105+
}
106+
func (emptyPullHook) AfterPullLayer(desc ocispec.Descriptor, skipped bool, err error) {}

0 commit comments

Comments
 (0)