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
13 changes: 7 additions & 6 deletions cmd/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,11 +35,12 @@ import (
)

var (
hwLoglevel = flag.Int("hw_loglevel", 0, "huawei log level, -1-debug, 0-info, 1-warning, 2-error 3-critical default value: 0")
configFile = flag.String("config_file", "", "config file path")
nodeConfigFile = flag.String("node_config_file", "", "node specific config file path")
nodeName = flag.String("node_name", os.Getenv("NODE_NAME"), "node name")
checkIdleVNPUInterval = flag.Int("check_idle_vnpu_interval", 60, "the interval (in seconds) to check idle vNPU and release them")
hwLoglevel = flag.Int("hw_loglevel", 0, "huawei log level, -1-debug, 0-info, 1-warning, 2-error 3-critical default value: 0")
configFile = flag.String("config_file", "", "config file path")
nodeConfigFile = flag.String("node_config_file", "", "node specific config file path")
nodeName = flag.String("node_name", os.Getenv("NODE_NAME"), "node name")
checkIdleVNPUInterval = flag.Int("check_idle_vnpu_interval", 60, "the interval (in seconds) to check idle vNPU and release them")
enablePeriodicIdleVNPUCleanup = flag.Bool("enable_periodic_idle_vnpu_cleanup", false, "whether to enable the periodic idle vNPU cleanup goroutine; when disabled, the one-shot cleanup on restart still runs (default false: periodic cleanup disabled)")
)

func checkFlags() {
Expand Down Expand Up @@ -144,7 +145,7 @@ func main() {
klog.Errorf("load node config failed: %v", err)
}
}
server, err := server.NewPluginServer(mgr, *nodeName, *checkIdleVNPUInterval)
server, err := server.NewPluginServer(mgr, *nodeName, *checkIdleVNPUInterval, *enablePeriodicIdleVNPUCleanup)
if err != nil {
klog.Fatalf("init PluginServer failed, error is %v", err)
}
Expand Down
5 changes: 2 additions & 3 deletions internal/server/register.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,10 +35,9 @@ import (
"github.com/Project-HAMi/HAMi/pkg/util"
)

// watchAndRegister must be launched with ps.wg.Add(1) already called by the
// caller (see Start()); doing the Add here would race with Stop()'s wg.Wait().
// watchAndRegister is launched via ps.wg.Go in Start(), which owns the
// WaitGroup Add(1)/Done() pairing; this function must not call wg.Done itself.
func (ps *PluginServer) watchAndRegister() {
defer ps.wg.Done()
timer := time.After(1 * time.Second)
for {
select {
Expand Down
79 changes: 38 additions & 41 deletions internal/server/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,19 +58,20 @@ var (
type PluginServer struct {
v1beta1.UnimplementedDevicePluginServer

commonWord string
nodeName string
registerAnno string
handshakeAnno string
allocAnno string
toAllocDeviceAnno string
grpcServer *grpc.Server
mgr manager.Manager
socket string
stopCh chan interface{}
healthCh chan int32
checkIdleVNPUInterval int
wg sync.WaitGroup
commonWord string
nodeName string
registerAnno string
handshakeAnno string
allocAnno string
toAllocDeviceAnno string
grpcServer *grpc.Server
mgr manager.Manager
socket string
stopCh chan any
healthCh chan int32
checkIdleVNPUInterval int
enablePeriodicIdleVNPUCleanup bool
wg sync.WaitGroup

// test hooks — injected by tests to avoid real socket/kubelet dependencies
dialFunc func(unixSocketPath string, timeout time.Duration) (*grpc.ClientConn, error)
Expand All @@ -85,20 +86,21 @@ type RuntimeInfo struct {
Core *int32 `json:"core,omitempty"`
}

func NewPluginServer(mgr manager.Manager, nodeName string, checkIdleVNPUInterval int) (*PluginServer, error) {
func NewPluginServer(mgr manager.Manager, nodeName string, checkIdleVNPUInterval int, enablePeriodicIdleVNPUCleanup bool) (*PluginServer, error) {
commonWord := mgr.CommonWord()
server := &PluginServer{
commonWord: commonWord,
nodeName: nodeName,
registerAnno: fmt.Sprintf("hami.io/node-register-%s", commonWord),
handshakeAnno: fmt.Sprintf("hami.io/node-handshake-%s", commonWord),
allocAnno: fmt.Sprintf("huawei.com/%s", commonWord),
toAllocDeviceAnno: fmt.Sprintf("hami.io/%s-devices-to-allocate", commonWord),
mgr: mgr,
socket: path.Join(v1beta1.DevicePluginPath, fmt.Sprintf("%s.sock", commonWord)),
stopCh: make(chan interface{}),
healthCh: make(chan int32),
checkIdleVNPUInterval: checkIdleVNPUInterval,
commonWord: commonWord,
nodeName: nodeName,
registerAnno: fmt.Sprintf("hami.io/node-register-%s", commonWord),
handshakeAnno: fmt.Sprintf("hami.io/node-handshake-%s", commonWord),
allocAnno: fmt.Sprintf("huawei.com/%s", commonWord),
toAllocDeviceAnno: fmt.Sprintf("hami.io/%s-devices-to-allocate", commonWord),
mgr: mgr,
socket: path.Join(v1beta1.DevicePluginPath, fmt.Sprintf("%s.sock", commonWord)),
stopCh: make(chan any),
healthCh: make(chan int32),
checkIdleVNPUInterval: checkIdleVNPUInterval,
enablePeriodicIdleVNPUCleanup: enablePeriodicIdleVNPUCleanup,
}
// enable calling hami methods
device.InRequestDevices[commonWord] = server.toAllocDeviceAnno
Expand All @@ -121,7 +123,7 @@ func (ps *PluginServer) Start() error {
return err
}

ps.stopCh = make(chan interface{})
ps.stopCh = make(chan any)
ps.grpcServer = grpc.NewServer()

err := ps.mgr.UpdateDevice()
Expand All @@ -139,19 +141,16 @@ func (ps *PluginServer) Start() error {
if err != nil {
return err
}
// Add to the WaitGroup synchronously before launching the goroutines.
// sync.WaitGroup requires a positive Add (from a zero counter) to
// happen-before Wait; doing Add inside the goroutine races with Stop()'s
// Wait and panics with "WaitGroup is reused before previous Wait has returned".
ps.wg.Add(1)
go ps.startPeriodicCheckIdleVNPUs()
ps.wg.Add(1)
go ps.watchAndRegister()
// sync.WaitGroup.Go handles the Add(1)/Done() pairing internally, so the
// goroutine bodies no longer need their own defer ps.wg.Done().
if ps.enablePeriodicIdleVNPUCleanup {
ps.wg.Go(ps.startPeriodicCheckIdleVNPUs)
}
ps.wg.Go(ps.watchAndRegister)
return nil
}

func (ps *PluginServer) startPeriodicCheckIdleVNPUs() {
defer ps.wg.Done()
ticker := time.NewTicker(time.Duration(ps.checkIdleVNPUInterval) * time.Second)
defer ticker.Stop()
for {
Expand Down Expand Up @@ -185,7 +184,7 @@ func (ps *PluginServer) Stop() error {
return nil
}

func (ps *PluginServer) StopCh() <-chan interface{} {
func (ps *PluginServer) StopCh() <-chan any {
return ps.stopCh
}

Expand All @@ -201,9 +200,7 @@ func (ps *PluginServer) serve() error {
}
v1beta1.RegisterDevicePluginServer(ps.grpcServer, ps)
resourceName := ps.mgr.ResourceName()
ps.wg.Add(1)
go func() {
defer ps.wg.Done()
ps.wg.Go(func() {
lastCrashTime := time.Now()
restartCount := 0
for {
Expand Down Expand Up @@ -236,7 +233,7 @@ func (ps *PluginServer) serve() error {
restartCount++
}
}
}()
})

// Wait for server to start by launching a blocking connexion
conn, err := ps.dial(ps.socket, 5*time.Second)
Expand All @@ -257,7 +254,7 @@ func (ps *PluginServer) apiDevices() []*v1beta1.Device {
if dev.Health {
health = v1beta1.Healthy
}
for i := 0; i < vCount; i++ {
for i := range vCount {
device := v1beta1.Device{
ID: fmt.Sprintf("%s-%d", dev.UUID, i),
Health: health,
Expand Down
58 changes: 30 additions & 28 deletions internal/server/server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ import (
"fmt"
"os"
"path"
"slices"
"strings"
"testing"

Expand Down Expand Up @@ -128,8 +129,8 @@ func setupAllocateEnv(nodeName, podName, podNamespace string, numContainers int,
// composeCleanup combines multiple CleanupFuncs into one that runs in reverse order.
func composeCleanup(fns ...CleanupFunc) CleanupFunc {
return func() {
for i := len(fns) - 1; i >= 0; i-- {
fns[i]()
for _, fn := range slices.Backward(fns) {
fn()
}
}
}
Expand Down Expand Up @@ -534,7 +535,7 @@ func TestNewPluginServer(t *testing.T) {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
mgr := &FakeManager{CommonWordFunc: func() string { return tc.args.commonWord }}
ps, err := NewPluginServer(mgr, tc.args.nodeName, 60)
ps, err := NewPluginServer(mgr, tc.args.nodeName, 60, false)
Comment thread
archlitchi marked this conversation as resolved.
if (err != nil) != tc.wantErr {
t.Fatalf("NewPluginServer() error = %v, wantErr %v", err, tc.wantErr)
}
Expand Down Expand Up @@ -569,7 +570,7 @@ func TestNewPluginServer_RegistersInRequestDevices(t *testing.T) {
mgr := &FakeManager{
CommonWordFunc: func() string { return commonWord },
}
ps, err := NewPluginServer(mgr, "test-node", 60)
ps, err := NewPluginServer(mgr, "test-node", 60, false)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
Expand Down Expand Up @@ -790,32 +791,32 @@ func newPanicOnFatalLogger() *panicOnFatalLogger {

var _ grpclog.LoggerV2 = (*panicOnFatalLogger)(nil)

func (l *panicOnFatalLogger) Info(args ...interface{}) { l.inner.Info(args...) }
func (l *panicOnFatalLogger) Infoln(args ...interface{}) { l.inner.Infoln(args...) }
func (l *panicOnFatalLogger) Infof(format string, args ...interface{}) {
func (l *panicOnFatalLogger) Info(args ...any) { l.inner.Info(args...) }
func (l *panicOnFatalLogger) Infoln(args ...any) { l.inner.Infoln(args...) }
func (l *panicOnFatalLogger) Infof(format string, args ...any) {
l.inner.Infof(format, args...)
}
func (l *panicOnFatalLogger) Warning(args ...interface{}) { l.inner.Warning(args...) }
func (l *panicOnFatalLogger) Warningln(args ...interface{}) { l.inner.Warningln(args...) }
func (l *panicOnFatalLogger) Warningf(format string, args ...interface{}) {
func (l *panicOnFatalLogger) Warning(args ...any) { l.inner.Warning(args...) }
func (l *panicOnFatalLogger) Warningln(args ...any) { l.inner.Warningln(args...) }
func (l *panicOnFatalLogger) Warningf(format string, args ...any) {
l.inner.Warningf(format, args...)
}
func (l *panicOnFatalLogger) Error(args ...interface{}) { l.inner.Error(args...) }
func (l *panicOnFatalLogger) Errorln(args ...interface{}) { l.inner.Errorln(args...) }
func (l *panicOnFatalLogger) Errorf(format string, args ...interface{}) {
func (l *panicOnFatalLogger) Error(args ...any) { l.inner.Error(args...) }
func (l *panicOnFatalLogger) Errorln(args ...any) { l.inner.Errorln(args...) }
func (l *panicOnFatalLogger) Errorf(format string, args ...any) {
l.inner.Errorf(format, args...)
}
func (l *panicOnFatalLogger) V(level int) bool { return l.inner.V(level) }

func (l *panicOnFatalLogger) Fatalf(format string, args ...interface{}) {
func (l *panicOnFatalLogger) Fatalf(format string, args ...any) {
panic(fmt.Sprintf("grpc FATAL: "+format, args...))
}

func (l *panicOnFatalLogger) Fatalln(args ...interface{}) {
func (l *panicOnFatalLogger) Fatalln(args ...any) {
panic(fmt.Sprintf("grpc FATAL: %v", fmt.Sprintln(args...)))
}

func (l *panicOnFatalLogger) Fatal(args ...interface{}) {
func (l *panicOnFatalLogger) Fatal(args ...any) {
panic(fmt.Sprintf("grpc FATAL: %v", fmt.Sprint(args...)))
}

Expand All @@ -825,17 +826,18 @@ func setupRestartablePluginServer(t *testing.T) *PluginServer {
t.Helper()

ps := &PluginServer{
commonWord: "test-ascend",
registerAnno: "hami.io/node-register-test-ascend",
handshakeAnno: "hami.io/node-handshake-test-ascend",
allocAnno: "huawei.com/test-ascend",
toAllocDeviceAnno: "hami.io/test-ascend-devices-to-allocate",
mgr: &FakeManager{ResourceNameFunc: func() string { return "test-ascend" }},
socket: path.Join(t.TempDir(), "test-ascend.sock"),
stopCh: make(chan interface{}),
healthCh: make(chan int32),
checkIdleVNPUInterval: 3600,
dialFunc: nil,
commonWord: "test-ascend",
registerAnno: "hami.io/node-register-test-ascend",
handshakeAnno: "hami.io/node-handshake-test-ascend",
allocAnno: "huawei.com/test-ascend",
toAllocDeviceAnno: "hami.io/test-ascend-devices-to-allocate",
mgr: &FakeManager{ResourceNameFunc: func() string { return "test-ascend" }},
socket: path.Join(t.TempDir(), "test-ascend.sock"),
stopCh: make(chan any),
healthCh: make(chan int32),
checkIdleVNPUInterval: 3600,
enablePeriodicIdleVNPUCleanup: true,
dialFunc: nil,
registerKubeletFunc: func() error {
return nil
},
Expand Down Expand Up @@ -883,7 +885,7 @@ func TestGrpcServer_MultipleRestarts(t *testing.T) {

ps := setupRestartablePluginServer(t)

for i := 0; i < 5; i++ {
for i := range 5 {
if err := ps.Start(); err != nil {
t.Fatalf("Start() iteration %d failed: %v", i, err)
}
Expand Down
Loading