diff --git a/internal/strategy/codeartifact.go b/internal/strategy/codeartifact.go index 8218811..78da6b8 100644 --- a/internal/strategy/codeartifact.go +++ b/internal/strategy/codeartifact.go @@ -303,7 +303,7 @@ func (c *CodeArtifact) do(r *http.Request, token string) (*http.Response, error) } copyHeaders(upstream.Header, codeArtifactRequestHeaders(r.Header)) if shouldRewriteCodeArtifactMetadata(target.Path) { - stripCodeArtifactMetadataRequestHeaders(upstream.Header) + normalizeCodeArtifactMetadataRequestHeaders(upstream.Header, target.Path) } setCodeArtifactAuthorization(upstream, token) @@ -389,7 +389,7 @@ func (c *CodeArtifact) followCrossOriginRedirect( } copyHeaders(redirected.Header, codeArtifactRequestHeaders(r.Header)) if rewriteMetadata { - stripCodeArtifactMetadataRequestHeaders(redirected.Header) + normalizeCodeArtifactMetadataRequestHeaders(redirected.Header, c.originURL(r).Path) } resp, err := c.client.Do(redirected) diff --git a/internal/strategy/codeartifact_metadata.go b/internal/strategy/codeartifact_metadata.go index 157cfc2..bf16347 100644 --- a/internal/strategy/codeartifact_metadata.go +++ b/internal/strategy/codeartifact_metadata.go @@ -41,7 +41,7 @@ func codeArtifactPackageFormat(path string) string { return parts[0] } -func stripCodeArtifactMetadataRequestHeaders(headers http.Header) { +func normalizeCodeArtifactMetadataRequestHeaders(headers http.Header, path string) { for _, name := range []string{ "Accept-Encoding", "If-Match", @@ -53,6 +53,9 @@ func stripCodeArtifactMetadataRequestHeaders(headers http.Header) { } { headers.Del(name) } + if codeArtifactPackageFormat(path) == codeArtifactCargoFormat { + headers.Set("Accept", "application/json") + } } func isCodeArtifactJSONResponse(path string, headers http.Header) bool { diff --git a/internal/strategy/codeartifact_metadata_test.go b/internal/strategy/codeartifact_metadata_test.go index db36269..7ce3e40 100644 --- a/internal/strategy/codeartifact_metadata_test.go +++ b/internal/strategy/codeartifact_metadata_test.go @@ -15,11 +15,13 @@ import ( func TestCodeArtifactRewritesPackageMetadata(t *testing.T) { tests := []struct { - name string - path string - contentType string - body func(string) string - want func(string) string + name string + path string + contentType string + requestAccept string + wantOriginAccept string + body func(string) string + want func(string) string }{ { name: "npm tarball URL", @@ -32,9 +34,11 @@ func TestCodeArtifactRewritesPackageMetadata(t *testing.T) { }, }, { - name: "Cargo download template and anonymous access", - path: "/cargo/repository/config.json", - contentType: "application/octet-stream", + name: "Cargo download template and anonymous access", + path: "/cargo/repository/config.json", + contentType: "application/octet-stream", + requestAccept: "text/plain", + wantOriginAccept: "application/json", body: func(origin string) string { return `{"dl":"` + origin + `/cargo/repository/crates/{crate}/{version}","api":"` + origin + `/cargo/repository/-","auth-required":true}` }, @@ -91,6 +95,9 @@ func TestCodeArtifactRewritesPackageMetadata(t *testing.T) { originURL = originServer.URL proxyURL := "https://cachew.example.com/" + originServer.Listener.Addr().String() req := httptest.NewRequest(http.MethodGet, codeArtifactPath(originServer, test.path), nil).WithContext(ctx) + if test.requestAccept != "" { + req.Header.Set("Accept", test.requestAccept) + } req.Header.Set("Accept-Encoding", "gzip") req.Header.Set("If-None-Match", `"old-metadata"`) req.Header.Set("Range", "bytes=0-10") @@ -107,6 +114,7 @@ func TestCodeArtifactRewritesPackageMetadata(t *testing.T) { mu.Lock() headers := observedHeaders.Clone() mu.Unlock() + assert.Equal(t, test.wantOriginAccept, headers.Get("Accept")) assert.Equal(t, "", headers.Get("Accept-Encoding")) assert.Equal(t, "", headers.Get("If-None-Match")) assert.Equal(t, "", headers.Get("Range")) @@ -197,6 +205,7 @@ func TestCodeArtifactRejectsUnsafePackageMetadata(t *testing.T) { func TestCodeArtifactRewritesRedirectedCargoMetadata(t *testing.T) { type observedHeaders struct { + accept string acceptEncoding string ifNoneMatch string rangeHeader string @@ -205,6 +214,7 @@ func TestCodeArtifactRewritesRedirectedCargoMetadata(t *testing.T) { var originURL string download := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { observed = observedHeaders{ + accept: r.Header.Get("Accept"), acceptEncoding: r.Header.Get("Accept-Encoding"), ifNoneMatch: r.Header.Get("If-None-Match"), rangeHeader: r.Header.Get("Range"), @@ -228,6 +238,7 @@ func TestCodeArtifactRewritesRedirectedCargoMetadata(t *testing.T) { transport.DisableCompression = true strategy.client.Transport = transport req := httptest.NewRequest(http.MethodGet, codeArtifactPath(originServer, "/cargo/repository/config.json"), nil).WithContext(ctx) + req.Header.Set("Accept", "text/plain") req.Header.Set("Accept-Encoding", "gzip") req.Header.Set("If-None-Match", `"metadata"`) req.Header.Set("Range", "bytes=0-10") @@ -238,7 +249,7 @@ func TestCodeArtifactRewritesRedirectedCargoMetadata(t *testing.T) { assert.Equal(t, http.StatusOK, w.Code) proxyURL := "https://cachew.example.com/" + originServer.Listener.Addr().String() assertJSONEqual(t, `{"dl":"`+proxyURL+`/cargo/repository/crates/{crate}/{version}","auth-required":false}`, w.Body.String()) - assert.Equal(t, observedHeaders{}, observed) + assert.Equal(t, observedHeaders{accept: "application/json"}, observed) } func TestCodeArtifactValidatesProxyBaseURL(t *testing.T) {