Skip to content
This repository was archived by the owner on Apr 21, 2026. It is now read-only.
Merged
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
25 changes: 22 additions & 3 deletions handler/azure.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,31 +6,42 @@ import (
"time"

"github.com/go-resty/resty/v2"
"github.com/sirupsen/logrus"
)

// NewAzureInterruptChecker checks for azure spot interrupt event from metadata server.
// See https://docs.microsoft.com/en-us/azure/virtual-machines/linux/scheduled-events#endpoint-discovery
func NewAzureInterruptChecker() MetadataChecker {
func NewAzureInterruptChecker(log logrus.FieldLogger) MetadataChecker {
client := resty.New()
// Times out if set to 1 second, after 2 we will try again soon anyway
client.SetTimeout(time.Second * 2)

return &azureInterruptChecker{
client: client,
metadataServerURL: "http://169.254.169.254",
log: log,
}
}

type azureInterruptChecker struct {
client *resty.Client
metadataServerURL string
log logrus.FieldLogger
}

// azureSpotScheduledEvent is a single event schema, not all fields are necessarily mapped
// see https://learn.microsoft.com/en-us/azure/virtual-machines/linux/scheduled-events#the-basics for details
type azureSpotScheduledEvent struct {
EventType string
EventId string
EventType string
EventStatus string
EventSource string
Description string
NotBefore string
}
type azureSpotScheduledEvents struct {
Events []azureSpotScheduledEvent
DocumentIncarnation int
Events []azureSpotScheduledEvent
}

func (c *azureInterruptChecker) CheckInterrupt(ctx context.Context) (bool, error) {
Expand All @@ -47,7 +58,15 @@ func (c *azureInterruptChecker) CheckInterrupt(ctx context.Context) (bool, error
return false, fmt.Errorf("received unexpected status code: %d", resp.StatusCode())
}

if len(responseBody.Events) == 0 {
return false, nil
}

c.log.Debugf("Received %d scheduled events with incarnation %d", len(responseBody.Events), responseBody.DocumentIncarnation)

for _, e := range responseBody.Events {
c.log.Debugf("Scheduled event seen: EventId=%s, EventType=%s, EventStatus=%s, EventSource=%s, NotBefore=%s, Description=%s",
e.EventId, e.EventType, e.EventStatus, e.EventSource, e.NotBefore, e.Description)
if e.EventType == "Preempt" {
return true, nil
}
Expand Down
3 changes: 3 additions & 0 deletions handler/azure_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"testing"

"github.com/go-resty/resty/v2"
"github.com/sirupsen/logrus"
"github.com/stretchr/testify/require"
)

Expand All @@ -32,9 +33,11 @@ func TestAzureInterruptChecker(t *testing.T) {
}))
defer s.Close()

log := logrus.New()
checker := azureInterruptChecker{
client: resty.New(),
metadataServerURL: s.URL,
log: log,
}

interrupted, err := checker.CheckInterrupt(context.Background())
Expand Down
4 changes: 4 additions & 0 deletions handler/handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,7 @@ func (g *SpotHandler) Run(ctx context.Context) error {
return err
}
// Stop after ACK.
g.log.Infof("stopping poll ticker")
t.Stop()
}

Expand All @@ -124,11 +125,14 @@ func (g *SpotHandler) Run(ctx context.Context) error {
if err != nil {
g.log.Errorf("checking for cloud events: %v", err)
}
g.log.Debugf("poll tick completed")
case <-deadline.C:
g.log.Infof("grace period elapsed, exiting")
return nil
case <-ctx.Done():
// Signal received, starting countdown until exiting the loop.
once.Do(func() {
g.log.Infof("termination signal received, waiting grace period of %s before exit", g.gracePeriod)
deadline.Reset(g.gracePeriod)
})
}
Expand Down
9 changes: 5 additions & 4 deletions main.go
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,8 @@ func main() {
cfg := config.Get()

logger := logrus.New()
log := logrus.WithFields(logrus.Fields{})
logger.SetLevel(logrus.Level(cfg.LogLevel))
log := logger.WithFields(logrus.Fields{})

kubeconfig, err := retrieveKubeConfig(log, cfg)
if err != nil {
Expand Down Expand Up @@ -66,7 +67,7 @@ func main() {
"k8s_version": k8sVersionField,
})

interruptChecker, err := buildInterruptChecker(cfg.Provider)
interruptChecker, err := buildInterruptChecker(cfg.Provider, log)
if err != nil {
log.Fatalf("interrupt checker: %v", err)
}
Expand Down Expand Up @@ -115,10 +116,10 @@ func main() {
}
}

func buildInterruptChecker(provider string) (handler.MetadataChecker, error) {
func buildInterruptChecker(provider string, log logrus.FieldLogger) (handler.MetadataChecker, error) {
switch provider {
case "azure":
return handler.NewAzureInterruptChecker(), nil
return handler.NewAzureInterruptChecker(log), nil
case "gcp":
return handler.NewGCPChecker(), nil
case "aws":
Expand Down
Loading