summaryrefslogtreecommitdiff
path: root/internal/image/download_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/image/download_test.go')
-rw-r--r--internal/image/download_test.go77
1 files changed, 76 insertions, 1 deletions
diff --git a/internal/image/download_test.go b/internal/image/download_test.go
index 84e8f5f..073e0d8 100644
--- a/internal/image/download_test.go
+++ b/internal/image/download_test.go
@@ -216,10 +216,85 @@ func TestDownloadImageRejectsSymlinkAncestor(t *testing.T) {
if err == nil {
t.Fatal("expected error for symlink ancestor")
}
- if !strings.Contains(err.Error(), "symlink") {
+ if !strings.Contains(err.Error(), "symlink") && !strings.Contains(err.Error(), "not a directory") {
t.Fatalf("DownloadImage() error = %v, want symlink rejection", err)
}
if _, statErr := os.Stat(filepath.Join(outsideDir, "escape.png")); !os.IsNotExist(statErr) {
t.Fatalf("unexpected file write outside output dir: %v", statErr)
}
}
+
+func TestDownloadImageRejectsSymlinkSwapDuringDownload(t *testing.T) {
+ t.Parallel()
+
+ baseDir := t.TempDir()
+ outsideDir := t.TempDir()
+ nestedDir := filepath.Join(baseDir, "nested")
+ if err := os.Mkdir(nestedDir, 0o755); err != nil {
+ t.Fatalf("Mkdir() error = %v", err)
+ }
+
+ downloadReady := make(chan struct{})
+ releaseDownload := make(chan struct{})
+ provider := &swapBlockingProvider{
+ payload: "image-bytes",
+ downloadReady: downloadReady,
+ releaseDownload: releaseDownload,
+ }
+ d := NewDownloader(provider, &DownloadOptions{
+ OutputDir: baseDir,
+ CreateDir: true,
+ OverwriteExisting: true,
+ })
+
+ result := &SearchResult{
+ URL: "https://example.com/image.png",
+ Source: Gemini,
+ Attribution: "safe attribution",
+ }
+ outputPath := filepath.Join(nestedDir, "escape.png")
+ errCh := make(chan error, 1)
+ go func() {
+ errCh <- d.DownloadImage(context.Background(), result, outputPath)
+ }()
+
+ <-downloadReady
+ if err := os.RemoveAll(nestedDir); err != nil {
+ t.Fatalf("RemoveAll() error = %v", err)
+ }
+ if err := os.Symlink(outsideDir, nestedDir); err != nil {
+ t.Fatalf("Symlink() error = %v", err)
+ }
+ close(releaseDownload)
+
+ err := <-errCh
+ if err == nil {
+ t.Fatal("expected error for symlink swap during download")
+ }
+ if !strings.Contains(err.Error(), "symlink") && !strings.Contains(err.Error(), "open dir") {
+ t.Fatalf("DownloadImage() error = %v, want symlink rejection", err)
+ }
+ if _, statErr := os.Stat(filepath.Join(outsideDir, "escape.png")); !os.IsNotExist(statErr) {
+ t.Fatalf("unexpected file write outside output dir: %v", statErr)
+ }
+}
+
+type swapBlockingProvider struct {
+ payload string
+ downloadReady chan struct{}
+ releaseDownload chan struct{}
+}
+
+func (p *swapBlockingProvider) Name() string { return "swap-blocking" }
+
+func (p *swapBlockingProvider) Search(context.Context, *SearchOptions) ([]SearchResult, error) {
+ return nil, nil
+}
+
+func (p *swapBlockingProvider) Download(_ context.Context, _ string) (io.ReadCloser, error) {
+ close(p.downloadReady)
+ <-p.releaseDownload
+ return io.NopCloser(strings.NewReader(p.payload)), nil
+}
+
+func (p *swapBlockingProvider) GetAttribution(*SearchResult) string { return "" }