@@ -23,6 +23,7 @@ import (
2323 "github.com/NVIDIA/go-nvml/pkg/nvml"
2424 mocknvml "github.com/NVIDIA/go-nvml/pkg/nvml/mock"
2525 "github.com/NVIDIA/go-nvml/pkg/nvml/mock/dgxa100"
26+ mockserver "github.com/NVIDIA/go-nvml/pkg/nvml/mock/server"
2627 "github.com/stretchr/testify/require"
2728
2829 "github.com/NVIDIA/go-nvlib/pkg/nvlib/device"
@@ -32,7 +33,7 @@ func TestNvmllibGetDeviceSpecGeneratorsForIDs(t *testing.T) {
3233 testCases := []struct {
3334 name string
3435 ids []string
35- setupMock func (* dgxa100 .Server )
36+ setupMock func (* mockserver .Server )
3637 expectedError error
3738 expectedLength int
3839 expectedGenerators DeviceSpecGenerators
@@ -46,10 +47,10 @@ func TestNvmllibGetDeviceSpecGeneratorsForIDs(t *testing.T) {
4647 {
4748 name : "single GPU index" ,
4849 ids : []string {"0" },
49- setupMock : func (server * dgxa100 .Server ) {
50+ setupMock : func (server * mockserver .Server ) {
5051 for _ , d := range server .Devices {
5152 // TODO: This is not implemented in the mock.
52- (d .(* dgxa100 .Device )).IsMigDeviceHandleFunc = func () (bool , nvml.Return ) {
53+ (d .(* mockserver .Device )).IsMigDeviceHandleFunc = func () (bool , nvml.Return ) {
5354 return false , nvml .SUCCESS
5455 }
5556 }
@@ -60,10 +61,10 @@ func TestNvmllibGetDeviceSpecGeneratorsForIDs(t *testing.T) {
6061 {
6162 name : "single UUID" ,
6263 ids : []string {"GPU-12345678-1234-1234-1234-123456789abc" },
63- setupMock : func (server * dgxa100 .Server ) {
64+ setupMock : func (server * mockserver .Server ) {
6465 for _ , d := range server .Devices {
6566 // TODO: This is not implemented in the mock.
66- (d .(* dgxa100 .Device )).IsMigDeviceHandleFunc = func () (bool , nvml.Return ) {
67+ (d .(* mockserver .Device )).IsMigDeviceHandleFunc = func () (bool , nvml.Return ) {
6768 return false , nvml .SUCCESS
6869 }
6970 }
@@ -80,7 +81,7 @@ func TestNvmllibGetDeviceSpecGeneratorsForIDs(t *testing.T) {
8081 {
8182 name : "MIG device index" ,
8283 ids : []string {"0:0" },
83- setupMock : func (server * dgxa100 .Server ) {
84+ setupMock : func (server * mockserver .Server ) {
8485 mig := & mocknvml.Device {
8586 IsMigDeviceHandleFunc : func () (bool , nvml.Return ) {
8687 return true , nvml .SUCCESS
@@ -96,7 +97,7 @@ func TestNvmllibGetDeviceSpecGeneratorsForIDs(t *testing.T) {
9697 },
9798 }
9899
99- server .Devices [0 ].(* dgxa100 .Device ).GetMigDeviceHandleByIndexFunc = func (n int ) (nvml.Device , nvml.Return ) {
100+ server .Devices [0 ].(* mockserver .Device ).GetMigDeviceHandleByIndexFunc = func (n int ) (nvml.Device , nvml.Return ) {
100101 if n != 0 {
101102 return nil , nvml .ERROR_INVALID_ARGUMENT
102103 }
@@ -143,17 +144,17 @@ func TestNvmllibGetDeviceSpecGeneratorsForIDs(t *testing.T) {
143144}
144145
145146// TODO: These need to be implemented in go-nvlib
146- func mockOverrides (server * dgxa100 .Server ) {
147+ func mockOverrides (server * mockserver .Server ) {
147148 for i , d := range server .Devices {
148149 // TODO: This is not implemented in the mock.
149- (d .(* dgxa100 .Device )).GetMaxMigDeviceCountFunc = func () (int , nvml.Return ) {
150+ (d .(* mockserver .Device )).GetMaxMigDeviceCountFunc = func () (int , nvml.Return ) {
150151 return 0 , nvml .SUCCESS
151152 }
152- (d .(* dgxa100 .Device )).GetIndexFunc = func () (int , nvml.Return ) {
153+ (d .(* mockserver .Device )).GetIndexFunc = func () (int , nvml.Return ) {
153154 return i , nvml .SUCCESS
154155 }
155- (d .(* dgxa100 .Device )).GetUUIDFunc = func () (string , nvml.Return ) {
156- return d .(* dgxa100 .Device ).UUID , nvml .SUCCESS
156+ (d .(* mockserver .Device )).GetUUIDFunc = func () (string , nvml.Return ) {
157+ return d .(* mockserver .Device ).UUID , nvml .SUCCESS
157158 }
158159 }
159160}
0 commit comments