diff --git a/cmd/proxy/main.go b/cmd/proxy/main.go index 9266de53..5b625631 100644 --- a/cmd/proxy/main.go +++ b/cmd/proxy/main.go @@ -479,6 +479,10 @@ func runMirror() { fmt.Fprintf(os.Stderr, "invalid configuration: %v\n", err) os.Exit(1) } + if !cfg.Storage.CacheArtifacts { + fmt.Fprintf(os.Stderr, "error: mirror is not available with storage.cache_artifacts: false: mirrored artifacts would never be served\n") + os.Exit(1) + } logger := setupLogger("info", "text") diff --git a/config.example.yaml b/config.example.yaml index 58a9de79..96a5ba2d 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -52,6 +52,12 @@ storage: # Empty or "0" means unlimited max_size: "" + # Store fetched artifacts. Set to false to stream every download from + # upstream without storing it; metadata filtering, cooldown and the + # denylist still apply. Useful when another cache sits in front of the + # proxy. false is incompatible with direct_serve, scanning and mirror_api. + cache_artifacts: true + # Redirect cached artifact downloads to presigned storage URLs (HTTP 302) # instead of streaming through the proxy. Only effective for S3, GCS, and Azure. # Leave disabled if clients reach the proxy through an authenticating gateway, diff --git a/docs/configuration.md b/docs/configuration.md index f09e0108..1510e70b 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -43,9 +43,23 @@ storage: | `storage.url` | `PROXY_STORAGE_URL` | `-storage-url` | Storage URL (file:// or s3://) | | `storage.path` | `PROXY_STORAGE_PATH` | `-storage-path` | Local path (deprecated, use url) | | `storage.max_size` | `PROXY_STORAGE_MAX_SIZE` | - | Max cache size (e.g., "10GB") | +| `storage.cache_artifacts` | `PROXY_STORAGE_CACHE_ARTIFACTS` | - | Store fetched artifacts (default: true); `false` streams them from upstream | `storage.max_size` counts cached artifacts only. An artifact replaced by a refetch stays in storage for at least an hour, or `storage.direct_serve_ttl` if longer, so requests already reading it can finish, and storage use can exceed the limit by what was replaced in that time. +### Serving artifacts without storing them + +With `storage.cache_artifacts: false` the proxy streams every artifact download from upstream to the client and stores no artifacts. Metadata is still filtered and cached, so cooldown and the denylist apply as usual, and the database still holds the publish times cooldown needs. Use it when another caching layer, such as an Artifactory remote repository, sits in front of the proxy and caching artifacts twice only costs storage. + +```yaml +storage: + cache_artifacts: false +``` + +Every download is a fresh upstream fetch, and concurrent requests for the same artifact are not combined. Artifacts with a digest known up front (OCI blobs, Swift archives, Helm charts) are verified while streaming: the response is sent chunked, and on a mismatch the connection is aborted before the response completes so the client never receives a tampered artifact as a good one. The same happens when the upstream connection fails mid-download. + +`cache_artifacts: false` cannot be combined with `scanning.enabled`, `storage.direct_serve` or `mirror_api`, which all need stored artifacts, and the `mirror` command refuses to run with it. + ### Amazon S3 ```yaml diff --git a/internal/config/config.go b/internal/config/config.go index fcc41fa7..6f8eda9a 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -370,6 +370,14 @@ type StorageConfig struct { // storage at an internal address (e.g. 127.0.0.1 or a Docker hostname) // but clients must use a public one. DirectServeBaseURL string `json:"direct_serve_base_url" yaml:"direct_serve_base_url"` + + // CacheArtifacts stores fetched artifacts so later downloads are served + // from storage. When false, every download streams from upstream and + // nothing is stored; metadata is still cached, and cooldown and the + // denylist still apply. Useful when another caching layer sits in front + // of the proxy. False is incompatible with scanning, direct_serve and + // mirror_api, which all depend on stored artifacts. Default: true. + CacheArtifacts bool `json:"cache_artifacts" yaml:"cache_artifacts"` } // GradleConfig configures Gradle-specific features. @@ -757,8 +765,9 @@ func Default() *Config { Listen: ":8080", BaseURL: "http://localhost:8080", Storage: StorageConfig{ - Path: "./cache/artifacts", - MaxSize: "", + Path: "./cache/artifacts", + MaxSize: "", + CacheArtifacts: true, }, Database: DatabaseConfig{ Driver: "sqlite", @@ -893,6 +902,7 @@ func (c *Config) LoadFromEnv() { setEnvBool(&c.Storage.DirectServe, "PROXY_STORAGE_DIRECT_SERVE") setEnvString(&c.Storage.DirectServeTTL, "PROXY_STORAGE_DIRECT_SERVE_TTL") setEnvString(&c.Storage.DirectServeBaseURL, "PROXY_STORAGE_DIRECT_SERVE_BASE_URL") + setEnvBool(&c.Storage.CacheArtifacts, "PROXY_STORAGE_CACHE_ARTIFACTS") setEnvString(&c.Database.Driver, "PROXY_DATABASE_DRIVER") setEnvString(&c.Database.Path, "PROXY_DATABASE_PATH") setEnvString(&c.Database.URL, "PROXY_DATABASE_URL") @@ -1023,6 +1033,10 @@ func (c *Config) Validate() error { } } + if err := c.validateCacheArtifacts(); err != nil { + return err + } + // Validate metadata TTL if specified if c.MetadataTTL != "" && c.MetadataTTL != "0" { if _, err := time.ParseDuration(c.MetadataTTL); err != nil { @@ -1041,6 +1055,21 @@ func (c *Config) Validate() error { return c.validateComponents() } +func (c *Config) validateCacheArtifacts() error { + if c.Storage.CacheArtifacts { + return nil + } + switch { + case c.Scanning.Enabled: + return fmt.Errorf("storage.cache_artifacts: false cannot be combined with scanning.enabled: scanning needs stored artifacts") + case c.Storage.DirectServe: + return fmt.Errorf("storage.cache_artifacts: false cannot be combined with storage.direct_serve: no artifacts are stored to redirect to") + case c.MirrorAPI: + return fmt.Errorf("storage.cache_artifacts: false cannot be combined with mirror_api: mirrored artifacts would never be served") + } + return nil +} + func (c *Config) validateComponents() error { if _, err := denylist.New(c.Denylist.Packages); err != nil { return err diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 050b3ce1..20bd1038 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -1303,3 +1303,61 @@ func TestValidateNamedUpstreams(t *testing.T) { }) } } + +func TestValidateCacheArtifactsDisabled(t *testing.T) { + tests := []struct { + name string + modify func(*Config) + wantErr string + }{ + {"alone", func(*Config) {}, ""}, + {"with scanning", func(c *Config) { c.Scanning.Enabled = true }, "scanning.enabled"}, + {"with direct_serve", func(c *Config) { c.Storage.DirectServe = true }, "storage.direct_serve"}, + {"with mirror_api", func(c *Config) { c.MirrorAPI = true }, "mirror_api"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := Default() + cfg.Storage.CacheArtifacts = false + tt.modify(cfg) + err := cfg.Validate() + if tt.wantErr == "" { + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + return + } + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("Validate() = %v, want an error mentioning %q", err, tt.wantErr) + } + }) + } +} + +func TestCacheArtifactsDefaultsToTrue(t *testing.T) { + if !Default().Storage.CacheArtifacts { + t.Fatal("Default().Storage.CacheArtifacts = false, want true") + } + + path := filepath.Join(t.TempDir(), "config.yaml") + if err := os.WriteFile(path, []byte("storage:\n url: \"file:///tmp/cache\"\n"), 0o600); err != nil { + t.Fatal(err) + } + cfg, err := Load(path) + if err != nil { + t.Fatalf("Load: %v", err) + } + if !cfg.Storage.CacheArtifacts { + t.Error("a config file without storage.cache_artifacts disabled artifact caching") + } +} + +func TestLoadCacheArtifactsFromEnv(t *testing.T) { + cfg := Default() + t.Setenv("PROXY_STORAGE_CACHE_ARTIFACTS", "false") + cfg.LoadFromEnv() + + if cfg.Storage.CacheArtifacts { + t.Error("Storage.CacheArtifacts should be false") + } +} diff --git a/internal/handler/handler.go b/internal/handler/handler.go index b9ecc625..37e353ca 100644 --- a/internal/handler/handler.go +++ b/internal/handler/handler.go @@ -175,6 +175,11 @@ type Proxy struct { HTTPClient *http.Client AuthForURL func(string) (headerName, headerValue string) + // StreamArtifacts streams artifacts from upstream without storing them. + // Each request fetches its own copy: there is no cache to check and + // nothing for concurrent misses to share. + StreamArtifacts bool + // Scanners runs pre-cache artifact scanning (e.g. trivy, ClamAV, Wiz). // Nil or disabled means artifacts are cached without scanning. Scanners *scanner.Group @@ -242,6 +247,17 @@ func (p *Proxy) GetOrFetchArtifact(ctx context.Context, ecosystem, name, version return nil, errors.New("resolved artifact has no filename") } } + if p.StreamArtifacts { + if info == nil { + if info, err = p.resolveArtifact(ctx, ecosystem, name, version); err != nil { + return nil, err + } + } + return p.streamFromUpstream(ctx, ecosystem, name, version, filename, versionPURL, info.URL, "", + func(fetchCtx context.Context) (*fetch.Artifact, error) { + return p.Fetcher.Fetch(fetchCtx, info.URL) + }) + } if cached, err := p.checkCache(ctx, pkgPURL, versionPURL, filename); err != nil { return nil, err } else if cached != nil { @@ -289,10 +305,15 @@ func (p *Proxy) ClearCachedArtifact(ecosystem, name, version, filename string) e } // checkCache looks up an artifact in the cache. Returns nil if not cached. +// With StreamArtifacts set it always reports a miss, so entries stored before the +// mode was enabled are never served. func (p *Proxy) checkCache(ctx context.Context, pkgPURL, versionPURL, filename string) (*CacheResult, error) { if p.Denylist.Denied(versionPURL) { return nil, fmt.Errorf("%w: %s", ErrVersionDenied, versionPURL) } + if p.StreamArtifacts { + return nil, nil + } artifact, err := p.DB.GetCachedArtifact(pkgPURL, versionPURL, filename) if err != nil { return nil, fmt.Errorf("checking artifact cache: %w", err) @@ -764,7 +785,14 @@ func serveArtifact(w http.ResponseWriter, method string, result *CacheResult) { buffer := artifactCopyBufferPool.Get().(*[]byte) defer artifactCopyBufferPool.Put(buffer) // Hide optional ReaderFrom methods so io.CopyBuffer uses the pooled buffer. - _, _ = io.CopyBuffer(struct{ io.Writer }{w}, result.Reader, *buffer) + written, err := io.CopyBuffer(struct{ io.Writer }{w}, result.Reader, *buffer) + if err != nil || (result.Artifact.Size > 0 && written != result.Artifact.Size) { + // Headers are already committed, so an error status is no longer + // possible. Aborting leaves the response unterminated and the + // client discards it instead of keeping a truncated or unverified + // artifact. + panic(http.ErrAbortHandler) + } } } @@ -1320,6 +1348,12 @@ func (p *Proxy) getOrFetchArtifactFromURLWithCachePURLs(ctx context.Context, eco if p.versionDenied(ecosystem, name, version) { return nil, fmt.Errorf("%w: %s", ErrVersionDenied, canonicalVersionPURL(ecosystem, name, version)) } + if p.StreamArtifacts { + return p.streamFromUpstream(ctx, ecosystem, name, version, filename, versionPURL, downloadURL, upstreamHash, + func(fetchCtx context.Context) (*fetch.Artifact, error) { + return p.Fetcher.FetchWithHeaders(fetchCtx, downloadURL, headers) + }) + } if cached, err := p.getCachedArtifactWithUpstreamHash(ctx, pkgPURL, versionPURL, filename, upstreamHash); err != nil { return nil, err } else if cached != nil { @@ -1402,6 +1436,88 @@ func (p *Proxy) fetchAndCacheFromURL(ctx context.Context, ecosystem, name, versi return p.storeArtifact(ctx, ecosystem, name, version, filename, pkgPURL, versionPURL, downloadURL, upstreamHash, artifact) } +// streamFromUpstream fetches an artifact when StreamArtifacts is set and hands its +// body to the caller without storing it. +// +// With upstreamHash set the body is verified as it streams. The size is left +// unknown so the response goes out chunked: a mismatch only shows at EOF, +// after the bytes have been sent, and an unterminated chunked response is the +// only way left to make the client reject them (see serveArtifact). +func (p *Proxy) streamFromUpstream(ctx context.Context, ecosystem, name, version, filename, versionPURL, upstreamURL, upstreamHash string, fetchArtifact func(context.Context) (*fetch.Artifact, error)) (*CacheResult, error) { + p.Logger.Info("streaming from upstream", + "ecosystem", ecosystem, "name", name, "version", version, "url", upstreamURL) + + fetchStart := time.Now() + artifact, err := fetchArtifact(ctx) + metrics.RecordUpstreamFetch(ecosystem, time.Since(fetchStart)) + if err != nil { + metrics.RecordUpstreamError(ecosystem, "fetch_failed") + if errors.Is(err, fetch.ErrNotFound) { + return nil, ErrUpstreamNotFound + } + return nil, fmt.Errorf("fetching from upstream: %w", err) + } + + body := &streamErrorLogger{ + ReadCloser: artifact.Body, + onError: func(read int64, err error) { + p.Logger.Warn("streaming artifact from upstream failed", + "purl", versionPURL, "filename", filename, "url", upstreamURL, "bytes", read, "error", err) + metrics.RecordUpstreamError(ecosystem, "stream_failed") + }, + } + result := &CacheResult{ + Reader: body, + Artifact: artifacts.Artifact{ + PURL: versionPURL, + Size: artifact.Size, + Filename: filename, + MediaType: artifact.ContentType, + }, + } + if upstreamHash == "" { + return result, nil + } + + hash := strings.ToLower(upstreamHash) + checks, err := newIntegrityChecks(hash, "") + if err != nil { + _ = artifact.Body.Close() + return nil, fmt.Errorf("parsing upstream digest: %w", err) + } + result.Reader, err = checks.wrapFailOnMismatch(body, func(reason string) { + p.Logger.Error("streamed artifact failed integrity check", + "purl", versionPURL, "filename", filename, "url", upstreamURL, "reason", reason) + metrics.RecordIntegrityFailure(purl.NormalizeEcosystem(ecosystem)) + }) + if err != nil { + _ = artifact.Body.Close() + return nil, err + } + result.Artifact.Digest = digest.Digest("sha256:" + hash) + result.Artifact.Size = -1 + return result, nil +} + +// streamErrorLogger reports the first read error of a streamed upstream body. +// The error itself still reaches serveArtifact, which aborts the response. +type streamErrorLogger struct { + io.ReadCloser + onError func(read int64, err error) + read int64 + logged bool +} + +func (r *streamErrorLogger) Read(p []byte) (int, error) { + n, err := r.ReadCloser.Read(p) + r.read += int64(n) + if err != nil && err != io.EOF && !r.logged { + r.logged = true + r.onError(r.read, err) + } + return n, err +} + // ErrArtifactDigestMismatch indicates that fetched bytes did not match the // checksum the upstream declared and were not recorded in the cache database. var ErrArtifactDigestMismatch = errors.New("artifact digest mismatch") diff --git a/internal/handler/helm.go b/internal/handler/helm.go index 8cacdb67..c0f86fde 100644 --- a/internal/handler/helm.go +++ b/internal/handler/helm.go @@ -115,8 +115,14 @@ func (h *HelmHandler) handleChart(w http.ResponseWriter, r *http.Request) { return } - result, err := h.proxy.GetOrFetchArtifactFromURL( - r.Context(), helmMetadataEcosystem, repository, digest, filename, downloadURL) + // A streamed fetch is never stored, so serveChart's digest check has + // nothing to compare against: verify the stream itself instead. + expectedDigest := "" + if h.proxy.StreamArtifacts { + expectedDigest = "sha256:" + digest + } + result, err := h.proxy.GetOrFetchArtifactFromURLWithDigest( + r.Context(), helmMetadataEcosystem, repository, digest, filename, downloadURL, expectedDigest) if err != nil { h.proxy.serveArtifactError(w, err, "failed to fetch chart") return diff --git a/internal/handler/integrity.go b/internal/handler/integrity.go index 07963b9f..97f41d99 100644 --- a/internal/handler/integrity.go +++ b/internal/handler/integrity.go @@ -40,6 +40,17 @@ func newIntegrityChecks(contentHash, native string) (integrityChecks, error) { } func (c integrityChecks) wrap(source io.ReadCloser, onMismatch func(string)) (io.ReadCloser, error) { + return c.newVerifyingReader(source, onMismatch, false) +} + +// wrapFailOnMismatch is wrap for bytes that have not been checked anywhere +// else: a mismatch is also returned from Read as ErrArtifactDigestMismatch in +// place of io.EOF, so the caller can abort rather than complete the response. +func (c integrityChecks) wrapFailOnMismatch(source io.ReadCloser, onMismatch func(string)) (io.ReadCloser, error) { + return c.newVerifyingReader(source, onMismatch, true) +} + +func (c integrityChecks) newVerifyingReader(source io.ReadCloser, onMismatch func(string), failOnMismatch bool) (io.ReadCloser, error) { if len(c.algorithms) == 0 { return source, nil } @@ -48,10 +59,11 @@ func (c integrityChecks) wrap(source io.ReadCloser, onMismatch func(string)) (io return nil, fmt.Errorf("create integrity reader: %w", err) } return &verifyingReader{ - source: source, - reader: reader, - checks: c, - onMismatch: onMismatch, + source: source, + reader: reader, + checks: c, + onMismatch: onMismatch, + failOnMismatch: failOnMismatch, }, nil } @@ -63,12 +75,18 @@ type verifyingReader struct { checks integrityChecks onMismatch func(reason string) verified bool + + failOnMismatch bool + mismatched bool } func (r *verifyingReader) Read(p []byte) (int, error) { n, err := r.reader.Read(p) if err == io.EOF { r.verify() + if r.failOnMismatch && r.mismatched { + return n, ErrArtifactDigestMismatch + } } return n, err } @@ -89,11 +107,13 @@ func (r *verifyingReader) verify() { if len(r.checks.contentHash) > 0 { if err := result.Verify(r.checks.contentHash); err != nil { + r.mismatched = true r.onMismatch("content_hash: " + err.Error()) } } if len(r.checks.native) > 0 { if err := result.Verify(r.checks.native); err != nil { + r.mismatched = true r.onMismatch("integrity: " + err.Error()) } } diff --git a/internal/handler/stream_artifacts_test.go b/internal/handler/stream_artifacts_test.go new file mode 100644 index 00000000..5edc5dd6 --- /dev/null +++ b/internal/handler/stream_artifacts_test.go @@ -0,0 +1,411 @@ +package handler + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/git-pkgs/artifacts" + "github.com/git-pkgs/cooldown" + "github.com/git-pkgs/proxy/internal/packageurl" + "github.com/git-pkgs/registries/fetch" +) + +func newStreamingProxy(t *testing.T, body string) (*Proxy, *mockStorage, *mockFetcher) { + t.Helper() + proxy, _, store, fetcher := setupTestProxy(t) + proxy.StreamArtifacts = true + fetcher.artifact = &fetch.Artifact{ + Body: io.NopCloser(strings.NewReader(body)), + Size: int64(len(body)), + ContentType: "application/gzip", + } + return proxy, store, fetcher +} + +func TestStreamArtifactsFromURLStreamsWithoutStoring(t *testing.T) { + proxy, store, fetcher := newStreamingProxy(t, "fetched content") + + result, err := proxy.GetOrFetchArtifactFromURL(context.Background(), "pypi", "newpkg", "1.0.0", + "newpkg-1.0.0.tar.gz", "https://pypi.org/files/newpkg-1.0.0.tar.gz") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer func() { _ = result.Reader.Close() }() + + body, err := io.ReadAll(result.Reader) + if err != nil { + t.Fatalf("reading body: %v", err) + } + if string(body) != "fetched content" { + t.Errorf("body = %q, want %q", body, "fetched content") + } + if fetcher.fetchedURL != "https://pypi.org/files/newpkg-1.0.0.tar.gz" { + t.Errorf("fetched URL = %q", fetcher.fetchedURL) + } + if result.Cached { + t.Error("streaming result reported as cached") + } + if result.Artifact.Size != int64(len("fetched content")) { + t.Errorf("Size = %d, want the upstream size", result.Artifact.Size) + } + if result.Artifact.MediaType != "application/gzip" { + t.Errorf("MediaType = %q", result.Artifact.MediaType) + } + if len(store.files) != 0 { + t.Errorf("streaming stored %d files, want none", len(store.files)) + } + cached, err := proxy.DB.GetCachedArtifact("pkg:pypi/newpkg", "pkg:pypi/newpkg@1.0.0", "newpkg-1.0.0.tar.gz") + if err != nil { + t.Fatalf("GetCachedArtifact: %v", err) + } + if cached != nil { + t.Error("streaming recorded the artifact in the cache database") + } +} + +func TestStreamArtifactsGetOrFetchArtifactStreamsWithoutStoring(t *testing.T) { + proxy, store, fetcher := newStreamingProxy(t, "tarball data") + + result, err := proxy.GetOrFetchArtifact(context.Background(), "npm", "leftpad", testVersion100, "leftpad-1.0.0.tgz") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer func() { _ = result.Reader.Close() }() + + body, _ := io.ReadAll(result.Reader) + if string(body) != "tarball data" { + t.Errorf("body = %q", body) + } + if !fetcher.fetchCalled { + t.Error("upstream was not fetched") + } + if len(store.files) != 0 { + t.Errorf("streaming stored %d files, want none", len(store.files)) + } +} + +func TestStreamArtifactsIgnoresCachedArtifacts(t *testing.T) { + proxy, _, store, fetcher := setupTestProxy(t) + fetcher.artifact = &fetch.Artifact{Body: io.NopCloser(strings.NewReader("old bytes"))} + url := "https://pypi.org/files/newpkg-1.0.0.tar.gz" + + // Populate the cache in normal mode first. + result, err := proxy.GetOrFetchArtifactFromURL(context.Background(), "pypi", "newpkg", "1.0.0", "newpkg-1.0.0.tar.gz", url) + if err != nil { + t.Fatalf("priming cache: %v", err) + } + _ = result.Reader.Close() + if len(store.files) != 1 { + t.Fatalf("expected the cache to hold the artifact, got %d files", len(store.files)) + } + + proxy.StreamArtifacts = true + cached, err := proxy.GetCachedArtifact(context.Background(), "pypi", "newpkg", "1.0.0", "newpkg-1.0.0.tar.gz") + if err != nil || cached != nil { + t.Fatalf("GetCachedArtifact = %v, %v; want nil, nil while streaming", cached, err) + } + + fetcher.fetchCalled = false + fetcher.artifact = &fetch.Artifact{Body: io.NopCloser(strings.NewReader("new bytes"))} + result, err = proxy.GetOrFetchArtifactFromURL(context.Background(), "pypi", "newpkg", "1.0.0", "newpkg-1.0.0.tar.gz", url) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer func() { _ = result.Reader.Close() }() + body, _ := io.ReadAll(result.Reader) + if !fetcher.fetchCalled || string(body) != "new bytes" { + t.Errorf("streaming served %q (fetched=%v), want a fresh upstream fetch", body, fetcher.fetchCalled) + } +} + +func TestStreamArtifactsVerifiesUpstreamDigest(t *testing.T) { + const content = "blob bytes" + + t.Run("matching digest streams the body", func(t *testing.T) { + proxy, _, _ := newStreamingProxy(t, content) + result, err := proxy.GetOrFetchArtifactFromURLWithDigest(context.Background(), "oci", "library/app", "v1", + "blob", "https://registry.test/blob", "sha256:"+sha256Hex(content)) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer func() { _ = result.Reader.Close() }() + + body, err := io.ReadAll(result.Reader) + if err != nil { + t.Fatalf("reading verified body: %v", err) + } + if string(body) != content { + t.Errorf("body = %q", body) + } + if got := result.Artifact.Digest.Encoded(); got != sha256Hex(content) { + t.Errorf("Digest = %q, want the upstream digest", got) + } + if result.Artifact.Size >= 0 { + t.Errorf("Size = %d, want unknown so the response is chunked", result.Artifact.Size) + } + }) + + t.Run("mismatched digest fails the read", func(t *testing.T) { + proxy, _, _ := newStreamingProxy(t, content) + result, err := proxy.GetOrFetchArtifactFromURLWithDigest(context.Background(), "oci", "library/app", "v1", + "blob", "https://registry.test/blob", "sha256:"+sha256Hex("other bytes")) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + defer func() { _ = result.Reader.Close() }() + + if _, err := io.ReadAll(result.Reader); !errors.Is(err, ErrArtifactDigestMismatch) { + t.Errorf("read error = %v, want ErrArtifactDigestMismatch", err) + } + }) +} + +func TestServeArtifactAbortsOnStreamedDigestMismatch(t *testing.T) { + proxy, _, _ := newStreamingProxy(t, "blob bytes") + result, err := proxy.GetOrFetchArtifactFromURLWithDigest(context.Background(), "oci", "library/app", "v1", + "blob", "https://registry.test/blob", "sha256:"+sha256Hex("other bytes")) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + rec := httptest.NewRecorder() + defer func() { + if r := recover(); r != http.ErrAbortHandler { + t.Fatalf("recovered %v, want http.ErrAbortHandler", r) + } + if rec.Header().Get(headerContentLength) != "" { + t.Errorf("Content-Length = %q, want none so the client sees an unterminated response", + rec.Header().Get(headerContentLength)) + } + }() + ServeArtifact(rec, result) +} + +func TestNPMDownloadCooldownWhileStreaming(t *testing.T) { + now := time.Now() + packument := `{ + "name": "leftpad", + "dist-tags": {"latest": "2.0.0"}, + "time": { + "1.0.0": "` + now.Add(-30*24*time.Hour).Format(time.RFC3339) + `", + "2.0.0": "` + now.Add(-1*time.Hour).Format(time.RFC3339) + `" + }, + "versions": {"1.0.0": {}, "2.0.0": {}} + }` + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", contentTypeJSON) + _, _ = io.WriteString(w, packument) + })) + defer upstream.Close() + + tests := []struct { + name string + version string + wantStatus int + }{ + {"published before the window streams the tarball", testVersion100, http.StatusOK}, + {"published inside the window is withheld", "2.0.0", http.StatusNotFound}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + proxy, store, fetcher := newStreamingProxy(t, "tarball data") + proxy.HTTPClient = upstream.Client() + proxy.Cooldown = &cooldown.Config{Default: "7d"} + + srv := httptest.NewServer(NewNPMHandler(proxy, "http://proxy.test", upstream.URL).Routes()) + defer srv.Close() + + resp, err := http.Get(srv.URL + "/leftpad/-/leftpad-" + tt.version + ".tgz") + if err != nil { + t.Fatalf("request failed: %v", err) + } + body, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + + if resp.StatusCode != tt.wantStatus { + t.Fatalf("status = %d, want %d", resp.StatusCode, tt.wantStatus) + } + if tt.wantStatus == http.StatusOK && string(body) != "tarball data" { + t.Errorf("body = %q", body) + } + if tt.wantStatus == http.StatusNotFound && fetcher.fetchCalled { + t.Error("fetched a version that is still inside the cooldown window") + } + if len(store.files) != 0 { + t.Errorf("streaming stored %d files, want none", len(store.files)) + } + }) + } +} + +// truncatingUpstream answers every request matching match with raw, then +// closes the connection, so the response body ends early. +func truncatingUpstream(t *testing.T, match func(*http.Request) bool, raw string, fallback http.HandlerFunc) *httptest.Server { + t.Helper() + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !match(r) { + fallback(w, r) + return + } + conn, buf, err := w.(http.Hijacker).Hijack() + if err != nil { + t.Errorf("hijack: %v", err) + return + } + _, _ = buf.WriteString(raw) + _ = buf.Flush() + _ = conn.Close() + })) + t.Cleanup(upstream.Close) + return upstream +} + +// requireIncompleteResponse fails unless reading the response surfaces an error, +// i.e. the client cannot mistake the body for a complete download. +func requireIncompleteResponse(t *testing.T, url string) { + t.Helper() + resp, err := http.Get(url) + if err != nil { + return + } + defer func() { _ = resp.Body.Close() }() + body, err := io.ReadAll(resp.Body) + if err == nil { + t.Fatalf("client read a complete %d response (Content-Length %q, body %q), want an incomplete one", + resp.StatusCode, resp.Header.Get(headerContentLength), body) + } +} + +func useRealFetcher(t *testing.T, proxy *Proxy, upstream *httptest.Server) { + t.Helper() + proxy.HTTPClient = upstream.Client() + fetcher := fetch.NewFetcher(fetch.WithHTTPClient(upstream.Client()), fetch.WithMaxRetries(0)) + proxy.Fetcher = fetcher + t.Cleanup(func() { _ = fetcher.Close() }) +} + +func TestStreamArtifactsOCIBlobShorterThanContentLengthIsIncomplete(t *testing.T) { + digest := "sha256:" + sha256Hex(strings.Repeat("x", 100)) + upstream := truncatingUpstream(t, + func(r *http.Request) bool { return strings.Contains(r.URL.Path, "/blobs/") }, + "HTTP/1.1 200 OK\r\nContent-Type: application/octet-stream\r\nContent-Length: 100\r\n\r\nshort", + http.NotFound) + + proxy, _, store, _ := setupTestProxy(t) + proxy.StreamArtifacts = true + useRealFetcher(t, proxy, upstream) + h := NewContainerHandler(proxy, "http://proxy.example", map[string]string{"ghcr": upstream.URL}) + srv := httptest.NewServer(h.Routes()) + defer srv.Close() + + requireIncompleteResponse(t, srv.URL+"/upstream/ghcr/owner/demo/blobs/"+digest) + if len(store.files) != 0 { + t.Errorf("streaming stored %d files, want none", len(store.files)) + } +} + +func TestStreamArtifactsNPMTarballWithoutFinalChunkIsIncomplete(t *testing.T) { + upstream := truncatingUpstream(t, + func(r *http.Request) bool { return strings.HasSuffix(r.URL.Path, ".tgz") }, + "HTTP/1.1 200 OK\r\nContent-Type: application/octet-stream\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nshort\r\n", + http.NotFound) + + proxy, _, store, _ := setupTestProxy(t) + proxy.StreamArtifacts = true + useRealFetcher(t, proxy, upstream) + srv := httptest.NewServer(NewNPMHandler(proxy, "http://proxy.test", upstream.URL).Routes()) + defer srv.Close() + + requireIncompleteResponse(t, srv.URL+"/leftpad/-/leftpad-1.0.0.tgz") + if len(store.files) != 0 { + t.Errorf("streaming stored %d files, want none", len(store.files)) + } +} + +func TestServeArtifactAbortsOnShortRead(t *testing.T) { + tests := []struct { + name string + reader io.Reader + size int64 + }{ + {"read error", io.MultiReader(strings.NewReader("short"), iotestErrReader{io.ErrUnexpectedEOF}), -1}, + {"fewer bytes than the declared size", strings.NewReader("short"), 100}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := &CacheResult{Reader: io.NopCloser(tt.reader), Artifact: artifacts.Artifact{Size: tt.size}} + defer func() { + if r := recover(); r != http.ErrAbortHandler { + t.Fatalf("recovered %v, want http.ErrAbortHandler", r) + } + }() + ServeArtifact(httptest.NewRecorder(), result) + }) + } +} + +type iotestErrReader struct{ err error } + +func (r iotestErrReader) Read([]byte) (int, error) { return 0, r.err } + +func TestStreamArtifactsSwiftArchiveHeadIgnoresCachedEntry(t *testing.T) { + archive := []byte("cached archive") + checksum := sha256.Sum256(archive) + var archiveRequests int + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if strings.HasSuffix(r.URL.Path, ".zip") { + archiveRequests++ + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintf(w, `{"id":"apple.example","version":"1.2.3","resources":[{"name":"source-archive","type":"application/zip","checksum":%q}]}`, hex.EncodeToString(checksum[:])) + })) + defer upstream.Close() + + proxy, _, _, fetcher := setupTestProxy(t) + proxy.HTTPClient = upstream.Client() + fetcher.artifact = &fetch.Artifact{ + Body: io.NopCloser(strings.NewReader(string(archive))), + Size: int64(len(archive)), + ContentType: "application/zip", + } + packagePURL, versionPURL := packageurl.MakeCacheStrings("swift", "apple/example", "1.2.3") + cached, err := proxy.getOrFetchArtifactFromURLWithCachePURLs( + context.Background(), "swift", "apple/example", "1.2.3", "example-1.2.3.zip", + packagePURL, versionPURL, upstream.URL+"/apple/example/1.2.3.zip", nil, hex.EncodeToString(checksum[:]), + ) + if err != nil { + t.Fatalf("seeding cache in normal mode: %v", err) + } + _ = cached.Reader.Close() + + proxy.StreamArtifacts = true + fetcher.artifact = nil + fetcher.fetchErr = fetch.ErrNotFound + handler := NewSwiftHandler(proxy, "https://proxy.example", upstream.URL).Routes() + + w := httptest.NewRecorder() + handler.ServeHTTP(w, httptest.NewRequest(http.MethodHead, "/apple/example/1.2.3.zip", nil)) + if w.Code == http.StatusOK { + t.Fatalf("HEAD served the cached archive (Content-Length %q) while streaming artifacts", w.Header().Get(headerContentLength)) + } + if archiveRequests == 0 { + t.Error("HEAD did not ask the upstream archive") + } + + w = httptest.NewRecorder() + handler.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/apple/example/1.2.3.zip", nil)) + if w.Code == http.StatusOK { + t.Fatalf("GET served the cached archive while streaming artifacts") + } +} diff --git a/internal/server/server.go b/internal/server/server.go index 35066486..da75aa02 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -257,6 +257,7 @@ func (s *Server) serve(listener net.Listener) error { proxy.DirectServe = s.cfg.Storage.DirectServe proxy.DirectServeTTL = s.cfg.ParseDirectServeTTL() proxy.DirectServeBaseURL = s.cfg.Storage.DirectServeBaseURL + proxy.StreamArtifacts = !s.cfg.Storage.CacheArtifacts // Create router with Chi r := chi.NewRouter()