Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 38 additions & 9 deletions chromeshell/chromeshell.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,13 @@ package chromeshell

import (
"archive/zip"
"context"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"runtime"
"sync"
)

const (
Expand All @@ -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.
Expand Down Expand Up @@ -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
}

Expand Down Expand Up @@ -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
Expand All @@ -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
}

Expand All @@ -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
}
Expand Down
61 changes: 61 additions & 0 deletions chromeshell/chromeshell_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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()
}
Loading