From 068b73e632ded4229e9431c0b24b2df94dff659b Mon Sep 17 00:00:00 2001 From: Daniel Lipovetsky Date: Tue, 22 Sep 2026 09:57:30 -0700 Subject: [PATCH] feat: Add a CLI to help prepare a precompiled driver image --- Makefile | 3 + cmd/prepareimage/main.go | 112 +++++++++++++++++++ cmd/prepareimage/output.go | 186 ++++++++++++++++++++++++++++++++ cmd/prepareimage/output_test.go | 108 +++++++++++++++++++ 4 files changed, 409 insertions(+) create mode 100644 cmd/prepareimage/main.go create mode 100644 cmd/prepareimage/output.go create mode 100644 cmd/prepareimage/output_test.go diff --git a/Makefile b/Makefile index dcd56aa17..ec9aa30ef 100644 --- a/Makefile +++ b/Makefile @@ -361,6 +361,9 @@ check-nfd-device-ids: ## Verify the AMD GPU PCI device-ID lists agree across the manager: $(shell find -name "*.go") go.mod go.sum ## Build manager binary (honors GOOS/GOARCH from the environment). go build -ldflags="-X main.Version=$(PROJECT_VERSION) -X main.GitCommit=$(GIT_COMMIT) -X main.BuildTag=$(HOURLY_TAG_LABEL)" -o $@ ./cmd +prepareimage: $(shell find -name "*.go") go.mod go.sum ## Build prepareimage binary (honors GOOS/GOARCH from the environment). + go build -ldflags="-X main.Version=$(PROJECT_VERSION) -X main.GitCommit=$(GIT_COMMIT) -X main.BuildTag=$(HOURLY_TAG_LABEL)" -o $@ ./cmd/prepareimage + # Build platform, default amd64. Set to a list (linux/amd64,linux/arm64) for multi-arch. PLATFORM ?= linux/amd64 diff --git a/cmd/prepareimage/main.go b/cmd/prepareimage/main.go new file mode 100644 index 000000000..30c266e84 --- /dev/null +++ b/cmd/prepareimage/main.go @@ -0,0 +1,112 @@ +package main + +import ( + "flag" + "fmt" + "log" + "os" + + workflowv1alpha1 "github.com/argoproj/argo-workflows/v4/pkg/apis/workflow/v1alpha1" + monitoringv1 "github.com/prometheus-operator/prometheus-operator/pkg/apis/monitoring/v1" + kmmv1beta1 "github.com/rh-ecosystem-edge/kernel-module-management/api/v1beta1" + apiextensionsv1 "k8s.io/apiextensions-apiserver/pkg/apis/apiextensions/v1" + "k8s.io/apimachinery/pkg/runtime" + utilruntime "k8s.io/apimachinery/pkg/util/runtime" + clientgoscheme "k8s.io/client-go/kubernetes/scheme" + + gpuev1alpha1 "github.com/ROCm/gpu-operator/api/v1alpha1" + utils "github.com/ROCm/gpu-operator/internal" +) + +const ( + DefaultRepoURL = "https://repo.radeon.com" + DefaultOutputFormat = OutputFormatDockerfile +) + +var ( + GitCommit = "undefined" + Version = "undefined" + BuildTag = "undefined" + scheme = runtime.NewScheme() +) + +func init() { + // Initialize the scheme + utilruntime.Must(clientgoscheme.AddToScheme(scheme)) + utilruntime.Must(gpuev1alpha1.AddToScheme(scheme)) + utilruntime.Must(kmmv1beta1.AddToScheme(scheme)) + utilruntime.Must(apiextensionsv1.AddToScheme(scheme)) + utilruntime.Must(monitoringv1.AddToScheme(scheme)) + utilruntime.Must(workflowv1alpha1.AddToScheme(scheme)) +} + +type Config struct { + OutputFormat OutputFormat + driverType string + osPrettyName string + imageName string + buildArgs BuildArgs +} + +type BuildArgs struct { + kernelFullVersion string + driversVersion string + repoURL string + packageRepoURL string + gpgKeyURL string +} + +type OutputFormat string + +const ( + OutputFormatDockerfile OutputFormat = "dockerfile" + OutputFormatDockerScript OutputFormat = "docker-script" + OutputFormatBuildahScript OutputFormat = "buildah-script" + OutputFormatPodmanScript OutputFormat = "podman-script" +) + +func (f *OutputFormat) String() string { + return string(*f) +} + +func (f *OutputFormat) Set(val string) error { + switch OutputFormat(val) { + case OutputFormatDockerfile, OutputFormatDockerScript, OutputFormatBuildahScript, OutputFormatPodmanScript: + *f = OutputFormat(val) + return nil + default: + return fmt.Errorf("invalid format %q (must be dockerfile, docker-script, buildah-script, or podman-script)", val) + } +} + +func main() { + config := deriveConfigFromFlags() + + output, err := createOutput(config) + if err != nil { + log.Fatalf("Failed to create output: %v", err) + } + + _, err = fmt.Fprint(os.Stdout, output) + if err != nil { + log.Fatalf("Failed to print output: %v", err) + } +} + +func deriveConfigFromFlags() *Config { + config := &Config{} + + config.OutputFormat = DefaultOutputFormat // Set the default value, because flag.Var() does not. + flag.Var(&config.OutputFormat, "format", "The format of the output to generate. Scripts embed the Dockerfile and execute the command to build the image using the correct name. Valid values are: dockerfile, docker-script, buildah-script, podman-script.") + flag.StringVar(&config.driverType, "driver-type", utils.DriverTypeContainer, "The type of the driver. Valid values are: container, vf-passthrough, pf-passthrough.") + flag.StringVar(&config.osPrettyName, "os-pretty-name", "", "The 'pretty name' of the OS. See the PRETTY_NAME variable defined in /etc/os-release. Example: 'Ubuntu 24.04.5 LTS'.") + flag.StringVar(&config.buildArgs.kernelFullVersion, "kernel-full-version", "", "The full version of the kernel. See the output of the 'uname -r' command. Example: '6.8.0-139-generic'.") + flag.StringVar(&config.buildArgs.driversVersion, "drivers-version", "", "The version of the drivers. For supported versions, see AMD documentation.") + flag.StringVar(&config.buildArgs.repoURL, "repo-url", DefaultRepoURL, "The URL for fetching the amdgpu installer.") + flag.StringVar(&config.buildArgs.packageRepoURL, "package-repo-url", "", "The full URL to the package repository for driver packages. When specified, this overrides the --repo-url flag. This is useful when using custom mirrors or when repo.radeon.com changes structure. Example: 'https://custom-mirror.example.com/amdgpu/30.20.1/ubuntu jammy main'.") + flag.StringVar(&config.buildArgs.gpgKeyURL, "gpg-key-url", "", "The full URL to the GPG key for package verification. When specified, this overrides the constrult GPG key URL. Example: 'https://custom-mirror.example.com/rocm/rocm.gpg.key'.") + flag.StringVar(&config.imageName, "image-name", "", "The name of the image with the registry and repository path, but without the tag. Example: 'registry.example.com/project/amdgpu'. Required for docker-script, buildah-script, and podman-script output formats.") + flag.Parse() + + return config +} diff --git a/cmd/prepareimage/output.go b/cmd/prepareimage/output.go new file mode 100644 index 000000000..e6832a546 --- /dev/null +++ b/cmd/prepareimage/output.go @@ -0,0 +1,186 @@ +package main + +import ( + "fmt" + "strings" + + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime" + + gpuev1alpha1 "github.com/ROCm/gpu-operator/api/v1alpha1" + "github.com/ROCm/gpu-operator/internal/kmmmodule" +) + +func createOutput(config *Config) (string, error) { + deviceConfig := deriveDeviceConfig(config) + + osName, err := deriveOSName(config.osPrettyName, deviceConfig) + if err != nil { + return "", fmt.Errorf("failed to derive OS name from OS pretty name %q: %w", config.osPrettyName, err) + } + + buildCM := deriveBuildConfigMap(osName, deviceConfig) + + dockerfile, err := deriveDockerfile(config, buildCM, deviceConfig, scheme) + if err != nil { + return "", fmt.Errorf("failed to derive Dockerfile: %w", err) + } + + var output string + switch config.OutputFormat { + case OutputFormatDockerfile: + output = dockerfile + case OutputFormatDockerScript, OutputFormatBuildahScript, OutputFormatPodmanScript: + if config.imageName == "" { + return "", fmt.Errorf("image name is required for docker-script, buildah-script, and podman-script output formats") + } + imageName := deriveImageNameWithTag(config, osName) + output = deriveScript(config.OutputFormat, imageName, dockerfile) + default: + return "", fmt.Errorf("unknown output format %q; valid values are: dockerfile, docker-script, buildah-script, podman-script", config.OutputFormat) + } + return output, nil +} + +func deriveDeviceConfig(config *Config) *gpuev1alpha1.DeviceConfig { + return &gpuev1alpha1.DeviceConfig{ + ObjectMeta: metav1.ObjectMeta{ + Name: "build", + }, + Spec: gpuev1alpha1.DeviceConfigSpec{ + Driver: gpuev1alpha1.DriverSpec{ + DriverType: config.driverType, + Version: config.buildArgs.driversVersion, + AMDGPUInstallerRepoURL: config.buildArgs.repoURL, + ImageBuild: gpuev1alpha1.ImageBuildSpec{ + PackageRepoURL: config.buildArgs.packageRepoURL, + GPGKeyURL: config.buildArgs.gpgKeyURL, + }, + }, + }, + } +} + +func deriveBuildConfigMap(osName string, deviceConfig *gpuev1alpha1.DeviceConfig) *corev1.ConfigMap { + return &corev1.ConfigMap{ + ObjectMeta: metav1.ObjectMeta{ + Name: kmmmodule.GetCMName(osName, deviceConfig), + }, + } +} + +func deriveDockerfile(config *Config, buildCM *corev1.ConfigMap, deviceConfig *gpuev1alpha1.DeviceConfig, scheme *runtime.Scheme) (string, error) { + kmmModule := kmmmodule.NewKMMModule( + nil, // The client is not used by any of the kmmModule methods we call. + scheme, // A correctly initialized scheme is required by SetBuildConfigMapAsDesired. + false, // Dockerfiles for all supported OSes can be derived without this input. + ) + + err := kmmModule.SetBuildConfigMapAsDesired(buildCM, deviceConfig) + if err != nil { + return "", fmt.Errorf("failed to set BuildConfigMap as desired: %v", err) + } + + dockerfile, ok := buildCM.Data["dockerfile"] + if !ok { + return "", fmt.Errorf("failed to get Dockerfile from BuildConfigMap") + } + + return setBuildArgDefaults(config, dockerfile), nil +} + +func setBuildArgDefaults(config *Config, dockerfile string) string { + buildArgReplacements := []string{} + + if config.buildArgs.kernelFullVersion != "" { + buildArgReplacements = append( + buildArgReplacements, + "ARG KERNEL_FULL_VERSION", + fmt.Sprintf("ARG KERNEL_FULL_VERSION=%s", config.buildArgs.kernelFullVersion), + ) + + // OpenShift Dockerfile templates use KERNEL_VERSION instead of KERNEL_FULL_VERSION, + // but the accepted values should be identical. + buildArgReplacements = append( + buildArgReplacements, + "ARG KERNEL_VERSION", + fmt.Sprintf("ARG KERNEL_VERSION=%s", config.buildArgs.kernelFullVersion), + ) + } + if config.buildArgs.driversVersion != "" { + buildArgReplacements = append( + buildArgReplacements, + "ARG DRIVERS_VERSION", + fmt.Sprintf("ARG DRIVERS_VERSION=%s", config.buildArgs.driversVersion), + ) + } + if config.buildArgs.repoURL != "" { + buildArgReplacements = append( + buildArgReplacements, + "ARG REPO_URL", + fmt.Sprintf("ARG REPO_URL=%s", config.buildArgs.repoURL), + ) + } + if config.buildArgs.packageRepoURL != "" { + buildArgReplacements = append( + buildArgReplacements, + "ARG PACKAGE_REPO_URL", + fmt.Sprintf("ARG PACKAGE_REPO_URL=%s", config.buildArgs.packageRepoURL), + ) + } + if config.buildArgs.gpgKeyURL != "" { + buildArgReplacements = append( + buildArgReplacements, + "ARG GPG_KEY_URL", + fmt.Sprintf("ARG GPG_KEY_URL=%s", config.buildArgs.gpgKeyURL), + ) + } + + return strings.NewReplacer(buildArgReplacements...).Replace(dockerfile) +} + +func deriveOSName(osPrettyName string, deviceConfig *gpuev1alpha1.DeviceConfig) (string, error) { + node := corev1.Node{ + Status: corev1.NodeStatus{ + NodeInfo: corev1.NodeSystemInfo{ + OSImage: osPrettyName, + }, + }, + } + return kmmmodule.GetOSName(node, deviceConfig) +} + +func deriveImageTag(osName, kernelVersion, driversVersion string) string { + return fmt.Sprintf("%s-%s-%s", osName, kernelVersion, driversVersion) +} + +func deriveImageNameWithTag(config *Config, osName string) string { + tag := deriveImageTag(osName, config.buildArgs.kernelFullVersion, config.buildArgs.driversVersion) + return fmt.Sprintf("%s:%s", config.imageName, tag) +} + +func deriveScript(outputFormat OutputFormat, imageName string, dockerfile string) string { + switch outputFormat { + case OutputFormatDockerScript: + return fmt.Sprintf(`#!/bin/sh +docker buildx build -t %s - <<'EOF' +%s +EOF +`, imageName, dockerfile) + case OutputFormatBuildahScript: + return fmt.Sprintf(`#!/bin/sh +buildah bud -t %s - <<'EOF' +%s +EOF +`, imageName, dockerfile) + case OutputFormatPodmanScript: + return fmt.Sprintf(`#!/bin/sh +podman buildx build -t %s - <<'EOF' +%s +EOF +`, imageName, dockerfile) + default: + return "" + } +} diff --git a/cmd/prepareimage/output_test.go b/cmd/prepareimage/output_test.go new file mode 100644 index 000000000..fcaeeb2c0 --- /dev/null +++ b/cmd/prepareimage/output_test.go @@ -0,0 +1,108 @@ +package main + +import ( + "strings" + "testing" + + utils "github.com/ROCm/gpu-operator/internal" +) + +func TestCreateOutput(t *testing.T) { + tests := []struct { + name string + config *Config + wantErrSubstring string + wantSubstrings []string + unwantSubstrings []string + }{ + { + name: "dockerfile", + config: &Config{ + OutputFormat: OutputFormatDockerfile, + driverType: utils.DriverTypeContainer, + osPrettyName: "Ubuntu 24.04.5 LTS", + imageName: "registry.example.com/project/amdgpu", + buildArgs: BuildArgs{ + kernelFullVersion: "6.8.0-139-generic", + driversVersion: "31.40.1", + }, + }, + wantSubstrings: []string{ + "FROM docker.io/ubuntu:24.04", + "ARG KERNEL_FULL_VERSION=6.8.0-139-generic", + "ARG DRIVERS_VERSION=31.40.1", + }, + unwantSubstrings: []string{ + "#!/bin/sh", + }, + }, + { + name: "docker-script", + config: &Config{ + OutputFormat: OutputFormatDockerScript, + driverType: utils.DriverTypeContainer, + osPrettyName: "Ubuntu 24.04.5 LTS", + imageName: "registry.example.com/project/amdgpu", + buildArgs: BuildArgs{ + kernelFullVersion: "6.8.0-139-generic", + driversVersion: "31.40.1", + }, + }, + wantSubstrings: []string{ + "#!/bin/sh", + "docker buildx build -t registry.example.com/project/amdgpu:ubuntu-24.04-6.8.0-139-generic-31.40.1 - <<'EOF'", + "FROM docker.io/ubuntu:24.04", + "ARG KERNEL_FULL_VERSION=6.8.0-139-generic", + "ARG DRIVERS_VERSION=31.40.1", + "EOF", + }, + }, + { + name: "unknown OS", + config: &Config{ + OutputFormat: OutputFormatDockerfile, + driverType: utils.DriverTypeContainer, + osPrettyName: "foo", + imageName: "registry.example.com/project/amdgpu", + buildArgs: BuildArgs{ + kernelFullVersion: "6.8.0-139-generic", + driversVersion: "31.40.1", + }, + }, + wantErrSubstring: "not supported", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + output, err := createOutput(tt.config) + + if tt.wantErrSubstring != "" { + if err == nil { + t.Fatalf("createOutput returned nil error; got output:\n%s", output) + } + if !strings.Contains(err.Error(), tt.wantErrSubstring) { + t.Errorf("error %q does not contain %q", err, tt.wantErrSubstring) + } + if output != "" { + t.Errorf("expected empty output on error, got %q", output) + } + return + } + + if err != nil { + t.Fatalf("createOutput returned unexpected error: %v", err) + } + for _, want := range tt.wantSubstrings { + if !strings.Contains(output, want) { + t.Errorf("output missing %q; got:\n%s", want, output) + } + } + for _, unwant := range tt.unwantSubstrings { + if strings.Contains(output, unwant) { + t.Errorf("output unexpectedly contains %q; got:\n%s", unwant, output) + } + } + }) + } +}