diff options
Diffstat (limited to 'internal/image/download_test.go')
| -rw-r--r-- | internal/image/download_test.go | 77 |
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 "" } |
