diff --git a/pkg/config/engine/docker/docker.go b/pkg/config/engine/docker/docker.go index fb6f502f5..751df54bc 100644 --- a/pkg/config/engine/docker/docker.go +++ b/pkg/config/engine/docker/docker.go @@ -74,8 +74,8 @@ func (c *Config) AddRuntime(name string, path string, setAsDefault bool) error { // Read the existing runtimes runtimes := make(map[string]any) - if _, exists := config["runtimes"]; exists { - runtimes = config["runtimes"].(map[string]any) + if rt, ok := config["runtimes"].(map[string]any); ok { + runtimes = rt } // Add / update the runtime definitions @@ -128,16 +128,13 @@ func (c *Config) RemoveRuntime(name string) error { } config := *c - if _, exists := config["default-runtime"]; exists { - defaultRuntime := config["default-runtime"].(string) + if defaultRuntime, ok := config["default-runtime"].(string); ok { if defaultRuntime == name { config["default-runtime"] = defaultDockerRuntime } } - if _, exists := config["runtimes"]; exists { - runtimes := config["runtimes"].(map[string]any) - + if runtimes, ok := config["runtimes"].(map[string]any); ok { delete(runtimes, name) if len(runtimes) == 0 { @@ -171,8 +168,7 @@ func (c *Config) UpdateDefaultRuntime(name string, action string) error { if action == engine.UpdateActionSet { config["default-runtime"] = name } else { - if _, exists := config["default-runtime"]; exists { - defaultRuntime := config["default-runtime"].(string) + if defaultRuntime, ok := config["default-runtime"].(string); ok { if defaultRuntime == name { config["default-runtime"] = defaultDockerRuntime } @@ -202,11 +198,9 @@ func (c *Config) GetRuntimeConfig(name string) (engine.RuntimeConfig, error) { cfg := *c - var runtimes map[string]any - if _, ok := cfg["runtimes"]; ok { - runtimes = cfg["runtimes"].(map[string]any) - if r, ok := runtimes[name]; ok { - dr := dockerRuntime(r.(map[string]any)) + if runtimes, ok := cfg["runtimes"].(map[string]any); ok { + if r, ok := runtimes[name].(map[string]any); ok { + dr := dockerRuntime(r) return &dr, nil } } diff --git a/pkg/config/engine/docker/docker_test.go b/pkg/config/engine/docker/docker_test.go index b5ea75c03..12aa24fed 100644 --- a/pkg/config/engine/docker/docker_test.go +++ b/pkg/config/engine/docker/docker_test.go @@ -24,6 +24,8 @@ import ( "testing" "github.com/stretchr/testify/require" + + "github.com/NVIDIA/nvidia-container-toolkit/pkg/config/engine" ) func TestUpdateConfigDefaultRuntime(t *testing.T) { @@ -250,6 +252,88 @@ func TestGetRuntimeConfig(t *testing.T) { } } +func TestNullRuntimesFromFile(t *testing.T) { + for _, action := range []string{"add", "remove", "get"} { + t.Run(action, func(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "daemon.json") + require.NoError(t, os.WriteFile(configPath, []byte(`{"runtimes":null,"log-driver":"json-file"}`), 0600)) + config, err := New(WithPath(configPath)) + require.NoError(t, err) + + switch action { + case "add": + require.NoError(t, config.AddRuntime("nvidia", "/usr/bin/nvidia-container-runtime", false)) + runtime, err := config.GetRuntimeConfig("nvidia") + require.NoError(t, err) + require.Equal(t, "/usr/bin/nvidia-container-runtime", runtime.GetBinaryPath()) + case "remove": + require.NoError(t, config.RemoveRuntime("nvidia")) + case "get": + runtime, err := config.GetRuntimeConfig("nvidia") + require.NoError(t, err) + require.Empty(t, runtime.GetBinaryPath()) + } + + _, err = config.Save(configPath) + require.NoError(t, err) + contents, err := os.ReadFile(configPath) + require.NoError(t, err) + var saved map[string]any + require.NoError(t, json.Unmarshal(contents, &saved)) + require.Equal(t, "json-file", saved["log-driver"]) + }) + } +} + +func TestNullRuntimeFieldsFromFile(t *testing.T) { + testCases := map[string]struct { + input string + action string + expected string + }{ + "remove with null default runtime": { + input: `{"default-runtime":null,"runtimes":{"nvidia":{"path":"nvidia-container-runtime"}},"log-driver":"json-file"}`, + action: "remove", + expected: `{"default-runtime":null,"log-driver":"json-file"}`, + }, + "unset null default runtime": { + input: `{"default-runtime":null,"runtimes":{"nvidia":{"path":"nvidia-container-runtime"}},"log-driver":"json-file"}`, + action: "unset", + expected: `{"default-runtime":null,"runtimes":{"nvidia":{"path":"nvidia-container-runtime"}},"log-driver":"json-file"}`, + }, + "get null runtime definition": { + input: `{"runtimes":{"nvidia":null},"log-driver":"json-file"}`, + action: "get", + expected: `{"runtimes":{"nvidia":null},"log-driver":"json-file"}`, + }, + } + for name, tc := range testCases { + t.Run(name, func(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "daemon.json") + require.NoError(t, os.WriteFile(configPath, []byte(tc.input), 0o600)) + config, err := New(WithPath(configPath)) + require.NoError(t, err) + + switch tc.action { + case "remove": + require.NoError(t, config.RemoveRuntime("nvidia")) + case "unset": + require.NoError(t, config.UpdateDefaultRuntime("nvidia", engine.UpdateActionUnset)) + case "get": + runtime, err := config.GetRuntimeConfig("nvidia") + require.NoError(t, err) + require.Empty(t, runtime.GetBinaryPath()) + } + + _, err = config.Save(configPath) + require.NoError(t, err) + contents, err := os.ReadFile(configPath) + require.NoError(t, err) + require.JSONEq(t, tc.expected, string(contents)) + }) + } +} + func TestEnableCDIPreservesFeaturesFromFile(t *testing.T) { tests := []struct { name string