Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion app/artifact-cas/internal/server/http.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
32 changes: 32 additions & 0 deletions app/artifact-cas/internal/service/download.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ import (
"io"
"net/http"
"strconv"
"strings"

"code.cloudfoundry.org/bytefmt"
"github.com/chainloop-dev/chainloop/app/controlplane/pkg/auditor/events"
Expand Down Expand Up @@ -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 == "" {
Expand All @@ -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.
Expand All @@ -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
Expand Down
158 changes: 149 additions & 9 deletions app/artifact-cas/internal/service/download_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -79,15 +86,15 @@ 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,
wantBodyContains: "server error",
},
{
name: "internal control plane traffic emits no event",
content: "hello world",
content: downloadContent,
claims: downloaderClaims(true),
wantStatus: http.StatusOK,
},
Expand All @@ -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).
Expand Down Expand Up @@ -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)
}
Expand All @@ -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)
})
}
}
34 changes: 28 additions & 6 deletions pkg/middlewares/http/jwt.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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 <token>" 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
Expand Down
25 changes: 18 additions & 7 deletions pkg/middlewares/http/jwt_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
Loading