blob: b03fde33666d2dfd03b79a657816033db5ee2b0b [file]
package gpuconfig
import (
"errors"
"fmt"
"os"
"strings"
"testing"
"cos.googlesource.com/cos/tools.git/src/pkg/gpuconfig/pb"
"github.com/golang/protobuf/proto"
"github.com/google/go-cmp/cmp"
"google.golang.org/protobuf/encoding/prototext"
)
const testProtoDataPath = "./testdata/gpu_driver_versions_test.textproto"
func TestGetGPUDriverVersion(t *testing.T) {
textProtoData, err := os.ReadFile(testProtoDataPath)
if err != nil {
t.Errorf("Cannot read the test textproto data: %v", err)
}
var gpuDriverVersionInfoList = &pb.GPUDriverVersionInfoList{}
err = (prototext.UnmarshalOptions{DiscardUnknown: true, AllowPartial: true}).Unmarshal(textProtoData, gpuDriverVersionInfoList)
if err != nil {
t.Errorf("failed when parsing the GPU driver version: %v", err)
}
testData, err := proto.Marshal(gpuDriverVersionInfoList)
if err != nil {
t.Errorf("fail to encode to the binary proto data: %v", err)
}
var testCases = []struct {
gpuProtoContent []byte
gpuType string
input string
fallback bool
expectedDriverVersion string
err error
}{
{
nil,
"NVIDIA_L4",
"",
false,
"",
errors.New("the gpu proto content must not be empty"),
},
{
testData,
"",
"",
false,
"",
errors.New("the GPU type must not be empty"),
},
{
testData,
" nvidia_l4 ",
"",
false,
"535.154.05",
nil,
},
{
testData,
"NVIDIA_L4",
"560.38",
false,
"560.38",
nil,
},
{
testData,
"NVIDIA_L4",
"DEFAULT",
false,
"535.154.05",
nil,
},
{
testData,
"NVIDIA_L4",
"default",
false,
"535.154.05",
nil,
},
{
testData,
"NVIDIA_L4",
"latest",
false,
"535.154.05",
nil,
},
{
testData,
"NVIDIA_L4",
"R535",
false,
"535.154.05",
nil,
},
{
testData,
"NVIDIA_L4",
"r535",
false,
"535.154.05",
nil,
},
{
testData,
"NVIDIA_L4",
"535.129.03",
false,
"535.129.03",
nil,
},
{
testData,
"NVIDIA_L4",
"R470",
false,
"",
errors.New("not supported for GPU type"),
},
{
testData,
"NVIDIA_L4",
"R470",
true,
"535.154.05",
nil,
},
{
testData,
"NVIDIA_L4",
"R525",
true,
"535.154.05",
nil,
},
{
testData,
"invalidGPU",
"R525",
true,
"",
errors.New("no supported driver versions found for gpu"),
},
{
testData,
"NVIDIA_L4",
"invalidInput",
false,
"",
errors.New("the input is invalid"),
},
{
testData,
"NVIDIA_L4",
"invalidInput",
true,
"",
errors.New("the input is invalid"),
},
{
testData,
"no_gpu",
"r525",
true,
"535.161.08",
nil,
},
{
testData,
"no_gpu",
"r525",
true,
"535.161.08",
nil,
},
{
testData,
"Sample_GPU_without_default_label",
"r525",
true,
"",
errors.New("unable to parse gpu_type"),
},
}
for index, tc := range testCases {
t.Run(fmt.Sprintf("Test %d: GetGPUDriverVersion: with gpuType: %s, input: %s, fallback flag: %v", index, tc.gpuType, tc.input, tc.fallback), func(t *testing.T) {
driverVersion, err := GetGPUDriverVersion(tc.gpuProtoContent, tc.gpuType, tc.input, tc.fallback)
if tc.err == nil {
if err != nil {
t.Errorf("Failed to GetGPUDriverVersion with error: %v", err)
}
if !cmp.Equal(tc.expectedDriverVersion, driverVersion) {
t.Errorf("Test GetGPUDriverVersion failed: the expected driver version is: %s, but the actual driver version is: %s.", tc.expectedDriverVersion, driverVersion)
}
} else {
if err == nil || !strings.Contains(err.Error(), tc.err.Error()) {
t.Errorf("Test GetGPUDriverVersion failed: the error been sent out is not correct, the exected: %s, the actual: %s", tc.err.Error(), err.Error())
}
}
})
}
}
func TestGetGPUDriverVersions(t *testing.T) {
textProtoData, err := os.ReadFile(testProtoDataPath)
if err != nil {
t.Errorf("Cannot read the test textproto data: %v", err)
}
var gpuDriverVersionInfoList = &pb.GPUDriverVersionInfoList{}
err = (prototext.UnmarshalOptions{DiscardUnknown: true, AllowPartial: true}).Unmarshal(textProtoData, gpuDriverVersionInfoList)
if err != nil {
t.Errorf("failed when parsing the GPU driver version: %v", err)
}
testData, err := proto.Marshal(gpuDriverVersionInfoList)
if err != nil {
t.Errorf("fail to encode to the binary proto data: %v", err)
}
var testCases = []struct {
gpuProtoContent []byte
gpuType string
expectedDriverVersions []*pb.DriverVersion
err error
}{
{
nil,
"NVIDIA_L4",
nil,
errors.New("the gpu proto content must not be empty"),
},
{
testData,
"",
nil,
errors.New("the GPU type must not be empty"),
},
{
testData,
"NVIDIA_TESLA_V100",
[]*pb.DriverVersion{
{Version: "535.154.05", Label: "DEFAULT"},
{Version: "535.154.05", Label: "LATEST"},
{Version: "535.154.05", Label: "R535"},
{Version: "535.129.03"},
{Version: "535.104.12"},
{Version: "535.104.05"},
{Version: "470.223.02", Label: "R470"},
{Version: "470.199.02"},
},
nil,
},
{
testData,
"NVIDIA_TESLA_P100",
[]*pb.DriverVersion{
{Version: "535.154.05", Label: "DEFAULT"},
{Version: "535.154.05", Label: "LATEST"},
{Version: "535.154.05", Label: "R535"},
{Version: "535.129.03"},
{Version: "535.104.12"},
{Version: "535.104.05"},
{Version: "470.223.02", Label: "R470"},
{Version: "470.199.02"},
},
nil,
},
{
testData,
"InvalidGPU",
nil,
errors.New("no supported driver versions found for gpu"),
},
}
for index, tc := range testCases {
t.Run(fmt.Sprintf("Test %d: GetGPUDriverVersions: with gpuType: %s", index, tc.gpuType), func(t *testing.T) {
driverVersions, err := GetGPUDriverVersions(tc.gpuProtoContent, tc.gpuType)
if tc.err == nil {
if err != nil {
t.Errorf("Failed to GetGPUDriverVersions with error: %v", err)
}
if len(driverVersions) != len(tc.expectedDriverVersions) {
t.Errorf("Test GetGPUDriverVersions failed: the length of expected gpu drivers is %v, but the actual lenght is %v", len(tc.expectedDriverVersions), len(driverVersions))
}
for index, driverVersion := range driverVersions {
expectedDriverVersion := tc.expectedDriverVersions[index]
if expectedDriverVersion.Label != driverVersion.Label {
t.Errorf("Test GetGPUDriverVersions failed: the expected driver version list contains DriverVersion{Label: %s, Version: %s},"+
"but the actual driver version list contains {DriverVersion{Label: %s, Version: %s}}.", expectedDriverVersion.Label, expectedDriverVersion.Version,
driverVersion.Label, driverVersion.Version)
}
if expectedDriverVersion.Version != driverVersion.Version {
t.Errorf("Test GetGPUDriverVersions failed: the expected driver version list contains DriverVersion{Label: %s, Version: %s},"+
"but the actual driver version list contains {DriverVersion{Label: %s, Version: %s}}.", expectedDriverVersion.Label, expectedDriverVersion.Version,
driverVersion.Label, driverVersion.Version)
}
}
} else {
if err == nil || !strings.Contains(err.Error(), tc.err.Error()) {
t.Errorf("Test GetGPUDriverVersions failed: the error been sent out is not correct, the exected: %s, the actual: %s", tc.err.Error(), err.Error())
}
}
})
}
}
func TestGetVGPUDriverVersion(t *testing.T) {
textProtoData, err := os.ReadFile(testProtoDataPath)
if err != nil {
t.Errorf("Cannot read the test textproto data: %v", err)
}
var testCases = []struct {
name string
gpuType string
input string
fallback bool
hostDriverVersion string
isVGPU bool
expectedDriverVersion string
err error
}{
{
name: "missing host driver version",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "default",
fallback: false,
hostDriverVersion: "",
isVGPU: true,
expectedDriverVersion: "",
err: errors.New("host driver version missing"),
},
{
name: "precise version matching host branch",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "575.57.08",
fallback: false,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid-gcp",
err: nil,
},
{
name: "precise grid version matching host branch",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "575.57.08-grid",
fallback: false,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid",
err: nil,
},
{
name: "precise grid-gcp version matching host branch",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "575.57.08-grid-gcp",
fallback: false,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid-gcp",
err: nil,
},
{
name: "precise version incompatible with host, fallback disabled",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "575.00.00",
fallback: false,
hostDriverVersion: "550.10.10",
isVGPU: true,
expectedDriverVersion: "",
err: errors.New("is either not supported for vGPU type"),
},
{
name: "precise version only has grid-gcp",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "550.00.00",
fallback: true,
hostDriverVersion: "550.10.10",
isVGPU: true,
expectedDriverVersion: "550.00.00-grid-gcp",
err: nil,
},
{
name: "precise version only has grid",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "540.00.00",
fallback: false,
hostDriverVersion: "540.10.10",
isVGPU: true,
expectedDriverVersion: "540.00.00-grid",
err: nil,
},
{
name: "label matched, incompatible with host, fallback disabled",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "R550",
fallback: false,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "",
err: errors.New("and fallback is disabled"),
},
{
name: "label matched, incompatible, fallback enabled",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "R550",
fallback: true,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid-gcp",
err: nil,
},
{
name: "unknown label, fallback disabled",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "UNKNOWN",
fallback: false,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "",
err: errors.New("is not supported for vGPU type"),
},
{
name: "unknown label, fallback enabled",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "UNKNOWN",
fallback: true,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid-gcp",
err: nil,
},
{
name: "fallback finds exact host version as grid when host driver incompatible, only grid available",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "R550", // R550 is incompatible with 540 host branch
fallback: true,
hostDriverVersion: "540.00.00",
isVGPU: true,
expectedDriverVersion: "540.00.00-grid",
err: nil,
},
{
name: "fallback when no matching host branch, returns latest grid",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "default",
fallback: true,
hostDriverVersion: "999.99.99",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid-gcp",
err: nil,
},
{
name: "precise grid uppercase version matching host branch",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "575.57.08-GRID",
fallback: false,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid",
err: nil,
},
{
name: "precise grid-gcp uppercase version matching host branch",
gpuType: "NVIDIA_RTX_PRO_6000",
input: "575.57.08-GRID-GCP",
fallback: false,
hostDriverVersion: "575.10.10",
isVGPU: true,
expectedDriverVersion: "575.57.08-grid-gcp",
err: nil,
},
}
var gpuDriverVersionInfoList = &pb.GPUDriverVersionInfoList{}
err = (prototext.UnmarshalOptions{DiscardUnknown: true, AllowPartial: true}).Unmarshal(textProtoData, gpuDriverVersionInfoList)
if err != nil {
t.Fatalf("failed when parsing the test textproto data: %v", err)
}
testData, err := proto.Marshal(gpuDriverVersionInfoList)
if err != nil {
t.Fatalf("fail to encode to the binary proto data: %v", err)
}
for index, tc := range testCases {
t.Run(fmt.Sprintf("Test %d: %s", index, tc.name), func(t *testing.T) {
gpuDriverVersions, err := parseGPUDriverVersion(testData, tc.gpuType)
if err != nil {
t.Fatalf("Failed to parse GPU driver versions: %v", err)
}
gpuDriverVersions.hostDriverVersion = tc.hostDriverVersion
gpuDriverVersions.gpu.IsVGPU = tc.isVGPU
driverVersion, err := getVGPUVersion(gpuDriverVersions, tc.input, tc.fallback, gpuDriverVersions.gpu)
if tc.err == nil {
if err != nil {
t.Errorf("Failed to getVGPUVersion with error: %v", err)
}
if driverVersion != tc.expectedDriverVersion {
t.Errorf("Expected driver version %s, but got %s", tc.expectedDriverVersion, driverVersion)
}
} else {
if err == nil || !strings.Contains(err.Error(), tc.err.Error()) {
t.Errorf("Expected error containing '%s', but got: %v", tc.err.Error(), err)
}
}
})
}
}
func TestSortGridVersions(t *testing.T) {
testCases := []struct {
name string
input []string
expected []string
}{
{
name: "versions with leading zeroes sorted numerically",
input: []string{"580.126.09-grid", "580.95.05-grid", "580.159.03-grid-gcp", "580.82.07-grid"},
expected: []string{"580.82.07-grid", "580.95.05-grid", "580.126.09-grid", "580.159.03-grid-gcp"},
},
{
name: "same version prioritizes grid-gcp over grid",
input: []string{"580.126.09-grid-gcp", "580.126.09-grid"},
expected: []string{"580.126.09-grid", "580.126.09-grid-gcp"},
},
{
name: "case insensitivity with uppercase grid suffixes",
input: []string{"580.126.09-GRID", "580.95.05-GRID-GCP"},
expected: []string{"580.95.05-grid-gcp", "580.126.09-grid"},
},
{
name: "two-part and three-part version numbers",
input: []string{"560.38.01-grid", "560.38-grid"},
expected: []string{"560.38-grid", "560.38.01-grid"},
},
{
name: "filters out non-grid versions",
input: []string{"580.173.02", "580.126.09-grid-gcp", "580.159.04", "580.95.05-grid"},
expected: []string{"580.95.05-grid", "580.126.09-grid-gcp"},
},
}
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
actual := sortGridVersions(tc.input)
if diff := cmp.Diff(tc.expected, actual); diff != "" {
t.Errorf("sortGridVersions mismatch (-expected +actual):\n%s", diff)
}
})
}
}
func TestCompareNvidiaVersions(t *testing.T) {
testCases := []struct {
v1 string
v2 string
expected int
}{
{"580.126.09-grid", "580.95.05-grid", 1},
{"580.95.05-grid", "580.126.09-grid", -1},
{"580.126.09-grid", "580.126.09-grid", 0},
{"580.126.09-grid", "580.126.09-grid-gcp", -1},
{"580.126.09-grid-gcp", "580.126.09-grid", 1},
{"580.159.03-grid-gcp", "580.126.09-grid-gcp", 1},
{"580.82.07-grid", "580.95.05-grid", -1},
{"575.57.08-GRID", "575.57.08-grid", 0},
}
for _, tc := range testCases {
t.Run(fmt.Sprintf("%s_vs_%s", tc.v1, tc.v2), func(t *testing.T) {
actual := compareNvidiaVersions(tc.v1, tc.v2)
if actual != tc.expected {
t.Errorf("compareNvidiaVersions(%q, %q) = %d, expected %d", tc.v1, tc.v2, actual, tc.expected)
}
})
}
}