Skip to content
Open
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
74 changes: 74 additions & 0 deletions pkg/interruptionevent/draincordon/handler_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
package draincordon

import (
"context"
"errors"
"testing"

"github.com/aws/aws-node-termination-handler/pkg/config"
"github.com/aws/aws-node-termination-handler/pkg/ec2metadata"
"github.com/aws/aws-node-termination-handler/pkg/interruptioneventstore"
"github.com/aws/aws-node-termination-handler/pkg/monitor"
"github.com/aws/aws-node-termination-handler/pkg/node"
"github.com/aws/aws-node-termination-handler/pkg/observability"
"github.com/aws/aws-node-termination-handler/pkg/uptime"
corev1 "k8s.io/api/core/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/client-go/kubernetes/fake"
"k8s.io/kubectl/pkg/drain"
)

func TestHandleEvent_PreDrainFailureSkipsCordonAndPostDrainTask(t *testing.T) {
const nodeName = "test-node"
preDrainErr := errors.New("taint update failed")
postDrainCalls := 0

client := fake.NewSimpleClientset(&corev1.Node{
ObjectMeta: metav1.ObjectMeta{Name: nodeName},
})
testNode, err := node.NewWithValues(config.Config{}, &drain.Helper{
Ctx: context.TODO(),
Client: client,
}, uptime.Uptime)
if err != nil {
t.Fatalf("create test node: %v", err)
}

nthConfig := config.Config{CordonOnly: true}
handler := New(
interruptioneventstore.New(nthConfig),
*testNode,
nthConfig,
ec2metadata.NodeMetadata{},
observability.Metrics{},
observability.K8sEventRecorder{},
)
drainEvent := &monitor.InterruptionEvent{
EventID: "taint-failure-event",
Kind: monitor.RebalanceRecommendationKind,
NodeName: nodeName,
PreDrainTask: func(monitor.InterruptionEvent, node.Node) error {
return preDrainErr
},
PostDrainTask: func(monitor.InterruptionEvent, node.Node) error {
postDrainCalls++
return nil
},
}

err = handler.HandleEvent(drainEvent)
if !errors.Is(err, preDrainErr) {
t.Fatalf("expected pre-drain error %v, got %v", preDrainErr, err)
}
if postDrainCalls != 0 {
t.Fatalf("post-drain task was called %d times after pre-drain failure", postDrainCalls)
}

updatedNode, err := client.CoreV1().Nodes().Get(context.Background(), nodeName, metav1.GetOptions{})
if err != nil {
t.Fatalf("get test node: %v", err)
}
if updatedNode.Spec.Unschedulable {
t.Fatal("node was cordoned after pre-drain failure")
}
}
3 changes: 2 additions & 1 deletion pkg/monitor/sqsevent/asg-lifecycle-event.go
Original file line number Diff line number Diff line change
Expand Up @@ -112,7 +112,7 @@ func (m SQSMonitor) asgTerminationToInterruptionEvent(event *EventBridgeEvent, m
close(stopHeartbeatCh)
return m.deleteMessage(message)
}

interruptionEvent.CancelDrainTask = func(_ monitor.InterruptionEvent, _ node.Node) error {
close(cancelHeartbeatCh)
return nil
Expand Down Expand Up @@ -140,6 +140,7 @@ func (m SQSMonitor) asgTerminationToInterruptionEvent(event *EventBridgeEvent, m
err := n.TaintASGLifecycleTermination(nodeName, interruptionEvent.EventID)
if err != nil {
log.Err(err).Msgf("unable to taint node with taint %s:%s", node.ASGLifecycleTerminationTaint, interruptionEvent.EventID)
return err
}
return nil
}
Expand Down
125 changes: 125 additions & 0 deletions pkg/monitor/sqsevent/pre-drain-task_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
package sqsevent

import (
"context"
"errors"
"testing"

"github.com/aws/aws-node-termination-handler/pkg/config"
"github.com/aws/aws-node-termination-handler/pkg/monitor"
"github.com/aws/aws-node-termination-handler/pkg/node"
h "github.com/aws/aws-node-termination-handler/pkg/test"
"github.com/aws/aws-node-termination-handler/pkg/uptime"
"github.com/aws/aws-sdk-go/service/sqs"
corev1 "k8s.io/api/core/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime"
"k8s.io/client-go/kubernetes/fake"
k8stesting "k8s.io/client-go/testing"
"k8s.io/kubectl/pkg/drain"
)

const (
taintFailureTestInstanceID = "i-0123456789"
taintFailureTestNodeName = "ip-10-0-0-157.us-east-2.compute.internal"
)

func TestPreDrainTaskReturnsTaintError(t *testing.T) {
sqsMonitor := SQSMonitor{
EC2: h.MockedEC2{
DescribeInstancesResp: getDescribeInstancesResp(
taintFailureTestInstanceID,
taintFailureTestNodeName,
map[string]string{ASGTagName: "test-asg"},
),
},
ASG: &h.MockedASG{},
}
message := &sqs.Message{}

tests := []struct {
name string
build func() (*monitor.InterruptionEvent, error)
}{
{
name: "ASG lifecycle termination",
build: func() (*monitor.InterruptionEvent, error) {
return sqsMonitor.asgTerminationToInterruptionEvent(&EventBridgeEvent{
ID: "asg-event",
Detail: []byte(`{
"LifecycleActionToken":"token",
"AutoScalingGroupName":"test-asg",
"LifecycleHookName":"test-hook",
"EC2InstanceId":"i-0123456789",
"LifecycleTransition":"autoscaling:EC2_INSTANCE_TERMINATING"
}`),
}, message)
},
},
{
name: "rebalance recommendation",
build: func() (*monitor.InterruptionEvent, error) {
return sqsMonitor.rebalanceRecommendationToInterruptionEvent(&EventBridgeEvent{
ID: "rebalance-event",
Detail: []byte(`{"instance-id":"i-0123456789"}`),
}, message)
},
},
{
name: "scheduled change",
build: func() (*monitor.InterruptionEvent, error) {
wrappers := sqsMonitor.scheduledEventToInterruptionEvents(&EventBridgeEvent{
ID: "scheduled-event",
Detail: []byte(`{
"eventTypeCategory":"scheduledChange",
"service":"EC2",
"affectedEntities":[{"entityValue":"i-0123456789"}]
}`),
}, message)
if len(wrappers) != 1 {
return nil, errors.New("expected one scheduled change interruption event")
}
return wrappers[0].InterruptionEvent, wrappers[0].Err
},
},
}

for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
interruptionEvent, err := test.build()
if err != nil {
t.Fatalf("build interruption event: %v", err)
}

err = interruptionEvent.PreDrainTask(*interruptionEvent, failingTaintNode(t))
if err == nil {
t.Fatal("expected pre-drain task to return the taint update error")
}
})
}
}

func failingTaintNode(t *testing.T) node.Node {
t.Helper()

client := fake.NewSimpleClientset(&corev1.Node{
ObjectMeta: metav1.ObjectMeta{Name: taintFailureTestNodeName},
})
client.PrependReactor("update", "nodes", func(k8stesting.Action) (bool, runtime.Object, error) {
return true, nil, errors.New("forced taint update failure")
})

testNode, err := node.NewWithValues(config.Config{
TaintNode: true,
HeartbeatInterval: -1,
HeartbeatUntil: -1,
}, &drain.Helper{
Ctx: context.TODO(),
Client: client,
}, uptime.Uptime)
if err != nil {
t.Fatalf("create test node: %v", err)
}

return *testNode
}
1 change: 1 addition & 0 deletions pkg/monitor/sqsevent/rebalance-recommendation-event.go
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,7 @@ func (m SQSMonitor) rebalanceRecommendationToInterruptionEvent(event *EventBridg
err := n.TaintRebalanceRecommendation(nodeName, interruptionEvent.EventID)
if err != nil {
log.Err(err).Msgf("Unable to taint node with taint %s:%s", node.RebalanceRecommendationTaint, interruptionEvent.EventID)
return err
}
return nil
}
Expand Down
1 change: 1 addition & 0 deletions pkg/monitor/sqsevent/scheduled-change-event.go
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,7 @@ func (m SQSMonitor) scheduledEventToInterruptionEvents(event *EventBridgeEvent,

if err := n.TaintScheduledMaintenance(nodeName, interruptionEvent.EventID); err != nil {
log.Err(err).Msgf("Unable to taint node with taint %s:%s", node.ScheduledMaintenanceTaint, interruptionEvent.EventID)
return err
}
return nil
}
Expand Down