Unverified Commit c0643af1 authored by Lemon-miaow's avatar Lemon-miaow
Browse files

feat(imagepush): 从公共 registry 拷贝单平台镜像与 OCI 制品

parent 80c4ea89
Loading
Loading
Loading
Loading
+450 −0
Changes for internal/imagepush/mirror.go: 450 added lines, 0 removed lines.
Original line number Diff line number Diff line
package imagepush

import (
	"bytes"
	"context"
	"crypto/sha256"
	"encoding/hex"
	"encoding/json"
	"errors"
	"fmt"
	"io"
	"net/http"
	"net/url"
	"os"
	"path/filepath"
	"regexp"
	"runtime"
	"strings"
	"sync"
	"time"
)

// Mirroring copies one image or OCI artifact from a public registry into the
// platform registry: the build tools (kaniko, trivy) and Trivy's vulnerability
// DBs, which a build Job cannot fetch itself because its namespace has no
// internet egress. Blobs stream from the source straight into the upload, and
// the destination registry checks each against its digest.
//
// An image index is narrowed to the one platform the node runs: the copy is a
// single-platform manifest under the destination tag. A source pinned by digest
// is checked against it before anything is copied.

const (
	mediaOCIIndex        = "application/vnd.oci.image.index.v1+json"
	mediaDockerList      = "application/vnd.docker.distribution.manifest.list.v2+json"
	manifestAccept       = mediaOCIIndex + "," + mediaDockerList + "," + mediaOCIManifest + "," + mediaDockerManifest
	maxManifestBytes     = 4 << 20
	sourceRequestTimeout = 2 * time.Minute
)

// SourceRef is host/repository[:tag][@digest].
type SourceRef struct {
	Host, Repo, Tag, Digest string
}

func (r SourceRef) String() string {
	s := r.Host + "/" + r.Repo
	if r.Tag != "" {
		s += ":" + r.Tag
	}
	if r.Digest != "" {
		s += "@" + r.Digest
	}
	return s
}

// reference is what the manifest endpoint is asked for: the digest when pinned.
func (r SourceRef) reference() string {
	if r.Digest != "" {
		return r.Digest
	}
	return r.Tag
}

var sha256DigestRE = regexp.MustCompile(`^sha256:[0-9a-f]{64}$`)

// ParseSourceRef parses a fully qualified reference. docker.io is served from
// registry-1.docker.io, and its single-component names live under library/.
func ParseSourceRef(ref string) (SourceRef, error) {
	name, digest, _ := strings.Cut(ref, "@")
	if digest != "" && !sha256DigestRE.MatchString(digest) {
		return SourceRef{}, fmt.Errorf("imagepush: %q: bad digest", ref)
	}
	host, rest, ok := strings.Cut(name, "/")
	if !ok || rest == "" || !strings.ContainsAny(host, ".:") {
		return SourceRef{}, fmt.Errorf("imagepush: %q must be host/repository[:tag][@digest]", ref)
	}
	repo, tag := rest, ""
	if i := strings.LastIndexByte(rest, ':'); i > strings.LastIndexByte(rest, '/') {
		repo, tag = rest[:i], rest[i+1:]
	}
	if repo == "" || (tag == "" && strings.HasSuffix(rest, ":")) {
		return SourceRef{}, fmt.Errorf("imagepush: %q must be host/repository[:tag][@digest]", ref)
	}
	if tag == "" && digest == "" {
		tag = "latest"
	}
	if host == "docker.io" {
		host = "registry-1.docker.io"
		if !strings.Contains(repo, "/") {
			repo = "library/" + repo
		}
	}
	return SourceRef{Host: host, Repo: repo, Tag: tag, Digest: digest}, nil
}

// Source reads manifests and blobs anonymously from public registries, answering
// their bearer-token challenges (ghcr.io, gcr.io, mirror.gcr.io, Docker Hub).
type Source struct {
	// Client performs the requests; nil uses one with response-header timeouts.
	Client *http.Client
	// Platform selects the manifest of an index, as "os/arch" or
	// "os/arch/variant". Empty means linux and this binary's architecture.
	Platform string
	// Scheme is "https" unless a test serves plain HTTP.
	Scheme string

	mu     sync.Mutex
	tokens map[string]string // host/repo → bearer token
}

func (s *Source) client() *http.Client {
	if s.Client != nil {
		return s.Client
	}
	return &http.Client{Transport: &http.Transport{Proxy: http.ProxyFromEnvironment, ResponseHeaderTimeout: sourceRequestTimeout}}
}

func (s *Source) url(r SourceRef, tail string) string {
	scheme := s.Scheme
	if scheme == "" {
		scheme = "https"
	}
	return (&url.URL{Scheme: scheme, Host: r.Host, Path: "/v2/" + r.Repo + "/" + tail}).String()
}

// get issues a GET, fetching a bearer token once when the registry challenges.
func (s *Source) get(ctx context.Context, r SourceRef, target, accept string) (*http.Response, error) {
	key := r.Host + "/" + r.Repo
	for attempt := 0; ; attempt++ {
		req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
		if err != nil {
			return nil, err
		}
		if accept != "" {
			req.Header.Set("Accept", accept)
		}
		s.mu.Lock()
		tok := s.tokens[key]
		s.mu.Unlock()
		if tok != "" {
			req.Header.Set("Authorization", "Bearer "+tok)
		}
		resp, err := s.client().Do(req)
		if err != nil {
			return nil, err
		}
		if resp.StatusCode != http.StatusUnauthorized || attempt > 0 {
			return resp, nil
		}
		challenge := resp.Header.Get("WWW-Authenticate")
		resp.Body.Close()
		tok, err = s.token(ctx, challenge)
		if err != nil {
			return nil, fmt.Errorf("%s: %w", key, err)
		}
		s.mu.Lock()
		if s.tokens == nil {
			s.tokens = map[string]string{}
		}
		s.tokens[key] = tok
		s.mu.Unlock()
	}
}

// token answers a Bearer challenge anonymously.
func (s *Source) token(ctx context.Context, challenge string) (string, error) {
	scheme, params, _ := strings.Cut(challenge, " ")
	if !strings.EqualFold(scheme, "Bearer") {
		return "", fmt.Errorf("registry wants %q authentication; only anonymous bearer tokens are supported", scheme)
	}
	p := parseChallenge(params)
	if p["realm"] == "" {
		return "", errors.New("bearer challenge without a realm")
	}
	u, err := url.Parse(p["realm"])
	if err != nil {
		return "", fmt.Errorf("bearer realm: %w", err)
	}
	q := u.Query()
	for _, k := range []string{"service", "scope"} {
		if p[k] != "" {
			q.Set(k, p[k])
		}
	}
	u.RawQuery = q.Encode()
	req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
	if err != nil {
		return "", err
	}
	resp, err := s.client().Do(req)
	if err != nil {
		return "", err
	}
	defer resp.Body.Close()
	if resp.StatusCode != http.StatusOK {
		return "", statusError("token", resp)
	}
	var body struct {
		Token       string `json:"token"`
		AccessToken string `json:"access_token"`
	}
	if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&body); err != nil {
		return "", fmt.Errorf("token: %w", err)
	}
	if body.Token != "" {
		return body.Token, nil
	}
	if body.AccessToken != "" {
		return body.AccessToken, nil
	}
	return "", errors.New("token endpoint returned no token")
}

// parseChallenge splits key="value",key="value".
func parseChallenge(s string) map[string]string {
	out := map[string]string{}
	for s != "" {
		s = strings.TrimLeft(s, " ,")
		k, rest, ok := strings.Cut(s, "=")
		if !ok {
			break
		}
		var v string
		if strings.HasPrefix(rest, `"`) {
			end := strings.IndexByte(rest[1:], '"')
			if end < 0 {
				break
			}
			v, s = rest[1:1+end], rest[2+end:]
		} else {
			v, s, _ = strings.Cut(rest, ",")
		}
		out[strings.ToLower(strings.TrimSpace(k))] = v
	}
	return out
}

// manifest fetches one manifest by tag or digest and returns its bytes and
// media type. A digest reference is verified against the bytes.
func (s *Source) manifest(ctx context.Context, r SourceRef, reference string) ([]byte, string, error) {
	resp, err := s.get(ctx, r, s.url(r, "manifests/"+reference), manifestAccept)
	if err != nil {
		return nil, "", err
	}
	defer resp.Body.Close()
	if resp.StatusCode != http.StatusOK {
		return nil, "", statusError("get manifest "+reference, resp)
	}
	body, err := io.ReadAll(io.LimitReader(resp.Body, maxManifestBytes+1))
	if err != nil {
		return nil, "", err
	}
	if len(body) > maxManifestBytes {
		return nil, "", fmt.Errorf("manifest %s is larger than %d bytes", reference, maxManifestBytes)
	}
	if strings.HasPrefix(reference, "sha256:") {
		if got := digestOf(body); got != reference {
			return nil, "", fmt.Errorf("manifest %s hashes to %s", reference, got)
		}
	}
	mt, _, _ := strings.Cut(resp.Header.Get("Content-Type"), ";")
	var probe struct {
		MediaType string `json:"mediaType"`
	}
	if json.Unmarshal(body, &probe) == nil && probe.MediaType != "" {
		mt = probe.MediaType
	}
	return body, strings.TrimSpace(mt), nil
}

func (s *Source) blob(ctx context.Context, r SourceRef, digest string) (io.ReadCloser, error) {
	resp, err := s.get(ctx, r, s.url(r, "blobs/"+digest), "")
	if err != nil {
		return nil, err
	}
	if resp.StatusCode != http.StatusOK {
		defer resp.Body.Close()
		return nil, statusError("get blob "+digest, resp)
	}
	return resp.Body, nil
}

func (s *Source) platform() (goos, arch, variant string) {
	p := s.Platform
	if p == "" {
		p = "linux/" + runtime.GOARCH
	}
	parts := strings.SplitN(p, "/", 3)
	goos = parts[0]
	if len(parts) > 1 {
		arch = parts[1]
	}
	if len(parts) > 2 {
		variant = parts[2]
	}
	return goos, arch, variant
}

type indexEntry struct {
	MediaType string `json:"mediaType"`
	Digest    string `json:"digest"`
	Platform  *struct {
		OS           string `json:"os"`
		Architecture string `json:"architecture"`
		Variant      string `json:"variant"`
	} `json:"platform"`
}

// pick returns the digest of the index entry for the wanted platform.
func (s *Source) pick(body []byte) (string, error) {
	var idx struct {
		Manifests []indexEntry `json:"manifests"`
	}
	if err := json.Unmarshal(body, &idx); err != nil {
		return "", fmt.Errorf("index: %w", err)
	}
	wantOS, wantArch, wantVariant := s.platform()
	var fallback string
	for _, m := range idx.Manifests {
		if m.Platform == nil || m.Platform.OS != wantOS || m.Platform.Architecture != wantArch {
			continue
		}
		if wantVariant == "" || m.Platform.Variant == wantVariant {
			return m.Digest, nil
		}
		if fallback == "" {
			fallback = m.Digest
		}
	}
	if fallback != "" {
		return fallback, nil
	}
	return "", fmt.Errorf("the index has no %s/%s manifest", wantOS, wantArch)
}

// Mirror copies src (a SourceRef string) to dst (host/repository:tag on p's
// registry) and returns the digest of the manifest it wrote.
func (p *Pusher) Mirror(ctx context.Context, s *Source, src, dst string) (string, error) {
	sr, err := ParseSourceRef(src)
	if err != nil {
		return "", err
	}
	dr, err := ParseRef(dst)
	if err != nil {
		return "", err
	}
	var body []byte
	var mt string
	err = p.retry(ctx, func() error {
		body, mt, err = s.manifest(ctx, sr, sr.reference())
		return err
	})
	if err != nil {
		return "", fmt.Errorf("imagepush: %s: %w", sr, err)
	}
	if mt == mediaOCIIndex || mt == mediaDockerList {
		d, err := s.pick(body)
		if err != nil {
			return "", fmt.Errorf("imagepush: %s: %w", sr, err)
		}
		err = p.retry(ctx, func() error {
			body, mt, err = s.manifest(ctx, sr, d)
			return err
		})
		if err != nil {
			return "", fmt.Errorf("imagepush: %s: %w", sr, err)
		}
	}
	if mt != mediaOCIManifest && mt != mediaDockerManifest {
		return "", fmt.Errorf("imagepush: %s: unsupported manifest type %q", sr, mt)
	}
	var m struct {
		Config descriptor   `json:"config"`
		Layers []descriptor `json:"layers"`
	}
	if err := json.Unmarshal(body, &m); err != nil {
		return "", fmt.Errorf("imagepush: %s: manifest: %w", sr, err)
	}
	blobs := append([]descriptor{m.Config}, m.Layers...)
	for i, b := range blobs {
		if !sha256DigestRE.MatchString(b.Digest) || b.Size < 0 {
			return "", fmt.Errorf("imagepush: %s: blob %d has digest %q size %d", sr, i+1, b.Digest, b.Size)
		}
		err := p.retry(ctx, func() error {
			return p.uploadBlob(ctx, dr, b.Digest, b.Size, func() (io.ReadCloser, error) {
				return s.blob(ctx, sr, b.Digest)
			})
		})
		if err != nil {
			return "", fmt.Errorf("imagepush: %s: blob %d/%d (%s): %w", sr, i+1, len(blobs), b.Digest, err)
		}
	}
	var digest string
	err = p.retry(ctx, func() error {
		d, err := p.putManifest(ctx, dr, mt, bytes.Clone(body))
		digest = d
		return err
	})
	if err != nil {
		return "", fmt.Errorf("imagepush: %s: manifest: %w", dr, err)
	}
	p.logf("mirrored %s to %s@%s", sr, dr, digest)
	return digest, nil
}

func digestOf(b []byte) string {
	sum := sha256.Sum256(b)
	return "sha256:" + hex.EncodeToString(sum[:])
}

// MirrorStatus is what `felis mirror-build-tools --status` records, and the
// watchdog reads to tell a stale vulnerability DB.
type MirrorStatus struct {
	LastAttempt time.Time `json:"last_attempt"`
	LastSuccess time.Time `json:"last_success,omitzero"`
	LastError   string    `json:"last_error,omitempty"`
}

// ReadMirrorStatus reads a status file; a missing one is (nil, nil).
func ReadMirrorStatus(path string) (*MirrorStatus, error) {
	b, err := os.ReadFile(path)
	if errors.Is(err, os.ErrNotExist) {
		return nil, nil
	}
	if err != nil {
		return nil, err
	}
	var st MirrorStatus
	if err := json.Unmarshal(b, &st); err != nil {
		return nil, fmt.Errorf("%s: %w", path, err)
	}
	return &st, nil
}

// WriteMirrorStatus replaces the status file atomically.
func WriteMirrorStatus(path string, st MirrorStatus) error {
	if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
		return err
	}
	b, err := json.MarshalIndent(st, "", "  ")
	if err != nil {
		return err
	}
	tmp := path + ".tmp"
	if err := os.WriteFile(tmp, append(b, '\n'), 0o644); err != nil {
		return err
	}
	return os.Rename(tmp, path)
}
+203 −0
Changes for internal/imagepush/mirror_test.go: 203 added lines, 0 removed lines.
Original line number Diff line number Diff line
package imagepush

import (
	"context"
	"encoding/json"
	"fmt"
	"net/http"
	"net/http/httptest"
	"path/filepath"
	"strings"
	"testing"
	"time"

	"felis.lolicon.best/internal/registrygate"
)

// fakeSource is a public registry that answers anonymous requests with a bearer
// challenge, like ghcr.io: a token from /token is required for every read.
type fakeSource struct {
	manifests map[string][]byte // reference (tag or digest) → body
	types     map[string]string
	blobs     map[string][]byte
	tokens    int
	srv       *httptest.Server
}

func newFakeSource(t *testing.T) *fakeSource {
	f := &fakeSource{manifests: map[string][]byte{}, types: map[string]string{}, blobs: map[string][]byte{}}
	f.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		if r.URL.Path == "/token" {
			if r.URL.Query().Get("scope") != "repository:tools/thing:pull" || r.URL.Query().Get("service") != "fake" {
				http.Error(w, "bad scope", http.StatusBadRequest)
				return
			}
			f.tokens++
			fmt.Fprint(w, `{"token":"t0k"}`)
			return
		}
		if r.Header.Get("Authorization") != "Bearer t0k" {
			w.Header().Set("WWW-Authenticate", fmt.Sprintf(`Bearer realm="%s/token",service="fake",scope="repository:tools/thing:pull"`, f.srv.URL))
			w.WriteHeader(http.StatusUnauthorized)
			return
		}
		ref := r.URL.Path[strings.LastIndex(r.URL.Path, "/")+1:]
		switch {
		case strings.Contains(r.URL.Path, "/manifests/"):
			b, ok := f.manifests[ref]
			if !ok {
				http.NotFound(w, r)
				return
			}
			w.Header().Set("Content-Type", f.types[ref])
			w.Write(b)
		case strings.Contains(r.URL.Path, "/blobs/"):
			b, ok := f.blobs[ref]
			if !ok {
				http.NotFound(w, r)
				return
			}
			w.Write(b)
		default:
			http.NotFound(w, r)
		}
	}))
	t.Cleanup(f.srv.Close)
	return f
}

func (f *fakeSource) host() string { return strings.TrimPrefix(f.srv.URL, "http://") }

func (f *fakeSource) addBlob(b []byte) descriptor {
	d := digestOf(b)
	f.blobs[d] = b
	return descriptor{MediaType: mediaOCILayerGz, Size: int64(len(b)), Digest: d}
}

// addImage stores a single-platform manifest and returns its digest.
func (f *fakeSource) addImage(arch string) string {
	cfg := f.addBlob([]byte(`{"architecture":"` + arch + `","os":"linux"}`))
	cfg.MediaType = mediaOCIConfig
	layer := f.addBlob([]byte("layer for " + arch))
	body, _ := json.Marshal(manifest{SchemaVersion: 2, MediaType: mediaOCIManifest, Config: cfg, Layers: []descriptor{layer}})
	d := digestOf(body)
	f.manifests[d] = body
	f.types[d] = mediaOCIManifest
	return d
}

// addIndex stores an index over amd64 and arm64 under tag and returns the
// index digest and the arm64 manifest digest.
func (f *fakeSource) addIndex(tag string) (index, arm64 string) {
	amd := f.addImage("amd64")
	arm := f.addImage("arm64")
	body := fmt.Sprintf(`{"schemaVersion":2,"mediaType":%q,"manifests":[`+
		`{"mediaType":%q,"digest":%q,"size":1,"platform":{"os":"linux","architecture":"amd64"}},`+
		`{"mediaType":%q,"digest":%q,"size":1,"platform":{"os":"linux","architecture":"arm64","variant":"v8"}}]}`,
		mediaOCIIndex, mediaOCIManifest, amd, mediaOCIManifest, arm)
	d := digestOf([]byte(body))
	f.manifests[tag], f.manifests[d] = []byte(body), []byte(body)
	f.types[tag], f.types[d] = mediaOCIIndex, mediaOCIIndex
	return d, arm
}

func TestMirrorCopiesThePlatformManifest(t *testing.T) {
	src := newFakeSource(t)
	idx, arm := src.addIndex("v1")
	reg, host := startStack(t)
	p := &Pusher{Scheme: "http", Username: registrygate.PrincipalPlatform, Password: "plat-secret", Attempts: 1}
	s := &Source{Scheme: "http", Platform: "linux/arm64"}

	got, err := p.Mirror(context.Background(), s, src.host()+"/tools/thing:v1@"+idx, host+"/mirror/thing:v1")
	if err != nil {
		t.Fatal(err)
	}
	if got != arm {
		t.Errorf("mirrored digest = %s, want the arm64 manifest %s", got, arm)
	}
	if string(reg.manifests["mirror/thing:v1"]) != string(src.manifests[arm]) {
		t.Error("the destination manifest differs from the source's bytes")
	}
	if reg.puts != 2 {
		t.Errorf("uploaded %d blobs, want the config and one layer", reg.puts)
	}
	if src.tokens != 1 {
		t.Errorf("fetched %d tokens, want one reused for every read", src.tokens)
	}

	// A second run finds every blob in place and only re-puts the manifest.
	if _, err := p.Mirror(context.Background(), s, src.host()+"/tools/thing:v1@"+idx, host+"/mirror/thing:v1"); err != nil {
		t.Fatal(err)
	}
	if reg.puts != 2 {
		t.Errorf("second run uploaded blobs again: %d", reg.puts)
	}
}

func TestMirrorRefusesAMovedPin(t *testing.T) {
	src := newFakeSource(t)
	src.addIndex("v1")
	other, _ := src.addIndex("v2")
	// Different bytes under the pinned digest: a tampered or re-pushed source.
	src.manifests[other] = []byte(`{"schemaVersion":2,"mediaType":"` + mediaOCIIndex + `","manifests":[]}`)
	reg, host := startStack(t)
	p := &Pusher{Scheme: "http", Username: registrygate.PrincipalPlatform, Password: "plat-secret", Attempts: 1}
	_, err := p.Mirror(context.Background(), &Source{Scheme: "http"}, src.host()+"/tools/thing:v1@"+other, host+"/mirror/thing:v1")
	if err == nil || !strings.Contains(err.Error(), "hashes to") {
		t.Fatalf("Mirror = %v, want a digest mismatch", err)
	}
	if len(reg.manifests) != 0 || reg.puts != 0 {
		t.Error("a refused source still wrote to the registry")
	}
}

func TestMirrorFollowsATagForAnArtifact(t *testing.T) {
	src := newFakeSource(t)
	d := src.addImage("amd64")
	src.manifests["2"], src.types["2"] = src.manifests[d], mediaOCIManifest
	reg, host := startStack(t)
	p := &Pusher{Scheme: "http", Username: registrygate.PrincipalPlatform, Password: "plat-secret", Attempts: 1}
	got, err := p.Mirror(context.Background(), &Source{Scheme: "http"}, src.host()+"/tools/thing:2", host+"/mirror/thing-db:2")
	if err != nil {
		t.Fatal(err)
	}
	if got != d || reg.manifests["mirror/thing-db:2"] == nil {
		t.Errorf("Mirror = %s, manifests %v", got, reg.manifests)
	}
}

func TestParseSourceRef(t *testing.T) {
	for in, want := range map[string]SourceRef{
		"ghcr.io/aquasecurity/trivy:0.74.0":                    {Host: "ghcr.io", Repo: "aquasecurity/trivy", Tag: "0.74.0"},
		"mirror.gcr.io/aquasec/trivy-db:2":                     {Host: "mirror.gcr.io", Repo: "aquasec/trivy-db", Tag: "2"},
		"docker.io/registry:2":                                 {Host: "registry-1.docker.io", Repo: "library/registry", Tag: "2"},
		"gcr.io/x/y":                                           {Host: "gcr.io", Repo: "x/y", Tag: "latest"},
		"localhost:5000/a/b@sha256:" + strings.Repeat("a", 64): {Host: "localhost:5000", Repo: "a/b", Digest: "sha256:" + strings.Repeat("a", 64)},
		"gcr.io/k/e:v1@sha256:" + strings.Repeat("b", 64):      {Host: "gcr.io", Repo: "k/e", Tag: "v1", Digest: "sha256:" + strings.Repeat("b", 64)},
	} {
		got, err := ParseSourceRef(in)
		if err != nil || got != want {
			t.Errorf("ParseSourceRef(%q) = %+v, %v; want %+v", in, got, err, want)
		}
	}
	for _, bad := range []string{"trivy:latest", "ghcr.io/", "gcr.io/x@sha256:zz", "gcr.io/x:"} {
		if _, err := ParseSourceRef(bad); err == nil {
			t.Errorf("ParseSourceRef(%q) accepted", bad)
		}
	}
}

func TestMirrorStatusRoundTrip(t *testing.T) {
	path := filepath.Join(t.TempDir(), "sub", "status.json")
	if st, err := ReadMirrorStatus(path); st != nil || err != nil {
		t.Fatalf("missing file = %v, %v", st, err)
	}
	now := time.Date(2026, 9, 24, 12, 0, 0, 0, time.UTC)
	if err := WriteMirrorStatus(path, MirrorStatus{LastAttempt: now, LastSuccess: now}); err != nil {
		t.Fatal(err)
	}
	st, err := ReadMirrorStatus(path)
	if err != nil || !st.LastSuccess.Equal(now) || st.LastError != "" {
		t.Errorf("round trip = %+v, %v", st, err)
	}
}
+22 −7
Changes for internal/imagepush/push.go: 22 added lines, 7 removed lines.
Original line number Diff line number Diff line
@@ -211,13 +211,28 @@ func (p *Pusher) describe(tarPath string) (*manifest, error) {
}

func (p *Pusher) pushBlob(ctx context.Context, r Ref, tarPath string, b descriptor) error {
	head, err := p.do(ctx, http.MethodHead, p.url(r, "blobs/"+b.Digest), nil, 0, "")
	return p.uploadBlob(ctx, r, b.Digest, b.Size, func() (io.ReadCloser, error) {
		f, entry, err := openEntry(tarPath, b.file)
		if err != nil {
			return nil, err
		}
		return struct {
			io.Reader
			io.Closer
		}{entry, f}, nil
	})
}

// uploadBlob uploads the blob open returns, unless r's repository already holds
// digest. The registry checks the bytes against digest.
func (p *Pusher) uploadBlob(ctx context.Context, r Ref, digest string, size int64, open func() (io.ReadCloser, error)) error {
	head, err := p.do(ctx, http.MethodHead, p.url(r, "blobs/"+digest), nil, 0, "")
	if err != nil {
		return err
	}
	head.Body.Close()
	if head.StatusCode == http.StatusOK {
		p.logf("exists  %s", b.Digest)
		p.logf("exists  %s", digest)
		return nil
	}

@@ -234,15 +249,15 @@ func (p *Pusher) pushBlob(ctx context.Context, r Ref, tarPath string, b descript
		return fmt.Errorf("start upload: %w", err)
	}
	q := loc.Query()
	q.Set("digest", b.Digest)
	q.Set("digest", digest)
	loc.RawQuery = q.Encode()

	f, entry, err := openEntry(tarPath, b.file)
	body, err := open()
	if err != nil {
		return err
	}
	defer f.Close()
	put, err := p.do(ctx, http.MethodPut, loc.String(), entry, b.Size, "application/octet-stream")
	defer body.Close()
	put, err := p.do(ctx, http.MethodPut, loc.String(), body, size, "application/octet-stream")
	if err != nil {
		return err
	}
@@ -250,7 +265,7 @@ func (p *Pusher) pushBlob(ctx context.Context, r Ref, tarPath string, b descript
	if put.StatusCode != http.StatusCreated {
		return statusError("upload", put)
	}
	p.logf("pushed  %s (%d bytes)", b.Digest, b.Size)
	p.logf("pushed  %s (%d bytes)", digest, size)
	return nil
}