fix(backend): harden final MVP invariants
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -32,32 +33,36 @@ func (f *closeFailingFile) Close() error {
|
||||
return errors.New("close failed")
|
||||
}
|
||||
|
||||
func TestSaveStoresFileUnderProjectDirectory(t *testing.T) {
|
||||
func TestSaveStoresFileUnderProjectIdentityWithOpaqueKey(t *testing.T) {
|
||||
service := NewService(t.TempDir())
|
||||
projectIdentity := "019b0000-0000-7000-8000-000000000012"
|
||||
|
||||
stored, err := service.Save(12, "brief.md", strings.NewReader("hello"))
|
||||
stored, err := service.Save(projectIdentity, "brief.md", strings.NewReader("hello"))
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "brief.md", stored.OriginalName)
|
||||
require.Contains(t, stored.RelativePath, "projects/12/")
|
||||
require.Contains(t, stored.RelativePath, "projects/"+projectIdentity+"/")
|
||||
require.NotContains(t, stored.RelativePath, "projects/12/")
|
||||
require.Regexp(t, `^[a-f0-9]{32}$`, stored.StorageKey)
|
||||
require.Equal(t, stored.StorageKey, filepath.Base(stored.RelativePath))
|
||||
require.FileExists(t, stored.AbsolutePath)
|
||||
}
|
||||
|
||||
func TestSaveNeutralizesPathTraversal(t *testing.T) {
|
||||
service := NewService(t.TempDir())
|
||||
|
||||
stored, err := service.Save(12, "..\\..\\secret.txt", strings.NewReader("hello"))
|
||||
stored, err := service.Save("019b0000-0000-7000-8000-000000000012", "..\\..\\secret.txt", strings.NewReader("hello"))
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "secret.txt", stored.OriginalName)
|
||||
require.Contains(t, stored.RelativePath, "projects/12/")
|
||||
require.NotContains(t, stored.RelativePath, "secret.txt")
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenCopyFails(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
|
||||
_, err := service.Save(12, "broken.bin", &failingReader{})
|
||||
_, err := service.Save("019b0000-0000-7000-8000-000000000012", "broken.bin", &failingReader{})
|
||||
|
||||
require.ErrorContains(t, err, "copy failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
@@ -74,24 +79,63 @@ func TestSaveRemovesTemporaryFileWhenCloseFails(t *testing.T) {
|
||||
return &closeFailingFile{File: file}, nil
|
||||
}
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
_, err := service.Save("019b0000-0000-7000-8000-000000000012", "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "close failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
}
|
||||
|
||||
func TestSaveRemovesTemporaryAndPartialFilesWhenRenameFails(t *testing.T) {
|
||||
func TestSaveCollisionNeverOverwritesOrDeletesExistingFile(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
service.rename = func(oldPath, newPath string) error {
|
||||
require.NoError(t, os.WriteFile(newPath, []byte("partial final"), 0o600))
|
||||
return errors.New("rename failed")
|
||||
service.newStorageKey = func() (string, error) { return "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", nil }
|
||||
projectIdentity := "019b0000-0000-7000-8000-000000000012"
|
||||
first, err := service.Save(projectIdentity, "first.bin", strings.NewReader("first owner"))
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = service.Save(projectIdentity, "second.bin", strings.NewReader("second owner"))
|
||||
|
||||
require.ErrorIs(t, err, ErrStorageKeyCollision)
|
||||
content, readErr := os.ReadFile(first.AbsolutePath)
|
||||
require.NoError(t, readErr)
|
||||
require.Equal(t, "first owner", string(content))
|
||||
requireOnlyStoredFile(t, storageRoot, first.AbsolutePath)
|
||||
}
|
||||
|
||||
func TestConcurrentSaveWithSameKeyPublishesExactlyOneOwner(t *testing.T) {
|
||||
storageRoot := t.TempDir()
|
||||
service := NewService(storageRoot)
|
||||
service.newStorageKey = func() (string, error) { return "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", nil }
|
||||
projectIdentity := "019b0000-0000-7000-8000-000000000012"
|
||||
start := make(chan struct{})
|
||||
type result struct {
|
||||
stored StoredFile
|
||||
err error
|
||||
}
|
||||
results := make(chan result, 2)
|
||||
var ready sync.WaitGroup
|
||||
ready.Add(2)
|
||||
for _, body := range []string{"alpha", "beta"} {
|
||||
go func(content string) {
|
||||
ready.Done()
|
||||
<-start
|
||||
stored, err := service.Save(projectIdentity, content+".bin", strings.NewReader(content))
|
||||
results <- result{stored: stored, err: err}
|
||||
}(body)
|
||||
}
|
||||
ready.Wait()
|
||||
close(start)
|
||||
firstResult, secondResult := <-results, <-results
|
||||
|
||||
_, err := service.Save(12, "broken.bin", strings.NewReader("content"))
|
||||
|
||||
require.ErrorContains(t, err, "rename failed")
|
||||
requireNoStoredFiles(t, storageRoot)
|
||||
errors := []error{firstResult.err, secondResult.err}
|
||||
require.Equal(t, 1, countNilErrors(errors))
|
||||
require.Equal(t, 1, countMatchingErrors(errors, ErrStorageKeyCollision))
|
||||
winner := firstResult.stored
|
||||
if firstResult.err != nil {
|
||||
winner = secondResult.stored
|
||||
}
|
||||
require.FileExists(t, winner.AbsolutePath)
|
||||
requireOnlyStoredFile(t, storageRoot, winner.AbsolutePath)
|
||||
}
|
||||
|
||||
func requireNoStoredFiles(t *testing.T, storageRoot string) {
|
||||
@@ -106,3 +150,38 @@ func requireNoStoredFiles(t *testing.T, storageRoot string) {
|
||||
return nil
|
||||
}))
|
||||
}
|
||||
|
||||
func requireOnlyStoredFile(t *testing.T, storageRoot, expectedPath string) {
|
||||
t.Helper()
|
||||
files := []string{}
|
||||
require.NoError(t, filepath.WalkDir(storageRoot, func(path string, entry fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !entry.IsDir() {
|
||||
files = append(files, path)
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
require.Equal(t, []string{expectedPath}, files)
|
||||
}
|
||||
|
||||
func countNilErrors(values []error) int {
|
||||
count := 0
|
||||
for _, err := range values {
|
||||
if err == nil {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func countMatchingErrors(values []error, target error) int {
|
||||
count := 0
|
||||
for _, err := range values {
|
||||
if errors.Is(err, target) {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user