diff --git a/app/artifact-cas/internal/server/http.go b/app/artifact-cas/internal/server/http.go index 634babe5b..d38873340 100644 --- a/app/artifact-cas/internal/server/http.go +++ b/app/artifact-cas/internal/server/http.go @@ -56,7 +56,7 @@ func NewHTTPServer(c *conf.Server, authConf *conf.Auth, downloadSvc *service.Dow srv := http.NewServer(opts...) - downloadHandler := middlewares_http.AuthFromQueryParam(loadPublicKey(publicKey), claimsFunc(), casJWT.SigningMethod, downloadSvc) + downloadHandler := middlewares_http.AuthFromHeaderOrQueryParam(loadPublicKey(publicKey), claimsFunc(), casJWT.SigningMethod, downloadSvc) srv.Handle(service.DownloadPath, CORSMiddleware(c.GetHttp().GetCors().GetAllowOrigins(), downloadHandler)) api.RegisterStatusServiceHTTPServer(srv, service.NewStatusService(Version, providers)) return srv, nil diff --git a/app/artifact-cas/internal/service/download.go b/app/artifact-cas/internal/service/download.go index 5264b4e4f..85b9bcf24 100644 --- a/app/artifact-cas/internal/service/download.go +++ b/app/artifact-cas/internal/service/download.go @@ -21,6 +21,7 @@ import ( "io" "net/http" "strconv" + "strings" "code.cloudfoundry.org/bytefmt" "github.com/chainloop-dev/chainloop/app/controlplane/pkg/auditor/events" @@ -94,6 +95,16 @@ func (s *DownloadService) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } + // The digest identifies the bytes, so a browser copy tagged with it never + // goes stale. The token and the object are already checked, so a browser + // that holds the content gets a 304 with no copy from the backend. + etag := strconv.Quote(wantChecksum.String()) + if etagMatches(r.Header.Get("If-None-Match"), etag) { + setCacheHeaders(w, etag) + w.WriteHeader(http.StatusNotModified) + return + } + // Override file nane if one is provided filename := r.URL.Query().Get("filename") if filename == "" { @@ -115,6 +126,7 @@ func (s *DownloadService) ServeHTTP(w http.ResponseWriter, r *http.Request) { // The content is verified: announce it to the browser with its exact size w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=%s", filename)) w.Header().Set("Content-Length", strconv.FormatInt(size, 10)) + setCacheHeaders(w, etag) // A plain io.Copy lets the response writer pull the file with sendfile, so // the verified bytes go kernel-to-kernel without a user-space buffer. @@ -134,6 +146,26 @@ func (s *DownloadService) ServeHTTP(w http.ResponseWriter, r *http.Request) { }, auth) } +// setCacheHeaders lets the browser keep a private copy of the content that it +// must revalidate with the CAS before each use, so access is checked every time. +func setCacheHeaders(w http.ResponseWriter, etag string) { + w.Header().Set("ETag", etag) + w.Header().Set("Cache-Control", "private, no-cache") +} + +// etagMatches reports whether an If-None-Match header value matches the etag. +// It uses the weak comparison that RFC 9110 requires for If-None-Match. +func etagMatches(ifNoneMatch, etag string) bool { + for candidate := range strings.SplitSeq(ifNoneMatch, ",") { + candidate = strings.TrimSpace(candidate) + if candidate == "*" || strings.TrimPrefix(candidate, "W/") == etag { + return true + } + } + + return false +} + // writeDownloadError maps a staging or streaming failure of a download to its // HTTP response. A mismatch means the backend holds content that does not hash // to its key, corrupt or tampered, so it is a 500 carrying both digests rather diff --git a/app/artifact-cas/internal/service/download_test.go b/app/artifact-cas/internal/service/download_test.go index 49324c736..a51e0fde6 100644 --- a/app/artifact-cas/internal/service/download_test.go +++ b/app/artifact-cas/internal/service/download_test.go @@ -35,17 +35,24 @@ import ( "github.com/stretchr/testify/require" ) +const ( + downloadBackendType = "backend-type" + downloadContent = "hello world" + // sha256 of downloadContent + downloadDigestHex = "b94d27b9934d3e08a52e52d7da7dabfac484efe37a5380ee9088f7ace2efcde9" + downloadFileName = "test.txt" +) + func TestDownloadServiceAuditEvents(t *testing.T) { const ( - backendType = "backend-type" - // sha256 of "hello world" - digestHex = "b94d27b9934d3e08a52e52d7da7dabfac484efe37a5380ee9088f7ace2efcde9" + backendType = downloadBackendType + digestHex = downloadDigestHex ) downloaderClaims := func(sourceInternal bool) *casJWT.Claims { return &casJWT.Claims{ Role: casJWT.Downloader, - StoredSecretID: "secret-id", + StoredSecretID: testStoredSecretID, BackendType: backendType, OrgID: testOrgID, SourceInternal: sourceInternal, @@ -65,7 +72,7 @@ func TestDownloadServiceAuditEvents(t *testing.T) { }{ { name: "successful download emits an event", - content: "hello world", + content: downloadContent, claims: downloaderClaims(false), wantStatus: http.StatusOK, wantEvents: 1, @@ -79,7 +86,7 @@ func TestDownloadServiceAuditEvents(t *testing.T) { }, { name: "staging failure is masked, sends no bytes and emits no event", - content: "hello world", + content: downloadContent, claims: downloaderClaims(false), removeStagingDir: true, wantStatus: http.StatusInternalServerError, @@ -87,7 +94,7 @@ func TestDownloadServiceAuditEvents(t *testing.T) { }, { name: "internal control plane traffic emits no event", - content: "hello world", + content: downloadContent, claims: downloaderClaims(true), wantStatus: http.StatusOK, }, @@ -99,7 +106,7 @@ func TestDownloadServiceAuditEvents(t *testing.T) { uploaderDownloader := mocks.NewUploaderDownloader(t) provider.On("FromCredentials", mock.Anything, mock.Anything).Return(uploaderDownloader, nil) uploaderDownloader.On("Describe", mock.Anything, digestHex).Return(&v1.CASResource{ - FileName: "test.txt", Digest: digestHex, Size: int64(len(tc.content)), + FileName: downloadFileName, Digest: digestHex, Size: int64(len(tc.content)), }, nil) if !tc.removeStagingDir { uploaderDownloader.On("Download", mock.Anything, mock.Anything, digestHex).Return(nil). @@ -134,7 +141,10 @@ func TestDownloadServiceAuditEvents(t *testing.T) { assert.Equal(t, tc.content, w.Body.String()) assert.Equal(t, strconv.Itoa(len(tc.content)), w.Header().Get("Content-Length")) assert.Equal(t, "attachment; filename=test.txt", w.Header().Get("Content-Disposition")) + assert.Equal(t, `"sha256:`+digestHex+`"`, w.Header().Get("ETag")) + assert.Equal(t, "private, no-cache", w.Header().Get("Cache-Control")) } else { + assert.Empty(t, w.Header().Get("ETag"), "a failed download must not be cached") assert.NotContains(t, w.Body.String(), tc.content, "no unverified byte may reach the client") assert.Contains(t, w.Body.String(), tc.wantBodyContains) } @@ -152,9 +162,139 @@ func TestDownloadServiceAuditEvents(t *testing.T) { info := decodeArtifactEvent(t, audit.published[0]) assert.Equal(t, digestHex, info.Digest) assert.Equal(t, int64(len(tc.content)), info.SizeBytes) - assert.Equal(t, "test.txt", info.FileName) + assert.Equal(t, downloadFileName, info.FileName) assert.Equal(t, backendType, info.BackendType) assert.False(t, info.Skipped) }) } } + +func TestDownloadServiceConditionalRequests(t *testing.T) { + const ( + backendType = downloadBackendType + content = downloadContent + digestHex = downloadDigestHex + etag = `"sha256:` + digestHex + `"` + ) + + claims := &casJWT.Claims{ + Role: casJWT.Downloader, + StoredSecretID: testStoredSecretID, + BackendType: backendType, + OrgID: testOrgID, + } + + tests := []struct { + name string + ifNoneMatch string + claims *casJWT.Claims + // describeErr is returned by the backend metadata call + describeErr error + // wantDownload means the backend object is copied + wantDownload bool + wantStatus int + wantEvents int + }{ + { + name: "matching etag answers not modified without copying the object", + ifNoneMatch: etag, + claims: claims, + wantStatus: http.StatusNotModified, + }, + { + name: "weak etag in a list matches", + ifNoneMatch: `"sha256:other", W/` + etag, + claims: claims, + wantStatus: http.StatusNotModified, + }, + { + name: "wildcard matches", + ifNoneMatch: "*", + claims: claims, + wantStatus: http.StatusNotModified, + }, + { + name: "unquoted digest does not match", + ifNoneMatch: "sha256:" + digestHex, + claims: claims, + wantDownload: true, + wantStatus: http.StatusOK, + wantEvents: 1, + }, + { + name: "different etag downloads the object", + ifNoneMatch: `"sha256:other"`, + claims: claims, + wantDownload: true, + wantStatus: http.StatusOK, + wantEvents: 1, + }, + { + name: "matching etag without a token is unauthorized", + ifNoneMatch: etag, + wantStatus: http.StatusUnauthorized, + }, + { + name: "matching etag for an object missing in the backend is not found", + ifNoneMatch: etag, + claims: claims, + describeErr: backend.NewErrNotFound("artifact"), + wantStatus: http.StatusNotFound, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + provider := mocks.NewProvider(t) + if tc.claims != nil { + uploaderDownloader := mocks.NewUploaderDownloader(t) + provider.On("FromCredentials", mock.Anything, mock.Anything).Return(uploaderDownloader, nil) + var resource *v1.CASResource + if tc.describeErr == nil { + resource = &v1.CASResource{FileName: downloadFileName, Digest: digestHex, Size: int64(len(content))} + } + uploaderDownloader.On("Describe", mock.Anything, digestHex).Return(resource, tc.describeErr) + // the mock fails the test on any Download call it does not expect + if tc.wantDownload { + uploaderDownloader.On("Download", mock.Anything, mock.Anything, digestHex).Return(nil). + Run(func(args mock.Arguments) { + _, err := io.WriteString(args.Get(1).(io.Writer), content) + require.NoError(t, err) + }) + } + } + + audit := &fakePublisher{} + svc := NewDownloadService( + backend.Providers{backendType: provider}, + WithLogger(log.DefaultLogger), + WithAuditDispatcher(newTestDispatcher(audit)), + WithStagingDir(t.TempDir()), + ) + + req := httptest.NewRequest(http.MethodGet, "/download/sha256:"+digestHex, nil) + req = mux.SetURLVars(req, map[string]string{"digest": "sha256:" + digestHex}) + req.Header.Set("If-None-Match", tc.ifNoneMatch) + if tc.claims != nil { + req = req.WithContext(jwtMiddleware.NewContext(req.Context(), tc.claims)) + } + + w := httptest.NewRecorder() + svc.ServeHTTP(w, req) + + assert.Equal(t, tc.wantStatus, w.Code) + switch tc.wantStatus { + case http.StatusNotModified: + assert.Empty(t, w.Body.String()) + assert.Equal(t, etag, w.Header().Get("ETag")) + assert.Equal(t, "private, no-cache", w.Header().Get("Cache-Control")) + case http.StatusOK: + assert.Equal(t, content, w.Body.String()) + assert.Equal(t, etag, w.Header().Get("ETag")) + default: + assert.Empty(t, w.Header().Get("ETag")) + } + assert.Len(t, audit.published, tc.wantEvents) + }) + } +} diff --git a/pkg/middlewares/http/jwt.go b/pkg/middlewares/http/jwt.go index 3ec09ee1b..fb240cf45 100644 --- a/pkg/middlewares/http/jwt.go +++ b/pkg/middlewares/http/jwt.go @@ -35,9 +35,23 @@ const ( // ClaimsFunc is a function that returns a jwt.Claims with the custom claims and correct type type ClaimsFunc func() jwt.Claims -// AuthFromQueryParam is a middleware that extracts the token from the query parameter and verifies it -func AuthFromQueryParam(keyFunc jwt.Keyfunc, claimsFunc ClaimsFunc, signingMethod jwt.SigningMethod, next nhttp.Handler) nhttp.Handler { +// AuthFromHeaderOrQueryParam is a middleware that extracts the token from the +// authorization header or, when that header is absent, from the "t" query +// parameter, and verifies it. A present header always wins: a bad header is +// rejected and never falls back to the query token. +func AuthFromHeaderOrQueryParam(keyFunc jwt.Keyfunc, claimsFunc ClaimsFunc, signingMethod jwt.SigningMethod, next nhttp.Handler) nhttp.Handler { return nhttp.HandlerFunc(func(w http.ResponseWriter, r *nhttp.Request) { + if r.Header.Get(authorizationKey) != "" { + token, ok := bearerToken(r) + if !ok { + nhttp.Error(w, "invalid authorization header", nhttp.StatusUnauthorized) + return + } + + verifyJWTAndServeNext(w, r, token, keyFunc, claimsFunc, signingMethod, next) + return + } + token := r.URL.Query().Get("t") if token == "" { nhttp.Error(w, "missing token", nhttp.StatusUnauthorized) @@ -68,18 +82,26 @@ func verifyJWTAndServeNext(w http.ResponseWriter, r *nhttp.Request, token string // AuthFromAuthorizationHeader is a middleware that extracts the token from the authorization header and verifies it func AuthFromAuthorizationHeader(keyFunc jwt.Keyfunc, claimsFunc ClaimsFunc, signingMethod jwt.SigningMethod, next nhttp.Handler) nhttp.Handler { return nhttp.HandlerFunc(func(w http.ResponseWriter, r *nhttp.Request) { - auths := strings.SplitN(r.Header.Get(authorizationKey), " ", 2) - if len(auths) != 2 || !strings.EqualFold(auths[0], bearerWord) { + jwtToken, ok := bearerToken(r) + if !ok { nhttp.Error(w, "JWT token is missing", nhttp.StatusUnauthorized) return } - jwtToken := auths[1] - verifyJWTAndServeNext(w, r, jwtToken, keyFunc, claimsFunc, signingMethod, next) }) } +// bearerToken returns the token of a "Bearer " authorization header +func bearerToken(r *nhttp.Request) (string, bool) { + auths := strings.SplitN(r.Header.Get(authorizationKey), " ", 2) + if len(auths) != 2 || !strings.EqualFold(auths[0], bearerWord) { + return "", false + } + + return auths[1], true +} + // verifyAndMarshalJWT verifies the token and returns the map claims func verifyAndMarshalJWT(token string, keyFunc jwt.Keyfunc, claimsFunc ClaimsFunc, signingMethod jwt.SigningMethod) (*jwt.Claims, error) { var tokenInfo *jwt.Token diff --git a/pkg/middlewares/http/jwt_test.go b/pkg/middlewares/http/jwt_test.go index 0dd1e3734..bc73a7f94 100644 --- a/pkg/middlewares/http/jwt_test.go +++ b/pkg/middlewares/http/jwt_test.go @@ -46,25 +46,36 @@ func genericClaimsFunc() ClaimsFunc { } } -func TestAuthFromQueryParam(t *testing.T) { +func TestAuthFromHeaderOrQueryParam(t *testing.T) { validToken := generateValidToken() + wantClaims := map[string]interface{}{"foo": "bar"} tests := []struct { name string - token string + header string + query string wantStatus int wantClaims map[string]interface{} }{ - {"Valid Token", validToken, http.StatusOK, map[string]interface{}{"foo": "bar"}}, - {"Missing Token", "", http.StatusUnauthorized, nil}, - {"Invalid Token", "invalidtoken", http.StatusUnauthorized, nil}, + {"valid query token", "", validToken, http.StatusOK, wantClaims}, + {"valid header token", bearerWord + " " + validToken, "", http.StatusOK, wantClaims}, + {"lowercase bearer keyword", "bearer " + validToken, "", http.StatusOK, wantClaims}, + {"header wins over an invalid query token", bearerWord + " " + validToken, "invalidtoken", http.StatusOK, wantClaims}, + {"invalid header token does not fall back to the query", bearerWord + " invalidtoken", validToken, http.StatusUnauthorized, nil}, + {"malformed header does not fall back to the query", bearerWord, validToken, http.StatusUnauthorized, nil}, + {"non bearer header does not fall back to the query", "Basic " + validToken, validToken, http.StatusUnauthorized, nil}, + {"missing token", "", "", http.StatusUnauthorized, nil}, + {"invalid query token", "", "invalidtoken", http.StatusUnauthorized, nil}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - req, _ := http.NewRequest("GET", "/?t="+tt.token, nil) + req, _ := http.NewRequest("GET", "/?t="+tt.query, nil) + if tt.header != "" { + req.Header.Set(authorizationKey, tt.header) + } rr := httptest.NewRecorder() - handler := AuthFromQueryParam(mockKeyFunc, genericClaimsFunc(), mockSigningMethod, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + handler := AuthFromHeaderOrQueryParam(mockKeyFunc, genericClaimsFunc(), mockSigningMethod, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { claims, ok := jwtmiddleware.FromContext(r.Context()) mapClaims := claims.(*jwt.MapClaims) if tt.wantClaims != nil {