Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 67 additions & 8 deletions workers/scaleset/scaleset.go
Original file line number Diff line number Diff line change
Expand Up @@ -155,6 +155,30 @@ func (w *Worker) ensureScaleSetInGitHub() error {
return nil
}

func (w *Worker) markScaleSetCreated() error {
if w.scaleSet.State == params.ScaleSetCreated {
return nil
}

entity, err := w.scaleSet.GetEntity()
if err != nil {
return fmt.Errorf("getting scale set entity: %w", err)
}
state := params.ScaleSetCreated
updated, err := w.store.UpdateEntityScaleSet(
w.ctx,
entity,
w.scaleSet.ID,
params.UpdateScaleSetParams{State: &state},
nil,
)
if err != nil {
return fmt.Errorf("updating scale set state: %w", err)
}
w.scaleSet = updated
return nil
}

func (w *Worker) Stop() error {
slog.DebugContext(w.ctx, "stopping scale set worker", "scale_set", w.consumerID)
w.mux.Lock()
Expand Down Expand Up @@ -303,6 +327,10 @@ func (w *Worker) Start() (err error) {
return fmt.Errorf("failed to ensure scale set: %w", err)
}

if err := w.markScaleSetCreated(); err != nil {
return fmt.Errorf("failed to mark scale set as created: %w", err)
}

consumer, err := watcher.RegisterConsumer(
w.ctx, w.consumerID,
watcher.WithAny(
Expand Down Expand Up @@ -688,6 +716,26 @@ func (w *Worker) handleInstanceCleanup(instance params.Instance) error {
return nil
}

func (w *Worker) reconcileRunners() error {
instances, err := w.store.ListScaleSetInstances(w.ctx, w.scaleSet.ID, false)
if err != nil {
return fmt.Errorf("listing scale set instances: %w", err)
}

runners := make(map[string]params.Instance, len(instances))
for _, instance := range instances {
runners[instance.ID] = instance
}
w.runners = runners

for _, instance := range instances {
if err := w.handleInstanceCleanup(instance); err != nil {
return err
}
}
return nil
}

func (w *Worker) handleInstanceEntityEvent(event dbCommon.ChangePayload) {
instance, ok := event.Payload.(params.Instance)
if !ok {
Expand Down Expand Up @@ -854,7 +902,8 @@ func (w *Worker) handleScaleUp() {
return
}

if w.targetRunners() <= w.runnerCount() {
runnersToAdd := w.runnersToAdd()
if runnersToAdd == 0 {
slog.DebugContext(w.ctx, "target is less than or equal to current; not scaling up")
return
}
Expand All @@ -870,7 +919,7 @@ func (w *Worker) handleScaleUp() {
slog.ErrorContext(w.ctx, "error getting scale set client", "error", err)
return
}
for i := w.runnerCount(); i < w.targetRunners(); i++ {
for range runnersToAdd {
newRunnerName := strings.ToLower(fmt.Sprintf("%s-%s", w.scaleSet.GetRunnerPrefix(), util.NewID()))
jitConfig, err := scaleSetCli.GenerateJitRunnerConfig(w.ctx, newRunnerName, w.scaleSet.ScaleSetID)
if err != nil {
Expand Down Expand Up @@ -1005,7 +1054,6 @@ func (w *Worker) handleScaleDown() {
removed++
case commonParams.InstancePendingDelete, commonParams.InstancePendingForceDelete,
commonParams.InstanceDeleting, commonParams.InstanceDeleted:
removed++
continue
default:
slog.WarnContext(w.ctx, "runner is not in a valid state; skipping", "runner_name", runner.Name, "runner_status", runner.Status)
Expand All @@ -1025,7 +1073,20 @@ func (w *Worker) targetRunners() int {
}

func (w *Worker) runnerCount() int {
return len(w.runners)
count := 0
for _, runner := range w.runners {
switch runner.Status {
case commonParams.InstancePendingDelete, commonParams.InstancePendingForceDelete,
commonParams.InstanceDeleting, commonParams.InstanceDeleted:
continue
}
count++
}
return count
}

func (w *Worker) runnersToAdd() int {
return max(w.targetRunners()-w.runnerCount(), 0)
}

func (w *Worker) handleAutoScale() {
Expand Down Expand Up @@ -1059,10 +1120,8 @@ func (w *Worker) handleAutoScale() {
return
case <-ticker.C:
w.mux.Lock()
for _, instance := range w.runners {
if err := w.handleInstanceCleanup(instance); err != nil {
slog.ErrorContext(w.ctx, "error cleaning up instance", "instance_id", instance.ID, "error", err)
}
if err := w.reconcileRunners(); err != nil {
slog.ErrorContext(w.ctx, "error reconciling scale set instances", "error", err)
}

if w.runnerCount() == w.targetRunners() {
Expand Down
263 changes: 263 additions & 0 deletions workers/scaleset/scaleset_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,263 @@
// Copyright 2026 Cloudbase Solutions SRL
//
// Licensed under the Apache License, Version 2.0 (the "License"); you may
// not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
// WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
// License for the specific language governing permissions and limitations
// under the License.

//go:build testing

package scaleset

import (
"context"
"encoding/base64"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"sync/atomic"
"testing"
"time"

"github.com/google/go-github/v84/github"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require"

commonParams "github.com/cloudbase/garm-provider-common/params"
"github.com/cloudbase/garm/cache"
storeMocks "github.com/cloudbase/garm/database/common/mocks"
"github.com/cloudbase/garm/params"
runnerMocks "github.com/cloudbase/garm/runner/common/mocks"
)

func TestRunnerCountExcludesInstancesBeingDeleted(t *testing.T) {
w := &Worker{
scaleSet: params.ScaleSet{MinIdleRunners: 1, MaxRunners: 1},
runners: map[string]params.Instance{
"pending-delete": {Status: commonParams.InstancePendingDelete},
"pending-force-delete": {Status: commonParams.InstancePendingForceDelete},
"deleting": {Status: commonParams.InstanceDeleting},
"deleted": {Status: commonParams.InstanceDeleted},
},
}

assert.Zero(t, w.runnerCount())
assert.Equal(t, 1, w.runnersToAdd())
}

func TestDeletingInstancesDoNotIncreaseScaleDownDelta(t *testing.T) {
w := &Worker{
scaleSet: params.ScaleSet{DesiredRunnerCount: 1, MaxRunners: 1},
runners: map[string]params.Instance{
"running-1": {Status: commonParams.InstanceRunning},
"running-2": {Status: commonParams.InstanceRunning},
"deleting": {Status: commonParams.InstanceDeleting},
},
}

assert.Equal(t, 1, w.runnerCount()-w.targetRunners())
}

func TestPendingCreatePreventsDuplicateReplacement(t *testing.T) {
w := &Worker{
scaleSet: params.ScaleSet{MinIdleRunners: 1, MaxRunners: 1},
runners: map[string]params.Instance{
"deleted": {Status: commonParams.InstanceDeleted},
"replacement": {Status: commonParams.InstancePendingCreate},
},
}

assert.Equal(t, 1, w.runnerCount())
assert.Equal(t, w.targetRunners(), w.runnerCount())
assert.Zero(t, w.runnersToAdd())
}

func TestTargetRunnersHonorsMaximum(t *testing.T) {
w := &Worker{
scaleSet: params.ScaleSet{
MinIdleRunners: 2,
DesiredRunnerCount: 3,
MaxRunners: 4,
},
}

assert.Equal(t, 4, w.targetRunners())
}

func TestReconcileRunnersCleansDeletedRowsAndReplacesStaleCache(t *testing.T) {
store := storeMocks.NewStore(t)
store.EXPECT().
ListScaleSetInstances(mock.Anything, uint(4), false).
Return([]params.Instance{
{ID: "deleted", Name: "deleted", Status: commonParams.InstanceDeleted, ScaleSetID: 4},
{ID: "replacement", Name: "replacement", Status: commonParams.InstancePendingCreate, ScaleSetID: 4},
}, nil)
store.EXPECT().
DeleteInstanceByName(mock.Anything, "deleted").
Return(nil)

w := &Worker{
ctx: context.Background(),
store: store,
scaleSet: params.ScaleSet{ID: 4, MinIdleRunners: 1, MaxRunners: 1},
runners: map[string]params.Instance{
"stale": {ID: "stale", Name: "stale", Status: commonParams.InstanceRunning, ScaleSetID: 4},
},
}

assert.NoError(t, w.reconcileRunners())
assert.Equal(t, 1, w.runnerCount())
assert.Len(t, w.runners, 1)
assert.Contains(t, w.runners, "replacement")
}

func TestReconcileDeletedRunnerCreatesOneReplacement(t *testing.T) {
const (
entityID = "repo-id"
scaleSetID = 4
)

var registrationRequests atomic.Int32
var jitRequests atomic.Int32
encodedJIT := base64.StdEncoding.EncodeToString([]byte(`{"key":"value"}`))

var server *httptest.Server
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/actions/runner-registration":
registrationRequests.Add(1)
assert.Equal(t, http.MethodPost, r.Method)
_, _ = fmt.Fprintf(w, `{"url":%q,"token":"eyJhbGciOiJub25lIn0.eyJleHAiOjQxNDk5MzYwMDB9."}`, server.URL)
case "/_apis/runtime/runnerscalesets/42/generatejitconfig":
jitRequests.Add(1)
assert.Equal(t, http.MethodPost, r.Method)
_, _ = fmt.Fprintf(w, `{"runner":{"id":99},"encodedJITConfig":%q}`, encodedJIT)
default:
http.NotFound(w, r)
t.Errorf("unexpected GitHub request: %s %s", r.Method, r.URL.Path)
}
}))
t.Cleanup(server.Close)

baseURL, err := url.Parse(server.URL)
require.NoError(t, err)
entity := params.ForgeEntity{
ID: entityID,
EntityType: params.ForgeEntityTypeRepository,
Owner: "owner",
Name: "repo",
}
githubClient := runnerMocks.NewGithubClient(t)
githubClient.EXPECT().
CreateEntityRegistrationToken(mock.Anything).
Return(&github.RegistrationToken{
Token: github.Ptr("registration-token"),
ExpiresAt: &github.Timestamp{Time: time.Now().Add(time.Hour)},
}, nil, nil).
Once()
githubClient.EXPECT().GithubBaseURL().Return(baseURL).Once()
githubClient.EXPECT().GetEntity().Return(entity).Maybe()
cache.SetGithubClient(entityID, githubClient)
t.Cleanup(func() { cache.DeleteGithubClient(entityID) })

store := storeMocks.NewStore(t)
store.EXPECT().
ListScaleSetInstances(mock.Anything, uint(scaleSetID), false).
Return([]params.Instance{{
ID: "deleted",
Name: "deleted",
Status: commonParams.InstanceDeleted,
RunnerStatus: params.RunnerIdle,
ScaleSetID: scaleSetID,
}}, nil).
Once()
store.EXPECT().DeleteInstanceByName(mock.Anything, "deleted").Return(nil).Once()
store.EXPECT().ControllerInfo().Return(params.ControllerInfo{
CallbackURL: "https://garm.example/callback",
MetadataURL: "https://garm.example/metadata",
}, nil).Once()
store.EXPECT().
CreateScaleSetInstance(
mock.Anything,
uint(scaleSetID),
mock.MatchedBy(func(create params.CreateInstanceParams) bool {
return create.Status == commonParams.InstancePendingCreate &&
create.RunnerStatus == params.RunnerPending &&
create.CallbackURL == "https://garm.example/callback" &&
create.MetadataURL == "https://garm.example/metadata" &&
create.AgentID == 99 &&
create.JitConfiguration["key"] == "value"
}),
).
Return(params.Instance{
ID: "replacement",
Name: "replacement",
Status: commonParams.InstancePendingCreate,
RunnerStatus: params.RunnerPending,
ScaleSetID: scaleSetID,
}, nil).
Once()

w := &Worker{
ctx: context.Background(),
store: store,
scaleSet: params.ScaleSet{
ID: scaleSetID,
ScaleSetID: 42,
RepoID: entityID,
Enabled: true,
MinIdleRunners: 1,
MaxRunners: 1,
},
runners: map[string]params.Instance{
"stale": {ID: "stale", Status: commonParams.InstanceRunning},
},
}

require.NoError(t, w.reconcileRunners())
w.handleScaleUp()
w.handleScaleUp()

assert.EqualValues(t, 1, registrationRequests.Load())
assert.EqualValues(t, 1, jitRequests.Load())
assert.Len(t, w.runners, 1)
assert.Equal(t, commonParams.InstancePendingCreate, w.runners["replacement"].Status)
}

func TestMarkScaleSetCreated(t *testing.T) {
store := storeMocks.NewStore(t)
entity := params.ForgeEntity{ID: "repo-id", EntityType: params.ForgeEntityTypeRepository}
store.EXPECT().
UpdateEntityScaleSet(
mock.Anything,
entity,
uint(4),
mock.MatchedBy(func(update params.UpdateScaleSetParams) bool {
return update.State != nil && *update.State == params.ScaleSetCreated
}),
mock.Anything,
).
Return(params.ScaleSet{ID: 4, RepoID: entity.ID, State: params.ScaleSetCreated}, nil)

w := &Worker{
ctx: context.Background(),
store: store,
scaleSet: params.ScaleSet{
ID: 4,
RepoID: entity.ID,
State: params.ScaleSetPendingCreate,
},
}

assert.NoError(t, w.markScaleSetCreated())
assert.Equal(t, params.ScaleSetCreated, w.scaleSet.State)
}