From d4d4497532c3ba3e2c28be433cf82f8dc93e8f69 Mon Sep 17 00:00:00 2001 From: Rohan Date: Mon, 31 Aug 2026 03:40:14 -0700 Subject: [PATCH 1/2] fix(virt-handler): re-register device plugins after kubelet restart Kubelet restarts exposed two independent failures. PCI and generic health checks only recognized an exact removal of their own socket, so a replaced device-plugin directory or a recreated kubelet socket could leave a plugin attached to a dead kubelet forever. PCI also reused lifecycle channels after teardown, causing its next ListAndWatch stream to deregister immediately and close an already-closed channel. Install the socket-directory watch before Register and hand that watcher to the health loop. Restart on plugin-socket removal, kubelet-socket creation, or device-plugin directory removal or rename, and use fsnotify bit membership so combined operation masks retain their meaning. Derive the registration endpoint from the plugin socket directory, bound Register to the existing five-second connection window, and cancel it when the cycle stops. Recreate PCI lifecycle channels for every Start. Close each cycle's done channel before stopping its server, and construct PCI and generic servers with WaitForHandlers so all old handlers are joined before the next cycle can replace their channels. This closes the cross-cycle race without changing the existing supervision and backoff model. Co-authored-by: Aseef Co-authored-by: RITANKAR SAHA --- pkg/virt-handler/device-manager/common.go | 19 ++++++ .../device-manager/device_plugin_base.go | 15 ++--- .../device-manager/generic_device.go | 62 +++++++++++-------- pkg/virt-handler/device-manager/pci_device.go | 50 +++++++++------ 4 files changed, 93 insertions(+), 53 deletions(-) diff --git a/pkg/virt-handler/device-manager/common.go b/pkg/virt-handler/device-manager/common.go index e3f86b117ac9..25620ab0492a 100644 --- a/pkg/virt-handler/device-manager/common.go +++ b/pkg/virt-handler/device-manager/common.go @@ -24,6 +24,7 @@ package device_manager import ( "bufio" "bytes" + "context" "fmt" "net" "os" @@ -211,6 +212,24 @@ func waitForGRPCServer(socketPath string, timeout time.Duration) error { return nil } +func newDevicePluginGRPCServer() *grpc.Server { + return grpc.NewServer(grpc.WaitForHandlers(true)) +} + +func registrationContext(stop <-chan struct{}) (context.Context, context.CancelFunc) { + ctx, cancel := context.WithTimeout(context.Background(), connectionTimeout) + if stop != nil { + go func() { + select { + case <-stop: + cancel() + case <-ctx.Done(): + } + }() + } + return ctx, cancel +} + // dial establishes the gRPC communication with the registered device plugin. func gRPCConnect(socketPath string, timeout time.Duration) (*grpc.ClientConn, error) { c, err := grpc.Dial(socketPath, diff --git a/pkg/virt-handler/device-manager/device_plugin_base.go b/pkg/virt-handler/device-manager/device_plugin_base.go index c1a212f19b0f..867a0127f901 100644 --- a/pkg/virt-handler/device-manager/device_plugin_base.go +++ b/pkg/virt-handler/device-manager/device_plugin_base.go @@ -177,11 +177,9 @@ func (dpi *DevicePluginBase) Allocate(ctx context.Context, r *pluginapi.Allocate } func (dpi *DevicePluginBase) stopDevicePlugin() error { - defer func() { - if !IsChanClosed(dpi.done) { - close(dpi.done) - } - }() + if !IsChanClosed(dpi.done) { + close(dpi.done) + } ticker := time.NewTicker(1 * time.Second) defer ticker.Stop() select { @@ -214,7 +212,8 @@ func (dpi *DevicePluginBase) setInitialized(initialized bool) { } func (dpi *DevicePluginBase) register() error { - conn, err := gRPCConnect(pluginapi.KubeletSocket, connectionTimeout) + kubeletSocket := filepath.Join(filepath.Dir(dpi.socketPath), filepath.Base(pluginapi.KubeletSocket)) + conn, err := gRPCConnect(kubeletSocket, connectionTimeout) if err != nil { return err } @@ -227,7 +226,9 @@ func (dpi *DevicePluginBase) register() error { ResourceName: dpi.resourceName, } - _, err = client.Register(context.Background(), reqt) + ctx, cancel := registrationContext(dpi.stop) + defer cancel() + _, err = client.Register(ctx, reqt) if err != nil { return err } diff --git a/pkg/virt-handler/device-manager/generic_device.go b/pkg/virt-handler/device-manager/generic_device.go index 7b2390c033ae..13c258a32111 100644 --- a/pkg/virt-handler/device-manager/generic_device.go +++ b/pkg/virt-handler/device-manager/generic_device.go @@ -130,7 +130,7 @@ func (dpi *GenericDevicePlugin) Start(stop <-chan struct{}) (err error) { return fmt.Errorf("error creating GRPC server socket: %v", err) } - dpi.server = grpc.NewServer([]grpc.ServerOption{}...) + dpi.server = newDevicePluginGRPCServer() defer dpi.stopDevicePlugin() pluginapi.RegisterDevicePluginServer(dpi.server, dpi) @@ -146,13 +146,27 @@ func (dpi *GenericDevicePlugin) Start(stop <-chan struct{}) (err error) { return fmt.Errorf("error starting the GRPC server: %v", err) } + socketDir := filepath.Dir(dpi.socketPath) + watcher, err := fsnotify.NewWatcher() + if err != nil { + return fmt.Errorf("failed to creating a fsnotify watcher: %v", err) + } + if err = watcher.Add(socketDir); err != nil { + _ = watcher.Close() + return fmt.Errorf("failed to add the device-plugin kubelet path to the watcher: %v", err) + } + if _, err = os.Stat(dpi.socketPath); err != nil { + _ = watcher.Close() + return fmt.Errorf("failed to stat the device-plugin socket: %v", err) + } err = dpi.register() if err != nil { + _ = watcher.Close() return fmt.Errorf("error registering with device plugin manager: %v", err) } go func() { - errChan <- dpi.healthCheck() + errChan <- dpi.healthCheck(watcher) }() dpi.setInitialized(true) @@ -164,11 +178,9 @@ func (dpi *GenericDevicePlugin) Start(stop <-chan struct{}) (err error) { // Stop stops the gRPC server func (dpi *GenericDevicePlugin) stopDevicePlugin() error { - defer func() { - if !IsChanClosed(dpi.done) { - close(dpi.done) - } - }() + if !IsChanClosed(dpi.done) { + close(dpi.done) + } // Give the device plugin one second to properly deregister ticker := time.NewTicker(1 * time.Second) @@ -184,7 +196,8 @@ func (dpi *GenericDevicePlugin) stopDevicePlugin() error { // Register registers the device plugin for the given resourceName with Kubelet. func (dpi *GenericDevicePlugin) register() error { - conn, err := gRPCConnect(pluginapi.KubeletSocket, connectionTimeout) + kubeletSocket := filepath.Join(filepath.Dir(dpi.socketPath), filepath.Base(pluginapi.KubeletSocket)) + conn, err := gRPCConnect(kubeletSocket, connectionTimeout) if err != nil { return err } @@ -197,7 +210,9 @@ func (dpi *GenericDevicePlugin) register() error { ResourceName: dpi.resourceName, } - _, err = client.Register(context.Background(), reqt) + ctx, cancel := registrationContext(dpi.stop) + defer cancel() + _, err = client.Register(ctx, reqt) if err != nil { return err } @@ -273,12 +288,8 @@ func (dpi *GenericDevicePlugin) PreStartContainer(_ context.Context, _ *pluginap return res, nil } -func (dpi *GenericDevicePlugin) healthCheck() error { +func (dpi *GenericDevicePlugin) healthCheck(watcher *fsnotify.Watcher) error { logger := log.DefaultLogger() - watcher, err := fsnotify.NewWatcher() - if err != nil { - return fmt.Errorf("failed to creating a fsnotify watcher: %v", err) - } defer watcher.Close() // This way we don't have to mount /dev from the node @@ -286,7 +297,7 @@ func (dpi *GenericDevicePlugin) healthCheck() error { // Start watching the files before we check for their existence to avoid races dirName := filepath.Dir(devicePath) - err = watcher.Add(dirName) + err := watcher.Add(dirName) if err != nil { return fmt.Errorf("failed to add the device root path to the watcher: %v", err) @@ -302,16 +313,8 @@ func (dpi *GenericDevicePlugin) healthCheck() error { } logger.Infof("device '%s' is present.", dpi.devicePath) - dirName = filepath.Dir(dpi.socketPath) - err = watcher.Add(dirName) - - if err != nil { - return fmt.Errorf("failed to add the device-plugin kubelet path to the watcher: %v", err) - } - _, err = os.Stat(dpi.socketPath) - if err != nil { - return fmt.Errorf("failed to stat the device-plugin socket: %v", err) - } + socketDir := filepath.Dir(dpi.socketPath) + kubeletSocketPath := filepath.Join(socketDir, filepath.Base(pluginapi.KubeletSocket)) for { select { @@ -330,9 +333,16 @@ func (dpi *GenericDevicePlugin) healthCheck() error { logger.Infof("monitored device %s disappeared", dpi.deviceName) dpi.health <- deviceHealth{Health: pluginapi.Unhealthy} } - } else if event.Name == dpi.socketPath && event.Op == fsnotify.Remove { + } else if event.Name == dpi.socketPath && event.Op.Has(fsnotify.Remove) { logger.Infof("device socket file for device %s was removed, kubelet probably restarted.", dpi.deviceName) return nil + } else if event.Name == kubeletSocketPath && event.Op.Has(fsnotify.Create) { + logger.Infof("kubelet socket %s was recreated, kubelet probably restarted. Restarting %s device plugin to re-register.", kubeletSocketPath, dpi.deviceName) + return nil + } else if event.Name == socketDir && + (event.Op.Has(fsnotify.Remove) || event.Op.Has(fsnotify.Rename)) { + logger.Infof("device plugin socket directory %s was removed or replaced, kubelet probably restarted. Restarting %s device plugin to re-register.", socketDir, dpi.deviceName) + return nil } } } diff --git a/pkg/virt-handler/device-manager/pci_device.go b/pkg/virt-handler/device-manager/pci_device.go index bcb6be4511d6..860a0482c10e 100644 --- a/pkg/virt-handler/device-manager/pci_device.go +++ b/pkg/virt-handler/device-manager/pci_device.go @@ -32,7 +32,6 @@ import ( "sync" "github.com/fsnotify/fsnotify" - "google.golang.org/grpc" v1 "kubevirt.io/api/core/v1" "kubevirt.io/client-go/log" @@ -63,6 +62,8 @@ type PCIDevicePlugin struct { func (dpi *PCIDevicePlugin) Start(stop <-chan struct{}) (err error) { logger := log.DefaultLogger() dpi.stop = stop + dpi.done = make(chan struct{}) + dpi.deregistered = make(chan struct{}) err = dpi.cleanup() if err != nil { @@ -74,7 +75,7 @@ func (dpi *PCIDevicePlugin) Start(stop <-chan struct{}) (err error) { return fmt.Errorf("error creating GRPC server socket: %v", err) } - dpi.server = grpc.NewServer([]grpc.ServerOption{}...) + dpi.server = newDevicePluginGRPCServer() defer dpi.stopDevicePlugin() pluginapi.RegisterDevicePluginServer(dpi.server, dpi) @@ -90,13 +91,27 @@ func (dpi *PCIDevicePlugin) Start(stop <-chan struct{}) (err error) { return fmt.Errorf("error starting the GRPC server: %v", err) } + socketDir := filepath.Dir(dpi.socketPath) + watcher, err := fsnotify.NewWatcher() + if err != nil { + return fmt.Errorf("failed to creating a fsnotify watcher: %v", err) + } + if err = watcher.Add(socketDir); err != nil { + _ = watcher.Close() + return fmt.Errorf("failed to add the device-plugin kubelet path to the watcher: %v", err) + } + if _, err = os.Stat(dpi.socketPath); err != nil { + _ = watcher.Close() + return fmt.Errorf("failed to stat the device-plugin socket: %v", err) + } err = dpi.register() if err != nil { + _ = watcher.Close() return fmt.Errorf("error registering with device plugin manager: %v", err) } go func() { - errChan <- dpi.healthCheck() + errChan <- dpi.healthCheck(watcher) }() dpi.setInitialized(true) @@ -180,13 +195,9 @@ func (dpi *PCIDevicePlugin) Allocate(_ context.Context, r *pluginapi.AllocateReq return resp, nil } -func (dpi *PCIDevicePlugin) healthCheck() error { +func (dpi *PCIDevicePlugin) healthCheck(watcher *fsnotify.Watcher) error { logger := log.DefaultLogger() monitoredDevices := make(map[string]string) - watcher, err := fsnotify.NewWatcher() - if err != nil { - return fmt.Errorf("failed to creating a fsnotify watcher: %v", err) - } defer watcher.Close() // This way we don't have to mount /dev from the node @@ -194,7 +205,7 @@ func (dpi *PCIDevicePlugin) healthCheck() error { // Start watching the files before we check for their existence to avoid races dirName := filepath.Dir(devicePath) - err = watcher.Add(dirName) + err := watcher.Add(dirName) if err != nil { return fmt.Errorf("failed to add the device root path to the watcher: %v", err) } @@ -216,16 +227,8 @@ func (dpi *PCIDevicePlugin) healthCheck() error { monitoredDevices[vfioDevice] = dev.ID } - dirName = filepath.Dir(dpi.socketPath) - err = watcher.Add(dirName) - - if err != nil { - return fmt.Errorf("failed to add the device-plugin kubelet path to the watcher: %v", err) - } - _, err = os.Stat(dpi.socketPath) - if err != nil { - return fmt.Errorf("failed to stat the device-plugin socket: %v", err) - } + socketDir := filepath.Dir(dpi.socketPath) + kubeletSocketPath := filepath.Join(socketDir, filepath.Base(pluginapi.KubeletSocket)) for { select { @@ -250,9 +253,16 @@ func (dpi *PCIDevicePlugin) healthCheck() error { Health: pluginapi.Unhealthy, } } - } else if event.Name == dpi.socketPath && event.Op == fsnotify.Remove { + } else if event.Name == dpi.socketPath && event.Op.Has(fsnotify.Remove) { logger.Infof("device socket file for device %s was removed, kubelet probably restarted.", dpi.resourceName) return nil + } else if event.Name == kubeletSocketPath && event.Op.Has(fsnotify.Create) { + logger.Infof("kubelet socket %s was recreated, kubelet probably restarted. Restarting %s device plugin to re-register.", kubeletSocketPath, dpi.resourceName) + return nil + } else if event.Name == socketDir && + (event.Op.Has(fsnotify.Remove) || event.Op.Has(fsnotify.Rename)) { + logger.Infof("device plugin socket directory %s was removed or replaced, kubelet probably restarted. Restarting %s device plugin to re-register.", socketDir, dpi.resourceName) + return nil } } } From bb5f66d365b61350cff2194c447cf0ff60264346 Mon Sep 17 00:00:00 2001 From: Rohan Date: Mon, 31 Aug 2026 03:40:19 -0700 Subject: [PATCH 2/2] test(virt-handler): cover device plugin restart lifecycle Drive PCI and generic plugins through the real controlled-device loop against a fake kubelet and verify that both register again and continue advertising devices after a restart. Cover kubelet-socket recreation, device-plugin directory replacement, combined fsnotify masks, and rejection of unrelated events. Exercise the watch-before-register ordering with an in-flight kubelet replacement, prove the production gRPC server waits for handlers while teardown wakes active ListAndWatch streams, and verify hung Register calls are cancelled on shutdown or fail at the deadline. --- pkg/virt-handler/device-manager/BUILD.bazel | 3 + .../device-manager/generic_device_test.go | 4 +- .../device-manager/kubelet_restart_test.go | 718 ++++++++++++++++++ 3 files changed, 723 insertions(+), 2 deletions(-) create mode 100644 pkg/virt-handler/device-manager/kubelet_restart_test.go diff --git a/pkg/virt-handler/device-manager/BUILD.bazel b/pkg/virt-handler/device-manager/BUILD.bazel index 322ce4ece096..bc7321846a7d 100644 --- a/pkg/virt-handler/device-manager/BUILD.bazel +++ b/pkg/virt-handler/device-manager/BUILD.bazel @@ -46,6 +46,7 @@ go_test( "device_controller_test.go", "device_manager_suite_test.go", "generic_device_test.go", + "kubelet_restart_test.go", "mediated_device_test.go", "mediated_devices_types_test.go", "pci_device_test.go", @@ -58,6 +59,7 @@ go_test( "//staging/src/kubevirt.io/api/core/v1:go_default_library", "//staging/src/kubevirt.io/client-go/log:go_default_library", "//staging/src/kubevirt.io/client-go/testutils:go_default_library", + "//vendor/github.com/fsnotify/fsnotify:go_default_library", "//vendor/github.com/onsi/ginkgo/v2:go_default_library", "//vendor/github.com/onsi/gomega:go_default_library", "//vendor/k8s.io/api/core/v1:go_default_library", @@ -68,5 +70,6 @@ go_test( "//vendor/k8s.io/client-go/testing:go_default_library", "//vendor/k8s.io/client-go/tools/cache:go_default_library", "//vendor/k8s.io/client-go/tools/cache/testing:go_default_library", + "@org_golang_google_grpc//:go_default_library", ], ) diff --git a/pkg/virt-handler/device-manager/generic_device_test.go b/pkg/virt-handler/device-manager/generic_device_test.go index 9cc011e7e953..73793f989aab 100644 --- a/pkg/virt-handler/device-manager/generic_device_test.go +++ b/pkg/virt-handler/device-manager/generic_device_test.go @@ -49,7 +49,7 @@ var _ = Describe("Generic Device", func() { errChan := make(chan error, 1) go func(errChan chan error) { - errChan <- dpi.healthCheck() + errChan <- dpi.healthCheck(newTestHealthWatcher(dpi.socketPath)) }(errChan) Consistently(func() string { return dpi.devs[0].Health @@ -63,7 +63,7 @@ var _ = Describe("Generic Device", func() { os.OpenFile(dpi.socketPath, os.O_RDONLY|os.O_CREATE, 0666) - go dpi.healthCheck() + go dpi.healthCheck(newTestHealthWatcher(dpi.socketPath)) Expect(dpi.devs[0].Health).To(Equal(pluginapi.Healthy)) time.Sleep(1 * time.Second) diff --git a/pkg/virt-handler/device-manager/kubelet_restart_test.go b/pkg/virt-handler/device-manager/kubelet_restart_test.go new file mode 100644 index 000000000000..1340644dbec8 --- /dev/null +++ b/pkg/virt-handler/device-manager/kubelet_restart_test.go @@ -0,0 +1,718 @@ +/* + * This file is part of the KubeVirt project + * + * 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. + * + * Copyright the KubeVirt Authors. + * + */ + +package device_manager + +import ( + "context" + "net" + "os" + "path/filepath" + "sync" + "time" + + "github.com/fsnotify/fsnotify" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc" + + pluginapi "kubevirt.io/kubevirt/pkg/virt-handler/device-manager/deviceplugin/v1beta1" +) + +type fakeKubeletRegistrationServer struct { + pluginDir string + + lock sync.Mutex + registrations []*pluginapi.RegisterRequest + deviceLists map[int][][]*pluginapi.Device + pluginConnections []*grpc.ClientConn + receivers sync.WaitGroup + beforeRegisterReturn func() +} + +func newFakeKubeletRegistrationServer(pluginDir string) *fakeKubeletRegistrationServer { + return &fakeKubeletRegistrationServer{ + pluginDir: pluginDir, + deviceLists: make(map[int][][]*pluginapi.Device), + } +} + +func (f *fakeKubeletRegistrationServer) Register(_ context.Context, request *pluginapi.RegisterRequest) (*pluginapi.Empty, error) { + conn, err := gRPCConnect(filepath.Join(f.pluginDir, request.Endpoint), connectionTimeout) + if err != nil { + return nil, err + } + + stream, err := pluginapi.NewDevicePluginClient(conn).ListAndWatch(context.Background(), &pluginapi.Empty{}) + if err != nil { + conn.Close() + return nil, err + } + + requestCopy := *request + f.lock.Lock() + f.registrations = append(f.registrations, &requestCopy) + registration := len(f.registrations) + f.pluginConnections = append(f.pluginConnections, conn) + f.receivers.Add(1) + f.lock.Unlock() + + go func() { + defer f.receivers.Done() + for { + response, err := stream.Recv() + if err != nil { + return + } + devices := make([]*pluginapi.Device, len(response.Devices)) + copy(devices, response.Devices) + f.lock.Lock() + f.deviceLists[registration] = append(f.deviceLists[registration], devices) + f.lock.Unlock() + } + }() + + if f.beforeRegisterReturn != nil { + f.beforeRegisterReturn() + } + return &pluginapi.Empty{}, nil +} + +func (f *fakeKubeletRegistrationServer) registrationCount() int { + f.lock.Lock() + defer f.lock.Unlock() + return len(f.registrations) +} + +func (f *fakeKubeletRegistrationServer) lastListLength(registration int) int { + f.lock.Lock() + defer f.lock.Unlock() + lists := f.deviceLists[registration] + if len(lists) == 0 { + return -1 + } + return len(lists[len(lists)-1]) +} + +func (f *fakeKubeletRegistrationServer) closePluginStreams() { + f.lock.Lock() + connections := append([]*grpc.ClientConn(nil), f.pluginConnections...) + f.pluginConnections = nil + f.lock.Unlock() + + for _, conn := range connections { + _ = conn.Close() + } + f.receivers.Wait() +} + +type fakeKubelet struct { + socketPath string + registration *fakeKubeletRegistrationServer + server *grpc.Server + serveDone chan struct{} +} + +func newFakeKubelet(socketPath, pluginDir string) *fakeKubelet { + return &fakeKubelet{ + socketPath: socketPath, + registration: newFakeKubeletRegistrationServer(pluginDir), + } +} + +func (f *fakeKubelet) start() error { + if err := os.Remove(f.socketPath); err != nil && !os.IsNotExist(err) { + return err + } + listener, err := net.Listen("unix", f.socketPath) + if err != nil { + return err + } + + f.server = grpc.NewServer() + f.serveDone = make(chan struct{}) + pluginapi.RegisterRegistrationServer(f.server, f.registration) + go func() { + _ = f.server.Serve(listener) + close(f.serveDone) + }() + + return waitForGRPCServer(f.socketPath, connectionTimeout) +} + +func (f *fakeKubelet) stop() { + f.registration.closePluginStreams() + if f.server != nil { + f.server.Stop() + <-f.serveDone + f.server = nil + } + _ = os.Remove(f.socketPath) +} + +type blockingFakeKubelet struct { + socketPath string + server *grpc.Server + serveDone chan struct{} + entered chan *pluginapi.RegisterRequest + release chan struct{} + releaseOne sync.Once +} + +func newBlockingFakeKubelet(socketPath string) *blockingFakeKubelet { + return &blockingFakeKubelet{ + socketPath: socketPath, + entered: make(chan *pluginapi.RegisterRequest, 2), + release: make(chan struct{}), + } +} + +func (f *blockingFakeKubelet) Register(ctx context.Context, request *pluginapi.RegisterRequest) (*pluginapi.Empty, error) { + requestCopy := *request + select { + case f.entered <- &requestCopy: + case <-ctx.Done(): + return nil, ctx.Err() + } + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-f.release: + return &pluginapi.Empty{}, nil + } +} + +func (f *blockingFakeKubelet) start() error { + if err := os.Remove(f.socketPath); err != nil && !os.IsNotExist(err) { + return err + } + listener, err := net.Listen("unix", f.socketPath) + if err != nil { + return err + } + + f.server = grpc.NewServer() + f.serveDone = make(chan struct{}) + pluginapi.RegisterRegistrationServer(f.server, f) + go func() { + _ = f.server.Serve(listener) + close(f.serveDone) + }() + + return waitForGRPCServer(f.socketPath, connectionTimeout) +} + +func (f *blockingFakeKubelet) releaseRegistrations() { + f.releaseOne.Do(func() { close(f.release) }) +} + +func (f *blockingFakeKubelet) stop() { + f.releaseRegistrations() + if f.server != nil { + f.server.Stop() + <-f.serveDone + f.server = nil + } + _ = os.Remove(f.socketPath) +} + +type blockingDevicePluginServer struct { + pluginapi.UnimplementedDevicePluginServer + entered chan struct{} + release <-chan struct{} +} + +func (s *blockingDevicePluginServer) ListAndWatch(_ *pluginapi.Empty, _ pluginapi.DevicePlugin_ListAndWatchServer) error { + close(s.entered) + <-s.release + return nil +} + +type restartTestPlugin struct { + Device + grpcPlugin pluginapi.DevicePluginServer + socketPath string + healthCheck func(*fsnotify.Watcher) error + setServer func(*grpc.Server) + stopPlugin func() error + resetChannels func() + wakeHandlers func() +} + +func newTestHealthWatcher(socketPath string) *fsnotify.Watcher { + watcher, err := fsnotify.NewWatcher() + ExpectWithOffset(1, err).ToNot(HaveOccurred()) + ExpectWithOffset(1, watcher.Add(filepath.Dir(socketPath))).To(Succeed()) + return watcher +} + +func verifyHealthCheckReturnsForEvent(plugin *restartTestPlugin, event fsnotify.Event) { + watcher := newTestHealthWatcher(plugin.socketPath) + result := make(chan error, 1) + go func() { result <- plugin.healthCheck(watcher) }() + + delivered := make(chan struct{}) + go func() { + watcher.Events <- event + close(delivered) + }() + + Eventually(result, 5*time.Second).Should(Receive(BeNil())) + Eventually(delivered, 5*time.Second).Should(BeClosed()) +} + +func verifyHealthCheckIgnoresEvent(plugin *restartTestPlugin, ignored, restart fsnotify.Event) { + watcher := newTestHealthWatcher(plugin.socketPath) + result := make(chan error, 1) + go func() { result <- plugin.healthCheck(watcher) }() + + delivered := make(chan struct{}) + go func() { + watcher.Events <- ignored + close(delivered) + }() + Eventually(delivered, 5*time.Second).Should(BeClosed()) + Consistently(result, 200*time.Millisecond).ShouldNot(Receive()) + + go func() { watcher.Events <- restart }() + Eventually(result, 5*time.Second).Should(Receive(BeNil())) +} + +func verifyActiveListAndWatchHandlerIsWoken(plugin *restartTestPlugin) { + listener, err := net.Listen("unix", plugin.socketPath) + Expect(err).ToNot(HaveOccurred()) + + server := newDevicePluginGRPCServer() + plugin.setServer(server) + pluginapi.RegisterDevicePluginServer(server, plugin.grpcPlugin) + go func() { _ = server.Serve(listener) }() + Expect(waitForGRPCServer(plugin.socketPath, connectionTimeout)).To(Succeed()) + DeferCleanup(func() { + plugin.wakeHandlers() + server.Stop() + }) + + conn, err := gRPCConnect(plugin.socketPath, connectionTimeout) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(conn.Close) + streamContext, cancelStream := context.WithTimeout(context.Background(), 10*time.Second) + DeferCleanup(cancelStream) + stream, err := pluginapi.NewDevicePluginClient(conn).ListAndWatch(streamContext, &pluginapi.Empty{}) + Expect(err).ToNot(HaveOccurred()) + response, err := stream.Recv() + Expect(err).ToNot(HaveOccurred()) + Expect(response.Devices).ToNot(BeEmpty()) + + stopResult := make(chan error, 1) + go func() { stopResult <- plugin.stopPlugin() }() + Eventually(stopResult, 5*time.Second).Should(Receive(Succeed())) +} + +var _ = Describe("Device plugin re-registration after kubelet restart", func() { + var workDir string + var socketDir string + var kubeletSocketPath string + var stop chan struct{} + + touch := func(path string) { + file, err := os.Create(path) + ExpectWithOffset(1, err).ToNot(HaveOccurred()) + ExpectWithOffset(1, file.Close()).To(Succeed()) + } + + BeforeEach(func() { + var err error + workDir, err = os.MkdirTemp("", "kubevirt-kubelet-restart") + Expect(err).ToNot(HaveOccurred()) + + Expect(os.MkdirAll(filepath.Join(workDir, "dev", "vfio"), 0755)).To(Succeed()) + touch(filepath.Join(workDir, "dev", "vfio", "42")) + + socketDir = filepath.Join(workDir, "device-plugins") + Expect(os.MkdirAll(socketDir, 0755)).To(Succeed()) + kubeletSocketPath = filepath.Join(socketDir, filepath.Base(pluginapi.KubeletSocket)) + touch(kubeletSocketPath) + stop = make(chan struct{}) + }) + + AfterEach(func() { + close(stop) + Expect(os.RemoveAll(workDir)).To(Succeed()) + }) + + newTestPCIPlugin := func() *restartTestPlugin { + dpi := NewPCIDevicePlugin([]*PCIDevice{ + {pciID: "dead:beef", pciAddress: "0000:00:00.0", iommuGroup: "42", numaNode: -1}, + }, "vendor.example.org/fake-nvme") + dpi.socketPath = filepath.Join(socketDir, "kubevirt-fake-nvme.sock") + dpi.deviceRoot = workDir + dpi.devicePath = filepath.Join("dev", "vfio") + dpi.stop = stop + return &restartTestPlugin{ + Device: dpi, + grpcPlugin: dpi, + socketPath: dpi.socketPath, + healthCheck: dpi.healthCheck, + setServer: func(server *grpc.Server) { dpi.server = server }, + stopPlugin: dpi.stopDevicePlugin, + resetChannels: func() { + dpi.done = make(chan struct{}) + dpi.deregistered = make(chan struct{}) + }, + wakeHandlers: func() { + if !IsChanClosed(dpi.done) { + close(dpi.done) + } + }, + } + } + + newTestGenericPlugin := func() *restartTestPlugin { + devicePath := filepath.Join(workDir, "dev", "kvm") + touch(devicePath) + dpi := NewGenericDevicePlugin("fake-kvm", devicePath, 1, "rw", false) + dpi.socketPath = filepath.Join(socketDir, "kubevirt-fake-kvm.sock") + dpi.deviceRoot = "/" + dpi.stop = stop + return &restartTestPlugin{ + Device: dpi, + grpcPlugin: dpi, + socketPath: dpi.socketPath, + healthCheck: dpi.healthCheck, + setServer: func(server *grpc.Server) { dpi.server = server }, + stopPlugin: dpi.stopDevicePlugin, + resetChannels: func() { + dpi.done = make(chan struct{}) + dpi.deregistered = make(chan struct{}) + }, + wakeHandlers: func() { + if !IsChanClosed(dpi.done) { + close(dpi.done) + } + }, + } + } + + newTestDevicePlugin := func(kind string) *restartTestPlugin { + if kind == "PCI" { + return newTestPCIPlugin() + } + return newTestGenericPlugin() + } + + It("waits for gRPC handlers when stopping the production server", func() { + socketPath := filepath.Join(socketDir, "blocking-handler.sock") + listener, err := net.Listen("unix", socketPath) + Expect(err).ToNot(HaveOccurred()) + + release := make(chan struct{}) + var releaseOnce sync.Once + releaseHandler := func() { releaseOnce.Do(func() { close(release) }) } + server := newDevicePluginGRPCServer() + serveDone := make(chan struct{}) + plugin := &blockingDevicePluginServer{ + entered: make(chan struct{}), + release: release, + } + pluginapi.RegisterDevicePluginServer(server, plugin) + go func() { + _ = server.Serve(listener) + close(serveDone) + }() + DeferCleanup(func() { + releaseHandler() + server.Stop() + Eventually(serveDone, 5*time.Second).Should(BeClosed()) + }) + Expect(waitForGRPCServer(socketPath, connectionTimeout)).To(Succeed()) + + conn, err := gRPCConnect(socketPath, connectionTimeout) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(conn.Close) + streamContext, cancelStream := context.WithCancel(context.Background()) + DeferCleanup(cancelStream) + _, err = pluginapi.NewDevicePluginClient(conn).ListAndWatch(streamContext, &pluginapi.Empty{}) + Expect(err).ToNot(HaveOccurred()) + Eventually(plugin.entered, 5*time.Second).Should(BeClosed()) + + stopped := make(chan struct{}) + go func() { + server.Stop() + close(stopped) + }() + Consistently(stopped, 300*time.Millisecond).ShouldNot(BeClosed()) + + releaseHandler() + Eventually(stopped, 5*time.Second).Should(BeClosed()) + }) + + DescribeTable("watches for restarts before registering", + func(kind string) { + plugin := newTestDevicePlugin(kind) + oldKubelet := newFakeKubelet(kubeletSocketPath, socketDir) + registerEntered := make(chan struct{}) + releaseRegister := make(chan struct{}) + var enterOnce sync.Once + var releaseOnce sync.Once + release := func() { releaseOnce.Do(func() { close(releaseRegister) }) } + oldKubelet.registration.beforeRegisterReturn = func() { + enterOnce.Do(func() { close(registerEntered) }) + <-releaseRegister + } + Expect(oldKubelet.start()).To(Succeed()) + + newKubelet := newFakeKubelet(kubeletSocketPath, socketDir) + controlled := &controlledDevice{ + devicePlugin: plugin.Device, + backoff: []time.Duration{10 * time.Millisecond, 20 * time.Millisecond}, + } + controlled.Start() + controlledStopped := false + DeferCleanup(func() { + release() + if !controlledStopped { + controlled.Stop() + Eventually(plugin.GetInitialized, 5*time.Second).Should(BeFalse()) + } + newKubelet.stop() + oldKubelet.stop() + }) + + Eventually(registerEntered, 5*time.Second).Should(BeClosed()) + Expect(newKubelet.start()).To(Succeed()) + release() + + Eventually(newKubelet.registration.registrationCount, 10*time.Second).Should(Equal(1)) + Eventually(plugin.GetInitialized, 5*time.Second).Should(BeTrue()) + + controlled.Stop() + controlledStopped = true + Eventually(plugin.GetInitialized, 5*time.Second).Should(BeFalse()) + }, + Entry("for PCI plugins", "PCI"), + Entry("for generic plugins", "generic"), + ) + + DescribeTable("re-registers and advertises devices on the second cycle", + func(kind string) { + plugin := newTestDevicePlugin(kind) + kubelet := newFakeKubelet(kubeletSocketPath, socketDir) + Expect(kubelet.start()).To(Succeed()) + DeferCleanup(kubelet.stop) + + controlled := &controlledDevice{ + devicePlugin: plugin.Device, + backoff: []time.Duration{10 * time.Millisecond, 20 * time.Millisecond}, + } + controlled.Start() + controlledStopped := false + DeferCleanup(func() { + if !controlledStopped { + controlled.Stop() + Eventually(plugin.GetInitialized, 5*time.Second).Should(BeFalse()) + } + }) + + Eventually(kubelet.registration.registrationCount, 5*time.Second).Should(Equal(1)) + Eventually(func() int { + return kubelet.registration.lastListLength(1) + }, 5*time.Second).Should(Equal(1)) + + kubelet.stop() + Expect(os.Remove(plugin.socketPath)).To(Succeed()) + Expect(kubelet.start()).To(Succeed()) + + Eventually(kubelet.registration.registrationCount, 10*time.Second).Should(Equal(2)) + Eventually(func() int { + return kubelet.registration.lastListLength(2) + }, 5*time.Second).Should(Equal(1)) + Consistently(func() int { + return kubelet.registration.lastListLength(2) + }, 300*time.Millisecond).Should(Equal(1)) + + controlled.Stop() + controlledStopped = true + Eventually(plugin.GetInitialized, 5*time.Second).Should(BeFalse()) + }, + Entry("PCI", "PCI"), + Entry("generic", "generic"), + ) + + DescribeTable("wakes active ListAndWatch handlers before stopping the server", + func(kind string) { + plugin := newTestDevicePlugin(kind) + Expect(os.RemoveAll(plugin.socketPath)).To(Succeed()) + plugin.resetChannels() + verifyActiveListAndWatchHandlerIsWoken(plugin) + }, + Entry("PCI", "PCI"), + Entry("generic", "generic"), + ) + + DescribeTable("detects kubelet restart filesystem events", + func(kind, restart string) { + plugin := newTestDevicePlugin(kind) + watcher := newTestHealthWatcher(plugin.socketPath) + result := make(chan error, 1) + go func() { result <- plugin.healthCheck(watcher) }() + + switch restart { + case "kubelet socket": + Expect(os.Remove(kubeletSocketPath)).To(Succeed()) + time.Sleep(100 * time.Millisecond) + touch(kubeletSocketPath) + case "plugin directory": + Expect(os.Rename(socketDir, socketDir+".gone")).To(Succeed()) + Expect(os.MkdirAll(socketDir, 0755)).To(Succeed()) + touch(kubeletSocketPath) + } + + Eventually(result, 5*time.Second).Should(Receive(BeNil())) + }, + Entry("PCI kubelet socket recreation", "PCI", "kubelet socket"), + Entry("generic kubelet socket recreation", "generic", "kubelet socket"), + Entry("PCI plugin directory replacement", "PCI", "plugin directory"), + Entry("generic plugin directory replacement", "generic", "plugin directory"), + ) + + Context("registration RPC supervision", func() { + It("cancels hung PCI and generic registrations on shutdown", func() { + kubelet := newBlockingFakeKubelet(kubeletSocketPath) + Expect(kubelet.start()).To(Succeed()) + DeferCleanup(kubelet.stop) + + pci := newTestPCIPlugin() + generic := newTestGenericPlugin() + pciStop := make(chan struct{}) + genericStop := make(chan struct{}) + var stopPCIOnce sync.Once + var stopGenericOnce sync.Once + stopPCI := func() { stopPCIOnce.Do(func() { close(pciStop) }) } + stopGeneric := func() { stopGenericOnce.Do(func() { close(genericStop) }) } + DeferCleanup(stopPCI) + DeferCleanup(stopGeneric) + + pciResult := make(chan error, 1) + genericResult := make(chan error, 1) + go func() { pciResult <- pci.Start(pciStop) }() + go func() { genericResult <- generic.Start(genericStop) }() + + for range 2 { + Eventually(kubelet.entered, 5*time.Second).Should(Receive()) + } + + started := time.Now() + stopPCI() + stopGeneric() + + var pciErr error + var genericErr error + Eventually(pciResult, 3*time.Second).Should(Receive(&pciErr)) + Eventually(genericResult, 3*time.Second).Should(Receive(&genericErr)) + Expect(pciErr).To(MatchError(ContainSubstring("Canceled"))) + Expect(genericErr).To(MatchError(ContainSubstring("Canceled"))) + Expect(time.Since(started)).To(BeNumerically("<", connectionTimeout)) + }) + + It("times out hung PCI and generic registrations", func() { + kubelet := newBlockingFakeKubelet(kubeletSocketPath) + Expect(kubelet.start()).To(Succeed()) + DeferCleanup(kubelet.stop) + + pci := newTestPCIPlugin() + generic := newTestGenericPlugin() + pciStop := make(chan struct{}) + genericStop := make(chan struct{}) + DeferCleanup(func() { close(pciStop) }) + DeferCleanup(func() { close(genericStop) }) + + pciResult := make(chan error, 1) + genericResult := make(chan error, 1) + started := time.Now() + go func() { pciResult <- pci.Start(pciStop) }() + go func() { genericResult <- generic.Start(genericStop) }() + + for range 2 { + Eventually(kubelet.entered, 5*time.Second).Should(Receive()) + } + + var pciErr error + var genericErr error + Eventually(pciResult, connectionTimeout+2*time.Second).Should(Receive(&pciErr)) + Eventually(genericResult, connectionTimeout+2*time.Second).Should(Receive(&genericErr)) + Expect(pciErr).To(MatchError(ContainSubstring("DeadlineExceeded"))) + Expect(genericErr).To(MatchError(ContainSubstring("DeadlineExceeded"))) + Expect(time.Since(started)).To(BeNumerically(">=", connectionTimeout)) + }) + }) + + DescribeTable("handles combined fsnotify operation masks", + func(kind, target string, operation fsnotify.Op) { + plugin := newTestDevicePlugin(kind) + eventName := plugin.socketPath + if target == "kubelet socket" { + eventName = kubeletSocketPath + } else if target == "plugin directory" { + eventName = socketDir + } + verifyHealthCheckReturnsForEvent(plugin, fsnotify.Event{Name: eventName, Op: operation}) + }, + Entry("PCI own socket Remove|Chmod", "PCI", "plugin socket", fsnotify.Remove|fsnotify.Chmod), + Entry("generic own socket Remove|Chmod", "generic", "plugin socket", fsnotify.Remove|fsnotify.Chmod), + Entry("PCI kubelet socket Create|Chmod", "PCI", "kubelet socket", fsnotify.Create|fsnotify.Chmod), + Entry("generic kubelet socket Create|Chmod", "generic", "kubelet socket", fsnotify.Create|fsnotify.Chmod), + Entry("PCI directory Rename|Chmod", "PCI", "plugin directory", fsnotify.Rename|fsnotify.Chmod), + Entry("generic directory Rename|Chmod", "generic", "plugin directory", fsnotify.Rename|fsnotify.Chmod), + Entry("PCI directory Remove|Chmod", "PCI", "plugin directory", fsnotify.Remove|fsnotify.Chmod), + Entry("generic directory Remove|Chmod", "generic", "plugin directory", fsnotify.Remove|fsnotify.Chmod), + ) + + DescribeTable("ignores restart events with the wrong path or operation", + func(kind, target, wrongField string) { + plugin := newTestDevicePlugin(kind) + restart := fsnotify.Event{Name: plugin.socketPath, Op: fsnotify.Remove} + if target == "kubelet socket" { + restart = fsnotify.Event{Name: kubeletSocketPath, Op: fsnotify.Create} + } else if target == "plugin directory" { + restart = fsnotify.Event{Name: socketDir, Op: fsnotify.Rename} + } + + ignored := restart + if wrongField == "path" { + ignored.Name += ".unrelated" + } else { + ignored.Op = fsnotify.Chmod + } + verifyHealthCheckIgnoresEvent(plugin, ignored, restart) + }, + Entry("PCI ignores a kubelet create on the wrong path", "PCI", "kubelet socket", "path"), + Entry("generic ignores a kubelet create on the wrong path", "generic", "kubelet socket", "path"), + Entry("PCI ignores a plugin remove on the wrong path", "PCI", "plugin socket", "path"), + Entry("generic ignores a plugin remove on the wrong path", "generic", "plugin socket", "path"), + Entry("PCI ignores a directory rename on the wrong path", "PCI", "plugin directory", "path"), + Entry("generic ignores a directory rename on the wrong path", "generic", "plugin directory", "path"), + Entry("PCI ignores the wrong kubelet operation", "PCI", "kubelet socket", "operation"), + Entry("generic ignores the wrong kubelet operation", "generic", "kubelet socket", "operation"), + Entry("PCI ignores the wrong plugin operation", "PCI", "plugin socket", "operation"), + Entry("generic ignores the wrong plugin operation", "generic", "plugin socket", "operation"), + Entry("PCI ignores the wrong directory operation", "PCI", "plugin directory", "operation"), + Entry("generic ignores the wrong directory operation", "generic", "plugin directory", "operation"), + ) +})