Skip to content
Merged
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
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -322,6 +322,7 @@ These are methods on a pipe that change its configuration:

| Source | Modifies |
| -------- | ------------- |
| [`WithContext`](https://pkg.go.dev/github.com/bitfield/script#Pipe.WithContext) | context used in commands or HTTP requests |
| [`WithEnv`](https://pkg.go.dev/github.com/bitfield/script#Pipe.WithEnv) | environment for commands |
| [`WithError`](https://pkg.go.dev/github.com/bitfield/script#Pipe.WithError) | pipe error status |
| [`WithHTTPClient`](https://pkg.go.dev/github.com/bitfield/script#Pipe.WithHTTPClient) | client for HTTP requests |
Expand Down
63 changes: 56 additions & 7 deletions script.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package script
import (
"bufio"
"container/ring"
"context"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
Expand Down Expand Up @@ -38,6 +39,7 @@ type Pipe struct {
err error
stderr io.Writer
env []string
ctx context.Context
}

// Args creates a pipe containing the program's command-line arguments from
Expand Down Expand Up @@ -174,6 +176,7 @@ func NewPipe() *Pipe {
stdout: os.Stdout,
httpClient: http.DefaultClient,
env: nil,
ctx: context.Background(),
}
}

Expand All @@ -198,6 +201,11 @@ func Stdin() *Pipe {
return NewPipe().WithReader(os.Stdin)
}

// WithContext creates a pipe with the context ctx.
func WithContext(ctx context.Context) *Pipe {
return NewPipe().WithContext(ctx)
}

// AppendFile appends the contents of the pipe to the file path, creating it if
// necessary, and returns the number of bytes successfully written, or an
// error.
Expand Down Expand Up @@ -326,6 +334,12 @@ func (p *Pipe) Dirname() *Pipe {
// set by [Pipe.WithHTTPClient], or [http.DefaultClient] otherwise. The
// response body is streamed concurrently to the pipe's output. If the response
// status is anything other than HTTP 200-299, the pipe's error status is set.
//
// # Context
//
// The request does not inherit the pipe's context (if any was set by
// [Pipe.WithContext]). If you want the request to be cancellable, you will need to set
// your own context on it.
func (p *Pipe) Do(req *http.Request) *Pipe {
return p.Filter(func(r io.Reader, w io.Writer) error {
resp, err := p.httpClient.Do(req)
Expand Down Expand Up @@ -413,6 +427,11 @@ func (p *Pipe) Error() error {
// The command inherits the current process's environment, optionally modified
// by [Pipe.WithEnv].
//
// # Context
//
// The command inherits the pipe's context (if any was set by [Pipe.WithContext]), and
// will be cancelled if the context is cancelled or times out.
//
// # Error handling
//
// If the command had a non-zero exit status, the pipe's error status will also
Expand All @@ -432,7 +451,7 @@ func (p *Pipe) Exec(cmdLine string) *Pipe {
if err != nil {
return err
}
cmd := exec.Command(args[0], args[1:]...)
cmd := exec.CommandContext(p.ctx, args[0], args[1:]...)
cmd.Stdin = r
cmd.Stdout = w
cmd.Stderr = w
Expand Down Expand Up @@ -462,6 +481,16 @@ func (p *Pipe) Exec(cmdLine string) *Pipe {
// syntax. For example:
//
// ListFiles("*").ExecForEach("touch {{.}}").Wait()
//
// # Environment
//
// The command inherits the current process's environment, optionally modified
// by [Pipe.WithEnv].
//
// # Context
//
// The command inherits the pipe's context (if any was set by [Pipe.WithContext]), and
// will be cancelled if the context is cancelled or times out.
func (p *Pipe) ExecForEach(cmdLine string) *Pipe {
tpl, err := template.New("").Parse(cmdLine)
if err != nil {
Expand All @@ -470,6 +499,9 @@ func (p *Pipe) ExecForEach(cmdLine string) *Pipe {
return p.Filter(func(r io.Reader, w io.Writer) error {
scanner := newScanner(r)
for scanner.Scan() {
if p.ctx.Err() != nil {
return p.ctx.Err()
}
cmdLine := new(strings.Builder)
err := tpl.Execute(cmdLine, scanner.Text())
if err != nil {
Expand All @@ -479,7 +511,7 @@ func (p *Pipe) ExecForEach(cmdLine string) *Pipe {
if err != nil {
return err
}
cmd := exec.Command(args[0], args[1:]...)
cmd := exec.CommandContext(p.ctx, args[0], args[1:]...)
cmd.Stdout = w
cmd.Stderr = w
pipeStderr := p.stdErr()
Expand Down Expand Up @@ -530,9 +562,10 @@ func (p *Pipe) ExitStatus() int {
// [io.Writer] to write its output to, and returns an error, which will be set
// on the pipe.
//
// filter runs concurrently, so its goroutine will not exit until the pipe has
// been fully read. Use [Pipe.Wait] to wait for all concurrent filters to
// complete.
// filter runs concurrently, so its goroutine will not exit until the pipe has been
// fully read. Use [Pipe.Wait] to wait for all concurrent filters to complete. These
// goroutines are not automatically cancelled by the pipe's context, if one was set by
// [Pipe.WithContext].
func (p *Pipe) Filter(filter func(io.Reader, io.Writer) error) *Pipe {
if p.Error() != nil {
return p
Expand Down Expand Up @@ -651,8 +684,13 @@ func (p *Pipe) Freq() *Pipe {
// Get makes an HTTP GET request to url, sending the contents of the pipe as
// the request body, and produces the server's response. See [Pipe.Do] for how
// the HTTP response status is interpreted.
//
// # Context
//
// The request inherits the pipe's context (if any was set by [Pipe.WithContext]), and
// will be cancelled if the context is cancelled or times out.
func (p *Pipe) Get(url string) *Pipe {
req, err := http.NewRequest(http.MethodGet, url, p.Reader)
req, err := http.NewRequestWithContext(p.ctx, http.MethodGet, url, p.Reader)
if err != nil {
return p.WithError(err)
}
Expand Down Expand Up @@ -806,8 +844,13 @@ func (p *Pipe) MatchRegexp(re *regexp.Regexp) *Pipe {
// Post makes an HTTP POST request to url, using the contents of the pipe as
// the request body, and produces the server's response. See [Pipe.Do] for how
// the HTTP response status is interpreted.
//
// # Context
//
// The request inherits the pipe's context (if any was set by [Pipe.WithContext]), and
// will be cancelled if the context is cancelled or times out.
func (p *Pipe) Post(url string) *Pipe {
req, err := http.NewRequest(http.MethodPost, url, p.Reader)
req, err := http.NewRequestWithContext(p.ctx, http.MethodPost, url, p.Reader)
if err != nil {
return p.WithError(err)
}
Expand Down Expand Up @@ -963,6 +1006,12 @@ func (p *Pipe) Wait() error {
return p.Error()
}

// WithContext sets the context for subsequent [Pipe.Exec], [Pipe.ExecForEach], [Pipe.Get] and [Pipe.Post] commands.
func (p *Pipe) WithContext(ctx context.Context) *Pipe {
p.ctx = ctx
return p
}

// WithEnv sets the environment for subsequent [Pipe.Exec] and [Pipe.ExecForEach]
// commands to the string slice env, using the same format as [os/exec.Cmd.Env].
// An empty slice unsets all existing environment variables.
Expand Down
67 changes: 67 additions & 0 deletions script_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package script_test
import (
"bufio"
"bytes"
"context"
"crypto/sha256"
"crypto/sha512"
"errors"
Expand All @@ -18,6 +19,7 @@ import (
"strings"
"testing"
"testing/iotest"
"time"

"github.com/bitfield/script"
"github.com/google/go-cmp/cmp"
Expand Down Expand Up @@ -2165,6 +2167,7 @@ func TestHash_ReturnsErrorGivenReadErrorOnPipe(t *testing.T) {
}

func TestHashSums_OutputsEmptyStringForFileThatCannotBeHashed(t *testing.T) {
t.Parallel()
got, err := script.Echo("file_does_not_exist.txt").HashSums(sha256.New()).String()
if err != nil {
t.Fatal(err)
Expand All @@ -2175,6 +2178,55 @@ func TestHashSums_OutputsEmptyStringForFileThatCannotBeHashed(t *testing.T) {
}
}

func TestWithContext_SkipsExecWithCancelledContext(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
cancel()
output, err := script.WithContext(ctx).Exec("echo Hello, World").String()
if err == nil {
t.Error("expected error due to context cancellation, got nil")
}
if strings.Contains(output, "Hello") {
t.Fatal("command ran even though context cancelled beforehand")
}
}

func TestWithContext_SkipsGETRequestWithCancelledContext(t *testing.T) {
t.Parallel()
handler_called := false
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
handler_called = true
}))
defer ts.Close()
ctx, cancel := context.WithCancel(context.Background())
cancel()
_, err := script.WithContext(ctx).Echo("request data").Get(ts.URL).String()
if err == nil {
t.Fatal("expected error from cancelled context")
}
if handler_called {
t.Fatal("request made even though context cancelled beforehand")
}
}

func TestWithContext_SkipsPOSTRequestWithCancelledContext(t *testing.T) {
t.Parallel()
handler_called := false
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
handler_called = true
}))
defer ts.Close()
ctx, cancel := context.WithCancel(context.Background())
cancel()
_, err := script.WithContext(ctx).Echo("request data").Post(ts.URL).String()
if err == nil {
t.Fatal("expected error from cancelled context")
}
if handler_called {
t.Fatal("request made even though context cancelled beforehand")
}
}

func ExampleArgs() {
script.Args().Stdout()
// prints command-line arguments
Expand Down Expand Up @@ -2622,6 +2674,21 @@ func ExamplePipe_Tee_writers() {
// hello
}

func ExamplePipe_WithContext() {
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
time.Sleep(20 * time.Millisecond)
}))
defer ts.Close()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
defer cancel()
_, err := script.NewPipe().WithContext(ctx).Get(ts.URL).String()
if err != nil {
fmt.Println("cancelled")
}
// Output:
// cancelled
}

func ExamplePipe_WithStderr() {
buf := new(bytes.Buffer)
script.NewPipe().WithStderr(buf).Exec("go").Wait()
Expand Down
Loading