Use retryablehttp for token calls too (#35)
This commit is contained in:
@@ -12,7 +12,6 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"maps"
|
"maps"
|
||||||
"math/rand"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
@@ -227,12 +226,20 @@ func newClient(githubConfigURL string, creds *actionsAuth, options ...Option) (*
|
|||||||
option(ac)
|
option(ac)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
retryClient, err := ac.newRetryableHTTPClient()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create retryable HTTP client: %w", err)
|
||||||
|
}
|
||||||
|
ac.httpClient = retryClient.StandardClient()
|
||||||
|
|
||||||
|
return ac, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) newRetryableHTTPClient() (*retryablehttp.Client, error) {
|
||||||
retryClient := retryablehttp.NewClient()
|
retryClient := retryablehttp.NewClient()
|
||||||
retryClient.Logger = ac.logger
|
retryClient.Logger = c.logger
|
||||||
|
retryClient.RetryMax = c.retryMax
|
||||||
retryClient.RetryMax = ac.retryMax
|
retryClient.RetryWaitMax = c.retryWaitMax
|
||||||
retryClient.RetryWaitMax = ac.retryWaitMax
|
|
||||||
|
|
||||||
retryClient.HTTPClient.Timeout = 5 * time.Minute // timeout must be > 1m to accomodate long polling
|
retryClient.HTTPClient.Timeout = 5 * time.Minute // timeout must be > 1m to accomodate long polling
|
||||||
|
|
||||||
transport, ok := retryClient.HTTPClient.Transport.(*http.Transport)
|
transport, ok := retryClient.HTTPClient.Transport.(*http.Transport)
|
||||||
@@ -245,20 +252,19 @@ func newClient(githubConfigURL string, creds *actionsAuth, options ...Option) (*
|
|||||||
transport.TLSClientConfig = &tls.Config{}
|
transport.TLSClientConfig = &tls.Config{}
|
||||||
}
|
}
|
||||||
|
|
||||||
if ac.rootCAs != nil {
|
if c.rootCAs != nil {
|
||||||
transport.TLSClientConfig.RootCAs = ac.rootCAs
|
transport.TLSClientConfig.RootCAs = c.rootCAs
|
||||||
}
|
}
|
||||||
|
|
||||||
if ac.tlsInsecureSkipVerify {
|
if c.tlsInsecureSkipVerify {
|
||||||
transport.TLSClientConfig.InsecureSkipVerify = true
|
transport.TLSClientConfig.InsecureSkipVerify = true
|
||||||
}
|
}
|
||||||
|
|
||||||
transport.Proxy = ac.proxyFunc
|
transport.Proxy = c.proxyFunc
|
||||||
|
|
||||||
retryClient.HTTPClient.Transport = transport
|
retryClient.HTTPClient.Transport = transport
|
||||||
ac.httpClient = retryClient.StandardClient()
|
|
||||||
|
|
||||||
return ac, nil
|
return retryClient, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetUserAgent updates the user agent
|
// SetUserAgent updates the user agent
|
||||||
@@ -1108,52 +1114,35 @@ func (c *Client) getActionsServiceAdminConnection(ctx context.Context, rt *regis
|
|||||||
|
|
||||||
c.logger.Info("getting Actions tenant URL and JWT", "registrationURL", req.URL.String())
|
c.logger.Info("getting Actions tenant URL and JWT", "registrationURL", req.URL.String())
|
||||||
|
|
||||||
retry := 0
|
|
||||||
for {
|
|
||||||
adminConnection, err := c.getActionsServiceAdminConnectionRequest(req)
|
adminConnection, err := c.getActionsServiceAdminConnectionRequest(req)
|
||||||
if err == nil {
|
if err != nil {
|
||||||
return adminConnection, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
retry++
|
|
||||||
if retry > 5 {
|
|
||||||
return nil, fmt.Errorf("unable to register runner after 5 retries: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var ghErr *GitHubAPIError
|
|
||||||
if !errors.As(err, &ghErr) {
|
|
||||||
return nil, fmt.Errorf("failed to get actions service admin connection: %w", err)
|
return nil, fmt.Errorf("failed to get actions service admin connection: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if ghErr.StatusCode != http.StatusUnauthorized && ghErr.StatusCode != http.StatusForbidden {
|
return adminConnection, nil
|
||||||
return nil, fmt.Errorf("failed to get actions service admin connection: %w", ghErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
c.logger.Debug("received unauthorized or forbidden response, retrying", "retryAttempt", retry, "statusCode", ghErr.StatusCode)
|
|
||||||
|
|
||||||
// Add exponential backoff + jitter to avoid thundering herd
|
|
||||||
// This will generate a backoff schedule:
|
|
||||||
// 1: 1s
|
|
||||||
// 2: 3s
|
|
||||||
// 3: 4s
|
|
||||||
// 4: 8s
|
|
||||||
// 5: 17s
|
|
||||||
baseDelay := 500 * time.Millisecond
|
|
||||||
jitter := time.Duration(rand.Intn(1000))
|
|
||||||
maxDelay := 20 * time.Second
|
|
||||||
delay := min(baseDelay*(1<<retry)+jitter, maxDelay)
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return nil, fmt.Errorf("context cancelled while waiting to retry: %w", ctx.Err())
|
|
||||||
case <-time.After(delay):
|
|
||||||
// continue to next retry
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) getActionsServiceAdminConnectionRequest(req *http.Request) (*actionsServiceAdminConnection, error) {
|
func (c *Client) getActionsServiceAdminConnectionRequest(req *http.Request) (*actionsServiceAdminConnection, error) {
|
||||||
resp, err := c.do(req)
|
retryableClient, err := c.newRetryableHTTPClient()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create retryable HTTP client: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
retryableClient.CheckRetry = func(ctx context.Context, resp *http.Response, err error) (bool, error) {
|
||||||
|
if resp != nil && (resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden) {
|
||||||
|
// Retry on 401 Unauthorized and 403 Forbidden
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return retryablehttp.DefaultRetryPolicy(ctx, resp, err)
|
||||||
|
}
|
||||||
|
// Adding custom error handler to also return response in case of error
|
||||||
|
retryableClient.ErrorHandler = func(resp *http.Response, err error, numTries int) (*http.Response, error) {
|
||||||
|
return resp, err
|
||||||
|
}
|
||||||
|
httpClient := retryableClient.StandardClient()
|
||||||
|
|
||||||
|
resp, err := httpClient.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to issue the request: %w", err)
|
return nil, fmt.Errorf("failed to issue the request: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
+18
-4
@@ -159,7 +159,11 @@ func TestNewActionsServiceRequest(t *testing.T) {
|
|||||||
testserver.WithActionsToken(newToken),
|
testserver.WithActionsToken(newToken),
|
||||||
testserver.WithActionsRegistrationTokenHandler(unauthorizedHandler),
|
testserver.WithActionsRegistrationTokenHandler(unauthorizedHandler),
|
||||||
)
|
)
|
||||||
client, err := newClient(server.ConfigURLForOrg("my-org"), defaultCreds)
|
client, err := newClient(
|
||||||
|
server.ConfigURLForOrg("my-org"),
|
||||||
|
defaultCreds,
|
||||||
|
WithRetryWaitMax(1*time.Millisecond),
|
||||||
|
)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
expiringToken := "expiring-token"
|
expiringToken := "expiring-token"
|
||||||
expiresAt := time.Now().Add(59 * time.Second)
|
expiresAt := time.Now().Add(59 * time.Second)
|
||||||
@@ -183,7 +187,7 @@ func TestNewActionsServiceRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
failures := 0
|
failures := 0
|
||||||
unauthorizedHandler := func(w http.ResponseWriter, r *http.Request) {
|
unauthorizedHandler := func(w http.ResponseWriter, r *http.Request) {
|
||||||
if failures < 5 {
|
if failures < 4 {
|
||||||
failures++
|
failures++
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
w.WriteHeader(http.StatusUnauthorized)
|
w.WriteHeader(http.StatusUnauthorized)
|
||||||
@@ -194,8 +198,18 @@ func TestNewActionsServiceRequest(t *testing.T) {
|
|||||||
w.WriteHeader(http.StatusCreated)
|
w.WriteHeader(http.StatusCreated)
|
||||||
_ = json.NewEncoder(w).Encode(resp)
|
_ = json.NewEncoder(w).Encode(resp)
|
||||||
}
|
}
|
||||||
server := testserver.New(t, nil, testserver.WithActionsToken("random-token"), testserver.WithActionsToken(newToken), testserver.WithActionsRegistrationTokenHandler(unauthorizedHandler))
|
server := testserver.New(
|
||||||
client, err := newClient(server.ConfigURLForOrg("my-org"), defaultCreds)
|
t,
|
||||||
|
nil,
|
||||||
|
testserver.WithActionsToken("random-token"),
|
||||||
|
testserver.WithActionsToken(newToken),
|
||||||
|
testserver.WithActionsRegistrationTokenHandler(unauthorizedHandler),
|
||||||
|
)
|
||||||
|
client, err := newClient(
|
||||||
|
server.ConfigURLForOrg("my-org"),
|
||||||
|
defaultCreds,
|
||||||
|
WithRetryWaitMax(1*time.Millisecond),
|
||||||
|
)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
expiringToken := "expiring-token"
|
expiringToken := "expiring-token"
|
||||||
expiresAt := time.Now().Add(59 * time.Second)
|
expiresAt := time.Now().Add(59 * time.Second)
|
||||||
|
|||||||
Reference in New Issue
Block a user