diff --git a/chromeshell/chromeshell.go b/chromeshell/chromeshell.go index 624f5f2..ae398f2 100644 --- a/chromeshell/chromeshell.go +++ b/chromeshell/chromeshell.go @@ -2,13 +2,13 @@ package chromeshell import ( "archive/zip" + "context" "fmt" "io" "net/http" "os" "path/filepath" "runtime" - "sync" ) const ( @@ -21,7 +21,23 @@ const ( revision = 1520797742 ) -var ensureMu sync.Mutex +var ensureLock = func() chan struct{} { + lock := make(chan struct{}, 1) + lock <- struct{}{} + return lock +}() + +// acquireEnsureLock waits for the single download slot. A channel is used +// instead of a sync.Mutex so a caller whose ctx expires stops waiting on a +// download it cannot cancel. The returned release must always be called. +func acquireEnsureLock(ctx context.Context) (func(), error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-ensureLock: + return func() { ensureLock <- struct{}{} }, nil + } +} // Supported reports whether chrome-headless-shell auto-download is available // for the current platform. @@ -58,18 +74,27 @@ func BinPath() string { // returns the executable path. It is a no-op download when the binary already // exists. Callers should only invoke this when Supported() is true. func Ensure() (string, error) { + return EnsureContext(context.Background()) +} + +// EnsureContext is like Ensure, but stops waiting for the shared download lock +// and cancels the browser download when ctx is done. +func EnsureContext(ctx context.Context) (string, error) { if !Supported() { return "", fmt.Errorf("chrome-headless-shell auto-download is only supported on linux/amd64") } - ensureMu.Lock() - defer ensureMu.Unlock() + release, err := acquireEnsureLock(ctx) + if err != nil { + return "", err + } + defer release() if p := findBin(Dir()); p != "" { return p, nil } - if err := downloadAndExtract(Host(), Dir()); err != nil { + if err := downloadAndExtract(ctx, Host(), Dir()); err != nil { return "", err } @@ -112,7 +137,7 @@ func defaultBrowserDir() string { return filepath.Join(home, ".cache", "rod", "browser") } -func downloadAndExtract(url, destDir string) error { +func downloadAndExtract(ctx context.Context, url, destDir string) error { tmpParent := filepath.Dir(destDir) if err := os.MkdirAll(tmpParent, 0o755); err != nil { return err @@ -125,7 +150,7 @@ func downloadAndExtract(url, destDir string) error { defer func() { _ = os.RemoveAll(tmpDir) }() zipPath := filepath.Join(tmpDir, "chrome-headless-shell.zip") - if err := downloadFile(url, zipPath); err != nil { + if err := downloadFile(ctx, url, zipPath); err != nil { return err } @@ -147,8 +172,12 @@ func downloadAndExtract(url, destDir string) error { return nil } -func downloadFile(url, dest string) error { - resp, err := http.Get(url) //nolint:noctx // one-shot browser binary fetch +func downloadFile(ctx context.Context, url, dest string) error { + request, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return err + } + resp, err := http.DefaultClient.Do(request) if err != nil { return err } diff --git a/chromeshell/chromeshell_test.go b/chromeshell/chromeshell_test.go index c3fe6e9..0132216 100644 --- a/chromeshell/chromeshell_test.go +++ b/chromeshell/chromeshell_test.go @@ -2,10 +2,15 @@ package chromeshell import ( "archive/zip" + "context" + "errors" + "net/http" + "net/http/httptest" "os" "path/filepath" "strings" "testing" + "time" ) func writeTestZip(path string, entries map[string]string) error { @@ -94,3 +99,59 @@ func TestEnsureUnsupported(t *testing.T) { t.Fatal("expected unsupported error") } } + +func TestDownloadFileContextCancellation(t *testing.T) { + requestDone := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer close(requestDone) + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("partial")) + w.(http.Flusher).Flush() + <-r.Context().Done() + })) + t.Cleanup(server.Close) + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + + started := time.Now() + err := downloadFile(ctx, server.URL, filepath.Join(t.TempDir(), "browser.zip")) + + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("expected context deadline exceeded, got %v", err) + } + if elapsed := time.Since(started); elapsed > time.Second { + t.Fatalf("download cancellation took %s", elapsed) + } + select { + case <-requestDone: + case <-time.After(time.Second): + t.Fatal("server request did not observe cancellation") + } +} + +func TestAcquireEnsureLockGivesUpOnContext(t *testing.T) { + release, err := acquireEnsureLock(context.Background()) + if err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + + started := time.Now() + if _, err := acquireEnsureLock(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("expected context deadline exceeded while lock is held, got %v", err) + } + if elapsed := time.Since(started); elapsed > time.Second { + t.Fatalf("waiting for the lock took %s", elapsed) + } + + release() + + next, err := acquireEnsureLock(context.Background()) + if err != nil { + t.Fatalf("lock not reusable after release: %v", err) + } + next() +}