diff --git a/pkg/api/cache.go b/pkg/api/cache.go index 94a5e28..09545ea 100644 --- a/pkg/api/cache.go +++ b/pkg/api/cache.go @@ -4,6 +4,7 @@ import ( "bufio" "bytes" "crypto/sha256" + "encoding/hex" "errors" "fmt" "io" @@ -143,34 +144,67 @@ 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 { @@ -178,6 +212,7 @@ func (fs *fileStorage) store(key string, res *http.Response) error { 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 @@ -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 { diff --git a/pkg/api/cache_test.go b/pkg/api/cache_test.go index b2051ff..2b3fefe 100644 --- a/pkg/api/cache_test.go +++ b/pkg/api/cache_test.go @@ -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. @@ -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{ @@ -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{