Skip to content
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
21 changes: 14 additions & 7 deletions cmd/launcher/signal_listener.go
Original file line number Diff line number Diff line change
Expand Up @@ -14,24 +14,31 @@ type signalListener struct {
sigChannel chan os.Signal
cancel context.CancelFunc
slogger *slog.Logger
interrupt chan struct{}
interrupted atomic.Bool
}

func newSignalListener(sigChannel chan os.Signal, cancel context.CancelFunc, slogger *slog.Logger) *signalListener {
signal.Notify(sigChannel, os.Interrupt, syscall.SIGTERM)
return &signalListener{
sigChannel: sigChannel,
cancel: cancel,
slogger: slogger.With("component", "signal_listener"),
interrupt: make(chan struct{}, 1),
}
}

func (s *signalListener) Execute() error {
signal.Notify(s.sigChannel, os.Interrupt, syscall.SIGTERM)
sig := <-s.sigChannel
s.slogger.Log(context.TODO(), slog.LevelInfo,
"beginning shutdown via signal",
"signal_received", sig,
)
select {
case sig := <-s.sigChannel:
s.slogger.Log(context.TODO(), slog.LevelInfo,
"beginning shutdown via signal",
"signal_received", sig,
)
case <-s.interrupt:
// Rungroup shutdown
}

return nil
}

Expand All @@ -44,6 +51,6 @@ func (s *signalListener) Interrupt(_ error) {
// tell sender in `os/signal` package to stop sending on `s.sigChannel`
// to avoid panics for sending on a closed channel
signal.Stop(s.sigChannel)
close(s.interrupt)
s.cancel()
close(s.sigChannel)
}
40 changes: 40 additions & 0 deletions cmd/launcher/signal_listener_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"errors"
"log/slog"
"os"
"syscall"
"testing"
"time"

Expand All @@ -19,6 +20,45 @@ func TestMain(m *testing.M) {
goleak.VerifyTestMain(m, goleak.IgnoreAnyFunction("github.com/Microsoft/go-winio.ioCompletionProcessor"))
}

func TestInterruptBeforeExecute(t *testing.T) {
t.Parallel()

sigChannel := make(chan os.Signal, 1)
_, cancel := context.WithCancel(t.Context())
var logBytes threadsafebuffer.ThreadSafeBuffer
slogger := slog.New(slog.NewTextHandler(&logBytes, &slog.HandlerOptions{
Level: slog.LevelDebug,
}))
sigListener := newSignalListener(sigChannel, cancel, slogger)

// Call Interrupt and Execute out of order -- Execute should immediately return
sigListener.Interrupt(errors.New("test error"))
err := sigListener.Execute()
require.NoError(t, err)

// Send a sigterm, confirm no panic
sigChannel <- syscall.SIGTERM
}

func TestSigtermBeforeExecute(t *testing.T) {
t.Parallel()

sigChannel := make(chan os.Signal, 1)
_, cancel := context.WithCancel(t.Context())
var logBytes threadsafebuffer.ThreadSafeBuffer
slogger := slog.New(slog.NewTextHandler(&logBytes, &slog.HandlerOptions{
Level: slog.LevelDebug,
}))
sigListener := newSignalListener(sigChannel, cancel, slogger)

// Send a sigterm, confirm no panic
sigChannel <- syscall.SIGTERM

// Start up the listener, then shut it down
go sigListener.Execute()
sigListener.Interrupt(errors.New("test error"))
}

func TestInterrupt_Multiple(t *testing.T) {
t.Parallel()

Expand Down