Skip to content
Closed
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
3 changes: 3 additions & 0 deletions internal/pkg/appconfig/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,9 @@ type Config struct {
EnableGPUBindUnbindWatch bool // Enable GPU bind/unbind event monitoring
GPUBindUnbindPollInterval time.Duration // Poll interval for GPU bind/unbind events
EnablePprof bool // Enable /debug/pprof/ HTTP endpoints
NVMLInitRetryAttempts int // Max attempts to initialize NVML before giving up
NVMLInitRetryBaseWait time.Duration // Initial backoff wait between NVML init retries
NVMLInitRetryMaxWait time.Duration // Cap on backoff wait between NVML init retries
}

// Clone returns a copy of Config with slices duplicated for reload snapshots.
Expand Down
126 changes: 118 additions & 8 deletions internal/pkg/nvmlprovider/provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,14 @@
package nvmlprovider

import (
"context"
"errors"
"fmt"
"log/slog"
"math/rand/v2"
"strconv"
"strings"
"time"

"github.com/NVIDIA/go-nvml/pkg/nvml"
)
Expand All @@ -32,16 +35,123 @@ type MIGDeviceInfo struct {
ComputeInstanceID int
}

// Default retry parameters for Initialize. The GPU driver installer on GKE
// (and similar node-bootstrap setups) can still be running when the exporter
// pod starts, so nvml.Init briefly returns ERROR_LIBRARY_NOT_FOUND. These
// defaults give the driver installer up to ~1 minute to finish before giving up.
const (
DefaultNVMLInitRetryAttempts = 5
DefaultNVMLInitRetryBaseWait = 2 * time.Second
DefaultNVMLInitRetryMaxWait = 20 * time.Second
)

// maxJitterFraction is the maximum fraction of the backoff duration added as
// random jitter, so pods restarting together (e.g. after a node reboot) don't
// retry in lockstep against the same node.
const maxJitterFraction = 0.30

var nvmlInterface NVML

// Initialize sets up the Singleton NVML interface.
// nvmlInitFunc is a package-level indirection over nvml.Init so tests can
// substitute a fake without requiring real GPU hardware.
var nvmlInitFunc = nvml.Init

// nvmlInitError wraps the raw NVML return code from nvmlInitFunc so callers
// can distinguish transient failures (e.g. ERROR_LIBRARY_NOT_FOUND) from
// permanent ones without resorting to string matching.
type nvmlInitError struct {
ret nvml.Return
}

func (e *nvmlInitError) Error() string {
return nvml.ErrorString(e.ret)
}

// isLibraryNotFoundErr reports whether err represents nvml.ERROR_LIBRARY_NOT_FOUND,
// the transient error returned while the GPU driver installer hasn't finished yet.
func isLibraryNotFoundErr(err error) bool {
var initErr *nvmlInitError
if errors.As(err, &initErr) {
return initErr.ret == nvml.ERROR_LIBRARY_NOT_FOUND
}
return false
}

// Initialize sets up the Singleton NVML interface, retrying on the transient
// ERROR_LIBRARY_NOT_FOUND error using sane default retry parameters.
func Initialize() error {
var err error
nvmlInterface, err = newNVMLProvider()
if err != nil {
return err
return InitializeWithRetry(context.Background(), DefaultNVMLInitRetryAttempts, DefaultNVMLInitRetryBaseWait, DefaultNVMLInitRetryMaxWait)
}

// InitializeWithRetry sets up the Singleton NVML interface, retrying up to
// attempts times with exponential backoff (base, 2x, 4x, ... capped at
// maxWait, plus jitter) when nvml.Init fails with ERROR_LIBRARY_NOT_FOUND.
// Any other error is returned immediately without retrying, since retrying a
// permanent failure only delays an unavoidable error. attempts < 1 is treated
// as 1, so at least one init attempt always happens. If ctx is cancelled
// while waiting between attempts, InitializeWithRetry returns ctx.Err()
// immediately instead of sleeping out the remaining backoff, so a shutdown
// signal during startup isn't ignored until retries are exhausted.
func InitializeWithRetry(ctx context.Context, attempts int, baseWait, maxWait time.Duration) error {
if attempts < 1 {
attempts = 1
}
return nil

var lastErr error
for attempt := 0; attempt < attempts; attempt++ {
var err error
nvmlInterface, err = newNVMLProvider()
if err == nil {
return nil
}
lastErr = err

if !isLibraryNotFoundErr(err) {
return fmt.Errorf("failed to initialize NVML library: %w", err)
}

if attempt < attempts-1 {
wait := backoffDuration(attempt, baseWait, maxWait)
slog.Warn("NVML library not found yet (GPU driver may still be installing); retrying",
slog.Int("attempt", attempt+1),
slog.Int("maxAttempts", attempts),
slog.Duration("wait", wait))

timer := time.NewTimer(wait)
select {
case <-timer.C:
case <-ctx.Done():
timer.Stop()
return fmt.Errorf("NVML initialization cancelled after %d attempt(s): %w", attempt+1, ctx.Err())
}
}
}

return fmt.Errorf("failed to initialize NVML library after %d attempts, last error: %w", attempts, lastErr)
}

// backoffDuration computes the exponential backoff wait for the given attempt
// (0-indexed): baseWait * 2^attempt, capped at maxWait, plus up to
// maxJitterFraction of additional random jitter on top of the cap. The
// doubling is computed via repeated, overflow-checked multiplication rather
// than a bit shift so an unbounded attempt count (attempts is a
// user-configurable CLI value) can never wrap into a bogus small positive
// duration that defeats the cap.
func backoffDuration(attempt int, baseWait, maxWait time.Duration) time.Duration {
wait := baseWait
for i := 0; i < attempt; i++ {
if wait >= maxWait || wait > maxWait/2 {
wait = maxWait
break
}
wait *= 2
}
if wait > maxWait {
wait = maxWait
}

jitter := time.Duration(rand.Float64() * maxJitterFraction * float64(wait)) //nolint:gosec // #nosec G404 -- jitter only needs to desynchronize retries, not be cryptographically secure
return wait + jitter
}

// reset clears the current NVML interface instance.
Expand Down Expand Up @@ -77,9 +187,9 @@ func newNVMLProvider() (NVML, error) {
}

slog.Info("Attempting to initialize NVML library.")
ret := nvml.Init()
ret := nvmlInitFunc()
if ret != nvml.SUCCESS {
err := errors.New(nvml.ErrorString(ret))
err := &nvmlInitError{ret: ret}
slog.Error(fmt.Sprintf("Cannot init NVML library; err: %v", err))
return nvmlProvider{initialized: false}, err
}
Expand Down
42 changes: 37 additions & 5 deletions internal/pkg/nvmlprovider/provider_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,10 +19,22 @@ package nvmlprovider
import (
"testing"

"github.com/NVIDIA/go-nvml/pkg/nvml"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

// mockSuccessfulInit stubs nvmlInitFunc to succeed without a real NVML
// library, and returns a restore func the caller should defer. Only safe for
// tests that don't go on to call any other real nvml.* function: those
// still need the actual library loaded to resolve, and will crash the test
// binary with a symbol lookup error rather than return a Go error if it
// isn't. Use the Initialize()-then-t.Skip pattern elsewhere in this file for
// tests that do.
func mockSuccessfulInit() func() {
nvmlInitFunc = func() nvml.Return { return nvml.SUCCESS }
return func() { nvmlInitFunc = nvml.Init }
}

func TestGetMIGDeviceInfoByID_When_NVML_Not_Initialized(t *testing.T) {
validMIGUUID := "MIG-GPU-b8ea3855-276c-c9cb-b366-c6fa655957c5/1/5"
newNvmlProvider := nvmlProvider{}
Expand Down Expand Up @@ -56,7 +68,14 @@ func TestGetAllMIGDevicesProcessMemory_When_NVML_Not_Initialized(t *testing.T) {
}

func TestGetMIGDeviceInfoByID_When_DriverVersion_Below_R470(t *testing.T) {
_ = Initialize()
// This exercises nvml.DeviceGetHandleByUUID beyond just Initialize, so a
// stubbed nvmlInitFunc isn't enough: the real NVML library still needs to
// be loaded for that symbol to resolve, or the test binary crashes
// outright instead of returning a Go error. Skip like its sibling tests
// when no real NVML is present.
if err := Initialize(); err != nil {
t.Skip("NVML not available, skipping test")
}
assert.NotNil(t, Client(), "expected NVML Client to be not nil")
assert.True(t, Client().(nvmlProvider).initialized, "expected Client to be initialized")
defer Client().Cleanup()
Expand Down Expand Up @@ -112,6 +131,8 @@ func TestGetMIGDeviceInfoByID_When_DriverVersion_Below_R470(t *testing.T) {
}

func Test_newNVMLProvider(t *testing.T) {
defer mockSuccessfulInit()()

tests := []struct {
name string
preRunFunc func() NVML
Expand Down Expand Up @@ -195,9 +216,13 @@ func TestCleanup_WhenNotInitialized(t *testing.T) {

// TestCleanup_WhenInitialized tests cleanup when NVML is initialized
func TestCleanup_WhenInitialized(t *testing.T) {
// Initialize NVML
// Cleanup calls the real nvml.Shutdown, which needs the real library
// loaded to resolve, so this can't be satisfied by stubbing
// nvmlInitFunc alone. Skip like its sibling tests when unavailable.
err := Initialize()
assert.NoError(t, err)
if err != nil {
t.Skip("NVML not available, skipping test")
}

provider := Client()
assert.NotNil(t, provider)
Expand Down Expand Up @@ -239,9 +264,16 @@ func TestPreCheck(t *testing.T) {
errorContains string
}{
{
// This subtest's assertion below only cares that GetMIGDeviceInfoByID
// doesn't fail with "NVML not initialized" - it tolerates any other
// error - but reaching that call still needs Initialize to have
// actually succeeded, which needs the real NVML library. Skip like
// the other hardware-dependent tests in this file when unavailable.
name: "Initialized provider",
setupFunc: func(t *testing.T) {
require.NoError(t, Initialize())
if err := Initialize(); err != nil {
t.Skip("NVML not available, skipping test")
}
},
expectError: false,
},
Expand Down
Loading