Hide session handling and generate session client from the main client (#42)

This commit is contained in:
Nikola Jokic
2026-01-05 18:19:41 +01:00
committed by GitHub
parent ebe67dba45
commit 1208e1216e
11 changed files with 1563 additions and 2092 deletions
+103 -400
View File
@@ -4,25 +4,19 @@ package scaleset
import (
"bytes"
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"maps"
"net/http"
"net/url"
"runtime/debug"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/golang-jwt/jwt/v4"
"github.com/google/uuid"
"github.com/hashicorp/go-retryablehttp"
)
@@ -31,13 +25,14 @@ const (
scaleSetEndpoint = "_apis/runtime/runnerscalesets"
)
var (
packageVersion string
commitSHA string
)
var buildInfo clientBuildInfo
func init() {
packageVersion, commitSHA = detectModuleVersionAndCommit()
packageVersion, commitSHA := detectModuleVersionAndCommit()
buildInfo = clientBuildInfo{
version: packageVersion,
commitSHA: commitSHA,
}
}
// HeaderScaleSetMaxCapacity is used to propagate the scale set max
@@ -46,37 +41,17 @@ const HeaderScaleSetMaxCapacity = "X-ScaleSetMaxCapacity"
// Client implements a GitHub Actions Scale Set client.
type Client struct {
httpClient *http.Client
mu sync.Mutex // guards every public call
actionsMu sync.Mutex // guards actionsService fields
// admin session info
actionsServiceAdminToken string
actionsServiceAdminTokenExpiresAt time.Time
actionsServiceURL string
retryMax int
retryWaitMax time.Duration
creds actionsAuth
config gitHubConfig
creds *actionsAuth
config *gitHubConfig
logger *slog.Logger
buildInfo clientBuildInfo
// systemInfoMu guards setting system info.
systemInfoMu sync.Mutex
systemInfo SystemInfo
// userAgent is computed based on buildInfo and systemInfo.
// userAgent should be re-computed every time client.SetSystemInfo
// is called.
//
// On every call, load the userAgent first locally so we can
// avoid lock-unlock on every call.
userAgent atomic.Pointer[string]
rootCAs *x509.CertPool
tlsInsecureSkipVerify bool
proxyFunc ProxyFunc
commonClient
}
type clientBuildInfo struct {
@@ -94,10 +69,12 @@ type debugInfo struct {
// including whether a proxy or custom root CA is configured, and the current system info.
// This method is intended for diagnostic and troubleshooting purposes.
func (c *Client) DebugInfo() string {
c.mu.Lock()
defer c.mu.Unlock()
info := debugInfo{
HasProxy: c.proxyFunc != nil,
HasRootCA: c.rootCAs != nil,
SystemInfo: *c.userAgent.Load(),
SystemInfo: c.userAgent,
}
b, _ := json.Marshal(info)
@@ -138,9 +115,6 @@ type actionsAuth struct {
// ProxyFunc defines the function signature for a proxy function.
type ProxyFunc func(req *http.Request) (*url.URL, error)
// Option defines a functional option for configuring the Client.
type Option func(*Client)
// SystemInfo contains information about the system that uses the
// scaleset client.
//
@@ -164,48 +138,6 @@ type SystemInfo struct {
Subsystem string `json:"subsystem"`
}
// WithLogger sets a custom logger for the Client.
func WithLogger(logger slog.Logger) Option {
return func(c *Client) {
c.logger = &logger
}
}
// WithRetryMax sets the maximum number of retries for the Client.
func WithRetryMax(retryMax int) Option {
return func(c *Client) {
c.retryMax = retryMax
}
}
// WithRetryWaitMax sets the maximum wait time between retries for the Client.
func WithRetryWaitMax(retryWaitMax time.Duration) Option {
return func(c *Client) {
c.retryWaitMax = retryWaitMax
}
}
// WithRootCAs sets custom root certificate authorities for the Client.
func WithRootCAs(rootCAs *x509.CertPool) Option {
return func(c *Client) {
c.rootCAs = rootCAs
}
}
// WithoutTLSVerify disables TLS certificate verification for the Client.
func WithoutTLSVerify() Option {
return func(c *Client) {
c.tlsInsecureSkipVerify = true
}
}
// WithProxy sets a custom proxy function for the Client.
func WithProxy(proxyFunc ProxyFunc) Option {
return func(c *Client) {
c.proxyFunc = proxyFunc
}
}
type ClientWithGitHubAppConfig struct {
GitHubConfigURL string
GitHubAppAuth GitHubAppAuth
@@ -213,18 +145,18 @@ type ClientWithGitHubAppConfig struct {
}
// NewClientWithGitHubApp creates a new Client using GitHub App credentials.
func NewClientWithGitHubApp(config ClientWithGitHubAppConfig, options ...Option) (*Client, error) {
creds := &actionsAuth{
app: &config.GitHubAppAuth,
}
func NewClientWithGitHubApp(config ClientWithGitHubAppConfig, options ...HTTPOption) (*Client, error) {
return newClient(
config.SystemInfo,
config.GitHubConfigURL,
creds,
actionsAuth{
app: &config.GitHubAppAuth,
},
options...,
)
}
// NewClientWithPersonalAccessTokenConfig contains the configuration for creating a new Client using a personal access token.
type NewClientWithPersonalAccessTokenConfig struct {
GitHubConfigURL string
PersonalAccessToken string
@@ -232,99 +164,58 @@ type NewClientWithPersonalAccessTokenConfig struct {
}
// NewClientWithPersonalAccessToken creates a new Client using a personal access token.
func NewClientWithPersonalAccessToken(config NewClientWithPersonalAccessTokenConfig, options ...Option) (*Client, error) {
creds := &actionsAuth{
token: config.PersonalAccessToken,
}
func NewClientWithPersonalAccessToken(config NewClientWithPersonalAccessTokenConfig, options ...HTTPOption) (*Client, error) {
return newClient(
config.SystemInfo,
config.GitHubConfigURL,
creds,
actionsAuth{
token: config.PersonalAccessToken,
},
options...,
)
}
func newClient(systemInfo SystemInfo, githubConfigURL string, creds *actionsAuth, options ...Option) (*Client, error) {
func newClient(systemInfo SystemInfo, githubConfigURL string, creds actionsAuth, options ...HTTPOption) (*Client, error) {
config, err := parseGitHubConfigFromURL(githubConfigURL)
if err != nil {
return nil, fmt.Errorf("failed to parse githubConfigURL: %w", err)
}
ac := &Client{
creds: creds,
config: config,
logger: slog.New(slog.DiscardHandler),
// retryablehttp defaults
httpClientOption := httpClientOption{
retryMax: 4,
retryWaitMax: 30 * time.Second,
buildInfo: clientBuildInfo{
version: packageVersion,
commitSHA: commitSHA,
},
}
ac.SetSystemInfo(systemInfo)
httpClientOption.defaults()
for _, option := range options {
option(ac)
option(&httpClientOption)
}
retryClient, err := ac.newRetryableHTTPClient()
if err != nil {
return nil, fmt.Errorf("failed to create retryable HTTP client: %w", err)
commonClient := newCommonClient(
systemInfo,
httpClientOption,
)
ac := &Client{
creds: creds,
config: *config,
commonClient: *commonClient,
}
ac.httpClient = retryClient.StandardClient()
return ac, nil
}
func (c *Client) newRetryableHTTPClient() (*retryablehttp.Client, error) {
retryClient := retryablehttp.NewClient()
retryClient.Logger = c.logger
retryClient.RetryMax = c.retryMax
retryClient.RetryWaitMax = c.retryWaitMax
retryClient.HTTPClient.Timeout = 5 * time.Minute // timeout must be > 1m to accomodate long polling
transport, ok := retryClient.HTTPClient.Transport.(*http.Transport)
if !ok {
// this should always be true, because retryablehttp.NewClient() uses
// cleanhttp.DefaultPooledTransport()
return nil, fmt.Errorf("failed to get http transport from retryablehttp client")
}
if transport.TLSClientConfig == nil {
transport.TLSClientConfig = &tls.Config{}
}
if c.rootCAs != nil {
transport.TLSClientConfig.RootCAs = c.rootCAs
}
if c.tlsInsecureSkipVerify {
transport.TLSClientConfig.InsecureSkipVerify = true
}
transport.Proxy = c.proxyFunc
retryClient.HTTPClient.Transport = transport
return retryClient, nil
}
// SetSystemInfo updates the information about the system.
func (c *Client) SetSystemInfo(info SystemInfo) {
c.systemInfoMu.Lock()
defer c.systemInfoMu.Unlock()
c.systemInfo = info
c.setUserAgent()
c.mu.Lock()
defer c.mu.Unlock()
c.setSystemInfo(info)
}
// SystemInfo returns the current system info that the client
// has configured.
func (c *Client) SystemInfo() SystemInfo {
c.systemInfoMu.Lock()
defer c.systemInfoMu.Unlock()
c.mu.Lock()
defer c.mu.Unlock()
return c.systemInfo
}
@@ -334,36 +225,6 @@ type userAgent struct {
BuildCommitSHA string `json:"build_commit_sha"`
}
func (c *Client) setUserAgent() {
b, _ := json.Marshal(userAgent{
SystemInfo: c.systemInfo,
BuildVersion: c.buildInfo.version,
BuildCommitSHA: c.buildInfo.commitSHA,
})
userAgent := string(b)
c.userAgent.Store(&userAgent)
}
func (c *Client) do(req *http.Request) (*http.Response, error) {
resp, err := c.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("client request failed: %w", err)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("failed to read the response body: %w", err)
}
err = resp.Body.Close()
if err != nil {
return nil, fmt.Errorf("failed to close the response body: %w", err)
}
body = trimByteOrderMark(body)
resp.Body = io.NopCloser(bytes.NewReader(body))
return resp, nil
}
func (c *Client) newGitHubAPIRequest(ctx context.Context, method, path string, body io.Reader) (*http.Request, error) {
u := c.config.gitHubAPIURL(path)
req, err := http.NewRequestWithContext(ctx, method, u.String(), body)
@@ -371,7 +232,7 @@ func (c *Client) newGitHubAPIRequest(ctx context.Context, method, path string, b
return nil, fmt.Errorf("failed to create new GitHub API request: %w", err)
}
req.Header.Set("User-Agent", *c.userAgent.Load())
req.Header.Set("User-Agent", c.userAgent)
return req, nil
}
@@ -412,13 +273,16 @@ func (c *Client) newActionsServiceRequest(ctx context.Context, method, path stri
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", c.actionsServiceAdminToken))
req.Header.Set("User-Agent", *c.userAgent.Load())
req.Header.Set("User-Agent", c.userAgent)
return req, nil
}
// GetRunnerScaleSet fetches a runner scale set by its name within a runner group.
func (c *Client) GetRunnerScaleSet(ctx context.Context, runnerGroupID int, runnerScaleSetName string) (*RunnerScaleSet, error) {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/%s?runnerGroupId=%d&name=%s", scaleSetEndpoint, runnerGroupID, runnerScaleSetName)
req, err := c.newActionsServiceRequest(ctx, http.MethodGet, path, nil)
if err != nil {
@@ -460,6 +324,9 @@ func (c *Client) GetRunnerScaleSet(ctx context.Context, runnerGroupID int, runne
// GetRunnerScaleSetByID fetches a runner scale set by its ID.
func (c *Client) GetRunnerScaleSetByID(ctx context.Context, runnerScaleSetID int) (*RunnerScaleSet, error) {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/%s/%d", scaleSetEndpoint, runnerScaleSetID)
req, err := c.newActionsServiceRequest(ctx, http.MethodGet, path, nil)
if err != nil {
@@ -489,6 +356,9 @@ func (c *Client) GetRunnerScaleSetByID(ctx context.Context, runnerScaleSetID int
// GetRunnerGroupByName fetches a runner group by its name.
func (c *Client) GetRunnerGroupByName(ctx context.Context, runnerGroup string) (*RunnerGroup, error) {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/_apis/runtime/runnergroups/?groupName=%s", runnerGroup)
req, err := c.newActionsServiceRequest(ctx, http.MethodGet, path, nil)
if err != nil {
@@ -547,6 +417,9 @@ func (c *Client) GetRunnerGroupByName(ctx context.Context, runnerGroup string) (
// CreateRunnerScaleSet creates a new runner scale set. Note that runner scale set names must be unique within a runner group.
func (c *Client) CreateRunnerScaleSet(ctx context.Context, runnerScaleSet *RunnerScaleSet) (*RunnerScaleSet, error) {
c.mu.Lock()
defer c.mu.Unlock()
body, err := json.Marshal(runnerScaleSet)
if err != nil {
return nil, fmt.Errorf("failed to marshal runner scale set: %w", err)
@@ -578,6 +451,9 @@ func (c *Client) CreateRunnerScaleSet(ctx context.Context, runnerScaleSet *Runne
// UpdateRunnerScaleSet updates an existing runner scale set.
func (c *Client) UpdateRunnerScaleSet(ctx context.Context, runnerScaleSetID int, runnerScaleSet *RunnerScaleSet) (*RunnerScaleSet, error) {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("%s/%d", scaleSetEndpoint, runnerScaleSetID)
body, err := json.Marshal(runnerScaleSet)
@@ -612,6 +488,9 @@ func (c *Client) UpdateRunnerScaleSet(ctx context.Context, runnerScaleSetID int,
// DeleteRunnerScaleSet deletes a runner scale set by its ID.
func (c *Client) DeleteRunnerScaleSet(ctx context.Context, runnerScaleSetID int) error {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/%s/%d", scaleSetEndpoint, runnerScaleSetID)
req, err := c.newActionsServiceRequest(ctx, http.MethodDelete, path, nil)
if err != nil {
@@ -631,78 +510,7 @@ func (c *Client) DeleteRunnerScaleSet(ctx context.Context, runnerScaleSetID int)
return nil
}
// GetMessage fetches a message from the runner scale set message queue. If there are no messages available, it returns (nil, nil).
// Unless a message is deleted after being processed (using DeleteMessage), it will be returned again in subsequent calls.
// If the current session token is expired, it returns a MessageQueueTokenExpiredError.
// In these cases the caller should refresh the session with RefreshMessageSession.
func (c *Client) GetMessage(ctx context.Context, messageQueueURL, messageQueueAccessToken string, lastMessageID int, maxCapacity int) (*RunnerScaleSetMessage, error) {
u, err := url.Parse(messageQueueURL)
if err != nil {
return nil, fmt.Errorf("failed to parse message queue url: %w", err)
}
if lastMessageID > 0 {
q := u.Query()
q.Set("lastMessageId", strconv.Itoa(lastMessageID))
u.RawQuery = q.Encode()
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
if err != nil {
return nil, fmt.Errorf("failed to create new request with context: %w", err)
}
req.Header.Set("Accept", "application/json; api-version=6.0-preview")
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", messageQueueAccessToken))
req.Header.Set("User-Agent", *c.userAgent.Load())
req.Header.Set(HeaderScaleSetMaxCapacity, strconv.Itoa(maxCapacity))
resp, err := c.do(req)
if err != nil {
return nil, fmt.Errorf("failed to issue the request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusAccepted {
return nil, nil
}
if resp.StatusCode != http.StatusOK {
if resp.StatusCode != http.StatusUnauthorized {
return nil, ParseActionsErrorFromResponse(resp)
}
body, err := io.ReadAll(resp.Body)
body = trimByteOrderMark(body)
if err != nil {
return nil, &ActionsError{
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: err,
}
}
return nil, &ActionsError{
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: &messageQueueTokenExpiredError{
message: string(body),
},
}
}
message, err := c.parseRunnerScaleSetMessageResponse(resp.Body)
if err != nil {
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
return message, nil
}
func (c *Client) parseRunnerScaleSetMessageResponse(respBody io.Reader) (*RunnerScaleSetMessage, error) {
func parseRunnerScaleSetMessageResponse(respBody io.Reader) (*RunnerScaleSetMessage, error) {
var messageResponse runnerScaleSetMessageResponse
if err := json.NewDecoder(respBody).Decode(&messageResponse); err != nil {
return nil, fmt.Errorf("failed to decode runner scale set message response: %w", err)
@@ -762,157 +570,46 @@ func (c *Client) parseRunnerScaleSetMessageResponse(respBody io.Reader) (*Runner
return message, nil
}
// DeleteMessage deletes a message from the runner scale set message queue.
// This should typically be done after processing the message and acts as an acknowledgment.
// If the current session token is expired, it returns a MessageQueueTokenExpiredError.
// In these cases the caller should refresh the session with RefreshMessageSession.
func (c *Client) DeleteMessage(ctx context.Context, messageQueueURL, messageQueueAccessToken string, messageID int) error {
u, err := url.Parse(messageQueueURL)
if err != nil {
return fmt.Errorf("failed to parse message queue url: %w", err)
// MessageSessionClient creates a new MessageSessionClient for the specified runner scale set ID and owner.
//
// It exposes client options that could be overwritten, providing ability to specify different retry policies or TLS settings, proxy, etc.
func (c *Client) MessageSessionClient(ctx context.Context, runnerScaleSetID int, owner string, options ...HTTPOption) (*MessageSessionClient, error) {
c.mu.Lock()
defer c.mu.Unlock()
// Copy original options
httpClientOption := c.httpClientOption
// Apply overwrites
for _, option := range options {
option(&httpClientOption)
}
// Instantiate a new common client
commonClient := newCommonClient(
c.systemInfo,
httpClientOption,
)
client := &MessageSessionClient{
innerClient: c,
commonClient: commonClient,
owner: owner,
scaleSetID: runnerScaleSetID,
session: nil,
}
u.Path = fmt.Sprintf("%s/%d", u.Path, messageID)
req, err := http.NewRequestWithContext(ctx, http.MethodDelete, u.String(), nil)
if err != nil {
return fmt.Errorf("failed to create new request with context: %w", err)
if err := client.createMessageSession(ctx); err != nil {
return nil, fmt.Errorf("failed to create message session: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", messageQueueAccessToken))
req.Header.Set("User-Agent", *c.userAgent.Load())
resp, err := c.do(req)
if err != nil {
return fmt.Errorf("failed to issue the request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusNoContent {
return nil
}
if resp.StatusCode != http.StatusUnauthorized {
return ParseActionsErrorFromResponse(resp)
}
body, err := io.ReadAll(resp.Body)
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),
},
}
}
// CreateMessageSession creates a new message session for the specified runner scale set.
// The resulting session contains the message queue URL and access token used to GetMessage.
func (c *Client) CreateMessageSession(ctx context.Context, runnerScaleSetID int, owner string) (*RunnerScaleSetSession, error) {
path := fmt.Sprintf("/%s/%d/sessions", scaleSetEndpoint, runnerScaleSetID)
newSession := &RunnerScaleSetSession{
OwnerName: owner,
}
requestData, err := json.Marshal(newSession)
if err != nil {
return nil, fmt.Errorf("failed to marshal new session: %w", err)
}
var createdSession RunnerScaleSetSession
if err = c.doSessionRequest(
ctx,
http.MethodPost,
path,
bytes.NewBuffer(requestData),
http.StatusOK,
&createdSession,
); err != nil {
return nil, fmt.Errorf("failed to do the session request: %w", err)
}
return &createdSession, nil
}
// DeleteMessageSession deletes a message session for the specified runner scale set.
func (c *Client) DeleteMessageSession(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) error {
path := fmt.Sprintf("/%s/%d/sessions/%s", scaleSetEndpoint, runnerScaleSetID, sessionID.String())
return c.doSessionRequest(ctx, http.MethodDelete, path, nil, http.StatusNoContent, nil)
}
// RefreshMessageSession refreshes a message session for the specified runner scale set.
// This should be used when a MessageQueueTokenExpiredError is encountered.
func (c *Client) RefreshMessageSession(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) (*RunnerScaleSetSession, error) {
path := fmt.Sprintf("/%s/%d/sessions/%s", scaleSetEndpoint, runnerScaleSetID, sessionID.String())
refreshedSession := &RunnerScaleSetSession{}
if err := c.doSessionRequest(ctx, http.MethodPatch, path, nil, http.StatusOK, refreshedSession); err != nil {
return nil, fmt.Errorf("failed to do the session request: %w", err)
}
return refreshedSession, nil
}
func (c *Client) doSessionRequest(ctx context.Context, method, path string, requestData io.Reader, expectedResponseStatusCode int, responseUnmarshalTarget any) error {
req, err := c.newActionsServiceRequest(ctx, method, path, requestData)
if err != nil {
return fmt.Errorf("failed to create new actions service request: %w", err)
}
resp, err := c.do(req)
if err != nil {
return fmt.Errorf("failed to issue the request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == expectedResponseStatusCode {
if responseUnmarshalTarget == nil {
return nil
}
if err := json.NewDecoder(resp.Body).Decode(responseUnmarshalTarget); err != nil {
return &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
return nil
}
if resp.StatusCode >= 400 && resp.StatusCode < 500 {
return ParseActionsErrorFromResponse(resp)
}
body, err := io.ReadAll(resp.Body)
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)),
})
return client, nil
}
// GenerateJitRunnerConfig generates a JIT runner configuration for the specified runner scale set. This returns an encoded
// configuration that can be used to directly start a new runner.
func (c *Client) GenerateJitRunnerConfig(ctx context.Context, jitRunnerSetting *RunnerScaleSetJitRunnerSetting, scaleSetID int) (*RunnerScaleSetJitRunnerConfig, error) {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/%s/%d/generatejitconfig", scaleSetEndpoint, scaleSetID)
body, err := json.Marshal(jitRunnerSetting)
@@ -948,6 +645,9 @@ func (c *Client) GenerateJitRunnerConfig(ctx context.Context, jitRunnerSetting *
// GetRunner fetches a runner by its ID. This can be used to check if a runner exists.
func (c *Client) GetRunner(ctx context.Context, runnerID int) (*RunnerReference, error) {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/%s/%d", runnerEndpoint, runnerID)
req, err := c.newActionsServiceRequest(ctx, http.MethodGet, path, nil)
@@ -979,6 +679,9 @@ func (c *Client) GetRunner(ctx context.Context, runnerID int) (*RunnerReference,
// GetRunnerByName fetches a runner by its name. This can be used to check if a runner exists.
func (c *Client) GetRunnerByName(ctx context.Context, runnerName string) (*RunnerReference, error) {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/%s?agentName=%s", runnerEndpoint, runnerName)
req, err := c.newActionsServiceRequest(ctx, http.MethodGet, path, nil)
@@ -1022,6 +725,9 @@ func (c *Client) GetRunnerByName(ctx context.Context, runnerName string) (*Runne
// RemoveRunner removes a runner by its ID.
func (c *Client) RemoveRunner(ctx context.Context, runnerID int64) error {
c.mu.Lock()
defer c.mu.Unlock()
path := fmt.Sprintf("/%s/%d", runnerEndpoint, runnerID)
req, err := c.newActionsServiceRequest(ctx, http.MethodDelete, path, nil)
@@ -1048,7 +754,7 @@ type registrationToken struct {
}
func (c *Client) getRunnerRegistrationToken(ctx context.Context) (*registrationToken, error) {
path, err := createRegistrationTokenPath(c.config)
path, err := createRegistrationTokenPath(&c.config)
if err != nil {
return nil, fmt.Errorf("failed to create registration token path: %w", err)
}
@@ -1323,9 +1029,6 @@ func actionsServiceAdminTokenExpiresAt(jwtToken string) (time.Time, error) {
}
func (c *Client) updateTokenIfNeeded(ctx context.Context) error {
c.actionsMu.Lock()
defer c.actionsMu.Unlock()
aboutToExpire := time.Now().Add(60 * time.Second).After(c.actionsServiceAdminTokenExpiresAt)
if !aboutToExpire && !c.actionsServiceAdminTokenExpiresAt.IsZero() {
return nil
+52 -656
View File
@@ -14,7 +14,6 @@ import (
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"testing"
"time"
@@ -24,7 +23,6 @@ import (
"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/http/httpproxy"
)
const exampleRequestID = "5ddf2050-dae0-013c-9159-04421ad31b68"
@@ -77,7 +75,7 @@ func TestNewGitHubAPIRequest(t *testing.T) {
client, err := newClient(
testSystemInfo,
scenario.configURL,
nil,
actionsAuth{token: "token"},
)
require.NoError(t, err)
@@ -91,7 +89,7 @@ func TestNewGitHubAPIRequest(t *testing.T) {
client, err := newClient(
testSystemInfo,
"http://localhost/my-org",
nil,
actionsAuth{token: "token"},
)
require.NoError(t, err)
@@ -111,7 +109,7 @@ func TestNewGitHubAPIRequest(t *testing.T) {
func TestNewActionsServiceRequest(t *testing.T) {
ctx := context.Background()
defaultCreds := &actionsAuth{token: "token"}
defaultCreds := actionsAuth{token: "token"}
t.Run("manages authentication", func(t *testing.T) {
t.Run("client is brand new", func(t *testing.T) {
@@ -289,14 +287,14 @@ func TestNewActionsServiceRequest(t *testing.T) {
req, err := client.newActionsServiceRequest(ctx, http.MethodGet, "/my/path", nil)
require.NoError(t, err)
assert.Equal(t, *client.userAgent.Load(), req.Header.Get("User-Agent"))
assert.Equal(t, client.userAgent, req.Header.Get("User-Agent"))
assert.Equal(t, "application/json", req.Header.Get("Content-Type"))
})
}
func TestGetRunner(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -354,7 +352,7 @@ func TestGetRunner(t *testing.T) {
func TestGetRunnerByName(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -434,7 +432,7 @@ func TestGetRunnerByName(t *testing.T) {
func TestDeleteRunner(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -487,7 +485,7 @@ func TestDeleteRunner(t *testing.T) {
func TestGetRunnerGroupByName(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -539,7 +537,7 @@ func TestGetRunnerGroupByName(t *testing.T) {
func TestGetRunnerScaleSet(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -694,7 +692,7 @@ func TestGetRunnerScaleSet(t *testing.T) {
func TestGetRunnerScaleSetByID(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -826,7 +824,7 @@ func TestGetRunnerScaleSetByID(t *testing.T) {
func TestCreateRunnerScaleSet(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -924,7 +922,7 @@ func TestCreateRunnerScaleSet(t *testing.T) {
func TestUpdateRunnerScaleSet(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -977,7 +975,7 @@ func TestUpdateRunnerScaleSet(t *testing.T) {
func TestDeleteRunnerScaleSet(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -1019,526 +1017,9 @@ func TestDeleteRunnerScaleSet(t *testing.T) {
})
}
func TestCreateMessageSession(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
token: "token",
}
t.Run("CreateMessageSession unmarshals correctly", func(t *testing.T) {
owner := "foo"
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
RunnerSetting: RunnerSetting{},
}
want := &RunnerScaleSetSession{
OwnerName: "foo",
RunnerScaleSet: &RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
},
MessageQueueURL: "http://fake.github.com/123",
MessageQueueAccessToken: "fake.jwt.here",
}
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
resp := []byte(`{
"ownerName": "foo",
"runnerScaleSet": {
"id": 1,
"name": "ScaleSet"
},
"messageQueueUrl": "http://fake.github.com/123",
"messageQueueAccessToken": "fake.jwt.here"
}`)
w.Write(resp)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
got, err := client.CreateMessageSession(ctx, runnerScaleSet.ID, owner)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("CreateMessageSession unmarshals errors into ActionsError", func(t *testing.T) {
owner := "foo"
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
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) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set(headerActionsActivityID, exampleRequestID)
w.WriteHeader(http.StatusBadRequest)
resp := []byte(`{"typeName": "CSharpExceptionNameHere","message": "could not do something"}`)
w.Write(resp)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
_, err = client.CreateMessageSession(ctx, runnerScaleSet.ID, owner)
require.NotNil(t, err)
errorTypeForComparison := &ActionsError{}
assert.True(
t,
errors.As(err, &errorTypeForComparison),
"CreateMessageSession expected to be able to parse the error into ActionsError type: %v",
err,
)
assert.Equal(t, want, errorTypeForComparison)
})
t.Run("CreateMessageSession call is retried the correct amount of times", func(t *testing.T) {
owner := "foo"
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
RunnerSetting: RunnerSetting{},
}
gotRetries := 0
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
gotRetries++
}))
retryMax := 3
retryWaitMax := 1 * time.Microsecond
wantRetries := retryMax + 1
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(retryWaitMax),
)
require.NoError(t, err)
_, err = client.CreateMessageSession(ctx, runnerScaleSet.ID, owner)
assert.NotNil(t, err)
assert.Equalf(t, gotRetries, wantRetries, "CreateMessageSession got unexpected retry count: got=%v, want=%v", gotRetries, wantRetries)
})
}
func TestDeleteMessageSession(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
token: "token",
}
t.Run("DeleteMessageSession call is retried the correct amount of times", func(t *testing.T) {
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
RunnerSetting: RunnerSetting{},
}
gotRetries := 0
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
gotRetries++
}))
retryMax := 3
retryWaitMax := 1 * time.Microsecond
wantRetries := retryMax + 1
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(retryWaitMax),
)
require.NoError(t, err)
sessionID := uuid.New()
err = client.DeleteMessageSession(ctx, runnerScaleSet.ID, sessionID)
assert.NotNil(t, err)
assert.Equalf(t, gotRetries, wantRetries, "CreateMessageSession got unexpected retry count: got=%v, want=%v", gotRetries, wantRetries)
})
}
func TestRefreshMessageSession(t *testing.T) {
auth := &actionsAuth{
token: "token",
}
t.Run("RefreshMessageSession call is retried the correct amount of times", func(t *testing.T) {
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
RunnerSetting: RunnerSetting{},
}
gotRetries := 0
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
gotRetries++
}))
retryMax := 3
retryWaitMax := 1 * time.Microsecond
wantRetries := retryMax + 1
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(retryWaitMax),
)
require.NoError(t, err)
sessionID := uuid.New()
_, err = client.RefreshMessageSession(context.Background(), runnerScaleSet.ID, sessionID)
assert.NotNil(t, err)
assert.Equalf(t, gotRetries, wantRetries, "CreateMessageSession got unexpected retry count: got=%v, want=%v", gotRetries, wantRetries)
})
}
func TestGetMessage(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
token: "token",
}
token := "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwiaWF0IjoxNTE2MjM5MDIyLCJleHAiOjI1MTYyMzkwMjJ9.tlrHslTmDkoqnc4Kk9ISoKoUNDfHo-kjlH-ByISBqzE"
runnerScaleSetMessage := &RunnerScaleSetMessage{
MessageID: 1,
}
t.Run("Get Runner Scale Set Message", func(t *testing.T) {
want := runnerScaleSetMessage
response := []byte(`{"messageId":1,"messageType":"RunnerScaleSetJobMessages"}`)
s := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Write(response)
}))
client, err := newClient(
testSystemInfo,
s.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
got, err := client.GetMessage(ctx, s.URL, token, 0, 10)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("GetMessage sets the last message id if not 0", func(t *testing.T) {
want := runnerScaleSetMessage
response := []byte(`{"messageId":1,"messageType":"RunnerScaleSetJobMessages"}`)
s := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
q := r.URL.Query()
assert.Equal(t, "1", q.Get("lastMessageId"))
w.Write(response)
}))
client, err := newClient(
testSystemInfo,
s.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
got, err := client.GetMessage(ctx, s.URL, token, 1, 10)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("Default retries on server error", func(t *testing.T) {
retryMax := 1
actualRetry := 0
expectedRetry := retryMax + 1
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
actualRetry++
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(1*time.Millisecond),
)
require.NoError(t, err)
_, err = client.GetMessage(ctx, server.URL, token, 0, 10)
assert.NotNil(t, err)
assert.Equalf(t, actualRetry, expectedRetry, "A retry was expected after the first request but got: %v", actualRetry)
})
t.Run("Message token expired", func(t *testing.T) {
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusUnauthorized)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
_, err = client.GetMessage(ctx, server.URL, token, 0, 10)
require.NotNil(t, err)
var expectedErr *ActionsError
require.True(t, errors.As(err, &expectedErr))
assert.True(t, expectedErr.IsMessageQueueTokenExpired())
})
t.Run("Status code not found", func(t *testing.T) {
want := ActionsError{
Err: errors.New("unknown exception"),
StatusCode: 404,
}
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNotFound)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
_, err = client.GetMessage(ctx, server.URL, token, 0, 10)
require.NotNil(t, err)
assert.Equal(t, want.Error(), err.Error())
})
t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain")
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
_, err = client.GetMessage(ctx, server.URL, token, 0, 10)
assert.NotNil(t, err)
})
t.Run("Capacity error handling", func(t *testing.T) {
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
hc := r.Header.Get(HeaderScaleSetMaxCapacity)
c, err := strconv.Atoi(hc)
require.NoError(t, err)
assert.GreaterOrEqual(t, c, 0)
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain")
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
_, err = client.GetMessage(ctx, server.URL, token, 0, 0)
assert.Error(t, err)
var expectedErr *ActionsError
assert.ErrorAs(t, err, &expectedErr)
assert.Equal(t, http.StatusBadRequest, expectedErr.StatusCode)
})
}
func TestDeleteMessage(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
token: "token",
}
token := "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwiaWF0IjoxNTE2MjM5MDIyLCJleHAiOjI1MTYyMzkwMjJ9.tlrHslTmDkoqnc4Kk9ISoKoUNDfHo-kjlH-ByISBqzE"
runnerScaleSetMessage := &RunnerScaleSetMessage{
MessageID: 1,
}
t.Run("Delete existing message", func(t *testing.T) {
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
err = client.DeleteMessage(ctx, server.URL, token, runnerScaleSetMessage.MessageID)
assert.Nil(t, err)
})
t.Run("Message token expired", func(t *testing.T) {
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusUnauthorized)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
err = client.DeleteMessage(ctx, server.URL, token, 0)
require.NotNil(t, err)
var expectedErr *ActionsError
require.ErrorAs(t, err, &expectedErr)
assert.True(t, expectedErr.IsMessageQueueTokenExpired())
})
t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain")
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
err = client.DeleteMessage(ctx, server.URL, token, runnerScaleSetMessage.MessageID)
require.NotNil(t, err)
var expectedErr *ActionsError
assert.True(t, errors.As(err, &expectedErr))
},
)
t.Run("Default retries on server error", func(t *testing.T) {
actualRetry := 0
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
actualRetry++
}))
retryMax := 1
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(1*time.Nanosecond),
)
require.NoError(t, err)
err = client.DeleteMessage(ctx, server.URL, token, runnerScaleSetMessage.MessageID)
assert.NotNil(t, err)
expectedRetry := retryMax + 1
assert.Equalf(t, actualRetry, expectedRetry, "A retry was expected after the first request but got: %v", actualRetry)
})
t.Run("No message found", func(t *testing.T) {
want := (*RunnerScaleSetMessage)(nil)
rsl, err := json.Marshal(want)
require.NoError(t, err)
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Write(rsl)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
err = client.DeleteMessage(ctx, server.URL, token, runnerScaleSetMessage.MessageID+1)
var expectedErr *ActionsError
require.True(t, errors.As(err, &expectedErr))
})
}
func TestClientProxy(t *testing.T) {
serverCalled := false
proxy := testserver.New(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
serverCalled = true
}))
proxyConfig := &httpproxy.Config{
HTTPProxy: proxy.URL,
}
proxyFunc := func(req *http.Request) (*url.URL, error) {
return proxyConfig.ProxyFunc()(req.URL)
}
c, err := newClient(
testSystemInfo,
"http://github.com/org/repo",
nil,
WithProxy(proxyFunc),
)
require.NoError(t, err)
req, err := http.NewRequest(http.MethodGet, "http://example.com", nil)
require.NoError(t, err)
_, err = c.do(req)
require.NoError(t, err)
assert.True(t, serverCalled)
}
func TestGenerateJitRunnerConfig(t *testing.T) {
ctx := context.Background()
auth := &actionsAuth{
auth := actionsAuth{
token: "token",
}
@@ -1590,62 +1071,9 @@ func TestGenerateJitRunnerConfig(t *testing.T) {
})
}
func TestClient_Do(t *testing.T) {
t.Run("trims byte order mark from response if present", func(t *testing.T) {
t.Run("when there is no body", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
}))
defer server.Close()
type serverCtxKey int
client, err := newClient(
testSystemInfo,
"https://localhost/org/repo",
&actionsAuth{token: "token"},
)
require.NoError(t, err)
req, err := http.NewRequest("GET", server.URL, nil)
require.NoError(t, err)
resp, err := client.do(req)
require.NoError(t, err)
body, err := io.ReadAll(resp.Body)
require.NoError(t, err)
assert.Empty(t, string(body))
})
responses := []string{
"\xef\xbb\xbf{\"foo\":\"bar\"}",
"{\"foo\":\"bar\"}",
}
for _, response := range responses {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.Write([]byte(response))
}))
defer server.Close()
client, err := newClient(
testSystemInfo,
"https://localhost/org/repo",
&actionsAuth{token: "token"},
)
require.NoError(t, err)
req, err := http.NewRequest("GET", server.URL, nil)
require.NoError(t, err)
resp, err := client.do(req)
require.NoError(t, err)
body, err := io.ReadAll(resp.Body)
require.NoError(t, err)
assert.Equal(t, "{\"foo\":\"bar\"}", string(body))
}
})
}
const ctxKeyServer serverCtxKey = iota
// newActionsServer returns a new httptest.Server that handles the
// authentication requests neeeded to create a new client. Any requests not
@@ -1667,6 +1095,7 @@ func newActionsServer(t *testing.T, handler http.Handler, options ...actionsServ
}
h := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
r = r.WithContext(context.WithValue(r.Context(), ctxKeyServer, server))
// handle getRunnerRegistrationToken
if strings.HasSuffix(r.URL.Path, "/runners/registration-token") {
w.WriteHeader(http.StatusCreated)
@@ -1700,6 +1129,29 @@ type actionsServer struct {
token string
}
func (s *actionsServer) testRunnerScaleSetSession() RunnerScaleSetSession {
session := RunnerScaleSetSession{
SessionID: uuid.New(),
OwnerName: "foo",
RunnerScaleSet: &RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
},
MessageQueueURL: s.URL,
MessageQueueAccessToken: s.token,
Statistics: &RunnerScaleSetStatistic{
TotalAvailableJobs: 0,
TotalAcquiredJobs: 0,
TotalAssignedJobs: 0,
TotalRunningJobs: 0,
TotalRegisteredRunners: 0,
TotalBusyRunners: 0,
TotalIdleRunners: 0,
},
}
return session
}
func (s *actionsServer) configURLForOrg(org string) string {
return s.URL + "/" + org
}
@@ -1761,13 +1213,12 @@ func TestServerWithSelfSignedCertificates(t *testing.T) {
u = server.URL
configURL := server.URL + "/my-org"
auth := &actionsAuth{
token: "token",
}
client, err := newClient(
testSystemInfo,
configURL,
auth,
actionsAuth{
token: "token",
},
)
require.NoError(t, err)
require.NotNil(t, client)
@@ -1796,10 +1247,6 @@ func TestServerWithSelfSignedCertificates(t *testing.T) {
u = server.URL
configURL := server.URL + "/my-org"
auth := &actionsAuth{
token: "token",
}
cert, err := os.ReadFile(filepath.Join("testdata", "rootCA.crt"))
require.NoError(t, err)
@@ -1809,7 +1256,9 @@ func TestServerWithSelfSignedCertificates(t *testing.T) {
client, err := newClient(
testSystemInfo,
configURL,
auth,
actionsAuth{
token: "token",
},
WithRootCAs(pool),
)
require.NoError(t, err)
@@ -1829,10 +1278,6 @@ func TestServerWithSelfSignedCertificates(t *testing.T) {
u = server.URL
configURL := server.URL + "/my-org"
auth := &actionsAuth{
token: "token",
}
cert, err := os.ReadFile(filepath.Join("testdata", "intermediate.crt"))
require.NoError(t, err)
@@ -1842,7 +1287,9 @@ func TestServerWithSelfSignedCertificates(t *testing.T) {
client, err := newClient(
testSystemInfo,
configURL,
auth,
actionsAuth{
token: "token",
},
WithRootCAs(pool),
WithRetryMax(0),
)
@@ -1857,14 +1304,12 @@ func TestServerWithSelfSignedCertificates(t *testing.T) {
server := startNewTLSTestServer(t, certPath, keyPath, http.HandlerFunc(h))
configURL := server.URL + "/my-org"
auth := &actionsAuth{
token: "token",
}
client, err := newClient(
testSystemInfo,
configURL,
auth,
actionsAuth{
token: "token",
},
WithoutTLSVerify(),
)
require.NoError(t, err)
@@ -1887,55 +1332,6 @@ func startNewTLSTestServer(t *testing.T, certPath, keyPath string, handler http.
return server
}
func TestUserAgent(t *testing.T) {
version, sha := detectModuleVersionAndCommit()
userAgentInfo := SystemInfo{
System: "actions-runner-controller",
Version: "0.1.0",
CommitSHA: "1234567890abcdef",
ScaleSetID: 10,
Subsystem: "test",
}
client, err := newClient(
testSystemInfo,
"https://github.com/org/repo",
&actionsAuth{token: "token"},
)
require.NoError(t, err, "failed to instantiate the client")
got := *client.userAgent.Load()
wantInfo := userAgent{
SystemInfo: testSystemInfo,
BuildCommitSHA: sha,
BuildVersion: version,
}
b, err := json.Marshal(wantInfo)
require.NoError(t, err, "failed to marshal expected user agent")
want := string(b)
assert.Equal(t, want, got)
client.SetSystemInfo(SystemInfo{
System: "actions-runner-controller",
Version: "0.1.0",
CommitSHA: "1234567890abcdef",
ScaleSetID: 10,
Subsystem: "test",
})
got = *client.userAgent.Load()
wantInfo = userAgent{
SystemInfo: userAgentInfo,
BuildCommitSHA: sha,
BuildVersion: version,
}
b, err = json.Marshal(wantInfo)
require.NoError(t, err, "failed to marshal expected user agent after SetSystemInfo")
want = string(b)
assert.Equal(t, want, got)
}
const samplePrivateKey = `-----BEGIN PRIVATE KEY-----
MIIEugIBADANBgkqhkiG9w0BAQEFAASCBKQwggSgAgEAAoIBAQC7tgquvNIp+Ik3
rRVZ9r0zJLsSzTHqr2dA6EUUmpRiQ25MzjMqKqu0OBwvh/pZyfjSIkKrhIridNK4
+175
View File
@@ -0,0 +1,175 @@
package scaleset
import (
"bytes"
"crypto/tls"
"crypto/x509"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"time"
"github.com/hashicorp/go-retryablehttp"
)
type commonClient struct {
httpClient *http.Client
systemInfo SystemInfo // never set directly, use setSystemInfoUnlocked
userAgent string
httpClientOption
}
func newCommonClient(systemInfo SystemInfo, httpClientOption httpClientOption) *commonClient {
c := &commonClient{
httpClientOption: httpClientOption,
}
c.setSystemInfo(systemInfo)
retryableHTTPClient, err := httpClientOption.newRetryableHTTPClient()
if err != nil {
panic(fmt.Sprintf("failed to create retryable HTTP client: %v", err))
}
c.httpClient = retryableHTTPClient.StandardClient()
return c
}
func (c *commonClient) newRetryableHTTPClient() (*retryablehttp.Client, error) {
return c.httpClientOption.newRetryableHTTPClient()
}
func (c *commonClient) do(req *http.Request) (*http.Response, error) {
resp, err := c.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("client request failed: %w", err)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("failed to read the response body: %w", err)
}
if err := resp.Body.Close(); err != nil {
return nil, fmt.Errorf("failed to close the response body: %w", err)
}
body = trimByteOrderMark(body)
resp.Body = io.NopCloser(bytes.NewReader(body))
return resp, nil
}
type httpClientOption struct {
logger *slog.Logger
retryMax int
retryWaitMax time.Duration
rootCAs *x509.CertPool
tlsInsecureSkipVerify bool
proxyFunc ProxyFunc
}
func (o *httpClientOption) defaults() {
if o.logger == nil {
o.logger = slog.New(slog.DiscardHandler)
}
if o.retryMax == 0 {
o.retryMax = 4
}
if o.retryWaitMax == 0 {
o.retryWaitMax = 30 * time.Second
}
}
func (o *httpClientOption) newRetryableHTTPClient() (*retryablehttp.Client, error) {
retryClient := retryablehttp.NewClient()
retryClient.Logger = o.logger
retryClient.RetryMax = o.retryMax
retryClient.RetryWaitMax = o.retryWaitMax
retryClient.HTTPClient.Timeout = 5 * time.Minute // timeout must be > 1m to accomodate long polling
transport, ok := retryClient.HTTPClient.Transport.(*http.Transport)
if !ok {
// this should always be true, because retryablehttp.NewClient() uses
// cleanhttp.DefaultPooledTransport()
return nil, fmt.Errorf("failed to get http transport from retryablehttp client")
}
if transport.TLSClientConfig == nil {
transport.TLSClientConfig = &tls.Config{}
}
if o.rootCAs != nil {
transport.TLSClientConfig.RootCAs = o.rootCAs
}
if o.tlsInsecureSkipVerify {
transport.TLSClientConfig.InsecureSkipVerify = true
}
transport.Proxy = o.proxyFunc
retryClient.HTTPClient.Transport = transport
return retryClient, nil
}
func (c *commonClient) setSystemInfo(info SystemInfo) {
c.systemInfo = info
c.setUserAgent()
}
func (c *commonClient) setUserAgent() {
b, _ := json.Marshal(userAgent{
SystemInfo: c.systemInfo,
BuildVersion: buildInfo.version,
BuildCommitSHA: buildInfo.commitSHA,
})
c.userAgent = string(b)
}
// HTTPOption defines a functional option for configuring the Client.
type HTTPOption func(*httpClientOption)
// WithLogger sets a custom logger for the Client.
func WithLogger(logger slog.Logger) HTTPOption {
return func(c *httpClientOption) {
c.logger = &logger
}
}
// WithRetryMax sets the maximum number of retries for the Client.
func WithRetryMax(retryMax int) HTTPOption {
return func(c *httpClientOption) {
c.retryMax = retryMax
}
}
// WithRetryWaitMax sets the maximum wait time between retries for the Client.
func WithRetryWaitMax(retryWaitMax time.Duration) HTTPOption {
return func(c *httpClientOption) {
c.retryWaitMax = retryWaitMax
}
}
// WithRootCAs sets custom root certificate authorities for the Client.
func WithRootCAs(rootCAs *x509.CertPool) HTTPOption {
return func(c *httpClientOption) {
c.rootCAs = rootCAs
}
}
// WithoutTLSVerify disables TLS certificate verification for the Client.
func WithoutTLSVerify() HTTPOption {
return func(c *httpClientOption) {
c.tlsInsecureSkipVerify = true
}
}
// WithProxy sets a custom proxy function for the Client.
func WithProxy(proxyFunc ProxyFunc) HTTPOption {
return func(c *httpClientOption) {
c.proxyFunc = proxyFunc
}
}
+153
View File
@@ -0,0 +1,153 @@
package scaleset
import (
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"github.com/actions/scaleset/internal/testserver"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/http/httpproxy"
)
func defaultHTTPClientOption() httpClientOption {
var opt httpClientOption
opt.defaults()
return opt
}
func TestClient_Do(t *testing.T) {
t.Run("trims byte order mark from response if present", func(t *testing.T) {
t.Run("when there is no body", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
}))
defer server.Close()
client := newCommonClient(
testSystemInfo,
defaultHTTPClientOption(),
)
req, err := http.NewRequest("GET", server.URL, nil)
require.NoError(t, err)
resp, err := client.do(req)
require.NoError(t, err)
body, err := io.ReadAll(resp.Body)
require.NoError(t, err)
assert.Empty(t, string(body))
})
responses := []string{
"\xef\xbb\xbf{\"foo\":\"bar\"}",
"{\"foo\":\"bar\"}",
}
for _, response := range responses {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.Write([]byte(response))
}))
defer server.Close()
client := newCommonClient(
testSystemInfo,
defaultHTTPClientOption(),
)
req, err := http.NewRequest("GET", server.URL, nil)
require.NoError(t, err)
resp, err := client.do(req)
require.NoError(t, err)
body, err := io.ReadAll(resp.Body)
require.NoError(t, err)
assert.Equal(t, "{\"foo\":\"bar\"}", string(body))
}
})
}
func TestClientProxy(t *testing.T) {
serverCalled := false
proxy := testserver.New(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
serverCalled = true
}))
proxyConfig := &httpproxy.Config{
HTTPProxy: proxy.URL,
}
proxyFunc := func(req *http.Request) (*url.URL, error) {
return proxyConfig.ProxyFunc()(req.URL)
}
opts := defaultHTTPClientOption()
WithProxy(proxyFunc)(&opts)
client := newCommonClient(
testSystemInfo,
opts,
)
req, err := http.NewRequest(http.MethodGet, "http://example.com", nil)
require.NoError(t, err)
_, err = client.do(req)
require.NoError(t, err)
assert.True(t, serverCalled)
}
func TestUserAgent(t *testing.T) {
version, sha := detectModuleVersionAndCommit()
userAgentInfo := SystemInfo{
System: "actions-runner-controller",
Version: "0.1.0",
CommitSHA: "1234567890abcdef",
ScaleSetID: 10,
Subsystem: "test",
}
client := newCommonClient(
testSystemInfo,
defaultHTTPClientOption(),
)
got := client.userAgent
wantInfo := userAgent{
SystemInfo: testSystemInfo,
BuildCommitSHA: sha,
BuildVersion: version,
}
b, err := json.Marshal(wantInfo)
require.NoError(t, err, "failed to marshal expected user agent")
want := string(b)
assert.Equal(t, want, got)
client.setSystemInfo(SystemInfo{
System: "actions-runner-controller",
Version: "0.1.0",
CommitSHA: "1234567890abcdef",
ScaleSetID: 10,
Subsystem: "test",
})
got = client.userAgent
wantInfo = userAgent{
SystemInfo: userAgentInfo,
BuildCommitSHA: sha,
BuildVersion: version,
}
b, err = json.Marshal(wantInfo)
require.NoError(t, err, "failed to marshal expected user agent after SetSystemInfo")
want = string(b)
assert.Equal(t, want, got)
}
+1 -1
View File
@@ -121,7 +121,7 @@ func TestActionsError(t *testing.T) {
t.Run("is message queue token expired exception", func(t *testing.T) {
err := &ActionsError{
ActivityID: "activity-id",
StatusCode: 404,
StatusCode: 401,
Err: &messageQueueTokenExpiredError{
message: "example error message",
},
+15 -1
View File
@@ -13,6 +13,7 @@ import (
"github.com/actions/scaleset/listener"
"github.com/docker/docker/api/types/image"
dockerclient "github.com/docker/docker/client"
"github.com/google/uuid"
"github.com/spf13/cobra"
)
@@ -131,8 +132,21 @@ func run(ctx context.Context, c Config) error {
return fmt.Errorf("failed to close image pull: %w", err)
}
// Get the name of the client which will be used as the owner
hostname, err := os.Hostname()
if err != nil {
hostname = uuid.NewString()
logger.Info("Failed to get hostname, fallback to uuid", "uuid", hostname, "error", err)
}
sessionClient, err := scalesetClient.MessageSessionClient(ctx, scaleSet.ID, hostname)
if err != nil {
return fmt.Errorf("failed to create message session client: %w", err)
}
defer sessionClient.Close(context.Background())
logger.Info("Initializing listener")
listener, err := listener.New(scalesetClient, listener.Config{
listener, err := listener.New(sessionClient, listener.Config{
ScaleSetID: scaleSet.ID,
MaxRunners: c.MaxRunners,
Logger: logger.WithGroup("listener"),
+30 -196
View File
@@ -3,24 +3,16 @@ package listener
import (
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
"math"
"net/http"
"os"
"sync/atomic"
"time"
"github.com/actions/scaleset"
"github.com/google/uuid"
)
const (
sessionCreationMaxRetries = 10
)
// Config holds the configuration for the Listener.
type Config struct {
// ScaleSetID is the ID of the runner scale set to listen to.
@@ -55,13 +47,13 @@ func (c *Config) Validate() error {
// This interface is defined to allow for easier testing and mocking, as well
// as allowing wrappers around the scaleset client if needed.
type Client interface {
CreateMessageSession(ctx context.Context, runnerScaleSetID int, owner string) (*scaleset.RunnerScaleSetSession, error)
GetMessage(ctx context.Context, messageQueueURL, messageQueueAccessToken string, lastMessageID int, maxCapacity int) (*scaleset.RunnerScaleSetMessage, error)
DeleteMessage(ctx context.Context, messageQueueURL, messageQueueAccessToken string, messageID int) error
RefreshMessageSession(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) (*scaleset.RunnerScaleSetSession, error)
DeleteMessageSession(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) error
GetMessage(ctx context.Context, lastMessageID, maxCapacity int) (*scaleset.RunnerScaleSetMessage, error)
DeleteMessage(ctx context.Context, messageID int) error
Session() scaleset.RunnerScaleSetSession
}
type Option func(*Listener)
// Listener listens for messages from the scaleset service and handles them. It automatically handles session
// creation/deletion/refreshing and message polling and acking.
type Listener struct {
@@ -72,13 +64,6 @@ type Listener struct {
scaleSetID int
maxRunners atomic.Uint32
// lastMessageID keeps track of the last processed message ID
lastMessageID int
// hostname of the current machine
hostname string
// session represents the current message session
session *scaleset.RunnerScaleSetSession
// configuration for the listener
logger *slog.Logger
}
@@ -90,7 +75,7 @@ func (l *Listener) SetMaxRunners(count int) {
}
// New creates a new Listener with the given configuration.
func New(client Client, config Config) (*Listener, error) {
func New(client Client, config Config, options ...Option) (*Listener, error) {
if client == nil {
return nil, errors.New("client is required")
}
@@ -99,16 +84,9 @@ func New(client Client, config Config) (*Listener, error) {
return nil, fmt.Errorf("invalid config: %w", err)
}
hostname, err := os.Hostname()
if err != nil {
hostname = uuid.NewString()
config.Logger.Info("Failed to get hostname, fallback to uuid", "uuid", hostname, "error", err)
}
listener := &Listener{
client: client,
scaleSetID: config.ScaleSetID,
hostname: hostname,
logger: config.Logger,
}
listener.SetMaxRunners(config.MaxRunners)
@@ -125,32 +103,27 @@ type Scaler interface {
// Run starts the listener and processes messages using the provided scaler.
func (l *Listener) Run(ctx context.Context, scaler Scaler) error {
l.logger.Info("Creating message session")
if err := l.createSession(ctx); err != nil {
return fmt.Errorf("failed to create session: %w", err)
}
{
initialSession := l.client.Session()
defer func() {
l.logger.Debug("Deleting message session")
if err := l.deleteMessageSession(); err != nil {
l.logger.Error(
"failed to delete message session",
slog.String("error", err.Error()),
)
if initialSession.SessionID == uuid.Nil {
return fmt.Errorf("initial session is nil")
}
}()
if l.session.Statistics == nil {
return fmt.Errorf("session statistics is nil")
}
if initialSession.Statistics == nil {
return fmt.Errorf("session statistics is nil")
}
l.logger.Info("Message session created; listening for messages", "sessionID", l.session.SessionID)
// Handle initial statistics
if _, err := scaler.HandleDesiredRunnerCount(ctx, l.session.Statistics.TotalAssignedJobs); err != nil {
return fmt.Errorf("handling initial message failed: %w", err)
l.logger.Info(
"Handling initial session statistics",
slog.Int("totalAssignedJobs", initialSession.Statistics.TotalAssignedJobs),
)
if _, err := scaler.HandleDesiredRunnerCount(ctx, initialSession.Statistics.TotalAssignedJobs); err != nil {
return fmt.Errorf("handling initial message failed: %w", err)
}
}
var lastMessageID int
for {
select {
case <-ctx.Done():
@@ -158,7 +131,12 @@ func (l *Listener) Run(ctx context.Context, scaler Scaler) error {
default:
}
msg, err := l.getMessage(ctx)
l.logger.Info("Getting next message", slog.Int("lastMessageID", lastMessageID))
msg, err := l.client.GetMessage(
ctx,
lastMessageID,
int(l.maxRunners.Load()),
)
if err != nil {
return fmt.Errorf("failed to get message: %w", err)
}
@@ -172,6 +150,8 @@ func (l *Listener) Run(ctx context.Context, scaler Scaler) error {
continue
}
lastMessageID = msg.MessageID
// Remove cancellation from the context to avoid cancelling the message handling.
if err := l.handleMessage(context.WithoutCancel(ctx), scaler, msg); err != nil {
return fmt.Errorf("failed to handle message: %w", err)
@@ -180,8 +160,7 @@ func (l *Listener) Run(ctx context.Context, scaler Scaler) error {
}
func (l *Listener) handleMessage(ctx context.Context, handler Scaler, msg *scaleset.RunnerScaleSetMessage) error {
l.lastMessageID = msg.MessageID
if err := l.deleteLastMessage(ctx); err != nil {
if err := l.client.DeleteMessage(ctx, msg.MessageID); err != nil {
return fmt.Errorf("failed to delete message: %w", err)
}
@@ -202,148 +181,3 @@ func (l *Listener) handleMessage(ctx context.Context, handler Scaler, msg *scale
return nil
}
func (l *Listener) createSession(ctx context.Context) error {
var session *scaleset.RunnerScaleSetSession
var retries int
for {
var err error
session, err = l.client.CreateMessageSession(ctx, l.scaleSetID, l.hostname)
if err == nil {
break
}
clientErr := &scaleset.ActionsError{}
if !errors.As(err, &clientErr) {
return fmt.Errorf("failed to create session: %w", err)
}
if clientErr.StatusCode != http.StatusConflict {
return fmt.Errorf("failed to create session: %w", err)
}
retries++
if retries >= sessionCreationMaxRetries {
return fmt.Errorf("failed to create session after %d retries: %w", retries, err)
}
l.logger.Info("Unable to create message session. Will try again in 30 seconds", "error", err.Error())
select {
case <-ctx.Done():
return fmt.Errorf("context cancelled: %w", ctx.Err())
case <-time.After(30 * time.Second):
}
}
statistics, err := json.Marshal(session.Statistics)
if err != nil {
return fmt.Errorf("failed to marshal statistics: %w", err)
}
l.logger.Info("Current runner scale set statistics.", "statistics", string(statistics))
l.session = session
return nil
}
func (l *Listener) getMessage(ctx context.Context) (*scaleset.RunnerScaleSetMessage, error) {
l.logger.Info("Getting next message", "lastMessageID", l.lastMessageID)
msg, err := l.client.GetMessage(
ctx,
l.session.MessageQueueURL,
l.session.MessageQueueAccessToken,
l.lastMessageID,
int(l.maxRunners.Load()),
)
if err == nil { // if NO error
return msg, nil
}
expiredError := &scaleset.ActionsError{}
if !errors.As(err, &expiredError) || !expiredError.IsMessageQueueTokenExpired() {
return nil, fmt.Errorf("failed to get next message: %w", err)
}
if err := l.refreshSession(ctx); err != nil {
return nil, fmt.Errorf("failed to refresh message session: %w", err)
}
l.logger.Info("Getting next message", "lastMessageID", l.lastMessageID)
msg, err = l.client.GetMessage(
ctx,
l.session.MessageQueueURL,
l.session.MessageQueueAccessToken,
l.lastMessageID,
int(l.maxRunners.Load()),
)
if err != nil { // if error
return nil, fmt.Errorf("failed to get next message after message session refresh: %w", err)
}
return msg, nil
}
func (l *Listener) deleteLastMessage(ctx context.Context) error {
l.logger.Info("Deleting last message", "lastMessageID", l.lastMessageID)
err := l.client.DeleteMessage(
ctx,
l.session.MessageQueueURL,
l.session.MessageQueueAccessToken,
l.lastMessageID,
)
if err == nil { // if NO error
return nil
}
expiredError := &scaleset.ActionsError{}
if !errors.As(err, &expiredError) || !expiredError.IsMessageQueueTokenExpired() {
return fmt.Errorf("failed to delete last message: %w", err)
}
if err := l.refreshSession(ctx); err != nil {
return fmt.Errorf("failed to refresh message session: %w", err)
}
err = l.client.DeleteMessage(
ctx,
l.session.MessageQueueURL,
l.session.MessageQueueAccessToken,
l.lastMessageID,
)
if err != nil {
return fmt.Errorf("failed to delete last message after message session refresh: %w", err)
}
return nil
}
func (l *Listener) refreshSession(ctx context.Context) error {
l.logger.Info("Message queue token is expired during GetNextMessage, refreshing...")
session, err := l.client.RefreshMessageSession(
ctx,
l.session.RunnerScaleSet.ID,
l.session.SessionID,
)
if err != nil {
return fmt.Errorf("refresh message session failed. %w", err)
}
l.session = session
return nil
}
func (l *Listener) deleteMessageSession() error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
l.logger.Info("Deleting message session")
if err := l.client.DeleteMessageSession(ctx, l.session.RunnerScaleSet.ID, l.session.SessionID); err != nil {
return fmt.Errorf("failed to delete message session: %w", err)
}
return nil
}
+5 -601
View File
@@ -2,11 +2,8 @@ package listener
import (
"context"
"errors"
"math"
"net/http"
"testing"
"time"
"github.com/actions/scaleset"
"github.com/google/uuid"
@@ -64,577 +61,9 @@ func TestNew(t *testing.T) {
})
}
func TestListener_createSession(t *testing.T) {
t.Parallel()
t.Run("fail once", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
MaxRunners: 10,
}
client := NewMockClient(t)
client.On(
"CreateMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(
nil,
assert.AnError,
).Once()
l, err := New(client, config)
require.Nil(t, err)
err = l.createSession(ctx)
assert.NotNil(t, err)
})
t.Run("fail context", func(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
client.On(
"CreateMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(
nil,
scaleset.ParseActionsErrorFromResponse(&http.Response{
StatusCode: http.StatusConflict,
}),
).Once()
l, err := New(client, config)
require.Nil(t, err)
err = l.createSession(ctx)
assert.True(t, errors.Is(err, context.DeadlineExceeded))
})
t.Run("sets session", func(t *testing.T) {
t.Parallel()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
uuid := uuid.New()
session := &scaleset.RunnerScaleSetSession{
SessionID: uuid,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
MessageQueueURL: "https://example.com",
MessageQueueAccessToken: "1234567890",
Statistics: nil,
}
client.On(
"CreateMessageSession",
mock.Anything,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
l, err := New(client, config)
require.Nil(t, err)
err = l.createSession(context.Background())
assert.Nil(t, err)
assert.Equal(t, session, l.session)
})
}
func TestListener_getMessage(t *testing.T) {
t.Parallel()
t.Run("receives message", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
MaxRunners: 10,
}
client := NewMockClient(t)
want := &scaleset.RunnerScaleSetMessage{
MessageID: 1,
}
client.On(
"GetMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
10,
).Return(
want,
nil,
).Once()
l, err := New(client, config)
require.Nil(t, err)
l.session = &scaleset.RunnerScaleSetSession{}
got, err := l.getMessage(ctx)
assert.Nil(t, err)
assert.Equal(t, want, got)
})
t.Run("not expired error", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
MaxRunners: 10,
}
client := NewMockClient(t)
client.On(
"GetMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
10,
).Return(
nil,
scaleset.ParseActionsErrorFromResponse(&http.Response{
StatusCode: http.StatusNotFound,
}),
).Once()
l, err := New(client, config)
require.Nil(t, err)
l.session = &scaleset.RunnerScaleSetSession{}
_, err = l.getMessage(ctx)
assert.NotNil(t, err)
})
t.Run("refresh and succeeds", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
MaxRunners: 10,
}
client := NewMockClient(t)
uuid := uuid.New()
session := &scaleset.RunnerScaleSetSession{
SessionID: uuid,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
MessageQueueURL: "https://example.com",
MessageQueueAccessToken: "1234567890",
Statistics: nil,
}
client.On(
"RefreshMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
client.On(
"GetMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
10,
).Return(
nil,
&scaleset.ActionsError{
StatusCode: http.StatusUnauthorized,
ActivityID: "1234",
Err: scaleset.NewMessageQueueTokenExpiredError("token expired"),
},
).Once()
want := &scaleset.RunnerScaleSetMessage{
MessageID: 1,
}
client.On(
"GetMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
10,
).Return(want, nil).Once()
l, err := New(client, config)
require.Nil(t, err)
l.session = &scaleset.RunnerScaleSetSession{
SessionID: uuid,
RunnerScaleSet: &scaleset.RunnerScaleSet{},
}
got, err := l.getMessage(ctx)
assert.Nil(t, err)
assert.Equal(t, want, got)
})
t.Run("refresh and fails", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
MaxRunners: 10,
}
client := NewMockClient(t)
uuid := uuid.New()
session := &scaleset.RunnerScaleSetSession{
SessionID: uuid,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
MessageQueueURL: "https://example.com",
MessageQueueAccessToken: "1234567890",
Statistics: nil,
}
client.On(
"RefreshMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
client.On(
"GetMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
10,
).Return(
nil,
&scaleset.ActionsError{
StatusCode: http.StatusUnauthorized,
ActivityID: "1234",
Err: scaleset.NewMessageQueueTokenExpiredError("token expired"),
},
).Twice()
l, err := New(client, config)
require.Nil(t, err)
l.session = &scaleset.RunnerScaleSetSession{
SessionID: uuid,
RunnerScaleSet: &scaleset.RunnerScaleSet{},
}
got, err := l.getMessage(ctx)
assert.NotNil(t, err)
assert.Nil(t, got)
})
}
func TestListener_refreshSession(t *testing.T) {
t.Parallel()
t.Run("successfully refreshes", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
newUUID := uuid.New()
session := &scaleset.RunnerScaleSetSession{
SessionID: newUUID,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
MessageQueueURL: "https://example.com",
MessageQueueAccessToken: "1234567890",
Statistics: nil,
}
client.On(
"RefreshMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
l, err := New(client, config)
require.Nil(t, err)
oldUUID := uuid.New()
l.session = &scaleset.RunnerScaleSetSession{
SessionID: oldUUID,
RunnerScaleSet: &scaleset.RunnerScaleSet{},
}
err = l.refreshSession(ctx)
assert.Nil(t, err)
assert.Equal(t, session, l.session)
})
t.Run("fails to refresh", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
client.On(
"RefreshMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(nil, errors.New("error")).Once()
l, err := New(client, config)
require.Nil(t, err)
oldUUID := uuid.New()
oldSession := &scaleset.RunnerScaleSetSession{
SessionID: oldUUID,
RunnerScaleSet: &scaleset.RunnerScaleSet{},
}
l.session = oldSession
err = l.refreshSession(ctx)
assert.NotNil(t, err)
assert.Equal(t, oldSession, l.session)
})
}
func TestListener_deleteLastMessage(t *testing.T) {
t.Parallel()
t.Run("successfully deletes", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
client.On(
"DeleteMessage",
ctx,
mock.Anything,
mock.Anything,
mock.MatchedBy(
func(lastMessageID any) bool {
return lastMessageID.(int) == 5
},
),
).Return(nil).Once()
l, err := New(client, config)
require.Nil(t, err)
l.session = &scaleset.RunnerScaleSetSession{}
l.lastMessageID = 5
err = l.deleteLastMessage(ctx)
assert.Nil(t, err)
})
t.Run("fails to delete", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
client.On(
"DeleteMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
).Return(errors.New("error")).Once()
l, err := New(client, config)
require.Nil(t, err)
l.session = &scaleset.RunnerScaleSetSession{}
l.lastMessageID = 5
err = l.deleteLastMessage(ctx)
assert.NotNil(t, err)
})
t.Run("refresh and succeeds", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
newUUID := uuid.New()
session := &scaleset.RunnerScaleSetSession{
SessionID: newUUID,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
MessageQueueURL: "https://example.com",
MessageQueueAccessToken: "1234567890",
Statistics: nil,
}
client.On(
"RefreshMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
client.On(
"DeleteMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
).Return(
&scaleset.ActionsError{
StatusCode: http.StatusUnauthorized,
ActivityID: "1234",
Err: scaleset.NewMessageQueueTokenExpiredError("token expired"),
},
).Once()
client.On(
"DeleteMessage",
ctx,
mock.Anything,
mock.Anything,
mock.MatchedBy(
func(lastMessageID any) bool {
return lastMessageID.(int) == 5
},
),
).Return(nil).Once()
l, err := New(client, config)
require.Nil(t, err)
oldUUID := uuid.New()
l.session = &scaleset.RunnerScaleSetSession{
SessionID: oldUUID,
RunnerScaleSet: &scaleset.RunnerScaleSet{},
}
l.lastMessageID = 5
err = l.deleteLastMessage(ctx)
assert.NoError(t, err)
})
t.Run("refresh and fails", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
newUUID := uuid.New()
session := &scaleset.RunnerScaleSetSession{
SessionID: newUUID,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
MessageQueueURL: "https://example.com",
MessageQueueAccessToken: "1234567890",
Statistics: nil,
}
client.On(
"RefreshMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
client.On(
"DeleteMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
).Return(
&scaleset.ActionsError{
StatusCode: http.StatusUnauthorized,
ActivityID: "1234",
Err: scaleset.NewMessageQueueTokenExpiredError("token expired"),
},
).Twice()
l, err := New(client, config)
require.Nil(t, err)
oldUUID := uuid.New()
l.session = &scaleset.RunnerScaleSetSession{
SessionID: oldUUID,
RunnerScaleSet: &scaleset.RunnerScaleSet{},
}
l.lastMessageID = 5
err = l.deleteLastMessage(ctx)
assert.Error(t, err)
})
}
func TestListener_Run(t *testing.T) {
t.Parallel()
t.Run("create session fails", func(t *testing.T) {
t.Parallel()
ctx := context.Background()
config := Config{
ScaleSetID: 1,
}
client := NewMockClient(t)
client.On(
"CreateMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(nil, assert.AnError).Once()
l, err := New(client, config)
require.Nil(t, err)
err = l.Run(ctx, nil)
assert.NotNil(t, err)
})
t.Run("call handle regardless of initial message", func(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
@@ -646,7 +75,7 @@ func TestListener_Run(t *testing.T) {
client := NewMockClient(t)
uuid := uuid.New()
session := &scaleset.RunnerScaleSetSession{
session := scaleset.RunnerScaleSetSession{
SessionID: uuid,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
@@ -654,18 +83,8 @@ func TestListener_Run(t *testing.T) {
MessageQueueAccessToken: "1234567890",
Statistics: &scaleset.RunnerScaleSetStatistic{},
}
client.On(
"CreateMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
client.On(
"DeleteMessageSession",
mock.Anything,
session.RunnerScaleSet.ID,
session.SessionID,
).Return(nil).Once()
client.On("Session").Return(session).Once()
l, err := New(client, config)
require.Nil(t, err)
@@ -703,7 +122,7 @@ func TestListener_Run(t *testing.T) {
client := NewMockClient(t)
uuid := uuid.New()
session := &scaleset.RunnerScaleSetSession{
session := scaleset.RunnerScaleSetSession{
SessionID: uuid,
OwnerName: "example",
RunnerScaleSet: &scaleset.RunnerScaleSet{},
@@ -711,29 +130,16 @@ func TestListener_Run(t *testing.T) {
MessageQueueAccessToken: "1234567890",
Statistics: &scaleset.RunnerScaleSetStatistic{},
}
client.On(
"CreateMessageSession",
ctx,
mock.Anything,
mock.Anything,
).Return(session, nil).Once()
client.On(
"DeleteMessageSession",
mock.Anything,
session.RunnerScaleSet.ID,
session.SessionID,
).Return(nil).Once()
msg := &scaleset.RunnerScaleSetMessage{
MessageID: 1,
Statistics: &scaleset.RunnerScaleSetStatistic{},
}
client.On("Session").Return(session).Once()
client.On(
"GetMessage",
ctx,
mock.Anything,
mock.Anything,
mock.Anything,
10,
).
Return(msg, nil).
@@ -749,8 +155,6 @@ func TestListener_Run(t *testing.T) {
"DeleteMessage",
context.WithoutCancel(ctx),
mock.Anything,
mock.Anything,
mock.Anything,
).Return(nil).Once()
handler := NewMockScaler(t)
+45 -237
View File
@@ -8,7 +8,6 @@ import (
"context"
"github.com/actions/scaleset"
"github.com/google/uuid"
mock "github.com/stretchr/testify/mock"
)
@@ -39,91 +38,17 @@ func (_m *MockClient) EXPECT() *MockClient_Expecter {
return &MockClient_Expecter{mock: &_m.Mock}
}
// CreateMessageSession provides a mock function for the type MockClient
func (_mock *MockClient) CreateMessageSession(ctx context.Context, runnerScaleSetID int, owner string) (*scaleset.RunnerScaleSetSession, error) {
ret := _mock.Called(ctx, runnerScaleSetID, owner)
if len(ret) == 0 {
panic("no return value specified for CreateMessageSession")
}
var r0 *scaleset.RunnerScaleSetSession
var r1 error
if returnFunc, ok := ret.Get(0).(func(context.Context, int, string) (*scaleset.RunnerScaleSetSession, error)); ok {
return returnFunc(ctx, runnerScaleSetID, owner)
}
if returnFunc, ok := ret.Get(0).(func(context.Context, int, string) *scaleset.RunnerScaleSetSession); ok {
r0 = returnFunc(ctx, runnerScaleSetID, owner)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*scaleset.RunnerScaleSetSession)
}
}
if returnFunc, ok := ret.Get(1).(func(context.Context, int, string) error); ok {
r1 = returnFunc(ctx, runnerScaleSetID, owner)
} else {
r1 = ret.Error(1)
}
return r0, r1
}
// MockClient_CreateMessageSession_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateMessageSession'
type MockClient_CreateMessageSession_Call struct {
*mock.Call
}
// CreateMessageSession is a helper method to define mock.On call
// - ctx context.Context
// - runnerScaleSetID int
// - owner string
func (_e *MockClient_Expecter) CreateMessageSession(ctx interface{}, runnerScaleSetID interface{}, owner interface{}) *MockClient_CreateMessageSession_Call {
return &MockClient_CreateMessageSession_Call{Call: _e.mock.On("CreateMessageSession", ctx, runnerScaleSetID, owner)}
}
func (_c *MockClient_CreateMessageSession_Call) Run(run func(ctx context.Context, runnerScaleSetID int, owner string)) *MockClient_CreateMessageSession_Call {
_c.Call.Run(func(args mock.Arguments) {
var arg0 context.Context
if args[0] != nil {
arg0 = args[0].(context.Context)
}
var arg1 int
if args[1] != nil {
arg1 = args[1].(int)
}
var arg2 string
if args[2] != nil {
arg2 = args[2].(string)
}
run(
arg0,
arg1,
arg2,
)
})
return _c
}
func (_c *MockClient_CreateMessageSession_Call) Return(runnerScaleSetSession *scaleset.RunnerScaleSetSession, err error) *MockClient_CreateMessageSession_Call {
_c.Call.Return(runnerScaleSetSession, err)
return _c
}
func (_c *MockClient_CreateMessageSession_Call) RunAndReturn(run func(ctx context.Context, runnerScaleSetID int, owner string) (*scaleset.RunnerScaleSetSession, error)) *MockClient_CreateMessageSession_Call {
_c.Call.Return(run)
return _c
}
// DeleteMessage provides a mock function for the type MockClient
func (_mock *MockClient) DeleteMessage(ctx context.Context, messageQueueURL string, messageQueueAccessToken string, messageID int) error {
ret := _mock.Called(ctx, messageQueueURL, messageQueueAccessToken, messageID)
func (_mock *MockClient) DeleteMessage(ctx context.Context, messageID int) error {
ret := _mock.Called(ctx, messageID)
if len(ret) == 0 {
panic("no return value specified for DeleteMessage")
}
var r0 error
if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, int) error); ok {
r0 = returnFunc(ctx, messageQueueURL, messageQueueAccessToken, messageID)
if returnFunc, ok := ret.Get(0).(func(context.Context, int) error); ok {
r0 = returnFunc(ctx, messageID)
} else {
r0 = ret.Error(0)
}
@@ -137,36 +62,24 @@ type MockClient_DeleteMessage_Call struct {
// DeleteMessage is a helper method to define mock.On call
// - ctx context.Context
// - messageQueueURL string
// - messageQueueAccessToken string
// - messageID int
func (_e *MockClient_Expecter) DeleteMessage(ctx interface{}, messageQueueURL interface{}, messageQueueAccessToken interface{}, messageID interface{}) *MockClient_DeleteMessage_Call {
return &MockClient_DeleteMessage_Call{Call: _e.mock.On("DeleteMessage", ctx, messageQueueURL, messageQueueAccessToken, messageID)}
func (_e *MockClient_Expecter) DeleteMessage(ctx interface{}, messageID interface{}) *MockClient_DeleteMessage_Call {
return &MockClient_DeleteMessage_Call{Call: _e.mock.On("DeleteMessage", ctx, messageID)}
}
func (_c *MockClient_DeleteMessage_Call) Run(run func(ctx context.Context, messageQueueURL string, messageQueueAccessToken string, messageID int)) *MockClient_DeleteMessage_Call {
func (_c *MockClient_DeleteMessage_Call) Run(run func(ctx context.Context, messageID int)) *MockClient_DeleteMessage_Call {
_c.Call.Run(func(args mock.Arguments) {
var arg0 context.Context
if args[0] != nil {
arg0 = args[0].(context.Context)
}
var arg1 string
var arg1 int
if args[1] != nil {
arg1 = args[1].(string)
}
var arg2 string
if args[2] != nil {
arg2 = args[2].(string)
}
var arg3 int
if args[3] != nil {
arg3 = args[3].(int)
arg1 = args[1].(int)
}
run(
arg0,
arg1,
arg2,
arg3,
)
})
return _c
@@ -177,77 +90,14 @@ func (_c *MockClient_DeleteMessage_Call) Return(err error) *MockClient_DeleteMes
return _c
}
func (_c *MockClient_DeleteMessage_Call) RunAndReturn(run func(ctx context.Context, messageQueueURL string, messageQueueAccessToken string, messageID int) error) *MockClient_DeleteMessage_Call {
_c.Call.Return(run)
return _c
}
// DeleteMessageSession provides a mock function for the type MockClient
func (_mock *MockClient) DeleteMessageSession(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) error {
ret := _mock.Called(ctx, runnerScaleSetID, sessionID)
if len(ret) == 0 {
panic("no return value specified for DeleteMessageSession")
}
var r0 error
if returnFunc, ok := ret.Get(0).(func(context.Context, int, uuid.UUID) error); ok {
r0 = returnFunc(ctx, runnerScaleSetID, sessionID)
} else {
r0 = ret.Error(0)
}
return r0
}
// MockClient_DeleteMessageSession_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteMessageSession'
type MockClient_DeleteMessageSession_Call struct {
*mock.Call
}
// DeleteMessageSession is a helper method to define mock.On call
// - ctx context.Context
// - runnerScaleSetID int
// - sessionID uuid.UUID
func (_e *MockClient_Expecter) DeleteMessageSession(ctx interface{}, runnerScaleSetID interface{}, sessionID interface{}) *MockClient_DeleteMessageSession_Call {
return &MockClient_DeleteMessageSession_Call{Call: _e.mock.On("DeleteMessageSession", ctx, runnerScaleSetID, sessionID)}
}
func (_c *MockClient_DeleteMessageSession_Call) Run(run func(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID)) *MockClient_DeleteMessageSession_Call {
_c.Call.Run(func(args mock.Arguments) {
var arg0 context.Context
if args[0] != nil {
arg0 = args[0].(context.Context)
}
var arg1 int
if args[1] != nil {
arg1 = args[1].(int)
}
var arg2 uuid.UUID
if args[2] != nil {
arg2 = args[2].(uuid.UUID)
}
run(
arg0,
arg1,
arg2,
)
})
return _c
}
func (_c *MockClient_DeleteMessageSession_Call) Return(err error) *MockClient_DeleteMessageSession_Call {
_c.Call.Return(err)
return _c
}
func (_c *MockClient_DeleteMessageSession_Call) RunAndReturn(run func(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) error) *MockClient_DeleteMessageSession_Call {
func (_c *MockClient_DeleteMessage_Call) RunAndReturn(run func(ctx context.Context, messageID int) error) *MockClient_DeleteMessage_Call {
_c.Call.Return(run)
return _c
}
// GetMessage provides a mock function for the type MockClient
func (_mock *MockClient) GetMessage(ctx context.Context, messageQueueURL string, messageQueueAccessToken string, lastMessageID int, maxCapacity int) (*scaleset.RunnerScaleSetMessage, error) {
ret := _mock.Called(ctx, messageQueueURL, messageQueueAccessToken, lastMessageID, maxCapacity)
func (_mock *MockClient) GetMessage(ctx context.Context, lastMessageID int, maxCapacity int) (*scaleset.RunnerScaleSetMessage, error) {
ret := _mock.Called(ctx, lastMessageID, maxCapacity)
if len(ret) == 0 {
panic("no return value specified for GetMessage")
@@ -255,18 +105,18 @@ func (_mock *MockClient) GetMessage(ctx context.Context, messageQueueURL string,
var r0 *scaleset.RunnerScaleSetMessage
var r1 error
if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, int, int) (*scaleset.RunnerScaleSetMessage, error)); ok {
return returnFunc(ctx, messageQueueURL, messageQueueAccessToken, lastMessageID, maxCapacity)
if returnFunc, ok := ret.Get(0).(func(context.Context, int, int) (*scaleset.RunnerScaleSetMessage, error)); ok {
return returnFunc(ctx, lastMessageID, maxCapacity)
}
if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, int, int) *scaleset.RunnerScaleSetMessage); ok {
r0 = returnFunc(ctx, messageQueueURL, messageQueueAccessToken, lastMessageID, maxCapacity)
if returnFunc, ok := ret.Get(0).(func(context.Context, int, int) *scaleset.RunnerScaleSetMessage); ok {
r0 = returnFunc(ctx, lastMessageID, maxCapacity)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*scaleset.RunnerScaleSetMessage)
}
}
if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, int, int) error); ok {
r1 = returnFunc(ctx, messageQueueURL, messageQueueAccessToken, lastMessageID, maxCapacity)
if returnFunc, ok := ret.Get(1).(func(context.Context, int, int) error); ok {
r1 = returnFunc(ctx, lastMessageID, maxCapacity)
} else {
r1 = ret.Error(1)
}
@@ -280,42 +130,30 @@ type MockClient_GetMessage_Call struct {
// GetMessage is a helper method to define mock.On call
// - ctx context.Context
// - messageQueueURL string
// - messageQueueAccessToken string
// - lastMessageID int
// - maxCapacity int
func (_e *MockClient_Expecter) GetMessage(ctx interface{}, messageQueueURL interface{}, messageQueueAccessToken interface{}, lastMessageID interface{}, maxCapacity interface{}) *MockClient_GetMessage_Call {
return &MockClient_GetMessage_Call{Call: _e.mock.On("GetMessage", ctx, messageQueueURL, messageQueueAccessToken, lastMessageID, maxCapacity)}
func (_e *MockClient_Expecter) GetMessage(ctx interface{}, lastMessageID interface{}, maxCapacity interface{}) *MockClient_GetMessage_Call {
return &MockClient_GetMessage_Call{Call: _e.mock.On("GetMessage", ctx, lastMessageID, maxCapacity)}
}
func (_c *MockClient_GetMessage_Call) Run(run func(ctx context.Context, messageQueueURL string, messageQueueAccessToken string, lastMessageID int, maxCapacity int)) *MockClient_GetMessage_Call {
func (_c *MockClient_GetMessage_Call) Run(run func(ctx context.Context, lastMessageID int, maxCapacity int)) *MockClient_GetMessage_Call {
_c.Call.Run(func(args mock.Arguments) {
var arg0 context.Context
if args[0] != nil {
arg0 = args[0].(context.Context)
}
var arg1 string
var arg1 int
if args[1] != nil {
arg1 = args[1].(string)
arg1 = args[1].(int)
}
var arg2 string
var arg2 int
if args[2] != nil {
arg2 = args[2].(string)
}
var arg3 int
if args[3] != nil {
arg3 = args[3].(int)
}
var arg4 int
if args[4] != nil {
arg4 = args[4].(int)
arg2 = args[2].(int)
}
run(
arg0,
arg1,
arg2,
arg3,
arg4,
)
})
return _c
@@ -326,81 +164,51 @@ func (_c *MockClient_GetMessage_Call) Return(runnerScaleSetMessage *scaleset.Run
return _c
}
func (_c *MockClient_GetMessage_Call) RunAndReturn(run func(ctx context.Context, messageQueueURL string, messageQueueAccessToken string, lastMessageID int, maxCapacity int) (*scaleset.RunnerScaleSetMessage, error)) *MockClient_GetMessage_Call {
func (_c *MockClient_GetMessage_Call) RunAndReturn(run func(ctx context.Context, lastMessageID int, maxCapacity int) (*scaleset.RunnerScaleSetMessage, error)) *MockClient_GetMessage_Call {
_c.Call.Return(run)
return _c
}
// RefreshMessageSession provides a mock function for the type MockClient
func (_mock *MockClient) RefreshMessageSession(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) (*scaleset.RunnerScaleSetSession, error) {
ret := _mock.Called(ctx, runnerScaleSetID, sessionID)
// Session provides a mock function for the type MockClient
func (_mock *MockClient) Session() scaleset.RunnerScaleSetSession {
ret := _mock.Called()
if len(ret) == 0 {
panic("no return value specified for RefreshMessageSession")
panic("no return value specified for Session")
}
var r0 *scaleset.RunnerScaleSetSession
var r1 error
if returnFunc, ok := ret.Get(0).(func(context.Context, int, uuid.UUID) (*scaleset.RunnerScaleSetSession, error)); ok {
return returnFunc(ctx, runnerScaleSetID, sessionID)
}
if returnFunc, ok := ret.Get(0).(func(context.Context, int, uuid.UUID) *scaleset.RunnerScaleSetSession); ok {
r0 = returnFunc(ctx, runnerScaleSetID, sessionID)
var r0 scaleset.RunnerScaleSetSession
if returnFunc, ok := ret.Get(0).(func() scaleset.RunnerScaleSetSession); ok {
r0 = returnFunc()
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*scaleset.RunnerScaleSetSession)
}
r0 = ret.Get(0).(scaleset.RunnerScaleSetSession)
}
if returnFunc, ok := ret.Get(1).(func(context.Context, int, uuid.UUID) error); ok {
r1 = returnFunc(ctx, runnerScaleSetID, sessionID)
} else {
r1 = ret.Error(1)
}
return r0, r1
return r0
}
// MockClient_RefreshMessageSession_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RefreshMessageSession'
type MockClient_RefreshMessageSession_Call struct {
// MockClient_Session_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Session'
type MockClient_Session_Call struct {
*mock.Call
}
// RefreshMessageSession is a helper method to define mock.On call
// - ctx context.Context
// - runnerScaleSetID int
// - sessionID uuid.UUID
func (_e *MockClient_Expecter) RefreshMessageSession(ctx interface{}, runnerScaleSetID interface{}, sessionID interface{}) *MockClient_RefreshMessageSession_Call {
return &MockClient_RefreshMessageSession_Call{Call: _e.mock.On("RefreshMessageSession", ctx, runnerScaleSetID, sessionID)}
// Session is a helper method to define mock.On call
func (_e *MockClient_Expecter) Session() *MockClient_Session_Call {
return &MockClient_Session_Call{Call: _e.mock.On("Session")}
}
func (_c *MockClient_RefreshMessageSession_Call) Run(run func(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID)) *MockClient_RefreshMessageSession_Call {
func (_c *MockClient_Session_Call) Run(run func()) *MockClient_Session_Call {
_c.Call.Run(func(args mock.Arguments) {
var arg0 context.Context
if args[0] != nil {
arg0 = args[0].(context.Context)
}
var arg1 int
if args[1] != nil {
arg1 = args[1].(int)
}
var arg2 uuid.UUID
if args[2] != nil {
arg2 = args[2].(uuid.UUID)
}
run(
arg0,
arg1,
arg2,
)
run()
})
return _c
}
func (_c *MockClient_RefreshMessageSession_Call) Return(runnerScaleSetSession *scaleset.RunnerScaleSetSession, err error) *MockClient_RefreshMessageSession_Call {
_c.Call.Return(runnerScaleSetSession, err)
func (_c *MockClient_Session_Call) Return(runnerScaleSetSession scaleset.RunnerScaleSetSession) *MockClient_Session_Call {
_c.Call.Return(runnerScaleSetSession)
return _c
}
func (_c *MockClient_RefreshMessageSession_Call) RunAndReturn(run func(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) (*scaleset.RunnerScaleSetSession, error)) *MockClient_RefreshMessageSession_Call {
func (_c *MockClient_Session_Call) RunAndReturn(run func() scaleset.RunnerScaleSetSession) *MockClient_Session_Call {
_c.Call.Return(run)
return _c
}
+322
View File
@@ -0,0 +1,322 @@
package scaleset
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"strconv"
"sync"
"github.com/google/uuid"
)
// MessageSessionClient is a client used to interact with a message session for a runner scale set.
// It provides methods to Get and Delete messages from the message queue associated with the session,
// handling session token expiration and refreshing as needed.
//
// It is safe for concurrent use by multiple goroutines.
// Please do not forget to call Close when done to clean up the session.
type MessageSessionClient struct {
mu sync.Mutex
// inner client is the parent of the message session, allowing session refreshing
// use this client to create (and potentially refresh the session) requests.
innerClient *Client
// commonClient uses different options than the original client
// use this client for message session requests
commonClient *commonClient
scaleSetID int
owner string
session *RunnerScaleSetSession
}
// Close deletes the message session associated with this client.
func (c *MessageSessionClient) Close(ctx context.Context) error {
c.mu.Lock()
defer c.mu.Unlock()
return c.deleteMessageSession(ctx, c.scaleSetID, c.session.SessionID)
}
func (c *MessageSessionClient) createMessageSession(ctx context.Context) error {
path := fmt.Sprintf("/%s/%d/sessions", scaleSetEndpoint, c.scaleSetID)
newSession := &RunnerScaleSetSession{
OwnerName: c.owner,
}
requestData, err := json.Marshal(newSession)
if err != nil {
return fmt.Errorf("failed to marshal new session: %w", err)
}
var createdSession RunnerScaleSetSession
if err = c.doSessionRequest(
ctx,
http.MethodPost,
path,
bytes.NewBuffer(requestData),
http.StatusOK,
&createdSession,
); err != nil {
return fmt.Errorf("failed to do the session request: %w", err)
}
c.session = &createdSession
return nil
}
// DeleteMessageSession deletes a message session for the specified runner scale set.
func (c *MessageSessionClient) deleteMessageSession(ctx context.Context, runnerScaleSetID int, sessionID uuid.UUID) error {
path := fmt.Sprintf("/%s/%d/sessions/%s", scaleSetEndpoint, runnerScaleSetID, sessionID.String())
return c.doSessionRequest(ctx, http.MethodDelete, path, nil, http.StatusNoContent, nil)
}
// RefreshMessageSession refreshes a message session for the specified runner scale set.
// This should be used when a MessageQueueTokenExpiredError is encountered.
func (c *MessageSessionClient) refreshMessageSession(ctx context.Context) error {
path := fmt.Sprintf("/%s/%d/sessions/%s", scaleSetEndpoint, c.scaleSetID, c.session.SessionID.String())
refreshedSession := &RunnerScaleSetSession{}
if err := c.doSessionRequest(ctx, http.MethodPatch, path, nil, http.StatusOK, refreshedSession); err != nil {
return fmt.Errorf("failed to do the session request: %w", err)
}
c.session = refreshedSession
return nil
}
// GetMessage fetches a message from the runner scale set message queue. If there are no messages available, it returns (nil, nil).
// Unless a message is deleted after being processed (using DeleteMessage), it will be returned again in subsequent calls.
// If the current session token is expired, it refreshes the session and tries one more time.
func (c *MessageSessionClient) GetMessage(ctx context.Context, lastMessageID int, maxCapacity int) (*RunnerScaleSetMessage, error) {
c.mu.Lock()
defer c.mu.Unlock()
message, err := c.getMessage(
ctx,
lastMessageID,
maxCapacity,
)
if err == nil {
return message, nil
}
expiredError := &ActionsError{}
if !errors.As(err, &expiredError) || !expiredError.IsMessageQueueTokenExpired() {
return nil, fmt.Errorf("failed to get next message: %w", err)
}
if err := c.refreshMessageSession(ctx); err != nil {
return nil, fmt.Errorf("failed to refresh message session: %w", err)
}
return c.getMessage(
ctx,
lastMessageID,
maxCapacity,
)
}
func (c *MessageSessionClient) getMessage(ctx context.Context, lastMessageID int, maxCapacity int) (*RunnerScaleSetMessage, error) {
u, err := url.Parse(c.session.MessageQueueURL)
if err != nil {
return nil, fmt.Errorf("failed to parse message queue url: %w", err)
}
if lastMessageID > 0 {
q := u.Query()
q.Set("lastMessageId", strconv.Itoa(lastMessageID))
u.RawQuery = q.Encode()
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
if err != nil {
return nil, fmt.Errorf("failed to create new request with context: %w", err)
}
req.Header.Set("Accept", "application/json; api-version=6.0-preview")
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", c.session.MessageQueueAccessToken))
req.Header.Set("User-Agent", c.commonClient.userAgent)
req.Header.Set(HeaderScaleSetMaxCapacity, strconv.Itoa(maxCapacity))
resp, err := c.commonClient.do(req)
if err != nil {
return nil, fmt.Errorf("failed to issue the request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusAccepted {
return nil, nil
}
if resp.StatusCode != http.StatusOK {
if resp.StatusCode != http.StatusUnauthorized {
return nil, ParseActionsErrorFromResponse(resp)
}
body, err := io.ReadAll(resp.Body)
body = trimByteOrderMark(body)
if err != nil {
return nil, &ActionsError{
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: err,
}
}
return nil, &ActionsError{
ActivityID: resp.Header.Get(headerActionsActivityID),
StatusCode: resp.StatusCode,
Err: &messageQueueTokenExpiredError{
message: string(body),
},
}
}
message, err := parseRunnerScaleSetMessageResponse(resp.Body)
if err != nil {
return nil, &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
return message, nil
}
// DeleteMessage deletes a message from the runner scale set message queue.
// This should typically be done after processing the message and acts as an acknowledgment.
// If the current session token is expired, it refreshes the session and tries one more time.
func (c *MessageSessionClient) DeleteMessage(ctx context.Context, messageID int) error {
c.mu.Lock()
defer c.mu.Unlock()
err := c.deleteMessage(ctx, messageID)
if err == nil {
return nil
}
expiredError := &ActionsError{}
if !errors.As(err, &expiredError) || !expiredError.IsMessageQueueTokenExpired() {
return fmt.Errorf("failed to delete message: %w", err)
}
if err := c.refreshMessageSession(ctx); err != nil {
return fmt.Errorf("failed to refresh message session: %w", err)
}
return c.deleteMessage(ctx, messageID)
}
func (c *MessageSessionClient) deleteMessage(ctx context.Context, messageID int) error {
u, err := url.Parse(c.session.MessageQueueURL)
if err != nil {
return fmt.Errorf("failed to parse message queue url: %w", err)
}
u.Path = fmt.Sprintf("%s/%d", u.Path, messageID)
req, err := http.NewRequestWithContext(ctx, http.MethodDelete, u.String(), nil)
if err != nil {
return fmt.Errorf("failed to create new request with context: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", c.session.MessageQueueAccessToken))
req.Header.Set("User-Agent", c.commonClient.userAgent)
resp, err := c.commonClient.do(req)
if err != nil {
return fmt.Errorf("failed to issue the request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusNoContent {
return nil
}
if resp.StatusCode != http.StatusUnauthorized {
return ParseActionsErrorFromResponse(resp)
}
body, err := io.ReadAll(resp.Body)
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 {
c.mu.Lock()
defer c.mu.Unlock()
if c.session == nil {
return RunnerScaleSetSession{}
}
return *c.session
}
func (c *MessageSessionClient) doSessionRequest(ctx context.Context, method, path string, requestData io.Reader, expectedResponseStatusCode int, responseUnmarshalTarget any) error {
req, err := c.innerClient.newActionsServiceRequest(ctx, method, path, requestData)
if err != nil {
return fmt.Errorf("failed to create new actions service request: %w", err)
}
// use potentially modified client to issue a request
resp, err := c.commonClient.do(req)
if err != nil {
return fmt.Errorf("failed to issue the request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode == expectedResponseStatusCode {
if responseUnmarshalTarget == nil {
return nil
}
if err := json.NewDecoder(resp.Body).Decode(responseUnmarshalTarget); err != nil {
return &ActionsError{
StatusCode: resp.StatusCode,
ActivityID: resp.Header.Get(headerActionsActivityID),
Err: err,
}
}
return nil
}
if resp.StatusCode >= 400 && resp.StatusCode < 500 {
return ParseActionsErrorFromResponse(resp)
}
body, err := io.ReadAll(resp.Body)
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)),
})
}
+662
View File
@@ -0,0 +1,662 @@
package scaleset
import (
"context"
"encoding/json"
"errors"
"net/http"
"strconv"
"strings"
"testing"
"time"
"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func newTestSessionRequestHandler(t *testing.T, session RunnerScaleSetSession) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
srv := r.Context().Value(ctxKeyServer).(*actionsServer)
session.MessageQueueURL = srv.URL
resp, err := json.Marshal(session)
require.NoError(t, err)
w.Header().Set("Content-Type", "application/json")
w.Write(resp)
}
}
func TestCreateMessageSession(t *testing.T) {
ctx := context.Background()
auth := actionsAuth{
token: "token",
}
t.Run("CreateMessageSession unmarshals correctly", func(t *testing.T) {
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
RunnerSetting: RunnerSetting{},
}
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
handleSessionRequest(w, r)
}))
want := server.testRunnerScaleSetSession()
handleSessionRequest = newTestSessionRequestHandler(t, want)
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, runnerScaleSet.ID, "my-org")
require.NoError(t, err)
session := sessionClient.Session()
require.NotEqual(t, session.SessionID, uuid.Nil)
assert.Equal(t, want, session)
})
t.Run("CreateMessageSession unmarshals errors into ActionsError", func(t *testing.T) {
owner := "foo"
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
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) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set(headerActionsActivityID, exampleRequestID)
w.WriteHeader(http.StatusBadRequest)
resp := []byte(`{"typeName": "CSharpExceptionNameHere","message": "could not do something"}`)
w.Write(resp)
}))
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(context.Background(), runnerScaleSet.ID, owner)
assert.Nil(t, sessionClient)
errorTypeForComparison := &ActionsError{}
assert.ErrorAs(t, err, &errorTypeForComparison)
assert.Equal(t, want, errorTypeForComparison)
})
t.Run("CreateMessageSession call is retried the correct amount of times", func(t *testing.T) {
owner := "foo"
runnerScaleSet := RunnerScaleSet{
ID: 1,
Name: "ScaleSet",
CreatedOn: time.Date(1, time.January, 1, 0, 0, 0, 0, time.UTC),
RunnerSetting: RunnerSetting{},
}
gotRetries := 0
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
gotRetries++
}))
retryMax := 3
retryWaitMax := 1 * time.Microsecond
wantRetries := retryMax + 1
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(retryWaitMax),
)
require.NoError(t, err)
_, err = client.MessageSessionClient(
ctx,
runnerScaleSet.ID,
owner,
WithRetryMax(retryMax),
WithRetryWaitMax(retryWaitMax),
)
assert.NotNil(t, err)
assert.Equalf(t, gotRetries, wantRetries, "CreateMessageSession got unexpected retry count: got=%v, want=%v", gotRetries, wantRetries)
})
}
func TestGetMessage(t *testing.T) {
ctx := context.Background()
auth := actionsAuth{
token: "token",
}
runnerScaleSetMessage := &RunnerScaleSetMessage{
MessageID: 1,
}
t.Run("Get Runner Scale Set Message", func(t *testing.T) {
want := runnerScaleSetMessage
response := []byte(`{"messageId":1,"messageType":"RunnerScaleSetJobMessages"}`)
var handleSessionRequest http.HandlerFunc
s := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.Write(response)
}))
handleSessionRequest = newTestSessionRequestHandler(t, s.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
s.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
got, err := sessionClient.GetMessage(ctx, 0, 10)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("GetMessage sets the last message id if not 0", func(t *testing.T) {
want := runnerScaleSetMessage
response := []byte(`{"messageId":1,"messageType":"RunnerScaleSetJobMessages"}`)
var handleSessionRequest http.HandlerFunc
s := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
q := r.URL.Query()
assert.Equal(t, "1", q.Get("lastMessageId"))
w.Write(response)
}))
handleSessionRequest = newTestSessionRequestHandler(t, s.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
s.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
got, err := sessionClient.GetMessage(ctx, 1, 10)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("Default retries on server error", func(t *testing.T) {
retryMax := 1
actualRetry := 0
expectedRetry := retryMax + 1
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusServiceUnavailable)
actualRetry++
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(1*time.Millisecond),
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(
ctx,
1,
"my-org",
WithRetryMax(retryMax),
WithRetryWaitMax(1*time.Millisecond),
)
require.NoError(t, err)
msg, err := sessionClient.GetMessage(ctx, 0, 10)
assert.Nil(t, msg)
assert.NotNil(t, err)
assert.Equalf(t, actualRetry, expectedRetry, "A retry was expected after the first request but got: %v", actualRetry)
})
t.Run("Message token expired", func(t *testing.T) {
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// create session
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
// refresh
if strings.Contains(r.URL.Path, "/sessions/") {
// just set the same session
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusUnauthorized)
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
msg, err := sessionClient.GetMessage(ctx, 0, 10)
assert.Nil(t, msg)
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) {
want := runnerScaleSetMessage
afterRefreshResponse := []byte(`{"messageId":1,"messageType":"RunnerScaleSetJobMessages"}`)
var handleSessionRequest http.HandlerFunc
type state int
const (
createSession state = iota
firstGetMessage
refreshToken
secondGetMessage
)
currentState := createSession
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// create session
if strings.HasSuffix(r.URL.Path, "sessions") {
require.Equal(t, createSession, currentState)
handleSessionRequest(w, r)
currentState = firstGetMessage
return
}
// refresh
if strings.Contains(r.URL.Path, "/sessions/") {
// just set the same session
require.Equal(t, refreshToken, currentState)
handleSessionRequest(w, r)
currentState = secondGetMessage
return
}
if currentState == firstGetMessage {
w.WriteHeader(http.StatusUnauthorized)
currentState = refreshToken
return
}
require.Equal(t, secondGetMessage, currentState)
w.Write(afterRefreshResponse)
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
got, err := sessionClient.GetMessage(ctx, 0, 10)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("Status code not found", func(t *testing.T) {
want := ActionsError{
Err: errors.New("unknown exception"),
StatusCode: 404,
}
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusNotFound)
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
msg, err := sessionClient.GetMessage(ctx, 0, 10)
assert.Nil(t, msg)
var got *ActionsError
require.ErrorAs(t, err, &got)
assert.Equal(t, want.StatusCode, got.StatusCode)
assert.Equal(t, want.Err.Error(), got.Err.Error())
})
t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain")
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
msg, err := sessionClient.GetMessage(ctx, 0, 10)
assert.Nil(t, msg)
assert.NotNil(t, err)
})
t.Run("Capacity error handling", func(t *testing.T) {
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
hc := r.Header.Get(HeaderScaleSetMaxCapacity)
c, err := strconv.Atoi(hc)
require.NoError(t, err)
assert.GreaterOrEqual(t, c, 0)
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain")
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
msg, err := sessionClient.GetMessage(ctx, 0, 0)
assert.Nil(t, msg)
assert.Error(t, err)
var expectedErr *ActionsError
assert.ErrorAs(t, err, &expectedErr)
assert.Equal(t, http.StatusBadRequest, expectedErr.StatusCode)
})
}
func TestDeleteMessage(t *testing.T) {
ctx := context.Background()
auth := actionsAuth{
token: "token",
}
runnerScaleSetMessage := &RunnerScaleSetMessage{
MessageID: 1,
}
t.Run("Delete existing message", func(t *testing.T) {
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusNoContent)
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID)
assert.Nil(t, err)
})
t.Run("Message token expired", func(t *testing.T) {
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// create session
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
// refresh
if strings.Contains(r.URL.Path, "/sessions/") {
// just set the same session
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusUnauthorized)
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
err = sessionClient.DeleteMessage(ctx, 0)
require.NotNil(t, err)
var expectedErr *ActionsError
require.ErrorAs(t, err, &expectedErr)
assert.True(t, expectedErr.IsMessageQueueTokenExpired())
})
t.Run("message token refreshed", func(t *testing.T) {
type state int
const (
createSession state = iota
firstDeleteMessage
refreshToken
secondDeleteMessage
)
currentState := createSession
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// create session
if strings.HasSuffix(r.URL.Path, "sessions") {
require.Equal(t, createSession, currentState)
handleSessionRequest(w, r)
currentState = firstDeleteMessage
return
}
// refresh
if strings.Contains(r.URL.Path, "/sessions/") {
// just set the same session
require.Equal(t, refreshToken, currentState)
handleSessionRequest(w, r)
currentState = secondDeleteMessage
return
}
if currentState == firstDeleteMessage {
w.WriteHeader(http.StatusUnauthorized)
currentState = refreshToken
return
}
require.Equal(t, secondDeleteMessage, currentState)
w.WriteHeader(http.StatusNoContent)
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
err = sessionClient.DeleteMessage(ctx, 0)
require.NoError(t, err)
})
t.Run("Error when Content-Type is text/plain", func(t *testing.T) {
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusBadRequest)
w.Header().Set("Content-Type", "text/plain")
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID)
require.NotNil(t, err)
var expectedErr *ActionsError
assert.True(t, errors.As(err, &expectedErr))
},
)
t.Run("Default retries on server error", func(t *testing.T) {
actualRetry := 0
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.WriteHeader(http.StatusServiceUnavailable)
actualRetry++
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
retryMax := 1
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
WithRetryMax(retryMax),
WithRetryWaitMax(1*time.Nanosecond),
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(
ctx,
1,
"my-org",
WithRetryMax(retryMax),
WithRetryWaitMax(1*time.Nanosecond),
)
require.NoError(t, err)
err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID)
assert.NotNil(t, err)
expectedRetry := retryMax + 1
assert.Equalf(t, actualRetry, expectedRetry, "A retry was expected after the first request but got: %v", actualRetry)
})
t.Run("No message found", func(t *testing.T) {
want := (*RunnerScaleSetMessage)(nil)
rsl, err := json.Marshal(want)
require.NoError(t, err)
var handleSessionRequest http.HandlerFunc
server := newActionsServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.HasSuffix(r.URL.Path, "sessions") {
handleSessionRequest(w, r)
return
}
w.Write(rsl)
}))
handleSessionRequest = newTestSessionRequestHandler(t, server.testRunnerScaleSetSession())
client, err := newClient(
testSystemInfo,
server.configURLForOrg("my-org"),
auth,
)
require.NoError(t, err)
sessionClient, err := client.MessageSessionClient(ctx, 1, "my-org")
require.NoError(t, err)
err = sessionClient.DeleteMessage(ctx, runnerScaleSetMessage.MessageID+1)
var expectedErr *ActionsError
require.True(t, errors.As(err, &expectedErr))
})
}