feat: add strict TOML configuration validation
This commit is contained in:
@@ -0,0 +1,288 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestAnnotatedExamples_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
serverData := readFixture(t, "server.toml")
|
||||
server, err := DecodeServer(serverData)
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeServer(example): %v", err)
|
||||
}
|
||||
if server.Queue.DefaultTTL != 15*time.Minute || server.Storage.TerminalRetention != 30*24*time.Hour {
|
||||
t.Fatalf("server defaults mismatch: TTL=%s retention=%s", server.Queue.DefaultTTL, server.Storage.TerminalRetention)
|
||||
}
|
||||
if server.JSONRPC.NonLoopbackBind {
|
||||
t.Fatal("disabled loopback JSON-RPC reported as external")
|
||||
}
|
||||
|
||||
clientData := readFixture(t, "client.toml")
|
||||
client, err := DecodeClient(clientData, ClientOptions{Platform: PlatformUnix, CheckFilesystem: true})
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeClient(example): %v", err)
|
||||
}
|
||||
if got := client.Shells.Advertised["bash"]; got != "/usr/bin/bash" && got != "/bin/bash" {
|
||||
t.Fatalf("advertised bash = %q", got)
|
||||
}
|
||||
if client.Client.MaxRunningCommands != 16 || client.Execution.WindowsTermGrace != 10*time.Second {
|
||||
t.Fatalf("client defaults mismatch: %+v", client.Client)
|
||||
}
|
||||
|
||||
windowsData := readFixture(t, "client.windows.toml")
|
||||
windowsClient, err := DecodeClient(windowsData, ClientOptions{Platform: PlatformWindows})
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeClient(windows example): %v", err)
|
||||
}
|
||||
if windowsClient.Client.StateDir != `C:\ProgramData\RVBox\state` {
|
||||
t.Fatalf("Windows state dir = %q", windowsClient.Client.StateDir)
|
||||
}
|
||||
if windowsClient.Observability.LivenessPath != "/livez" || windowsClient.Profiles.Light.CPUPercent != 50 {
|
||||
t.Fatal("omitted Windows entries did not receive compiled defaults")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStrictTOMLRejections_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
data []byte
|
||||
}{
|
||||
{"unknown key", []byte("[server]\nunknown = true\n")},
|
||||
{"duplicate key", []byte("[queue]\nmax_server = 2\nmax_server = 3\n")},
|
||||
{"duplicate table", []byte("[queue]\nmax_server = 2\n[queue]\nmax_per_client = 1\n")},
|
||||
{"bare duration", []byte("[queue]\ndefault_ttl = 15\n")},
|
||||
{"invalid utf8", []byte{0xff, 0xfe}},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
if _, err := DecodeServer(test.data); err == nil {
|
||||
t.Fatal("DecodeServer() succeeded")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExplicitFlagPrecedence_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server, err := DecodeServerWithOverrides(
|
||||
[]byte("[queue]\ndefault_ttl = \"20m\"\nmax_server = 9000\n"),
|
||||
map[string]string{"queue.default_ttl": "5m"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeServerWithOverrides(): %v", err)
|
||||
}
|
||||
if server.Queue.DefaultTTL != 5*time.Minute {
|
||||
t.Fatalf("explicit flag TTL = %s", server.Queue.DefaultTTL)
|
||||
}
|
||||
if server.Queue.MaxServer != 9000 {
|
||||
t.Fatalf("omitted flag overwrote TOML max_server: %d", server.Queue.MaxServer)
|
||||
}
|
||||
|
||||
client, err := DecodeClientWithOverrides(
|
||||
[]byte("[client]\nclient_id = \"from-file\"\nmax_running_commands = 12\n"),
|
||||
ClientOptions{Platform: PlatformUnix},
|
||||
map[string]string{"client.client_id": "from-flag"},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeClientWithOverrides(): %v", err)
|
||||
}
|
||||
if client.Client.ClientID != "from-flag" || client.Client.MaxRunningCommands != 12 {
|
||||
t.Fatalf("client precedence mismatch: %+v", client.Client)
|
||||
}
|
||||
|
||||
if _, err := DecodeServerWithOverrides(nil, map[string]string{"unknown": "x"}); err == nil {
|
||||
t.Fatal("unknown server override accepted")
|
||||
}
|
||||
if _, err := DecodeClientWithOverrides(nil, ClientOptions{Platform: PlatformUnix}, map[string]string{"client.max_running_commands": "not-a-number"}); err == nil {
|
||||
t.Fatal("invalid client override accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerBoundaryAndCrossFieldValidation_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
toml string
|
||||
}{
|
||||
{"negative duration", "[queue]\ndefault_ttl = \"-1s\"\n"},
|
||||
{"retry inversion", "[queue]\nretry_initial = \"31s\"\n"},
|
||||
{"queue inversion", "[queue]\nmax_per_client = 10001\n"},
|
||||
{"storage equality", "[storage]\ncommand_output_limit_bytes = 33554432\n"},
|
||||
{"closeout does not fit", "[storage]\ncommand_closeout_reserve_bytes = 33554432\n"},
|
||||
{"watermark inversion", "[flow]\nraw_output_low_bytes = 67108864\n"},
|
||||
{"heartbeat equality", "[protocol]\nheartbeat_idle = \"30s\"\n"},
|
||||
{"agent ceiling", "[protocol]\nmax_agent_envelope_bytes = 1048577\n"},
|
||||
{"zero nonoptional", "[storage]\naudit_limit_bytes = 0\n"},
|
||||
{"relative data path", "[server]\ndata_dir = \"relative\"\n"},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
if _, err := DecodeServer([]byte(test.toml)); err == nil {
|
||||
t.Fatal("DecodeServer() succeeded")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
for _, field := range []string{"default_ttl", "terminal_retention", "audit_retention"} {
|
||||
var source string
|
||||
switch field {
|
||||
case "default_ttl":
|
||||
source = "[queue]\ndefault_ttl = \"0s\"\n"
|
||||
case "terminal_retention":
|
||||
source = "[storage]\nterminal_retention = \"0s\"\n"
|
||||
default:
|
||||
source = "[storage]\naudit_retention = \"0s\"\n"
|
||||
}
|
||||
if _, err := DecodeServer([]byte(source)); err != nil {
|
||||
t.Errorf("documented zero for %s rejected: %v", field, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestJSONRPCExternalBindWarningState_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server, err := DecodeServer([]byte("[json_rpc]\nenabled = true\nlisten = \"0.0.0.0:6900\"\n"))
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeServer(): %v", err)
|
||||
}
|
||||
if !server.JSONRPC.NonLoopbackBind {
|
||||
t.Fatal("enabled external JSON-RPC did not request warning")
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientBoundaryShellAndProfileValidation_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
toml string
|
||||
platform Platform
|
||||
}{
|
||||
{"invalid URL", "[client]\nserver_url = \"http://example.test\"\n", PlatformUnix},
|
||||
{"invalid client id", "[client]\nclient_id = \"contains space\"\n", PlatformUnix},
|
||||
{"relative state", "[client]\nstate_dir = \"state\"\n", PlatformUnix},
|
||||
{"missing default shell", "[shells]\nsh = \"\"\n", PlatformUnix},
|
||||
{"inactive relative Windows shell", "[shells]\ncmd = \"cmd.exe\"\n", PlatformUnix},
|
||||
{"reconnect inversion", "[network]\nreconnect_initial = \"61s\"\n", PlatformUnix},
|
||||
{"command watermark inversion", "[flow]\nraw_output_command_low_bytes = 1048576\n", PlatformUnix},
|
||||
{"client watermark tier inversion", "[flow]\nraw_output_command_high_bytes = 8388608\n", PlatformUnix},
|
||||
{"execution ceiling", "[execution]\nmax_script_bytes = 10485761\n", PlatformUnix},
|
||||
{"profile missing control", "[profiles.cpu_medium]\nrequired_controls = [\"memory\"]\n", PlatformUnix},
|
||||
{"noncanonical device", "[profiles.disk_medium]\nlinux_io_read_bps = { \"08:0\" = 1 }\n", PlatformUnix},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
if _, err := DecodeClient([]byte(test.toml), ClientOptions{Platform: test.platform}); err == nil {
|
||||
t.Fatal("DecodeClient() succeeded")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestShellResolutionIgnoresRequestPATH_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, err := DecodeClient(nil, ClientOptions{Platform: PlatformUnix, CheckFilesystem: true})
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeClient(defaults): %v", err)
|
||||
}
|
||||
requestEnvironment := map[string]string{"PATH": "/tmp/untrusted"}
|
||||
_ = requestEnvironment
|
||||
resolved := client.Shells.Advertised["bash"]
|
||||
if resolved == "" || strings.HasPrefix(resolved, "/tmp/untrusted") {
|
||||
t.Fatalf("request PATH changed resolved shell to %q", resolved)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllowedCWDAndExecutableFilesystemChecks_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
regular := filepath.Join(tempDir, "not-executable")
|
||||
if err := os.WriteFile(regular, []byte("x"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
source := "[shells]\nsh = " + tomlString(regular) + "\nallowed_cwd_roots = [" + tomlString(regular) + "]\n"
|
||||
if _, err := DecodeClient([]byte(source), ClientOptions{Platform: PlatformUnix, CheckFilesystem: true}); err == nil {
|
||||
t.Fatal("non-executable shell/file CWD root accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileCombination_HP_CFG_01(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
valid := [][]string{{}, {"light"}, {"cpu_heavy", "mem_medium", "disk_medium"}}
|
||||
for _, names := range valid {
|
||||
if err := ValidateProfileCombination(names); err != nil {
|
||||
t.Errorf("ValidateProfileCombination(%v): %v", names, err)
|
||||
}
|
||||
}
|
||||
invalid := [][]string{{"light", "mem_heavy"}, {"cpu_medium", "cpu_heavy"}, {"mem_medium", "mem_medium"}, {"unknown"}}
|
||||
for _, names := range invalid {
|
||||
if err := ValidateProfileCombination(names); err == nil {
|
||||
t.Errorf("ValidateProfileCombination(%v) succeeded", names)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseDurationOverflowAndZero(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
if _, err := parseDuration("test", "999999999999999999999h", false); err == nil {
|
||||
t.Fatal("duration overflow accepted")
|
||||
}
|
||||
if value, err := parseDuration("test", "0s", true); err != nil || value != 0 {
|
||||
t.Fatalf("allowed zero = (%s, %v)", value, err)
|
||||
}
|
||||
if _, err := parseDuration("test", "0s", false); err == nil {
|
||||
t.Fatal("forbidden zero accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
if _, err := LoadServer(filepath.Join(t.TempDir(), "missing")); err == nil {
|
||||
t.Fatal("LoadServer missing file succeeded")
|
||||
}
|
||||
if _, err := LoadClient(filepath.Join(t.TempDir(), "missing"), PlatformUnix); err == nil {
|
||||
t.Fatal("LoadClient missing file succeeded")
|
||||
}
|
||||
}
|
||||
|
||||
func TestErrorValuesAreOrdinaryErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, err := DecodeServer([]byte("[server]\ndata_dir = \"relative\"\n"))
|
||||
if err == nil || errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("unexpected validation error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func readFixture(t *testing.T, name string) []byte {
|
||||
t.Helper()
|
||||
data, err := os.ReadFile(filepath.Join("..", "..", "docs", "examples", name))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
func tomlString(value string) string {
|
||||
return `"` + strings.ReplaceAll(value, `"`, `\"`) + `"`
|
||||
}
|
||||
Reference in New Issue
Block a user