Skip to content
Closed
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
77 changes: 68 additions & 9 deletions pkg/api/cache.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"bufio"
"bytes"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
Expand Down Expand Up @@ -143,41 +144,75 @@ func (fs *fileStorage) filePath(key string) string {

func (fs *fileStorage) read(key string) (*http.Response, error) {
cacheFile := fs.filePath(key)
checksumFile := cacheFile + ".sha256"

fs.mu.RLock()
defer fs.mu.RUnlock()

f, err := os.Open(cacheFile)
if err != nil {
fs.mu.RUnlock()
return nil, err
}
defer f.Close()

stat, err := f.Stat()
if err != nil {
_ = f.Close()
fs.mu.RUnlock()
return nil, err
}

age := time.Since(stat.ModTime())
if age > fs.ttl {
_ = f.Close()
fs.mu.RUnlock()
return nil, errors.New("cache expired")
}

body := &bytes.Buffer{}
_, err = io.Copy(body, f)
contents, err := io.ReadAll(f)
_ = f.Close()
if err != nil {
fs.mu.RUnlock()
return nil, err
}

res, err := http.ReadResponse(bufio.NewReader(body), nil)
return res, err
checksum, checksumErr := os.ReadFile(checksumFile)
expected := sha256.Sum256(contents)
validChecksum := checksumErr == nil && bytes.Equal(bytes.TrimSpace(checksum), []byte(hex.EncodeToString(expected[:])))
res, parseErr := http.ReadResponse(bufio.NewReader(bytes.NewReader(contents)), nil)
fs.mu.RUnlock()

if !validChecksum || parseErr != nil {
_ = fs.evict(key)
if checksumErr != nil {
return nil, errors.New("cache integrity check failed")
}
if parseErr != nil {
return nil, parseErr
}
return nil, errors.New("cache integrity check failed")
}

return res, nil
}

func (fs *fileStorage) evict(key string) error {
fs.mu.Lock()
defer fs.mu.Unlock()

var firstErr error
for _, path := range []string{fs.filePath(key), fs.filePath(key) + ".sha256"} {
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) && firstErr == nil {
firstErr = err
}
}
return firstErr
}

func (fs *fileStorage) store(key string, res *http.Response) error {
fs.mu.Lock()
defer fs.mu.Unlock()

cacheFilePath := fs.filePath(key)
checksumFilePath := cacheFilePath + ".sha256"
dir := filepath.Dir(cacheFilePath)
if err := os.MkdirAll(dir, 0755); err != nil {
return err
Expand All @@ -197,14 +232,38 @@ func (fs *fileStorage) store(key string, res *http.Response) error {
_ = os.Remove(tmpCacheFileName)
}()

if err := writeCacheResponse(tmpCacheFile, res); err != nil {
hash := sha256.New()
if err := writeCacheResponse(io.MultiWriter(tmpCacheFile, hash), res); err != nil {
return err
}
if err := tmpCacheFile.Close(); err != nil {
return err
}

return renameCacheFile(tmpCacheFileName, cacheFilePath)
tmpChecksumFile, err := os.CreateTemp(dir, ".gh-cache-checksum-*")
if err != nil {
return err
}
tmpChecksumFileName := tmpChecksumFile.Name()
defer func() {
_ = tmpChecksumFile.Close()
_ = os.Remove(tmpChecksumFileName)
}()
if _, err := fmt.Fprintf(tmpChecksumFile, "%x\n", hash.Sum(nil)); err != nil {
return err
}
if err := tmpChecksumFile.Close(); err != nil {
return err
}

if err := renameCacheFile(tmpChecksumFileName, checksumFilePath); err != nil {
return err
}
if err := renameCacheFile(tmpCacheFileName, cacheFilePath); err != nil {
_ = os.Remove(checksumFilePath)
return err
}
return nil
}

func writeCacheResponse(w io.Writer, res *http.Response) error {
Expand Down
57 changes: 53 additions & 4 deletions pkg/api/cache_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -349,10 +349,9 @@ func TestCachePublicationFailureDoesNotFailResponse(t *testing.T) {
dir := t.TempDir()
seed := cacheTestClient(t, dir, time.Hour, cacheTestResponse(strings.NewReader("previous response")))
assert.Equal(t, "previous response", cacheTestFetch(t, seed))
files := cacheTestFiles(t, dir)
require.Len(t, files, 1)
require.NoError(t, os.Remove(files[0]))
require.NoError(t, os.Mkdir(files[0], 0700))
cacheFile := cacheTestEntryFile(t, dir)
require.NoError(t, os.Remove(cacheFile))
require.NoError(t, os.Mkdir(cacheFile, 0700))
client := cacheTestClient(t, dir, time.Hour, cacheTestResponse(strings.NewReader("fresh response")))

// When the HTTP request succeeds but cache publication fails.
Expand All @@ -363,6 +362,45 @@ func TestCachePublicationFailureDoesNotFailResponse(t *testing.T) {
assert.Empty(t, cacheTestFiles(t, dir))
}

func TestCacheEvictsEntryWithCorruptBody(t *testing.T) {
t.Parallel()

dir := t.TempDir()
seed := cacheTestClient(t, dir, time.Hour, cacheTestResponse(strings.NewReader("cached response")))
assert.Equal(t, "cached response", cacheTestFetch(t, seed))

cacheFile := cacheTestEntryFile(t, dir)
contents, err := os.ReadFile(cacheFile)
require.NoError(t, err)
contents = bytes.Replace(contents, []byte("cached response"), []byte("corrupt response"), 1)
require.NoError(t, os.WriteFile(cacheFile, contents, 0600))

client := cacheTestClient(t, dir, time.Hour, cacheTestResponse(strings.NewReader("fresh response")))
assert.Equal(t, "fresh response", cacheTestFetch(t, client))

reader := cacheTestClient(t, dir, time.Hour, cacheTestTransport(func(*http.Request) (*http.Response, error) {
return nil, errors.New("unexpected cache miss")
}))
assert.Equal(t, "fresh response", cacheTestFetch(t, reader))
}

func TestCacheEvictsLegacyEntryWithoutChecksum(t *testing.T) {
t.Parallel()

dir := t.TempDir()
seed := cacheTestClient(t, dir, time.Hour, cacheTestResponse(strings.NewReader("legacy response")))
assert.Equal(t, "legacy response", cacheTestFetch(t, seed))
require.NoError(t, os.Remove(cacheTestEntryFile(t, dir)+".sha256"))

client := cacheTestClient(t, dir, time.Hour, cacheTestResponse(strings.NewReader("fresh response")))
assert.Equal(t, "fresh response", cacheTestFetch(t, client))

reader := cacheTestClient(t, dir, time.Hour, cacheTestTransport(func(*http.Request) (*http.Response, error) {
return nil, errors.New("unexpected cache miss")
}))
assert.Equal(t, "fresh response", cacheTestFetch(t, reader))
}

func cacheTestClient(t *testing.T, dir string, ttl time.Duration, transport http.RoundTripper) *http.Client {
t.Helper()
client, err := api.NewHTTPClient(api.ClientOptions{
Expand Down Expand Up @@ -395,6 +433,17 @@ func cacheTestFiles(t *testing.T, dir string) []string {
return files
}

func cacheTestEntryFile(t *testing.T, dir string) string {
t.Helper()
for _, path := range cacheTestFiles(t, dir) {
if !strings.HasSuffix(path, ".sha256") {
return path
}
}
t.Fatal("cache entry not found")
return ""
}

func cacheTestResponse(body io.Reader) http.RoundTripper {
return cacheTestTransport(func(*http.Request) (*http.Response, error) {
return &http.Response{
Expand Down