Add debug info while issuing request (#50)

* Add debug info while issuing request

* wip

* wip

* Rework tests
This commit is contained in:
Nikola Jokic
2026-01-30 13:03:22 +01:00
committed by GitHub
parent 145b7e382f
commit c841c96f1c
7 changed files with 423 additions and 612 deletions
+78 -165
View File
@@ -5,7 +5,6 @@ import (
"bytes" "bytes"
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"maps" "maps"
@@ -296,30 +295,22 @@ func (c *Client) GetRunnerScaleSet(ctx context.Context, runnerGroupID int, runne
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return nil, ParseActionsErrorFromResponse(resp) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
var runnerScaleSetList *runnerScaleSetsResponse var runnerScaleSetList runnerScaleSetsResponse
if err := json.NewDecoder(resp.Body).Decode(&runnerScaleSetList); err != nil { if err := json.NewDecoder(resp.Body).Decode(&runnerScaleSetList); err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode runner scale set list: %w", err))
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
if runnerScaleSetList.Count == 0 {
switch runnerScaleSetList.Count {
case 1:
return &runnerScaleSetList.RunnerScaleSets[0], nil
case 0:
return nil, nil return nil, nil
default:
return nil, newRequestResponseError(req, resp, fmt.Errorf("multiple runner scale sets found with name %q", runnerScaleSetName))
} }
if runnerScaleSetList.Count > 1 {
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: fmt.Errorf("multiple runner scale sets found with name %q", runnerScaleSetName),
}
}
return &runnerScaleSetList.RunnerScaleSets[0], nil
} }
// GetRunnerScaleSetByID fetches a runner scale set by its ID. // GetRunnerScaleSetByID fetches a runner scale set by its ID.
@@ -340,16 +331,12 @@ func (c *Client) GetRunnerScaleSetByID(ctx context.Context, runnerScaleSetID int
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return nil, ParseActionsErrorFromResponse(resp) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
var runnerScaleSet *RunnerScaleSet var runnerScaleSet *RunnerScaleSet
if err := json.NewDecoder(resp.Body).Decode(&runnerScaleSet); err != nil { if err := json.NewDecoder(resp.Body).Decode(&runnerScaleSet); err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode runner scale set: %w", err))
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
return runnerScaleSet, nil return runnerScaleSet, nil
} }
@@ -371,48 +358,22 @@ func (c *Client) GetRunnerGroupByName(ctx context.Context, runnerGroup string) (
} }
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
body, err := io.ReadAll(resp.Body) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
if err != nil {
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
return nil, fmt.Errorf("unexpected status code: %w", &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: errors.New(string(body)),
})
} }
var runnerGroupList *RunnerGroupList var runnerGroupList RunnerGroupList
err = json.NewDecoder(resp.Body).Decode(&runnerGroupList) if err := json.NewDecoder(resp.Body).Decode(&runnerGroupList); err != nil {
if err != nil { return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode runner group list: %w", err))
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
if runnerGroupList.Count == 0 { switch runnerGroupList.Count {
return nil, &ActionsError{ case 1:
StatusCode: resp.StatusCode, return &runnerGroupList.RunnerGroups[0], nil
ActivityID: resp.Header.Get(headerActionsActivityID), case 0:
Err: fmt.Errorf("no runner group found with name %q", runnerGroup), return nil, newRequestResponseError(req, resp, fmt.Errorf("no runner group found with name %q", runnerGroup))
} default:
return nil, newRequestResponseError(req, resp, fmt.Errorf("multiple runner group found with name %q", runnerGroup))
} }
if runnerGroupList.Count > 1 {
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: fmt.Errorf("multiple runner group found with name %q", runnerGroup),
}
}
return &runnerGroupList.RunnerGroups[0], nil
} }
// applyDefaultLabelTypes ensures that each label in the runner scale set has a Type set, // applyDefaultLabelTypes ensures that each label in the runner scale set has a Type set,
@@ -449,17 +410,15 @@ func (c *Client) CreateRunnerScaleSet(ctx context.Context, runnerScaleSet *Runne
} }
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return nil, ParseActionsErrorFromResponse(resp) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
var createdRunnerScaleSet *RunnerScaleSet
var createdRunnerScaleSet RunnerScaleSet
if err := json.NewDecoder(resp.Body).Decode(&createdRunnerScaleSet); err != nil { if err := json.NewDecoder(resp.Body).Decode(&createdRunnerScaleSet); err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode created runner scale set: %w", err))
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
return createdRunnerScaleSet, nil
return &createdRunnerScaleSet, nil
} }
// UpdateRunnerScaleSet updates an existing runner scale set. // UpdateRunnerScaleSet updates an existing runner scale set.
@@ -487,18 +446,14 @@ func (c *Client) UpdateRunnerScaleSet(ctx context.Context, runnerScaleSetID int,
} }
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return nil, ParseActionsErrorFromResponse(resp) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
var updatedRunnerScaleSet *RunnerScaleSet var updatedRunnerScaleSet RunnerScaleSet
if err := json.NewDecoder(resp.Body).Decode(&updatedRunnerScaleSet); err != nil { if err := json.NewDecoder(resp.Body).Decode(&updatedRunnerScaleSet); err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode updated runner scale set: %w", err))
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
return updatedRunnerScaleSet, nil return &updatedRunnerScaleSet, nil
} }
// DeleteRunnerScaleSet deletes a runner scale set by its ID. // DeleteRunnerScaleSet deletes a runner scale set by its ID.
@@ -519,7 +474,7 @@ func (c *Client) DeleteRunnerScaleSet(ctx context.Context, runnerScaleSetID int)
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusNoContent { if resp.StatusCode != http.StatusNoContent {
return ParseActionsErrorFromResponse(resp) return newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
return nil return nil
@@ -644,17 +599,14 @@ func (c *Client) GenerateJitRunnerConfig(ctx context.Context, jitRunnerSetting *
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return nil, ParseActionsErrorFromResponse(resp) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
var runnerJitConfig *RunnerScaleSetJitRunnerConfig var runnerJitConfig *RunnerScaleSetJitRunnerConfig
if err := json.NewDecoder(resp.Body).Decode(&runnerJitConfig); err != nil { if err := json.NewDecoder(resp.Body).Decode(&runnerJitConfig); err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode runner JIT config: %w", err))
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
return runnerJitConfig, nil return runnerJitConfig, nil
} }
@@ -677,16 +629,12 @@ func (c *Client) GetRunner(ctx context.Context, runnerID int) (*RunnerReference,
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return nil, ParseActionsErrorFromResponse(resp) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
var runnerReference *RunnerReference var runnerReference *RunnerReference
if err := json.NewDecoder(resp.Body).Decode(&runnerReference); err != nil { if err := json.NewDecoder(resp.Body).Decode(&runnerReference); err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode runner reference: %w", err))
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
return runnerReference, nil return runnerReference, nil
@@ -711,31 +659,22 @@ func (c *Client) GetRunnerByName(ctx context.Context, runnerName string) (*Runne
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
return nil, ParseActionsErrorFromResponse(resp) return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
var runnerList *RunnerReferenceList var runnerList *RunnerReferenceList
if err := json.NewDecoder(resp.Body).Decode(&runnerList); err != nil { if err := json.NewDecoder(resp.Body).Decode(&runnerList); err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode runner reference list: %w", err))
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
} }
if runnerList.Count == 0 { switch runnerList.Count {
case 1:
return &runnerList.RunnerReferences[0], nil
case 0:
return nil, nil return nil, nil
default:
return nil, fmt.Errorf("multiple runners found with name %q", runnerName)
} }
if runnerList.Count > 1 {
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: fmt.Errorf("multiple runner found with name %s", runnerName),
}
}
return &runnerList.RunnerReferences[0], nil
} }
// RemoveRunner removes a runner by its ID. // RemoveRunner removes a runner by its ID.
@@ -757,7 +696,7 @@ func (c *Client) RemoveRunner(ctx context.Context, runnerID int64) error {
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusNoContent { if resp.StatusCode != http.StatusNoContent {
return ParseActionsErrorFromResponse(resp) return newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
} }
return nil return nil
@@ -805,24 +744,12 @@ func (c *Client) getRunnerRegistrationToken(ctx context.Context) (*registrationT
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusCreated { if resp.StatusCode != http.StatusCreated {
body, err := io.ReadAll(resp.Body) return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to get runner registration token (%v)", resp.Status))
if err != nil {
return nil, fmt.Errorf("failed to read the body: %w", err)
}
return nil, &GitHubAPIError{
StatusCode: resp.StatusCode,
RequestID: resp.Header.Get(headerGitHubRequestID),
Err: errors.New(string(body)),
}
} }
var registrationToken *registrationToken var registrationToken *registrationToken
if err := json.NewDecoder(resp.Body).Decode(&registrationToken); err != nil { if err := json.NewDecoder(resp.Body).Decode(&registrationToken); err != nil {
return nil, &GitHubAPIError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode runner registration token: %w", err))
StatusCode: resp.StatusCode,
RequestID: resp.Header.Get(headerGitHubRequestID),
Err: err,
}
} }
return registrationToken, nil return registrationToken, nil
@@ -858,28 +785,15 @@ func (c *Client) fetchAccessToken(ctx context.Context, creds *GitHubAppAuth) (*a
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusCreated { if resp.StatusCode != http.StatusCreated {
errMsg := fmt.Sprintf("failed to get access token for GitHub App auth (%v)", resp.Status) return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to get access token for GitHub App auth (%v)", resp.Status))
if body, err := io.ReadAll(resp.Body); err == nil {
errMsg = fmt.Sprintf("%s: %s", errMsg, string(body))
}
return nil, &GitHubAPIError{
StatusCode: resp.StatusCode,
RequestID: resp.Header.Get(headerGitHubRequestID),
Err: errors.New(errMsg),
}
} }
// Format: https://docs.github.com/en/rest/apps/apps#create-an-installation-access-token-for-an-app // Format: https://docs.github.com/en/rest/apps/apps#create-an-installation-access-token-for-an-app
var accessToken *accessToken var accessToken accessToken
if err = json.NewDecoder(resp.Body).Decode(&accessToken); err != nil { if err := json.NewDecoder(resp.Body).Decode(&accessToken); err != nil {
return nil, &GitHubAPIError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode access token for GitHub App auth: %w", err))
StatusCode: resp.StatusCode,
RequestID: resp.Header.Get(headerGitHubRequestID),
Err: err,
}
} }
return accessToken, nil return &accessToken, nil
} }
type actionsServiceAdminConnection struct { type actionsServiceAdminConnection struct {
@@ -944,38 +858,28 @@ func (c *Client) getActionsServiceAdminConnectionRequest(req *http.Request) (*ac
} }
httpClient := retryableClient.StandardClient() httpClient := retryableClient.StandardClient()
resp, err := httpClient.Do(req) resp, err := sendRequest(httpClient, 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)
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode >= 200 && resp.StatusCode <= 299 { if resp.StatusCode < 200 || resp.StatusCode > 299 {
var actionsServiceAdminConnection *actionsServiceAdminConnection return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code: %d", resp.StatusCode))
if err := json.NewDecoder(resp.Body).Decode(&actionsServiceAdminConnection); err != nil {
return nil, &GitHubAPIError{
StatusCode: resp.StatusCode,
RequestID: resp.Header.Get(headerGitHubRequestID),
Err: err,
}
}
return actionsServiceAdminConnection, nil
} }
var innerErr error var actionsServiceAdminConnection actionsServiceAdminConnection
body, err := io.ReadAll(resp.Body) if err := json.NewDecoder(resp.Body).Decode(&actionsServiceAdminConnection); err != nil {
if err != nil { return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to decode actions service admin connection: %w", err))
innerErr = err }
} else { if actionsServiceAdminConnection.ActionsServiceURL == nil || *actionsServiceAdminConnection.ActionsServiceURL == "" {
innerErr = errors.New(string(body)) return nil, fmt.Errorf("actions service admin connection missing url")
}
if actionsServiceAdminConnection.AdminToken == nil || *actionsServiceAdminConnection.AdminToken == "" {
return nil, fmt.Errorf("actions service admin connection missing token")
} }
return nil, &GitHubAPIError{ return &actionsServiceAdminConnection, nil
StatusCode: resp.StatusCode,
RequestID: resp.Header.Get(headerGitHubRequestID),
Err: innerErr,
}
} }
func createRegistrationTokenPath(config *gitHubConfig) (string, error) { func createRegistrationTokenPath(config *gitHubConfig) (string, error) {
@@ -1049,7 +953,13 @@ func (c *Client) updateTokenIfNeeded(ctx context.Context) error {
return nil return nil
} }
c.logger.Info("refreshing token", "githubConfigUrl", c.config.configURL.String()) configURL := ""
if c.config.configURL != nil {
configURL = c.config.configURL.String()
}
if c.logger != nil {
c.logger.Info("refreshing token", "githubConfigUrl", configURL)
}
rt, err := c.getRunnerRegistrationToken(ctx) rt, err := c.getRunnerRegistrationToken(ctx)
if err != nil { if err != nil {
return fmt.Errorf("failed to get runner registration token on refresh: %w", err) return fmt.Errorf("failed to get runner registration token on refresh: %w", err)
@@ -1059,6 +969,9 @@ func (c *Client) updateTokenIfNeeded(ctx context.Context) error {
if err != nil { if err != nil {
return fmt.Errorf("failed to get actions service admin connection on refresh: %w", err) return fmt.Errorf("failed to get actions service admin connection on refresh: %w", err)
} }
if adminConnInfo == nil || adminConnInfo.ActionsServiceURL == nil || adminConnInfo.AdminToken == nil {
return fmt.Errorf("failed to get actions service admin connection on refresh: missing url or token")
}
c.actionsServiceURL = *adminConnInfo.ActionsServiceURL c.actionsServiceURL = *adminConnInfo.ActionsServiceURL
c.actionsServiceAdminToken = *adminConnInfo.AdminToken c.actionsServiceAdminToken = *adminConnInfo.AdminToken
+14 -12
View File
@@ -176,7 +176,7 @@ func TestNewActionsServiceRequest(t *testing.T) {
client.actionsServiceAdminTokenExpiresAt = expiresAt client.actionsServiceAdminTokenExpiresAt = expiresAt
_, err = client.newActionsServiceRequest(ctx, http.MethodGet, "my-path", nil) _, err = client.newActionsServiceRequest(ctx, http.MethodGet, "my-path", nil)
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), errMessage) assert.Contains(t, err.Error(), "test")
assert.Equal(t, client.actionsServiceAdminToken, expiringToken) assert.Equal(t, client.actionsServiceAdminToken, expiringToken)
assert.Equal(t, client.actionsServiceAdminTokenExpiresAt, expiresAt) assert.Equal(t, client.actionsServiceAdminTokenExpiresAt, expiresAt)
}) })
@@ -604,9 +604,11 @@ func TestGetRunnerScaleSet(t *testing.T) {
}) })
t.Run("Error when Content-Type is text/plain", func(t *testing.T) { t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
plainBody := "example plain text error"
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain") w.Header().Set("Content-Type", "text/plain")
w.WriteHeader(http.StatusBadRequest)
w.Write([]byte(plainBody))
})) }))
client, err := newClient( client, err := newClient(
@@ -666,11 +668,6 @@ func TestGetRunnerScaleSet(t *testing.T) {
t.Run("Multiple runner scale sets found", func(t *testing.T) { t.Run("Multiple runner scale sets found", func(t *testing.T) {
reqID := uuid.NewString() reqID := uuid.NewString()
wantErr := &ActionsError{
StatusCode: http.StatusOK,
ActivityID: reqID,
Err: fmt.Errorf("multiple runner scale sets found with name %q", scaleSetName),
}
runnerScaleSetsResp := []byte(`{"count":2,"value":[{"id":1,"name":"ScaleSet"}]}`) runnerScaleSetsResp := []byte(`{"count":2,"value":[{"id":1,"name":"ScaleSet"}]}`)
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set(headerActionsActivityID, reqID) w.Header().Set(headerActionsActivityID, reqID)
@@ -686,7 +683,8 @@ func TestGetRunnerScaleSet(t *testing.T) {
_, err = client.GetRunnerScaleSet(ctx, 1, scaleSetName) _, err = client.GetRunnerScaleSet(ctx, 1, scaleSetName)
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, wantErr.Error(), err.Error()) assert.Contains(t, err.Error(), "multiple runner scale sets found")
assert.Contains(t, err.Error(), "activity_id=\""+reqID+"\"")
}) })
} }
@@ -761,9 +759,11 @@ func TestGetRunnerScaleSetByID(t *testing.T) {
}) })
t.Run("Error when Content-Type is text/plain", func(t *testing.T) { t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
plainBody := "example plain text error"
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain") w.Header().Set("Content-Type", "text/plain")
w.WriteHeader(http.StatusBadRequest)
w.Write([]byte(plainBody))
})) }))
client, err := newClient( client, err := newClient(
@@ -876,9 +876,11 @@ func TestCreateRunnerScaleSet(t *testing.T) {
}) })
t.Run("Error when Content-Type is text/plain", func(t *testing.T) { t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
plainBody := "example plain text error"
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain") w.Header().Set("Content-Type", "text/plain")
w.WriteHeader(http.StatusBadRequest)
w.Write([]byte(plainBody))
})) }))
client, err := newClient( client, err := newClient(
@@ -890,8 +892,8 @@ func TestCreateRunnerScaleSet(t *testing.T) {
_, err = client.CreateRunnerScaleSet(ctx, &runnerScaleSet) _, err = client.CreateRunnerScaleSet(ctx, &runnerScaleSet)
require.NotNil(t, err) require.NotNil(t, err)
var expectedErr *ActionsError assert.Contains(t, err.Error(), "status=\"400 Bad Request\"")
assert.True(t, errors.As(err, &expectedErr)) assert.Contains(t, err.Error(), plainBody)
}) })
t.Run("Default retries on server error", func(t *testing.T) { t.Run("Default retries on server error", func(t *testing.T) {
+18 -6
View File
@@ -14,6 +14,11 @@ import (
"github.com/hashicorp/go-retryablehttp" "github.com/hashicorp/go-retryablehttp"
) )
const (
headerActionsActivityID = "ActivityId"
headerGitHubRequestID = "X-GitHub-Request-Id"
)
type commonClient struct { type commonClient struct {
httpClient *http.Client httpClient *http.Client
@@ -44,17 +49,24 @@ func (c *commonClient) newRetryableHTTPClient() (*retryablehttp.Client, error) {
} }
func (c *commonClient) do(req *http.Request) (*http.Response, error) { func (c *commonClient) do(req *http.Request) (*http.Response, error) {
resp, err := c.httpClient.Do(req) return sendRequest(c.httpClient, req)
if err != nil { }
return nil, fmt.Errorf("client request failed: %w", err)
}
// sendRequest ensures that the request is sent and the response body is fully read and closed.
// It trims the BOM when present in the response body.
//
// Make sure to use this function instead of http.Client.Do directly to avoid issues.
func sendRequest(c *http.Client, req *http.Request) (*http.Response, error) {
resp, err := c.Do(req)
if err != nil {
return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to send request: %w", err))
}
body, err := io.ReadAll(resp.Body) body, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to read the response body: %w", err) return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to read the response body: %w", err))
} }
if err := resp.Body.Close(); err != nil { if err := resp.Body.Close(); err != nil {
return nil, fmt.Errorf("failed to close the response body: %w", err) return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to close the response body: %w", err))
} }
body = trimByteOrderMark(body) body = trimByteOrderMark(body)
+64 -111
View File
@@ -9,137 +9,90 @@ import (
"strings" "strings"
) )
// Header names for request IDs type scalesetError string
const (
headerActionsActivityID = "ActivityId" func (e scalesetError) Error() string {
headerGitHubRequestID = "X-GitHub-Request-Id" return string(e)
}
var (
RunnerNotFoundError = scalesetError("runner not found")
RunnerExistsError = scalesetError("runner exists")
JobStillRunningError = scalesetError("job still running")
MessageQueueTokenExpiredError = scalesetError("message queue token expired")
) )
type GitHubAPIError struct {
StatusCode int
RequestID string
Err error
}
func (e *GitHubAPIError) Error() string {
return fmt.Sprintf("github api error: StatusCode %d, RequestID %q: %v", e.StatusCode, e.RequestID, e.Err)
}
func (e *GitHubAPIError) Unwrap() error {
return e.Err
}
type ActionsError struct {
ActivityID string
StatusCode int
Err error
}
func (e *ActionsError) Error() string {
return fmt.Sprintf("actions error: StatusCode %d, ActivityId %q: %v", e.StatusCode, e.ActivityID, e.Err)
}
func (e *ActionsError) Unwrap() error {
return e.Err
}
func (e *ActionsError) IsAgentNotFound() bool {
return e.isException("AgentNotFoundException")
}
func (e *ActionsError) IsJobStillRunning() bool {
return e.isException("JobStillRunningException")
}
func (e *ActionsError) IsMessageQueueTokenExpired() bool {
if e == nil {
return false
}
var err *messageQueueTokenExpiredError
return errors.As(e.Err, &err)
}
func (e *ActionsError) IsAgentExists() bool {
return e.isException("AgentExistsException")
}
func (e *ActionsError) isException(target string) bool {
if e == nil {
return false
}
if ex, ok := e.Err.(*actionsExceptionError); ok {
return strings.Contains(ex.ExceptionName, target)
}
return false
}
type actionsExceptionError struct { type actionsExceptionError struct {
ExceptionName string `json:"typeName,omitempty"` ExceptionName string `json:"typeName,omitempty"`
Message string `json:"message,omitempty"` Message string `json:"message,omitempty"`
} }
func (e *actionsExceptionError) Error() string { func (e actionsExceptionError) Error() string {
return fmt.Sprintf("%s: %s", e.ExceptionName, e.Message) return fmt.Sprintf("%s: %s", e.ExceptionName, e.Message)
} }
func ParseActionsErrorFromResponse(response *http.Response) error { // newRequestResponseError creates a detailed error message based on the HTTP request and response,
if response.ContentLength == 0 { // including parsing the response body for known error formats.
return &ActionsError{ //
ActivityID: response.Header.Get(headerActionsActivityID), // The sendRequest already parses errors using this method, so use this error if the client doesn't
StatusCode: response.StatusCode, // return an error, but the error is happening on the application logic level.
Err: errors.New("unknown exception"), //
} // Prefer creating errors using this function instead of manually constructing error messages since it automatically
// includes useful metadata like activity IDs and request IDs, and handles well-known error cases.
func newRequestResponseError(req *http.Request, resp *http.Response, err error) error {
var sb strings.Builder
fmt.Fprintf(&sb, "request %s %s failed", req.Method, req.URL.String())
if resp == nil {
return fmt.Errorf("%s: %w", sb.String(), err)
} }
body, err := io.ReadAll(response.Body) sb.WriteRune('(')
if err != nil { fmt.Fprintf(&sb, "status=%q", resp.Status)
return &ActionsError{ if resp.Header.Get(headerActionsActivityID) != "" {
ActivityID: response.Header.Get(headerActionsActivityID), fmt.Fprintf(&sb, ", activity_id=%q", resp.Header.Get(headerActionsActivityID))
StatusCode: response.StatusCode,
Err: err,
}
} }
body = trimByteOrderMark(body) if resp.Header.Get(headerGitHubRequestID) != "" {
contentType := response.Header.Get("Content-Type") fmt.Fprintf(&sb, ", github_request_id=%q", resp.Header.Get(headerGitHubRequestID))
}
sb.WriteRune(')')
if resp.Body == nil || resp.ContentLength == 0 {
return fmt.Errorf("%s: %w: unknown error", sb.String(), err)
}
body, bodyErr := io.ReadAll(resp.Body)
if bodyErr != nil {
return fmt.Errorf("%s: %w: failed to read error response body: %w", sb.String(), err, bodyErr)
}
if len(body) == 0 {
return fmt.Errorf("%s: %w: unknown error", sb.String(), err)
}
var scalesetErr scalesetError
if errors.As(err, &scalesetErr) {
return fmt.Errorf("%s: %w: %s", sb.String(), err, string(body))
}
contentType := resp.Header.Get("Content-Type")
if len(contentType) > 0 && strings.Contains(contentType, "text/plain") { if len(contentType) > 0 && strings.Contains(contentType, "text/plain") {
message := string(body) return fmt.Errorf("%s: %w: %s", sb.String(), err, string(body))
return &ActionsError{
ActivityID: response.Header.Get(headerActionsActivityID),
StatusCode: response.StatusCode,
Err: errors.New(message),
}
} }
var exception actionsExceptionError var exception actionsExceptionError
if err := json.Unmarshal(body, &exception); err != nil { if err := json.Unmarshal(body, &exception); err != nil {
return &ActionsError{ return fmt.Errorf("%s: %w: failed to unmarshal error response body: %q", sb.String(), err, string(body))
ActivityID: response.Header.Get(headerActionsActivityID),
StatusCode: response.StatusCode,
Err: err,
}
} }
return &ActionsError{ switch {
ActivityID: response.Header.Get(headerActionsActivityID), case strings.Contains(exception.ExceptionName, "AgentExistsException"):
StatusCode: response.StatusCode, return fmt.Errorf("%s: %w: %s", sb.String(), RunnerExistsError, exception.Message)
Err: &exception, case strings.Contains(exception.ExceptionName, "AgentNotFoundException"):
} return fmt.Errorf("%s: %w: %s", sb.String(), RunnerNotFoundError, exception.Message)
} case strings.Contains(exception.ExceptionName, "JobStillRunningException"):
return fmt.Errorf("%s: %w: %s", sb.String(), JobStillRunningError, exception.Message)
type messageQueueTokenExpiredError struct { default:
message string return fmt.Errorf("%s: %w: %w", sb.String(), err, exception)
}
func (e *messageQueueTokenExpiredError) Error() string {
return fmt.Sprintf("message queue token expired: %s", e.message)
}
// NewMessageQueueTokenExpiredError creates a new MessageQueueTokenExpiredError.
//
// This function is mostly used by tests.
func NewMessageQueueTokenExpiredError(message string) error {
return &messageQueueTokenExpiredError{
message: message,
} }
} }
+196 -195
View File
@@ -2,8 +2,10 @@ package scaleset
import ( import (
"errors" "errors"
"fmt"
"io" "io"
"net/http" "net/http"
"net/url"
"strings" "strings"
"testing" "testing"
@@ -11,129 +13,14 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
func TestActionsError(t *testing.T) { type readErrCloser struct{}
t.Run("contains the status code, activity ID, and error", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 404,
Err: errors.New("example error description"),
}
s := err.Error() func (readErrCloser) Read([]byte) (int, error) { return 0, fmt.Errorf("read failed") }
assert.Contains(t, s, "StatusCode 404") func (readErrCloser) Close() error { return nil }
assert.Contains(t, s, "ActivityId \"activity-id\"")
assert.Contains(t, s, "example error description")
})
t.Run("unwraps the error", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 404,
Err: &actionsExceptionError{
ExceptionName: "exception-name",
Message: "example error message",
},
}
assert.Equal(t, err.Unwrap(), err.Err)
})
t.Run("is exception is ok", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 404,
Err: &actionsExceptionError{
ExceptionName: "exception-name",
Message: "example error message",
},
}
var exception *actionsExceptionError
assert.True(t, errors.As(err, &exception))
assert.True(t, err.isException("exception-name"))
})
t.Run("is exception is not ok", func(t *testing.T) {
tt := map[string]*ActionsError{
"not an exception": {
ActivityID: "activity-id",
StatusCode: 404,
Err: errors.New("example error description"),
},
"not target exception": {
ActivityID: "activity-id",
StatusCode: 404,
Err: &actionsExceptionError{
ExceptionName: "exception-name",
Message: "example error message",
},
},
}
targetException := "target-exception"
for name, err := range tt {
t.Run(name, func(t *testing.T) {
assert.False(t, err.isException(targetException))
})
}
})
t.Run("is agent exists exception", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 404,
Err: &actionsExceptionError{
ExceptionName: "AgentExistsException",
Message: "example error message",
},
}
assert.True(t, err.IsAgentExists())
})
t.Run("is agent not found exception", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 404,
Err: &actionsExceptionError{
ExceptionName: "AgentNotFoundException",
Message: "example error message",
},
}
assert.True(t, err.IsAgentNotFound())
})
t.Run("is job still running exception", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 404,
Err: &actionsExceptionError{
ExceptionName: "JobStillRunningException",
Message: "example error message",
},
}
assert.True(t, err.IsJobStillRunning())
})
t.Run("is message queue token expired exception", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 401,
Err: &messageQueueTokenExpiredError{
message: "example error message",
},
}
assert.True(t, err.IsMessageQueueTokenExpired())
})
}
func TestActionsExceptionError(t *testing.T) { func TestActionsExceptionError(t *testing.T) {
t.Run("contains the exception name and message", func(t *testing.T) { t.Run("contains the exception name and message", func(t *testing.T) {
err := &actionsExceptionError{ err := actionsExceptionError{
ExceptionName: "exception-name", ExceptionName: "exception-name",
Message: "example error message", Message: "example error message",
} }
@@ -144,109 +31,223 @@ func TestActionsExceptionError(t *testing.T) {
}) })
} }
func TestGitHubAPIError(t *testing.T) { func TestNewRequestResponseError(t *testing.T) {
t.Run("contains the status code, request ID, and error", func(t *testing.T) { req := func(t *testing.T) *http.Request {
err := &GitHubAPIError{ t.Helper()
StatusCode: 404, u, err := url.Parse("https://example.com/org/repo")
RequestID: "request-id", require.NoError(t, err)
Err: errors.New("example error description"), return &http.Request{Method: http.MethodGet, URL: u}
} }
s := err.Error() t.Run("resp is nil", func(t *testing.T) {
assert.Contains(t, s, "StatusCode 404") base := errors.New("base")
assert.Contains(t, s, "RequestID \"request-id\"") err := newRequestResponseError(req(t), nil, base)
assert.Contains(t, s, "example error description") require.Error(t, err)
assert.Contains(t, err.Error(), "request GET https://example.com/org/repo failed")
assert.True(t, errors.Is(err, base))
}) })
t.Run("unwraps the error", func(t *testing.T) { t.Run("resp body is nil", func(t *testing.T) {
err := &GitHubAPIError{ base := errors.New("base")
StatusCode: 404, resp := &http.Response{
RequestID: "request-id", Status: "500 Internal Server Error",
Err: errors.New("example error description"), StatusCode: http.StatusInternalServerError,
ContentLength: 123,
Header: make(http.Header),
Body: nil,
} }
assert.Equal(t, err.Unwrap(), err.Err) err := newRequestResponseError(req(t), resp, base)
require.Error(t, err)
assert.Contains(t, err.Error(), "unknown error")
assert.True(t, errors.Is(err, base))
}) })
}
func TestParseActionsErrorFromResponse(t *testing.T) { t.Run("empty body returns unknown error", func(t *testing.T) {
t.Run("empty content length", func(t *testing.T) { base := errors.New("base")
response := &http.Response{ resp := &http.Response{
Status: "404 Not Found",
StatusCode: http.StatusNotFound,
ContentLength: 0, ContentLength: 0,
Header: http.Header{}, Header: make(http.Header),
StatusCode: 404,
} }
response.Header.Add(headerActionsActivityID, "activity-id") resp.Header.Set(headerActionsActivityID, "activity-id")
resp.Header.Set(headerGitHubRequestID, "request-id")
err := ParseActionsErrorFromResponse(response) err := newRequestResponseError(req(t), resp, base)
require.Error(t, err) require.Error(t, err)
assert.Equal(t, "activity-id", err.(*ActionsError).ActivityID) assert.Contains(t, err.Error(), "status=\"404 Not Found\"")
assert.Equal(t, 404, err.(*ActionsError).StatusCode) assert.Contains(t, err.Error(), "activity_id=\"activity-id\"")
assert.Equal(t, "unknown exception", err.(*ActionsError).Err.Error()) assert.Contains(t, err.Error(), "github_request_id=\"request-id\"")
assert.Contains(t, err.Error(), "unknown error")
assert.True(t, errors.Is(err, base))
}) })
t.Run("contains text plain error", func(t *testing.T) { t.Run("read body failure includes read error", func(t *testing.T) {
errorMessage := "example error message" base := errors.New("base")
response := &http.Response{ resp := &http.Response{
ContentLength: int64(len(errorMessage)), Status: "400 Bad Request",
StatusCode: 404, StatusCode: http.StatusBadRequest,
Header: http.Header{}, ContentLength: 1,
Body: io.NopCloser(strings.NewReader(errorMessage)), Header: make(http.Header),
Body: io.NopCloser(readErrCloser{}),
} }
response.Header.Add(headerActionsActivityID, "activity-id")
response.Header.Add("Content-Type", "text/plain")
err := ParseActionsErrorFromResponse(response) err := newRequestResponseError(req(t), resp, base)
require.Error(t, err) require.Error(t, err)
var actionsError *ActionsError assert.Contains(t, err.Error(), "failed to read error response body")
assert.ErrorAs(t, err, &actionsError) assert.True(t, errors.Is(err, base))
assert.Equal(t, "activity-id", actionsError.ActivityID) assert.Contains(t, err.Error(), "read failed")
assert.Equal(t, 404, actionsError.StatusCode)
assert.Equal(t, errorMessage, actionsError.Err.Error())
}) })
t.Run("contains json error", func(t *testing.T) { t.Run("unknown content length and empty body returns unknown error", func(t *testing.T) {
errorMessage := `{"typeName":"exception-name","message":"example error message"}` base := errors.New("base")
response := &http.Response{ resp := &http.Response{
ContentLength: int64(len(errorMessage)), Status: "400 Bad Request",
StatusCode: 404, StatusCode: http.StatusBadRequest,
Header: http.Header{}, ContentLength: -1,
Body: io.NopCloser(strings.NewReader(errorMessage)), Header: make(http.Header),
Body: io.NopCloser(strings.NewReader("")),
} }
response.Header.Add(headerActionsActivityID, "activity-id")
response.Header.Add("Content-Type", "application/json")
err := ParseActionsErrorFromResponse(response) err := newRequestResponseError(req(t), resp, base)
require.Error(t, err) require.Error(t, err)
var actionsError *ActionsError assert.Contains(t, err.Error(), "unknown error")
assert.ErrorAs(t, err, &actionsError) assert.True(t, errors.Is(err, base))
assert.Equal(t, "activity-id", actionsError.ActivityID)
assert.Equal(t, 404, actionsError.StatusCode)
inner, ok := actionsError.Err.(*actionsExceptionError)
require.True(t, ok)
assert.Equal(t, "exception-name", inner.ExceptionName)
assert.Equal(t, "example error message", inner.Message)
}) })
t.Run("wrapped exception error", func(t *testing.T) { t.Run("text/plain body is included", func(t *testing.T) {
errorMessage := `{"typeName":"exception-name","message":"example error message"}` base := errors.New("base")
response := &http.Response{ body := "example plain text error"
ContentLength: int64(len(errorMessage)), resp := &http.Response{
StatusCode: 404, Status: "400 Bad Request",
Header: http.Header{}, StatusCode: http.StatusBadRequest,
Body: io.NopCloser(strings.NewReader(errorMessage)), ContentLength: int64(len(body)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(body)),
} }
response.Header.Add(headerActionsActivityID, "activity-id") resp.Header.Set("Content-Type", "text/plain")
response.Header.Add("Content-Type", "application/json") resp.Header.Set(headerActionsActivityID, "activity-id")
err := ParseActionsErrorFromResponse(response) err := newRequestResponseError(req(t), resp, base)
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), body)
assert.True(t, errors.Is(err, base))
})
var actionsExceptionError *actionsExceptionError t.Run("scalesetError in error chain uses raw body (no JSON parsing)", func(t *testing.T) {
assert.ErrorAs(t, err, &actionsExceptionError) wrapped := fmt.Errorf("wrapped: %w", RunnerNotFoundError)
body := `{"typeName":"AgentExistsException","message":"should not be parsed"}`
resp := &http.Response{
Status: "404 Not Found",
StatusCode: http.StatusNotFound,
ContentLength: int64(len(body)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(body)),
}
resp.Header.Set("Content-Type", "application/json")
assert.Equal(t, "exception-name", actionsExceptionError.ExceptionName) err := newRequestResponseError(req(t), resp, wrapped)
assert.Equal(t, "example error message", actionsExceptionError.Message) require.Error(t, err)
assert.True(t, errors.Is(err, RunnerNotFoundError))
assert.Contains(t, err.Error(), body)
})
t.Run("known actions exception maps to sentinel error", func(t *testing.T) {
base := errors.New("base")
jsonBody := `{"typeName":"AgentExistsException","message":"runner already exists"}`
resp := &http.Response{
Status: "409 Conflict",
StatusCode: http.StatusConflict,
ContentLength: int64(len(jsonBody)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(jsonBody)),
}
resp.Header.Set("Content-Type", "application/json")
err := newRequestResponseError(req(t), resp, base)
require.Error(t, err)
assert.True(t, errors.Is(err, RunnerExistsError))
assert.False(t, errors.Is(err, base), "base error should not be wrapped for mapped exceptions")
assert.Contains(t, err.Error(), "runner already exists")
})
t.Run("agent not found exception maps to sentinel error", func(t *testing.T) {
base := errors.New("base")
jsonBody := `{"typeName":"AgentNotFoundException","message":"missing"}`
resp := &http.Response{
Status: "404 Not Found",
StatusCode: http.StatusNotFound,
ContentLength: int64(len(jsonBody)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(jsonBody)),
}
resp.Header.Set("Content-Type", "application/json")
err := newRequestResponseError(req(t), resp, base)
require.Error(t, err)
assert.True(t, errors.Is(err, RunnerNotFoundError))
assert.False(t, errors.Is(err, base))
assert.Contains(t, err.Error(), "missing")
})
t.Run("job still running exception maps to sentinel error", func(t *testing.T) {
base := errors.New("base")
jsonBody := `{"typeName":"JobStillRunningException","message":"still running"}`
resp := &http.Response{
Status: "409 Conflict",
StatusCode: http.StatusConflict,
ContentLength: int64(len(jsonBody)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(jsonBody)),
}
resp.Header.Set("Content-Type", "application/json")
err := newRequestResponseError(req(t), resp, base)
require.Error(t, err)
assert.True(t, errors.Is(err, JobStillRunningError))
assert.False(t, errors.Is(err, base))
assert.Contains(t, err.Error(), "still running")
})
t.Run("invalid json returns unmarshal error and includes body", func(t *testing.T) {
base := errors.New("base")
bad := "not-json"
resp := &http.Response{
Status: "400 Bad Request",
StatusCode: http.StatusBadRequest,
ContentLength: int64(len(bad)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(bad)),
}
resp.Header.Set("Content-Type", "application/json")
err := newRequestResponseError(req(t), resp, base)
require.Error(t, err)
assert.Contains(t, err.Error(), "failed to unmarshal error response body")
assert.Contains(t, err.Error(), "not-json")
assert.False(t, errors.Is(err, base), "base error is not wrapped on JSON unmarshal failures")
})
t.Run("unknown json error wraps exception", func(t *testing.T) {
base := errors.New("base")
jsonBody := `{"typeName":"SomeException","message":"example error message"}`
resp := &http.Response{
Status: "500 Internal Server Error",
StatusCode: http.StatusInternalServerError,
ContentLength: int64(len(jsonBody)),
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader(jsonBody)),
}
resp.Header.Set("Content-Type", "application/json")
err := newRequestResponseError(req(t), resp, base)
require.Error(t, err)
assert.True(t, errors.Is(err, base))
var ex actionsExceptionError
assert.True(t, errors.As(err, &ex))
assert.Equal(t, "SomeException", ex.ExceptionName)
assert.Equal(t, "example error message", ex.Message)
}) })
} }
+22 -81
View File
@@ -104,8 +104,7 @@ func (c *MessageSessionClient) GetMessage(ctx context.Context, lastMessageID int
return message, nil return message, nil
} }
expiredError := &ActionsError{} if !errors.Is(err, MessageQueueTokenExpiredError) {
if !errors.As(err, &expiredError) || !expiredError.IsMessageQueueTokenExpired() {
return nil, fmt.Errorf("failed to get next message: %w", err) return nil, fmt.Errorf("failed to get next message: %w", err)
} }
@@ -148,43 +147,23 @@ func (c *MessageSessionClient) getMessage(ctx context.Context, lastMessageID int
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode == http.StatusAccepted { switch resp.StatusCode {
case http.StatusAccepted:
return nil, nil return nil, nil
}
if resp.StatusCode != http.StatusOK { case http.StatusOK:
if resp.StatusCode != http.StatusUnauthorized { message, err := parseRunnerScaleSetMessageResponse(resp.Body)
return nil, ParseActionsErrorFromResponse(resp)
}
body, err := io.ReadAll(resp.Body)
body = trimByteOrderMark(body)
if err != nil { if err != nil {
return nil, &ActionsError{ return nil, newRequestResponseError(req, resp, fmt.Errorf("failed to parse message response: %w", err))
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: err,
}
} }
return nil, &ActionsError{ return message, nil
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: &messageQueueTokenExpiredError{
message: string(body),
},
}
}
message, err := parseRunnerScaleSetMessageResponse(resp.Body) case http.StatusUnauthorized:
if err != nil { return nil, newRequestResponseError(req, resp, MessageQueueTokenExpiredError)
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
return message, nil default:
return nil, newRequestResponseError(req, resp, fmt.Errorf("unexpected status code %s", resp.Status))
}
} }
// DeleteMessage deletes a message from the runner scale set message queue. // DeleteMessage deletes a message from the runner scale set message queue.
@@ -199,8 +178,7 @@ func (c *MessageSessionClient) DeleteMessage(ctx context.Context, messageID int)
return nil return nil
} }
expiredError := &ActionsError{} if !errors.Is(err, MessageQueueTokenExpiredError) {
if !errors.As(err, &expiredError) || !expiredError.IsMessageQueueTokenExpired() {
return fmt.Errorf("failed to delete message: %w", err) return fmt.Errorf("failed to delete message: %w", err)
} }
@@ -239,25 +217,10 @@ func (c *MessageSessionClient) deleteMessage(ctx context.Context, messageID int)
} }
if resp.StatusCode != http.StatusUnauthorized { if resp.StatusCode != http.StatusUnauthorized {
return ParseActionsErrorFromResponse(resp) return newRequestResponseError(req, resp, fmt.Errorf("unexpected status code %s", resp.Status))
} }
body, err := io.ReadAll(resp.Body) return newRequestResponseError(req, resp, MessageQueueTokenExpiredError)
body = trimByteOrderMark(body)
if err != nil {
return &ActionsError{
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: err,
}
}
return &ActionsError{
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: &messageQueueTokenExpiredError{
message: string(body),
},
}
} }
func (c *MessageSessionClient) Session() RunnerScaleSetSession { func (c *MessageSessionClient) Session() RunnerScaleSetSession {
@@ -284,39 +247,17 @@ func (c *MessageSessionClient) doSessionRequest(ctx context.Context, method, pat
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode == expectedResponseStatusCode { if resp.StatusCode != expectedResponseStatusCode {
if responseUnmarshalTarget == nil { return newRequestResponseError(req, resp, fmt.Errorf("unexpected status code %s", resp.Status))
return nil }
}
if err := json.NewDecoder(resp.Body).Decode(responseUnmarshalTarget); err != nil {
return &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
if responseUnmarshalTarget == nil {
return nil return nil
} }
if resp.StatusCode >= 400 && resp.StatusCode < 500 { if err := json.NewDecoder(resp.Body).Decode(responseUnmarshalTarget); err != nil {
return ParseActionsErrorFromResponse(resp) return newRequestResponseError(req, resp, fmt.Errorf("failed to unmarshal response body: %w", err))
} }
body, err := io.ReadAll(resp.Body) return nil
body = trimByteOrderMark(body)
if err != nil {
return &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
return fmt.Errorf("unexpected status code: %w", &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: errors.New(string(body)),
})
} }
+31 -42
View File
@@ -62,7 +62,7 @@ func TestCreateMessageSession(t *testing.T) {
assert.Equal(t, want, session) assert.Equal(t, want, session)
}) })
t.Run("CreateMessageSession unmarshals errors into ActionsError", func(t *testing.T) { t.Run("CreateMessageSession includes actions exception details", func(t *testing.T) {
owner := "foo" owner := "foo"
runnerScaleSet := RunnerScaleSet{ runnerScaleSet := RunnerScaleSet{
ID: 1, ID: 1,
@@ -71,15 +71,6 @@ func TestCreateMessageSession(t *testing.T) {
RunnerSetting: RunnerSetting{}, RunnerSetting: RunnerSetting{},
} }
want := &ActionsError{
ActivityID: exampleRequestID,
StatusCode: http.StatusBadRequest,
Err: &actionsExceptionError{
ExceptionName: "CSharpExceptionNameHere",
Message: "could not do something",
},
}
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/json") w.Header().Set("Content-Type", "application/json")
w.Header().Set(headerActionsActivityID, exampleRequestID) w.Header().Set(headerActionsActivityID, exampleRequestID)
@@ -97,11 +88,14 @@ func TestCreateMessageSession(t *testing.T) {
sessionClient, err := client.MessageSessionClient(context.Background(), runnerScaleSet.ID, owner) sessionClient, err := client.MessageSessionClient(context.Background(), runnerScaleSet.ID, owner)
assert.Nil(t, sessionClient) assert.Nil(t, sessionClient)
require.Error(t, err)
assert.Contains(t, err.Error(), "status=\"400 Bad Request\"")
assert.Contains(t, err.Error(), "activity_id=\""+exampleRequestID+"\"")
errorTypeForComparison := &ActionsError{} var ex actionsExceptionError
assert.ErrorAs(t, err, &errorTypeForComparison) assert.True(t, errors.As(err, &ex))
assert.Equal(t, "CSharpExceptionNameHere", ex.ExceptionName)
assert.Equal(t, want, errorTypeForComparison) assert.Equal(t, "could not do something", ex.Message)
}) })
t.Run("CreateMessageSession call is retried the correct amount of times", func(t *testing.T) { t.Run("CreateMessageSession call is retried the correct amount of times", func(t *testing.T) {
@@ -285,10 +279,7 @@ func TestGetMessage(t *testing.T) {
msg, err := sessionClient.GetMessage(ctx, 0, 10) msg, err := sessionClient.GetMessage(ctx, 0, 10)
assert.Nil(t, msg) assert.Nil(t, msg)
assert.ErrorIs(t, err, MessageQueueTokenExpiredError, "expected error to be MessageQueueTokenExpiredError but got: %v", err)
var expectedErr *ActionsError
require.ErrorAs(t, err, &expectedErr)
assert.True(t, expectedErr.IsMessageQueueTokenExpired(), "expected error to be of type MessageQueueTokenExpiredError but got: %v", err)
}) })
t.Run("Message token refreshed", func(t *testing.T) { t.Run("Message token refreshed", func(t *testing.T) {
@@ -345,10 +336,6 @@ func TestGetMessage(t *testing.T) {
}) })
t.Run("Status code not found", func(t *testing.T) { t.Run("Status code not found", func(t *testing.T) {
want := ActionsError{
Err: errors.New("unknown exception"),
StatusCode: 404,
}
var handleSessionRequest http.HandlerFunc var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") { if strings.HasSuffix(r.URL.Path, "sessions") {
@@ -371,21 +358,22 @@ func TestGetMessage(t *testing.T) {
msg, err := sessionClient.GetMessage(ctx, 0, 10) msg, err := sessionClient.GetMessage(ctx, 0, 10)
assert.Nil(t, msg) assert.Nil(t, msg)
var got *ActionsError require.Error(t, err)
require.ErrorAs(t, err, &got) assert.Contains(t, err.Error(), "status=\"404 Not Found\"")
assert.Equal(t, want.StatusCode, got.StatusCode) assert.Contains(t, err.Error(), "unknown error")
assert.Equal(t, want.Err.Error(), got.Err.Error())
}) })
t.Run("Error when Content-Type is text/plain", func(t *testing.T) { t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
plainBody := "example plain text error"
var handleSessionRequest http.HandlerFunc var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") { if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r) handleSessionRequest(w, r)
return return
} }
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain") w.Header().Set("Content-Type", "text/plain")
w.WriteHeader(http.StatusBadRequest)
w.Write([]byte(plainBody))
})) }))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession()) handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
@@ -402,9 +390,12 @@ func TestGetMessage(t *testing.T) {
msg, err := sessionClient.GetMessage(ctx, 0, 10) msg, err := sessionClient.GetMessage(ctx, 0, 10)
assert.Nil(t, msg) assert.Nil(t, msg)
assert.NotNil(t, err) assert.NotNil(t, err)
assert.Contains(t, err.Error(), "status=\"400 Bad Request\"")
assert.Contains(t, err.Error(), plainBody)
}) })
t.Run("Capacity error handling", func(t *testing.T) { t.Run("Capacity error handling", func(t *testing.T) {
plainBody := "capacity error"
var handleSessionRequest http.HandlerFunc var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") { if strings.HasSuffix(r.URL.Path, "sessions") {
@@ -415,9 +406,9 @@ func TestGetMessage(t *testing.T) {
c, err := strconv.Atoi(hc) c, err := strconv.Atoi(hc)
require.NoError(t, err) require.NoError(t, err)
assert.GreaterOrEqual(t, c, 0) assert.GreaterOrEqual(t, c, 0)
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain") w.Header().Set("Content-Type", "text/plain")
w.WriteHeader(http.StatusBadRequest)
w.Write([]byte(plainBody))
})) }))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession()) handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
@@ -434,9 +425,8 @@ func TestGetMessage(t *testing.T) {
msg, err := sessionClient.GetMessage(ctx, 0, 0) msg, err := sessionClient.GetMessage(ctx, 0, 0)
assert.Nil(t, msg) assert.Nil(t, msg)
assert.Error(t, err) assert.Error(t, err)
var expectedErr *ActionsError assert.Contains(t, err.Error(), "status=\"400 Bad Request\"")
assert.ErrorAs(t, err, &expectedErr) assert.Contains(t, err.Error(), plainBody)
assert.Equal(t, http.StatusBadRequest, expectedErr.StatusCode)
}) })
} }
@@ -506,9 +496,7 @@ func TestDeleteMessage(t *testing.T) {
err = sessionClient.DeleteMessage(ctx, 0) err = sessionClient.DeleteMessage(ctx, 0)
require.NotNil(t, err) require.NotNil(t, err)
var expectedErr *ActionsError assert.ErrorIs(t, err, MessageQueueTokenExpiredError, "expected error to be MessageQueueTokenExpiredError but got: %v", err)
require.ErrorAs(t, err, &expectedErr)
assert.True(t, expectedErr.IsMessageQueueTokenExpired())
}) })
t.Run("message token refreshed", func(t *testing.T) { t.Run("message token refreshed", func(t *testing.T) {
@@ -563,14 +551,16 @@ func TestDeleteMessage(t *testing.T) {
}) })
t.Run("Error when Content-Type is text/plain", func(t *testing.T) { t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
plainBody := "example plain text error"
var handleSessionRequest http.HandlerFunc var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") { if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r) handleSessionRequest(w, r)
return return
} }
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain") w.Header().Set("Content-Type", "text/plain")
w.WriteHeader(http.StatusBadRequest)
w.Write([]byte(plainBody))
})) }))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession()) handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
@@ -586,10 +576,9 @@ func TestDeleteMessage(t *testing.T) {
err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID) err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID)
require.NotNil(t, err) require.NotNil(t, err)
var expectedErr *ActionsError assert.Contains(t, err.Error(), "status=\"400 Bad Request\"")
assert.True(t, errors.As(err, &expectedErr)) assert.Contains(t, err.Error(), plainBody)
}, })
)
t.Run("Default retries on server error", func(t *testing.T) { t.Run("Default retries on server error", func(t *testing.T) {
actualRetry := 0 actualRetry := 0
@@ -656,7 +645,7 @@ func TestDeleteMessage(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID+1) err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID+1)
var expectedErr *ActionsError require.Error(t, err)
require.True(t, errors.As(err, &expectedErr)) assert.Contains(t, err.Error(), "unexpected status code")
}) })
} }