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
5 changes: 5 additions & 0 deletions internal/volume/csi/controller.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
75 changes: 73 additions & 2 deletions internal/volume/csi/plugin.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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
}

Expand Down
114 changes: 114 additions & 0 deletions internal/volume/csi/plugin_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
}
}