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
32 changes: 27 additions & 5 deletions internal/ghclient/helpers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,21 +6,32 @@ import (
"path/filepath"
"sync/atomic"
"testing"
"time"
)

type testRoundTripper struct {
delay time.Duration
called atomic.Int32
resp *http.Response
err error
}

func (r *testRoundTripper) RoundTrip(_ *http.Request) (*http.Response, error) {
r.called.Add(1)
if r.err != nil {
return nil, r.err
func (tr *testRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
if tr.delay > 0 {
ctx := req.Context()
select {
case <-time.After(tr.delay):
case <-ctx.Done():
}
}

return r.resp, nil
tr.called.Add(1)

if tr.err != nil {
return nil, tr.err
}

return tr.resp, nil
}

func mustMkdirTemp(t *testing.T, dir, pattern string) string {
Expand All @@ -44,3 +55,14 @@ func mustReadAppPrivateKey(t *testing.T) []byte {

return privateKeyData
}

func mustCreateRequest(t *testing.T, method, url string) *http.Request {
t.Helper()

req, err := http.NewRequestWithContext(t.Context(), method, url, nil)
if err != nil {
t.Fatalf("failed to create request: %v", err)
}

return req
}
33 changes: 2 additions & 31 deletions internal/ghclient/throttle.go
Original file line number Diff line number Diff line change
@@ -1,29 +1,11 @@
package ghclient

import (
"io"
"net/http"
"sync"

"golang.org/x/sync/semaphore"
)

// throttlerReadCloser is a wrapper around an io.ReadCloser that releases a semaphore weight when the ReadCloser is closed. This is used to ensure that the semaphore controlling concurrent requests is properly released after the response body has been fully read and closed, preventing resource leaks and allowing other requests to proceed.
type throttlerReadCloser struct {
io.ReadCloser
sema *semaphore.Weighted
once sync.Once
}

// Close releases the semaphore weight when the ReadCloser is closed. It ensures that the release operation is only performed once, even if Close is called multiple times, preventing potential double-release issues that could lead to incorrect semaphore state.
func (c *throttlerReadCloser) Close() error {
err := c.ReadCloser.Close()
c.once.Do(func() {
c.sema.Release(1)
})
return err
}

// throttler is an HTTP RoundTripper that limits the number of concurrent requests to a specified maximum. It uses a weighted semaphore to control access to the underlying RoundTripper, ensuring that no more than the allowed number of requests are in flight at any given time. This is useful for preventing overwhelming a server or API with too many simultaneous requests.
type throttler struct {
sema *semaphore.Weighted
Expand All @@ -35,17 +17,6 @@ func (t *throttler) RoundTrip(req *http.Request) (*http.Response, error) {
if err := t.sema.Acquire(req.Context(), 1); err != nil {
return nil, err
}

res, err := t.inner.RoundTrip(req)
if err != nil {
t.sema.Release(1)
return nil, err
}

res.Body = &throttlerReadCloser{
ReadCloser: res.Body,
sema: t.sema,
}

return res, nil
defer t.sema.Release(1)
return t.inner.RoundTrip(req)
}
115 changes: 76 additions & 39 deletions internal/ghclient/throttle_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,32 +7,12 @@ import (
"net/http"
"strings"
"testing"
"testing/synctest"
"time"

"golang.org/x/sync/semaphore"
)

func Test_throttlerReadCloser_Close(t *testing.T) {
t.Parallel()

sema := semaphore.NewWeighted(1)
if err := sema.Acquire(t.Context(), 1); err != nil {
t.Fatalf("failed to acquire semaphore for setup: %v", err)
}

rc := &throttlerReadCloser{
ReadCloser: io.NopCloser(strings.NewReader("ok")),
sema: sema,
}

if err := rc.Close(); err != nil {
t.Fatalf("failed to close read closer: %v", err)
}

if err := rc.Close(); err != nil {
t.Fatalf("failed to close read closer on second close: %v", err)
}
}

func Test_throttler_RoundTrip(t *testing.T) {
t.Parallel()

Expand All @@ -53,7 +33,11 @@ func Test_throttler_RoundTrip(t *testing.T) {
t.Fatalf("failed to create request: %v", err)
}

_, err = tr.RoundTrip(req)
resp, err := tr.RoundTrip(req)
if resp != nil && resp.Body != nil {
defer resp.Body.Close()
}

if err == nil {
t.Fatal("expected acquire error from canceled context")
}
Expand All @@ -76,12 +60,13 @@ func Test_throttler_RoundTrip(t *testing.T) {
inner: inner,
}

req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "https://example.com", nil)
if err != nil {
t.Fatalf("failed to create request: %v", err)
req := mustCreateRequest(t, http.MethodGet, "https://example.com")

resp, err := tr.RoundTrip(req)
if resp != nil && resp.Body != nil {
defer resp.Body.Close()
}

_, err = tr.RoundTrip(req)
if err == nil {
t.Fatal("expected round trip to fail")
}
Expand All @@ -99,21 +84,18 @@ func Test_throttler_RoundTrip(t *testing.T) {
t.Parallel()

inner := &testRoundTripper{resp: &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader("ok"))}}
sema := semaphore.NewWeighted(1)
tr := &throttler{
sema: sema,
sema: semaphore.NewWeighted(1),
inner: inner,
}

req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "https://example.com", nil)
if err != nil {
t.Fatalf("failed to create request: %v", err)
}
req := mustCreateRequest(t, http.MethodGet, "https://example.com")

resp, err := tr.RoundTrip(req)
if err != nil {
t.Fatalf("expected round trip to succeed, got error: %v", err)
}
defer resp.Body.Close()

if resp.StatusCode != http.StatusOK {
t.Fatalf("expected status code 200 OK, got %d", resp.StatusCode)
Expand All @@ -122,17 +104,72 @@ func Test_throttler_RoundTrip(t *testing.T) {
if inner.called.Load() != 1 {
t.Fatalf("expected inner transport to be called once, got %d calls", inner.called.Load())
}
})

if ok := sema.TryAcquire(1); ok {
t.Fatal("expected semaphore to be held until response body is closed")
t.Run("success_after_no_body_close", func(t *testing.T) {
t.Parallel()

inner := &testRoundTripper{err: errors.New("boom")}
sema := semaphore.NewWeighted(1)
tr := &throttler{
sema: sema,
inner: inner,
}

if err := resp.Body.Close(); err != nil {
t.Fatalf("failed to close response body: %v", err)
req := mustCreateRequest(t, http.MethodGet, "https://example.com")

_, err := tr.RoundTrip(req)
if err == nil || !errors.Is(err, inner.err) {
t.Fatalf("expected round trip to error with %v, got %v", inner.err, err)
}

if ok := sema.TryAcquire(1); !ok {
t.Fatal("expected semaphore to be released after closing response body")
if !sema.TryAcquire(1) {
t.Fatal("semaphore permit leaked after body not closed")
}
})

t.Run("throttles_concurrent_requests", func(t *testing.T) {
t.Parallel()

synctest.Test(t, func(t *testing.T) {
Comment thread
stevehipwell marked this conversation as resolved.
reqs := 5
inner := &testRoundTripper{delay: 1 * time.Second, resp: &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader("ok"))}}
tr := &throttler{
sema: semaphore.NewWeighted(1),
inner: inner,
}

for i := range reqs {
req := mustCreateRequest(t, http.MethodGet, "https://example.com")
go func() {
req := req
resp, err := tr.RoundTrip(req)
if err != nil {
t.Errorf("request %d: %v", i, err)
return
}
resp.Body.Close()

if resp.StatusCode != http.StatusOK {
t.Errorf("request %d: expected status code 200 OK, got %d", i, resp.StatusCode)
}
}()
}

// Both goroutines are now durably blocked (one sleeping, others on sema.Acquire).
synctest.Wait()
if inner.called.Load() != 0 {
t.Fatal("expected no completions before time advances")
}

for i := range reqs {
// Clock jumps 1s: sleeping goroutine wakes, completes, releases semaphore; next goroutine acquires and sleeps.
time.Sleep(time.Second)
synctest.Wait()
if inner.called.Load() != int32(i+1) {
t.Fatal("expected completions to be throttled to one per second")
}
}
})
})
}
Loading