From 9d469c656738ae8b303ae7929a2083f6e3a998ec Mon Sep 17 00:00:00 2001 From: Sunny Song Date: Thu, 17 Sep 2026 21:17:11 +0000 Subject: [PATCH] Query CSI driver capabilities at plugin initialization --- internal/volume/csi/controller.go | 5 ++ internal/volume/csi/plugin.go | 75 ++++++++++++++++++- internal/volume/csi/plugin_test.go | 114 +++++++++++++++++++++++++++++ 3 files changed, 192 insertions(+), 2 deletions(-) diff --git a/internal/volume/csi/controller.go b/internal/volume/csi/controller.go index 39f6090877..d4abe274c9 100644 --- a/internal/volume/csi/controller.go +++ b/internal/volume/csi/controller.go @@ -39,3 +39,8 @@ func (c *Client) ControllerPublishVolume(ctx context.Context, req *csi.Controlle func (c *Client) ControllerUnpublishVolume(ctx context.Context, req *csi.ControllerUnpublishVolumeRequest) (*csi.ControllerUnpublishVolumeResponse, error) { return c.controller.ControllerUnpublishVolume(ctx, req) } + +// ControllerGetCapabilities queries the supported capabilities of the controller service. +func (c *Client) ControllerGetCapabilities(ctx context.Context, req *csi.ControllerGetCapabilitiesRequest) (*csi.ControllerGetCapabilitiesResponse, error) { + return c.controller.ControllerGetCapabilities(ctx, req) +} diff --git a/internal/volume/csi/plugin.go b/internal/volume/csi/plugin.go index ab1337efd9..9138368342 100644 --- a/internal/volume/csi/plugin.go +++ b/internal/volume/csi/plugin.go @@ -56,6 +56,10 @@ var defaultTLSPaths = tlsPaths{ type Plugin struct { client *Client stagingDirPrefix string + + capMu sync.RWMutex + hasCapabilities bool + supportsPublishUnpublish bool } // Ensure Plugin implements volume.VolumePluginControlPlane and VolumePluginWorkerPlane @@ -121,8 +125,63 @@ func (p *Plugin) DeleteVolume(ctx context.Context, volumeID string) error { return nil } +// InitControllerCapabilities queries the CSI Controller service for its supported capabilities +// and caches them on the Plugin. +func (p *Plugin) InitControllerCapabilities(ctx context.Context) error { + p.capMu.Lock() + defer p.capMu.Unlock() + if p.hasCapabilities { + return nil + } + + resp, err := p.client.ControllerGetCapabilities(ctx, &csi.ControllerGetCapabilitiesRequest{}) + if err != nil { + if status.Code(err) == codes.Unimplemented { + p.hasCapabilities = true + p.supportsPublishUnpublish = false + return nil + } + return fmt.Errorf("CSI ControllerGetCapabilities failed: %w", err) + } + + for _, cap := range resp.GetCapabilities() { + if rpc := cap.GetRpc(); rpc != nil { + if rpc.GetType() == csi.ControllerServiceCapability_RPC_PUBLISH_UNPUBLISH_VOLUME { + p.supportsPublishUnpublish = true + } + } + } + p.hasCapabilities = true + return nil +} + +// SupportsPublishUnpublish returns whether the CSI driver supports ControllerPublishVolume +// and ControllerUnpublishVolume. +func (p *Plugin) SupportsPublishUnpublish() bool { + p.capMu.RLock() + defer p.capMu.RUnlock() + return p.supportsPublishUnpublish +} + +func (p *Plugin) ensureControllerCapabilities(ctx context.Context) error { + p.capMu.RLock() + initialized := p.hasCapabilities + p.capMu.RUnlock() + if initialized { + return nil + } + return p.InitControllerCapabilities(ctx) +} + // AttachVolume maps to CSI Controller ControllerPublishVolume. func (p *Plugin) AttachVolume(ctx context.Context, volumeID string, node string) error { + if err := p.ensureControllerCapabilities(ctx); err != nil { + return err + } + if !p.SupportsPublishUnpublish() { + return nil + } + req := &csi.ControllerPublishVolumeRequest{ VolumeId: volumeID, NodeId: node, @@ -132,8 +191,6 @@ func (p *Plugin) AttachVolume(ctx context.Context, volumeID string, node string) resp, err := p.client.ControllerPublishVolume(ctx, req) if err != nil { - // TODO: Query CSI driver capabilities ahead of time (e.g. during plugin initialization) - // to avoid calling unimplemented methods and generating spammy logs. if status.Code(err) == codes.Unimplemented { slog.WarnContext(ctx, "CSI ControllerPublishVolume is unimplemented by driver; skipping attach", slog.String("volume_id", volumeID), slog.String("node", node)) return nil @@ -155,6 +212,13 @@ func (p *Plugin) AttachVolume(ctx context.Context, volumeID string, node string) // DetachVolume maps to CSI Controller ControllerUnpublishVolume. func (p *Plugin) DetachVolume(ctx context.Context, volumeID string, node string) error { + if err := p.ensureControllerCapabilities(ctx); err != nil { + return err + } + if !p.SupportsPublishUnpublish() { + return nil + } + req := &csi.ControllerUnpublishVolumeRequest{ VolumeId: volumeID, NodeId: node, @@ -322,6 +386,13 @@ func newCSIPlugin(ctx context.Context, lister listersv1alpha1.CSIDriverConfigLis return nil, fmt.Errorf("reported driver name %q does not match requested name %q", reportedName, driverName) } + if isController { + if err := csiPlugin.InitControllerCapabilities(ctx); err != nil { + csiClient.Close() + return nil, fmt.Errorf("failed to query controller capabilities for %q: %w", driverName, err) + } + } + return csiPlugin, nil } diff --git a/internal/volume/csi/plugin_test.go b/internal/volume/csi/plugin_test.go index 79b5a75d9d..86f5bdd0a3 100644 --- a/internal/volume/csi/plugin_test.go +++ b/internal/volume/csi/plugin_test.go @@ -36,6 +36,7 @@ type mockCSIDriver struct { deleteVolumeFunc func(context.Context, *csi.DeleteVolumeRequest) (*csi.DeleteVolumeResponse, error) controllerPublishVolumeFunc func(context.Context, *csi.ControllerPublishVolumeRequest) (*csi.ControllerPublishVolumeResponse, error) controllerUnpublishVolumeFunc func(context.Context, *csi.ControllerUnpublishVolumeRequest) (*csi.ControllerUnpublishVolumeResponse, error) + controllerGetCapabilitiesFunc func(context.Context, *csi.ControllerGetCapabilitiesRequest) (*csi.ControllerGetCapabilitiesResponse, error) nodeStageVolumeFunc func(context.Context, *csi.NodeStageVolumeRequest) (*csi.NodeStageVolumeResponse, error) nodeUnstageVolumeFunc func(context.Context, *csi.NodeUnstageVolumeRequest) (*csi.NodeUnstageVolumeResponse, error) nodePublishVolumeFunc func(context.Context, *csi.NodePublishVolumeRequest) (*csi.NodePublishVolumeResponse, error) @@ -109,6 +110,23 @@ func (m *mockCSIDriver) ControllerUnpublishVolume(ctx context.Context, req *csi. return &csi.ControllerUnpublishVolumeResponse{}, nil } +func (m *mockCSIDriver) ControllerGetCapabilities(ctx context.Context, req *csi.ControllerGetCapabilitiesRequest) (*csi.ControllerGetCapabilitiesResponse, error) { + if m.controllerGetCapabilitiesFunc != nil { + return m.controllerGetCapabilitiesFunc(ctx, req) + } + return &csi.ControllerGetCapabilitiesResponse{ + Capabilities: []*csi.ControllerServiceCapability{ + { + Type: &csi.ControllerServiceCapability_Rpc{ + Rpc: &csi.ControllerServiceCapability_RPC{ + Type: csi.ControllerServiceCapability_RPC_PUBLISH_UNPUBLISH_VOLUME, + }, + }, + }, + }, + }, nil +} + func (m *mockCSIDriver) NodeStageVolume(ctx context.Context, req *csi.NodeStageVolumeRequest) (*csi.NodeStageVolumeResponse, error) { if m.nodeStageVolumeFunc != nil { return m.nodeStageVolumeFunc(ctx, req) @@ -408,3 +426,99 @@ func TestClient_Identity(t *testing.T) { t.Fatalf("Probe failed: %v", err) } } + +func TestPlugin_Capabilities_SkipAttachDetach(t *testing.T) { + driver := &mockCSIDriver{} + endpoint, cleanup := startMockCSIDriver(t, driver) + defer cleanup() + + // Configure driver to report no PUBLISH_UNPUBLISH_VOLUME capability. + driver.controllerGetCapabilitiesFunc = func(ctx context.Context, req *csi.ControllerGetCapabilitiesRequest) (*csi.ControllerGetCapabilitiesResponse, error) { + return &csi.ControllerGetCapabilitiesResponse{ + Capabilities: []*csi.ControllerServiceCapability{ + { + Type: &csi.ControllerServiceCapability_Rpc{ + Rpc: &csi.ControllerServiceCapability_RPC{ + Type: csi.ControllerServiceCapability_RPC_CREATE_DELETE_VOLUME, + }, + }, + }, + }, + }, nil + } + + driver.controllerPublishVolumeFunc = func(ctx context.Context, req *csi.ControllerPublishVolumeRequest) (*csi.ControllerPublishVolumeResponse, error) { + t.Fatalf("ControllerPublishVolume should not be called when capability is missing") + return nil, nil + } + driver.controllerUnpublishVolumeFunc = func(ctx context.Context, req *csi.ControllerUnpublishVolumeRequest) (*csi.ControllerUnpublishVolumeResponse, error) { + t.Fatalf("ControllerUnpublishVolume should not be called when capability is missing") + return nil, nil + } + + client, err := NewCSIClient(endpoint, nil) + if err != nil { + t.Fatalf("failed to create CSI client: %v", err) + } + defer client.Close() + + plugin := NewPlugin(client) + ctx := context.Background() + + if err := plugin.InitControllerCapabilities(ctx); err != nil { + t.Fatalf("InitControllerCapabilities failed: %v", err) + } + + if plugin.SupportsPublishUnpublish() { + t.Errorf("expected SupportsPublishUnpublish to be false, got true") + } + + if err := plugin.AttachVolume(ctx, "test-vol", "node-1"); err != nil { + t.Errorf("AttachVolume failed: %v", err) + } + + if err := plugin.DetachVolume(ctx, "test-vol", "node-1"); err != nil { + t.Errorf("DetachVolume failed: %v", err) + } +} + +func TestPlugin_Capabilities_Unimplemented(t *testing.T) { + driver := &mockCSIDriver{} + endpoint, cleanup := startMockCSIDriver(t, driver) + defer cleanup() + + // Driver returns Unimplemented for ControllerGetCapabilities. + driver.controllerGetCapabilitiesFunc = func(ctx context.Context, req *csi.ControllerGetCapabilitiesRequest) (*csi.ControllerGetCapabilitiesResponse, error) { + return nil, status.Error(codes.Unimplemented, "unimplemented") + } + + driver.controllerPublishVolumeFunc = func(ctx context.Context, req *csi.ControllerPublishVolumeRequest) (*csi.ControllerPublishVolumeResponse, error) { + t.Fatalf("ControllerPublishVolume should not be called when capabilities are unimplemented") + return nil, nil + } + driver.controllerUnpublishVolumeFunc = func(ctx context.Context, req *csi.ControllerUnpublishVolumeRequest) (*csi.ControllerUnpublishVolumeResponse, error) { + t.Fatalf("ControllerUnpublishVolume should not be called when capabilities are unimplemented") + return nil, nil + } + + client, err := NewCSIClient(endpoint, nil) + if err != nil { + t.Fatalf("failed to create CSI client: %v", err) + } + defer client.Close() + + plugin := NewPlugin(client) + ctx := context.Background() + + if err := plugin.AttachVolume(ctx, "test-vol", "node-1"); err != nil { + t.Errorf("AttachVolume failed: %v", err) + } + + if plugin.SupportsPublishUnpublish() { + t.Errorf("expected SupportsPublishUnpublish to be false when ControllerGetCapabilities is unimplemented") + } + + if err := plugin.DetachVolume(ctx, "test-vol", "node-1"); err != nil { + t.Errorf("DetachVolume failed: %v", err) + } +}