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, `"`, `\"`) + `"`
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"os"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/pelletier/go-toml/v2"
|
||||
)
|
||||
|
||||
// ClientOptions controls platform-specific static validation. Production uses
|
||||
// CheckFilesystem; documentation tests may perform syntax-only validation for
|
||||
// a non-native target.
|
||||
type ClientOptions struct {
|
||||
Platform Platform
|
||||
CheckFilesystem bool
|
||||
}
|
||||
|
||||
func LoadServer(path string) (*Server, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read server config: %w", err)
|
||||
}
|
||||
return DecodeServer(data)
|
||||
}
|
||||
|
||||
func DecodeServer(data []byte) (*Server, error) {
|
||||
return DecodeServerWithOverrides(data, nil)
|
||||
}
|
||||
|
||||
// DecodeServerWithOverrides applies only explicitly present flag values after
|
||||
// TOML decoding. Keys come from the single registry in overrides.go.
|
||||
func DecodeServerWithOverrides(data []byte, overrides map[string]string) (*Server, error) {
|
||||
var raw = defaultServerFile()
|
||||
if err := strictDecode(data, &raw); err != nil {
|
||||
return nil, fmt.Errorf("decode server config: %w", err)
|
||||
}
|
||||
if err := applyServerOverrides(&raw, overrides); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return validateServer(raw)
|
||||
}
|
||||
|
||||
func LoadClient(path string, platform Platform) (*Client, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read client config: %w", err)
|
||||
}
|
||||
return DecodeClient(data, ClientOptions{Platform: platform, CheckFilesystem: true})
|
||||
}
|
||||
|
||||
func DecodeClient(data []byte, options ClientOptions) (*Client, error) {
|
||||
return DecodeClientWithOverrides(data, options, nil)
|
||||
}
|
||||
|
||||
// DecodeClientWithOverrides applies explicitly present flags after the file.
|
||||
func DecodeClientWithOverrides(data []byte, options ClientOptions, overrides map[string]string) (*Client, error) {
|
||||
var raw = defaultClientFile()
|
||||
if err := strictDecode(data, &raw); err != nil {
|
||||
return nil, fmt.Errorf("decode client config: %w", err)
|
||||
}
|
||||
if err := applyClientOverrides(&raw, overrides); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return validateClient(raw, options)
|
||||
}
|
||||
|
||||
func strictDecode(data []byte, target any) error {
|
||||
if !utf8.Valid(data) {
|
||||
return fmt.Errorf("configuration is not valid UTF-8")
|
||||
}
|
||||
decoder := toml.NewDecoder(bytes.NewReader(data))
|
||||
decoder.DisallowUnknownFields()
|
||||
if err := decoder.Decode(target); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseDuration(name, value string, allowZero bool) (time.Duration, error) {
|
||||
duration, err := time.ParseDuration(value)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
if duration < 0 || (!allowZero && duration == 0) {
|
||||
return 0, fmt.Errorf("%s must be %s", name, map[bool]string{true: "non-negative", false: "positive"}[allowZero])
|
||||
}
|
||||
return duration, nil
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package config
|
||||
|
||||
const (
|
||||
HardMaxAgentEnvelopeBytes uint64 = 1 << 20
|
||||
HardMaxExecutionSpecBytes uint64 = 768 << 10
|
||||
HardMaxRawChunkBytes uint64 = 64 << 10
|
||||
HardMaxScriptBytes uint64 = 10 << 20
|
||||
HardMaxControlRequestBytes uint64 = 16 << 20
|
||||
HardMaxJSONRPCBodyBytes uint64 = 24 << 20
|
||||
HardMaxProtocolDetailBytes uint64 = 4 << 10
|
||||
)
|
||||
|
||||
func defaultServerFile() serverFile {
|
||||
return serverFile{
|
||||
Server: serverCoreFile{
|
||||
DataDir: "/var/lib/rvbox-server", AgentListen: "127.0.0.1:6899",
|
||||
AgentPath: "/v1/agent", ControlSocket: "/run/rvbox/server.sock", ShutdownGrace: "30s",
|
||||
},
|
||||
JSONRPC: jsonRPCFile{Listen: "127.0.0.1:6900"},
|
||||
Queue: serverQueueFile{DefaultTTL: "15m", MaxPerClient: 1000, MaxServer: 10000, RetryInitial: "1s", RetryMax: "30s"},
|
||||
Storage: serverStorageFile{
|
||||
CommandOutputLimitBytes: 10 << 20, CommandTotalLimitBytes: 32 << 20,
|
||||
ClientTotalLimitBytes: 256 << 20, ServerTotalLimitBytes: 4 << 30,
|
||||
TerminalRetention: "720h", AuditLimitBytes: 100 << 20, AuditRetention: "0s",
|
||||
TombstoneMaxEntries: 1_000_000, CommandCloseoutReserveBytes: 64 << 10,
|
||||
FreeSpaceFloorBytes: 256 << 20, SegmentTargetBytes: 256 << 10,
|
||||
DurabilityInterval: "100ms", SQLiteBusyTimeout: "5s",
|
||||
IncidentNoteMaxBytes: 4 << 10, ProtocolDetailMaxBytes: 4 << 10,
|
||||
},
|
||||
Flow: serverFlowFile{
|
||||
RawOutputHighBytes: 64 << 20, RawOutputLowBytes: 32 << 20,
|
||||
UnacknowledgedPerCommandBytes: 1 << 20, UnacknowledgedPerSessionBytes: 8 << 20,
|
||||
WriteDeadline: "10s",
|
||||
},
|
||||
Protocol: serverProtocolFile{
|
||||
HeartbeatIdle: "10s", LivenessTimeout: "30s", TakeoverTTL: "5m",
|
||||
MaxAgentEnvelopeBytes: 1 << 20, MaxExecutionSpecBytes: 768 << 10,
|
||||
MaxRawChunkBytes: 64 << 10, MaxScriptBytes: 10 << 20,
|
||||
MaxControlRequestBytes: 16 << 20, MaxJSONRPCBodyBytes: 24 << 20,
|
||||
},
|
||||
Observability: observabilityFile{
|
||||
Listen: "127.0.0.1:6901", LivenessPath: "/livez", ReadinessPath: "/readyz",
|
||||
MetricsPath: "/metrics", LogLevel: "info", LogFormat: "json",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func defaultClientFile() clientFile {
|
||||
return clientFile{
|
||||
Client: clientCoreFile{
|
||||
ServerURL: "wss://rvbox.example.test/v1/agent", StateDir: "/var/lib/rvbox",
|
||||
DaemonCWD: "/", MaxRunningCommands: 16, MaxQueuedCommands: 100, ShutdownGrace: "30s",
|
||||
},
|
||||
Shells: shellsFile{
|
||||
DefaultUnix: "sh", DefaultWindows: "powershell", SH: "/bin/sh", Bash: "/bin/bash",
|
||||
CMD: `C:\Windows\System32\cmd.exe`,
|
||||
PowerShell: `C:\Windows\System32\WindowsPowerShell\v1.0\powershell.exe`,
|
||||
AllowedCWDRoots: []string{},
|
||||
},
|
||||
Network: clientNetworkFile{
|
||||
HeartbeatIdle: "10s", LivenessTimeout: "30s", ReconnectInitial: "1s", ReconnectMax: "60s",
|
||||
StableSessionReset: "60s", ConnectTimeout: "15s", WriteDeadline: "10s",
|
||||
},
|
||||
Storage: clientStorageFile{
|
||||
CommandOutputLimitBytes: 10 << 20, CommandTotalLimitBytes: 32 << 20,
|
||||
ClientTotalLimitBytes: 256 << 20, TombstoneMaxEntries: 1_000_000,
|
||||
CommandCloseoutReserveBytes: 64 << 10, FreeSpaceFloorBytes: 64 << 20,
|
||||
SegmentTargetBytes: 256 << 10, DurabilityInterval: "100ms",
|
||||
},
|
||||
Flow: clientFlowFile{
|
||||
RawOutputCommandHighBytes: 1 << 20, RawOutputCommandLowBytes: 256 << 10,
|
||||
RawOutputClientHighBytes: 8 << 20, RawOutputClientLowBytes: 4 << 20,
|
||||
UnacknowledgedPerCommandBytes: 1 << 20, UnacknowledgedPerSessionBytes: 8 << 20,
|
||||
},
|
||||
Execution: executionFile{
|
||||
DescendantDrainGrace: "5s", WindowsTermGrace: "10s", HungThreshold: "10m", DiagnosticInterval: "30s",
|
||||
MaxScriptBytes: 10 << 20, MaxExecutionSpecBytes: 768 << 10,
|
||||
MaxAgentEnvelopeBytes: 1 << 20, MaxRawChunkBytes: 64 << 10, ProtocolDetailMaxBytes: 4 << 10,
|
||||
},
|
||||
Observability: observabilityFile{
|
||||
Listen: "127.0.0.1:6902", LivenessPath: "/livez", ReadinessPath: "/readyz",
|
||||
MetricsPath: "/metrics", LogLevel: "info", LogFormat: "json",
|
||||
LogMaxBytes: 10 << 20, LogMaxFiles: 5,
|
||||
},
|
||||
Profiles: defaultProfilesFile(),
|
||||
}
|
||||
}
|
||||
|
||||
func defaultProfilesFile() profilesFile {
|
||||
return profilesFile{
|
||||
Light: profileFile{Enabled: true, RequiredControls: []string{"cpu", "memory", "pids"}, CPUPercent: 50, MemoryMaxBytes: 512 << 20, PIDsMax: 64, LinuxIOReadBPS: map[string]uint64{}, LinuxIOWriteBPS: map[string]uint64{}},
|
||||
CPUMedium: profileFile{Enabled: true, RequiredControls: []string{"cpu"}, CPUPercent: 200},
|
||||
CPUHeavy: profileFile{Enabled: true, RequiredControls: []string{"cpu"}, CPUPercent: 800},
|
||||
MemMedium: profileFile{Enabled: true, RequiredControls: []string{"memory"}, MemoryMaxBytes: 2 << 30},
|
||||
MemHeavy: profileFile{Enabled: true, RequiredControls: []string{"memory"}, MemoryMaxBytes: 8 << 30},
|
||||
DiskMedium: profileFile{RequiredControls: []string{"io"}, WindowsIOReadBPS: 100 << 20, WindowsIOWriteBPS: 50 << 20, LinuxIOReadBPS: map[string]uint64{"8:0": 100 << 20}, LinuxIOWriteBPS: map[string]uint64{"8:0": 50 << 20}},
|
||||
DiskHeavy: profileFile{RequiredControls: []string{"io"}, WindowsIOReadBPS: 500 << 20, WindowsIOWriteBPS: 250 << 20, LinuxIOReadBPS: map[string]uint64{"8:0": 500 << 20}, LinuxIOWriteBPS: map[string]uint64{"8:0": 250 << 20}},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
type serverOverride func(*serverFile, string) error
|
||||
type clientOverride func(*clientFile, string) error
|
||||
|
||||
// These registries are the authority used later to generate daemon flag
|
||||
// bindings. A value is applied only when its key is present in the caller map.
|
||||
var serverOverrideRegistry = map[string]serverOverride{
|
||||
"server.data_dir": func(c *serverFile, v string) error { c.Server.DataDir = v; return nil },
|
||||
"server.agent_listen": func(c *serverFile, v string) error { c.Server.AgentListen = v; return nil },
|
||||
"server.agent_path": func(c *serverFile, v string) error { c.Server.AgentPath = v; return nil },
|
||||
"server.control_socket": func(c *serverFile, v string) error { c.Server.ControlSocket = v; return nil },
|
||||
"server.shutdown_grace": func(c *serverFile, v string) error { c.Server.ShutdownGrace = v; return nil },
|
||||
"json_rpc.enabled": func(c *serverFile, v string) error { return parseBoolOverride(v, &c.JSONRPC.Enabled) },
|
||||
"json_rpc.listen": func(c *serverFile, v string) error { c.JSONRPC.Listen = v; return nil },
|
||||
"queue.default_ttl": func(c *serverFile, v string) error { c.Queue.DefaultTTL = v; return nil },
|
||||
"queue.max_per_client": func(c *serverFile, v string) error { return parseUint32Override(v, &c.Queue.MaxPerClient) },
|
||||
"queue.max_server": func(c *serverFile, v string) error { return parseUint32Override(v, &c.Queue.MaxServer) },
|
||||
"protocol.heartbeat_idle": func(c *serverFile, v string) error { c.Protocol.HeartbeatIdle = v; return nil },
|
||||
"protocol.liveness_timeout": func(c *serverFile, v string) error { c.Protocol.LivenessTimeout = v; return nil },
|
||||
"protocol.takeover_ttl": func(c *serverFile, v string) error { c.Protocol.TakeoverTTL = v; return nil },
|
||||
"observability.listen": func(c *serverFile, v string) error { c.Observability.Listen = v; return nil },
|
||||
"observability.log_level": func(c *serverFile, v string) error { c.Observability.LogLevel = v; return nil },
|
||||
"observability.log_format": func(c *serverFile, v string) error { c.Observability.LogFormat = v; return nil },
|
||||
}
|
||||
|
||||
var clientOverrideRegistry = map[string]clientOverride{
|
||||
"client.server_url": func(c *clientFile, v string) error { c.Client.ServerURL = v; return nil },
|
||||
"client.state_dir": func(c *clientFile, v string) error { c.Client.StateDir = v; return nil },
|
||||
"client.client_id": func(c *clientFile, v string) error { c.Client.ClientID = v; return nil },
|
||||
"client.daemon_cwd": func(c *clientFile, v string) error { c.Client.DaemonCWD = v; return nil },
|
||||
"client.max_running_commands": func(c *clientFile, v string) error { return parseUint32Override(v, &c.Client.MaxRunningCommands) },
|
||||
"client.max_queued_commands": func(c *clientFile, v string) error { return parseUint32Override(v, &c.Client.MaxQueuedCommands) },
|
||||
"client.shutdown_grace": func(c *clientFile, v string) error { c.Client.ShutdownGrace = v; return nil },
|
||||
"tls.ca_file": func(c *clientFile, v string) error { c.TLS.CAFile = v; return nil },
|
||||
"tls.server_name": func(c *clientFile, v string) error { c.TLS.ServerName = v; return nil },
|
||||
"shells.default_unix": func(c *clientFile, v string) error { c.Shells.DefaultUnix = v; return nil },
|
||||
"shells.default_windows": func(c *clientFile, v string) error { c.Shells.DefaultWindows = v; return nil },
|
||||
"network.reconnect_initial": func(c *clientFile, v string) error { c.Network.ReconnectInitial = v; return nil },
|
||||
"network.reconnect_max": func(c *clientFile, v string) error { c.Network.ReconnectMax = v; return nil },
|
||||
"observability.listen": func(c *clientFile, v string) error { c.Observability.Listen = v; return nil },
|
||||
"observability.log_level": func(c *clientFile, v string) error { c.Observability.LogLevel = v; return nil },
|
||||
"observability.log_format": func(c *clientFile, v string) error { c.Observability.LogFormat = v; return nil },
|
||||
"observability.log_file": func(c *clientFile, v string) error { c.Observability.LogFile = v; return nil },
|
||||
}
|
||||
|
||||
func applyServerOverrides(config *serverFile, overrides map[string]string) error {
|
||||
for key, value := range overrides {
|
||||
setter, exists := serverOverrideRegistry[key]
|
||||
if !exists {
|
||||
return fmt.Errorf("unknown server flag override %q", key)
|
||||
}
|
||||
if err := setter(config, value); err != nil {
|
||||
return fmt.Errorf("server flag override %s: %w", key, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func applyClientOverrides(config *clientFile, overrides map[string]string) error {
|
||||
for key, value := range overrides {
|
||||
setter, exists := clientOverrideRegistry[key]
|
||||
if !exists {
|
||||
return fmt.Errorf("unknown client flag override %q", key)
|
||||
}
|
||||
if err := setter(config, value); err != nil {
|
||||
return fmt.Errorf("client flag override %s: %w", key, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseBoolOverride(value string, target *bool) error {
|
||||
parsed, err := strconv.ParseBool(value)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*target = parsed
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseUint32Override(value string, target *uint32) error {
|
||||
parsed, err := strconv.ParseUint(value, 10, 32)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*target = uint32(parsed)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func validateAbsolutePath(name, value string, platform Platform, allowEmpty bool) (string, error) {
|
||||
if value == "" && allowEmpty {
|
||||
return "", nil
|
||||
}
|
||||
if value == "" || strings.IndexByte(value, 0) >= 0 {
|
||||
return "", fmt.Errorf("%s must be a nonempty absolute path", name)
|
||||
}
|
||||
if platform == PlatformWindows {
|
||||
if !isWindowsAbsolute(value) {
|
||||
return "", fmt.Errorf("%s must be an absolute Windows path", name)
|
||||
}
|
||||
return cleanWindowsPath(value), nil
|
||||
}
|
||||
if !filepath.IsAbs(value) {
|
||||
return "", fmt.Errorf("%s must be an absolute Unix path", name)
|
||||
}
|
||||
return filepath.Clean(value), nil
|
||||
}
|
||||
|
||||
func isWindowsAbsolute(value string) bool {
|
||||
if strings.HasPrefix(value, `\\`) {
|
||||
parts := strings.Split(strings.TrimPrefix(value, `\\`), `\`)
|
||||
return len(parts) >= 2 && parts[0] != "" && parts[1] != ""
|
||||
}
|
||||
return len(value) >= 3 && ((value[0] >= 'A' && value[0] <= 'Z') || (value[0] >= 'a' && value[0] <= 'z')) &&
|
||||
value[1] == ':' && (value[2] == '\\' || value[2] == '/')
|
||||
}
|
||||
|
||||
func cleanWindowsPath(value string) string {
|
||||
value = strings.ReplaceAll(value, "/", `\`)
|
||||
prefix := value[:3]
|
||||
rest := strings.TrimPrefix(value[3:], `\`)
|
||||
if strings.HasPrefix(value, `\\`) {
|
||||
components := strings.Split(strings.TrimPrefix(value, `\\`), `\`)
|
||||
prefix = `\\` + components[0] + `\` + components[1] + `\`
|
||||
rest = strings.Join(components[2:], `\`)
|
||||
}
|
||||
parts := make([]string, 0)
|
||||
for _, part := range strings.Split(rest, `\`) {
|
||||
switch part {
|
||||
case "", ".":
|
||||
continue
|
||||
case "..":
|
||||
if len(parts) > 0 {
|
||||
parts = parts[:len(parts)-1]
|
||||
}
|
||||
default:
|
||||
parts = append(parts, part)
|
||||
}
|
||||
}
|
||||
return prefix + strings.Join(parts, `\`)
|
||||
}
|
||||
|
||||
func checkNativeExecutable(name, value string, platform Platform) (string, error) {
|
||||
if (platform == PlatformWindows) != (runtime.GOOS == "windows") {
|
||||
return "", fmt.Errorf("%s filesystem validation requires a native %s process", name, map[bool]string{true: "Windows", false: "Unix"}[platform == PlatformWindows])
|
||||
}
|
||||
canonical, err := filepath.Abs(value)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
if evaluated, evalErr := filepath.EvalSymlinks(canonical); evalErr == nil {
|
||||
canonical = evaluated
|
||||
}
|
||||
info, err := os.Stat(canonical)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return "", fmt.Errorf("%s must name a regular file", name)
|
||||
}
|
||||
if platform == PlatformUnix && info.Mode().Perm()&0o111 == 0 {
|
||||
return "", fmt.Errorf("%s is not executable", name)
|
||||
}
|
||||
if platform == PlatformWindows {
|
||||
extension := strings.ToLower(filepath.Ext(canonical))
|
||||
if extension != ".exe" && extension != ".com" && extension != ".cmd" && extension != ".bat" {
|
||||
return "", fmt.Errorf("%s is not a recognized Windows executable", name)
|
||||
}
|
||||
}
|
||||
return canonical, nil
|
||||
}
|
||||
|
||||
func checkNativeDirectory(name, value string, platform Platform) (string, error) {
|
||||
if (platform == PlatformWindows) != (runtime.GOOS == "windows") {
|
||||
return "", fmt.Errorf("%s filesystem validation requires a native target", name)
|
||||
}
|
||||
canonical, err := filepath.Abs(value)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
if evaluated, evalErr := filepath.EvalSymlinks(canonical); evalErr == nil {
|
||||
canonical = evaluated
|
||||
}
|
||||
info, err := os.Stat(canonical)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
if !info.IsDir() {
|
||||
return "", fmt.Errorf("%s must name a directory", name)
|
||||
}
|
||||
return canonical, nil
|
||||
}
|
||||
|
||||
func checkNativeRegularFile(name, value string, platform Platform) (string, error) {
|
||||
if (platform == PlatformWindows) != (runtime.GOOS == "windows") {
|
||||
return "", fmt.Errorf("%s filesystem validation requires a native target", name)
|
||||
}
|
||||
canonical, err := filepath.Abs(value)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
if evaluated, evalErr := filepath.EvalSymlinks(canonical); evalErr == nil {
|
||||
canonical = evaluated
|
||||
}
|
||||
info, err := os.Stat(canonical)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return "", fmt.Errorf("%s must name a regular file", name)
|
||||
}
|
||||
return canonical, nil
|
||||
}
|
||||
@@ -0,0 +1,377 @@
|
||||
// Package config owns strict TOML decoding, defaulting, and static validation.
|
||||
package config
|
||||
|
||||
import "time"
|
||||
|
||||
type Platform uint8
|
||||
|
||||
const (
|
||||
PlatformUnix Platform = iota + 1
|
||||
PlatformWindows
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
Server ServerCore
|
||||
JSONRPC JSONRPC
|
||||
Queue ServerQueue
|
||||
Storage ServerStorage
|
||||
Flow ServerFlow
|
||||
Protocol ServerProtocol
|
||||
Observability Observability
|
||||
}
|
||||
|
||||
type ServerCore struct {
|
||||
DataDir string
|
||||
AgentListen string
|
||||
AgentPath string
|
||||
ControlSocket string
|
||||
ShutdownGrace time.Duration
|
||||
}
|
||||
|
||||
type JSONRPC struct {
|
||||
Enabled bool
|
||||
Listen string
|
||||
NonLoopbackBind bool
|
||||
}
|
||||
|
||||
type ServerQueue struct {
|
||||
DefaultTTL time.Duration
|
||||
MaxPerClient uint32
|
||||
MaxServer uint32
|
||||
RetryInitial time.Duration
|
||||
RetryMax time.Duration
|
||||
}
|
||||
|
||||
type ServerStorage struct {
|
||||
CommandOutputLimitBytes uint64
|
||||
CommandTotalLimitBytes uint64
|
||||
ClientTotalLimitBytes uint64
|
||||
ServerTotalLimitBytes uint64
|
||||
TerminalRetention time.Duration
|
||||
AuditLimitBytes uint64
|
||||
AuditRetention time.Duration
|
||||
TombstoneMaxEntries uint64
|
||||
CommandCloseoutReserveBytes uint64
|
||||
FreeSpaceFloorBytes uint64
|
||||
SegmentTargetBytes uint64
|
||||
DurabilityInterval time.Duration
|
||||
SQLiteBusyTimeout time.Duration
|
||||
IncidentNoteMaxBytes uint64
|
||||
ProtocolDetailMaxBytes uint64
|
||||
}
|
||||
|
||||
type ServerFlow struct {
|
||||
RawOutputHighBytes uint64
|
||||
RawOutputLowBytes uint64
|
||||
UnacknowledgedPerCommandBytes uint64
|
||||
UnacknowledgedPerSessionBytes uint64
|
||||
WriteDeadline time.Duration
|
||||
}
|
||||
|
||||
type ServerProtocol struct {
|
||||
HeartbeatIdle time.Duration
|
||||
LivenessTimeout time.Duration
|
||||
TakeoverTTL time.Duration
|
||||
MaxAgentEnvelopeBytes uint64
|
||||
MaxExecutionSpecBytes uint64
|
||||
MaxRawChunkBytes uint64
|
||||
MaxScriptBytes uint64
|
||||
MaxControlRequestBytes uint64
|
||||
MaxJSONRPCBodyBytes uint64
|
||||
}
|
||||
|
||||
type Observability struct {
|
||||
Listen string
|
||||
LivenessPath string
|
||||
ReadinessPath string
|
||||
MetricsPath string
|
||||
LogLevel string
|
||||
LogFormat string
|
||||
LogFile string
|
||||
LogMaxBytes uint64
|
||||
LogMaxFiles uint32
|
||||
}
|
||||
|
||||
type Client struct {
|
||||
Client ClientCore
|
||||
TLS TLS
|
||||
Shells Shells
|
||||
Network ClientNetwork
|
||||
Storage ClientStorage
|
||||
Flow ClientFlow
|
||||
Execution Execution
|
||||
Observability Observability
|
||||
Profiles Profiles
|
||||
}
|
||||
|
||||
type ClientCore struct {
|
||||
ServerURL string
|
||||
StateDir string
|
||||
ClientID string
|
||||
DaemonCWD string
|
||||
MaxRunningCommands uint32
|
||||
MaxQueuedCommands uint32
|
||||
ShutdownGrace time.Duration
|
||||
}
|
||||
|
||||
type TLS struct {
|
||||
CAFile string
|
||||
ServerName string
|
||||
}
|
||||
|
||||
type Shells struct {
|
||||
DefaultUnix string
|
||||
DefaultWindows string
|
||||
SH string
|
||||
Bash string
|
||||
CMD string
|
||||
PowerShell string
|
||||
AllowedCWDRoots []string
|
||||
Advertised map[string]string
|
||||
}
|
||||
|
||||
type ClientNetwork struct {
|
||||
HeartbeatIdle time.Duration
|
||||
LivenessTimeout time.Duration
|
||||
ReconnectInitial time.Duration
|
||||
ReconnectMax time.Duration
|
||||
StableSessionReset time.Duration
|
||||
ConnectTimeout time.Duration
|
||||
WriteDeadline time.Duration
|
||||
}
|
||||
|
||||
type ClientStorage struct {
|
||||
CommandOutputLimitBytes uint64
|
||||
CommandTotalLimitBytes uint64
|
||||
ClientTotalLimitBytes uint64
|
||||
TombstoneMaxEntries uint64
|
||||
CommandCloseoutReserveBytes uint64
|
||||
FreeSpaceFloorBytes uint64
|
||||
SegmentTargetBytes uint64
|
||||
DurabilityInterval time.Duration
|
||||
}
|
||||
|
||||
type ClientFlow struct {
|
||||
RawOutputCommandHighBytes uint64
|
||||
RawOutputCommandLowBytes uint64
|
||||
RawOutputClientHighBytes uint64
|
||||
RawOutputClientLowBytes uint64
|
||||
UnacknowledgedPerCommandBytes uint64
|
||||
UnacknowledgedPerSessionBytes uint64
|
||||
}
|
||||
|
||||
type Execution struct {
|
||||
DescendantDrainGrace time.Duration
|
||||
WindowsTermGrace time.Duration
|
||||
HungThreshold time.Duration
|
||||
DiagnosticInterval time.Duration
|
||||
MaxScriptBytes uint64
|
||||
MaxExecutionSpecBytes uint64
|
||||
MaxAgentEnvelopeBytes uint64
|
||||
MaxRawChunkBytes uint64
|
||||
ProtocolDetailMaxBytes uint64
|
||||
}
|
||||
|
||||
type Profiles struct {
|
||||
Light Profile
|
||||
CPUMedium Profile
|
||||
CPUHeavy Profile
|
||||
MemMedium Profile
|
||||
MemHeavy Profile
|
||||
DiskMedium Profile
|
||||
DiskHeavy Profile
|
||||
}
|
||||
|
||||
type Profile struct {
|
||||
Enabled bool
|
||||
RequiredControls []string
|
||||
CPUPercent uint64
|
||||
MemoryMaxBytes uint64
|
||||
PIDsMax uint64
|
||||
WindowsIOReadBPS uint64
|
||||
WindowsIOWriteBPS uint64
|
||||
LinuxIOReadBPS map[string]uint64
|
||||
LinuxIOWriteBPS map[string]uint64
|
||||
}
|
||||
|
||||
type serverFile struct {
|
||||
Server serverCoreFile `toml:"server"`
|
||||
JSONRPC jsonRPCFile `toml:"json_rpc"`
|
||||
Queue serverQueueFile `toml:"queue"`
|
||||
Storage serverStorageFile `toml:"storage"`
|
||||
Flow serverFlowFile `toml:"flow"`
|
||||
Protocol serverProtocolFile `toml:"protocol"`
|
||||
Observability observabilityFile `toml:"observability"`
|
||||
}
|
||||
|
||||
type serverCoreFile struct {
|
||||
DataDir string `toml:"data_dir"`
|
||||
AgentListen string `toml:"agent_listen"`
|
||||
AgentPath string `toml:"agent_path"`
|
||||
ControlSocket string `toml:"control_socket"`
|
||||
ShutdownGrace string `toml:"shutdown_grace"`
|
||||
}
|
||||
|
||||
type jsonRPCFile struct {
|
||||
Enabled bool `toml:"enabled"`
|
||||
Listen string `toml:"listen"`
|
||||
}
|
||||
|
||||
type serverQueueFile struct {
|
||||
DefaultTTL string `toml:"default_ttl"`
|
||||
MaxPerClient uint32 `toml:"max_per_client"`
|
||||
MaxServer uint32 `toml:"max_server"`
|
||||
RetryInitial string `toml:"retry_initial"`
|
||||
RetryMax string `toml:"retry_max"`
|
||||
}
|
||||
|
||||
type serverStorageFile struct {
|
||||
CommandOutputLimitBytes uint64 `toml:"command_output_limit_bytes"`
|
||||
CommandTotalLimitBytes uint64 `toml:"command_total_limit_bytes"`
|
||||
ClientTotalLimitBytes uint64 `toml:"client_total_limit_bytes"`
|
||||
ServerTotalLimitBytes uint64 `toml:"server_total_limit_bytes"`
|
||||
TerminalRetention string `toml:"terminal_retention"`
|
||||
AuditLimitBytes uint64 `toml:"audit_limit_bytes"`
|
||||
AuditRetention string `toml:"audit_retention"`
|
||||
TombstoneMaxEntries uint64 `toml:"tombstone_max_entries"`
|
||||
CommandCloseoutReserveBytes uint64 `toml:"command_closeout_reserve_bytes"`
|
||||
FreeSpaceFloorBytes uint64 `toml:"free_space_floor_bytes"`
|
||||
SegmentTargetBytes uint64 `toml:"segment_target_bytes"`
|
||||
DurabilityInterval string `toml:"durability_interval"`
|
||||
SQLiteBusyTimeout string `toml:"sqlite_busy_timeout"`
|
||||
IncidentNoteMaxBytes uint64 `toml:"incident_note_max_bytes"`
|
||||
ProtocolDetailMaxBytes uint64 `toml:"protocol_detail_max_bytes"`
|
||||
}
|
||||
|
||||
type serverFlowFile struct {
|
||||
RawOutputHighBytes uint64 `toml:"raw_output_high_bytes"`
|
||||
RawOutputLowBytes uint64 `toml:"raw_output_low_bytes"`
|
||||
UnacknowledgedPerCommandBytes uint64 `toml:"unacknowledged_per_command_bytes"`
|
||||
UnacknowledgedPerSessionBytes uint64 `toml:"unacknowledged_per_session_bytes"`
|
||||
WriteDeadline string `toml:"write_deadline"`
|
||||
}
|
||||
|
||||
type serverProtocolFile struct {
|
||||
HeartbeatIdle string `toml:"heartbeat_idle"`
|
||||
LivenessTimeout string `toml:"liveness_timeout"`
|
||||
TakeoverTTL string `toml:"takeover_ttl"`
|
||||
MaxAgentEnvelopeBytes uint64 `toml:"max_agent_envelope_bytes"`
|
||||
MaxExecutionSpecBytes uint64 `toml:"max_execution_spec_bytes"`
|
||||
MaxRawChunkBytes uint64 `toml:"max_raw_chunk_bytes"`
|
||||
MaxScriptBytes uint64 `toml:"max_script_bytes"`
|
||||
MaxControlRequestBytes uint64 `toml:"max_control_request_bytes"`
|
||||
MaxJSONRPCBodyBytes uint64 `toml:"max_json_rpc_body_bytes"`
|
||||
}
|
||||
|
||||
type clientFile struct {
|
||||
Client clientCoreFile `toml:"client"`
|
||||
TLS tlsFile `toml:"tls"`
|
||||
Shells shellsFile `toml:"shells"`
|
||||
Network clientNetworkFile `toml:"network"`
|
||||
Storage clientStorageFile `toml:"storage"`
|
||||
Flow clientFlowFile `toml:"flow"`
|
||||
Execution executionFile `toml:"execution"`
|
||||
Observability observabilityFile `toml:"observability"`
|
||||
Profiles profilesFile `toml:"profiles"`
|
||||
}
|
||||
|
||||
type clientCoreFile struct {
|
||||
ServerURL string `toml:"server_url"`
|
||||
StateDir string `toml:"state_dir"`
|
||||
ClientID string `toml:"client_id"`
|
||||
DaemonCWD string `toml:"daemon_cwd"`
|
||||
MaxRunningCommands uint32 `toml:"max_running_commands"`
|
||||
MaxQueuedCommands uint32 `toml:"max_queued_commands"`
|
||||
ShutdownGrace string `toml:"shutdown_grace"`
|
||||
}
|
||||
|
||||
type tlsFile struct {
|
||||
CAFile string `toml:"ca_file"`
|
||||
ServerName string `toml:"server_name"`
|
||||
}
|
||||
|
||||
type shellsFile struct {
|
||||
DefaultUnix string `toml:"default_unix"`
|
||||
DefaultWindows string `toml:"default_windows"`
|
||||
SH string `toml:"sh"`
|
||||
Bash string `toml:"bash"`
|
||||
CMD string `toml:"cmd"`
|
||||
PowerShell string `toml:"powershell"`
|
||||
AllowedCWDRoots []string `toml:"allowed_cwd_roots"`
|
||||
}
|
||||
|
||||
type clientNetworkFile struct {
|
||||
HeartbeatIdle string `toml:"heartbeat_idle"`
|
||||
LivenessTimeout string `toml:"liveness_timeout"`
|
||||
ReconnectInitial string `toml:"reconnect_initial"`
|
||||
ReconnectMax string `toml:"reconnect_max"`
|
||||
StableSessionReset string `toml:"stable_session_reset"`
|
||||
ConnectTimeout string `toml:"connect_timeout"`
|
||||
WriteDeadline string `toml:"write_deadline"`
|
||||
}
|
||||
|
||||
type clientStorageFile struct {
|
||||
CommandOutputLimitBytes uint64 `toml:"command_output_limit_bytes"`
|
||||
CommandTotalLimitBytes uint64 `toml:"command_total_limit_bytes"`
|
||||
ClientTotalLimitBytes uint64 `toml:"client_total_limit_bytes"`
|
||||
TombstoneMaxEntries uint64 `toml:"tombstone_max_entries"`
|
||||
CommandCloseoutReserveBytes uint64 `toml:"command_closeout_reserve_bytes"`
|
||||
FreeSpaceFloorBytes uint64 `toml:"free_space_floor_bytes"`
|
||||
SegmentTargetBytes uint64 `toml:"segment_target_bytes"`
|
||||
DurabilityInterval string `toml:"durability_interval"`
|
||||
}
|
||||
|
||||
type clientFlowFile struct {
|
||||
RawOutputCommandHighBytes uint64 `toml:"raw_output_command_high_bytes"`
|
||||
RawOutputCommandLowBytes uint64 `toml:"raw_output_command_low_bytes"`
|
||||
RawOutputClientHighBytes uint64 `toml:"raw_output_client_high_bytes"`
|
||||
RawOutputClientLowBytes uint64 `toml:"raw_output_client_low_bytes"`
|
||||
UnacknowledgedPerCommandBytes uint64 `toml:"unacknowledged_per_command_bytes"`
|
||||
UnacknowledgedPerSessionBytes uint64 `toml:"unacknowledged_per_session_bytes"`
|
||||
}
|
||||
|
||||
type executionFile struct {
|
||||
DescendantDrainGrace string `toml:"descendant_drain_grace"`
|
||||
WindowsTermGrace string `toml:"windows_term_grace"`
|
||||
HungThreshold string `toml:"hung_threshold"`
|
||||
DiagnosticInterval string `toml:"diagnostic_interval"`
|
||||
MaxScriptBytes uint64 `toml:"max_script_bytes"`
|
||||
MaxExecutionSpecBytes uint64 `toml:"max_execution_spec_bytes"`
|
||||
MaxAgentEnvelopeBytes uint64 `toml:"max_agent_envelope_bytes"`
|
||||
MaxRawChunkBytes uint64 `toml:"max_raw_chunk_bytes"`
|
||||
ProtocolDetailMaxBytes uint64 `toml:"protocol_detail_max_bytes"`
|
||||
}
|
||||
|
||||
type observabilityFile struct {
|
||||
Listen string `toml:"listen"`
|
||||
LivenessPath string `toml:"liveness_path"`
|
||||
ReadinessPath string `toml:"readiness_path"`
|
||||
MetricsPath string `toml:"metrics_path"`
|
||||
LogLevel string `toml:"log_level"`
|
||||
LogFormat string `toml:"log_format"`
|
||||
LogFile string `toml:"log_file"`
|
||||
LogMaxBytes uint64 `toml:"log_max_bytes"`
|
||||
LogMaxFiles uint32 `toml:"log_max_files"`
|
||||
}
|
||||
|
||||
type profilesFile struct {
|
||||
Light profileFile `toml:"light"`
|
||||
CPUMedium profileFile `toml:"cpu_medium"`
|
||||
CPUHeavy profileFile `toml:"cpu_heavy"`
|
||||
MemMedium profileFile `toml:"mem_medium"`
|
||||
MemHeavy profileFile `toml:"mem_heavy"`
|
||||
DiskMedium profileFile `toml:"disk_medium"`
|
||||
DiskHeavy profileFile `toml:"disk_heavy"`
|
||||
}
|
||||
|
||||
type profileFile struct {
|
||||
Enabled bool `toml:"enabled"`
|
||||
RequiredControls []string `toml:"required_controls"`
|
||||
CPUPercent uint64 `toml:"cpu_percent"`
|
||||
MemoryMaxBytes uint64 `toml:"memory_max_bytes"`
|
||||
PIDsMax uint64 `toml:"pids_max"`
|
||||
WindowsIOReadBPS uint64 `toml:"windows_io_read_bps"`
|
||||
WindowsIOWriteBPS uint64 `toml:"windows_io_write_bps"`
|
||||
LinuxIOReadBPS map[string]uint64 `toml:"linux_io_read_bps"`
|
||||
LinuxIOWriteBPS map[string]uint64 `toml:"linux_io_write_bps"`
|
||||
}
|
||||
@@ -0,0 +1,549 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
func validateServer(raw serverFile) (*Server, error) {
|
||||
dataDir, err := validateAbsolutePath("server.data_dir", raw.Server.DataDir, PlatformUnix, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
controlSocket, err := validateAbsolutePath("server.control_socket", raw.Server.ControlSocket, PlatformUnix, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataDir == "/" {
|
||||
return nil, fmt.Errorf("server.data_dir may not be the filesystem root")
|
||||
}
|
||||
if !pathWithin(controlSocket, "/run") && !pathWithin(controlSocket, "/var/run") {
|
||||
return nil, fmt.Errorf("server.control_socket must be below /run or /var/run")
|
||||
}
|
||||
shutdownGrace, err := parseDuration("server.shutdown_grace", raw.Server.ShutdownGrace, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defaultTTL, err := parseDuration("queue.default_ttl", raw.Queue.DefaultTTL, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
retryInitial, err := parseDuration("queue.retry_initial", raw.Queue.RetryInitial, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
retryMax, err := parseDuration("queue.retry_max", raw.Queue.RetryMax, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
terminalRetention, err := parseDuration("storage.terminal_retention", raw.Storage.TerminalRetention, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
auditRetention, err := parseDuration("storage.audit_retention", raw.Storage.AuditRetention, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
durabilityInterval, err := parseDuration("storage.durability_interval", raw.Storage.DurabilityInterval, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sqliteBusyTimeout, err := parseDuration("storage.sqlite_busy_timeout", raw.Storage.SQLiteBusyTimeout, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
writeDeadline, err := parseDuration("flow.write_deadline", raw.Flow.WriteDeadline, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
heartbeatIdle, err := parseDuration("protocol.heartbeat_idle", raw.Protocol.HeartbeatIdle, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
livenessTimeout, err := parseDuration("protocol.liveness_timeout", raw.Protocol.LivenessTimeout, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
takeoverTTL, err := parseDuration("protocol.takeover_ttl", raw.Protocol.TakeoverTTL, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := validateListener("server.agent_listen", raw.Server.AgentListen); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateHTTPPath("server.agent_path", raw.Server.AgentPath); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateListener("json_rpc.listen", raw.JSONRPC.Listen); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw.Queue.MaxPerClient == 0 || raw.Queue.MaxServer == 0 || raw.Queue.MaxPerClient > raw.Queue.MaxServer {
|
||||
return nil, fmt.Errorf("queue limits must be positive and max_per_client <= max_server")
|
||||
}
|
||||
if retryInitial > retryMax {
|
||||
return nil, fmt.Errorf("queue.retry_initial must be <= queue.retry_max")
|
||||
}
|
||||
if err := increasing("server storage tiers", raw.Storage.CommandOutputLimitBytes, raw.Storage.CommandTotalLimitBytes, raw.Storage.ClientTotalLimitBytes, raw.Storage.ServerTotalLimitBytes); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw.Storage.CommandCloseoutReserveBytes == 0 || raw.Storage.CommandCloseoutReserveBytes >= raw.Storage.CommandTotalLimitBytes {
|
||||
return nil, fmt.Errorf("storage.command_closeout_reserve_bytes must fit below command_total_limit_bytes")
|
||||
}
|
||||
for name, value := range map[string]uint64{
|
||||
"storage.audit_limit_bytes": raw.Storage.AuditLimitBytes, "storage.tombstone_max_entries": raw.Storage.TombstoneMaxEntries,
|
||||
"storage.free_space_floor_bytes": raw.Storage.FreeSpaceFloorBytes, "storage.segment_target_bytes": raw.Storage.SegmentTargetBytes,
|
||||
} {
|
||||
if value == 0 {
|
||||
return nil, fmt.Errorf("%s must be positive", name)
|
||||
}
|
||||
}
|
||||
if raw.Storage.IncidentNoteMaxBytes == 0 || raw.Storage.IncidentNoteMaxBytes > HardMaxProtocolDetailBytes || raw.Storage.ProtocolDetailMaxBytes == 0 || raw.Storage.ProtocolDetailMaxBytes > HardMaxProtocolDetailBytes {
|
||||
return nil, fmt.Errorf("incident and protocol detail limits must be in 1..%d", HardMaxProtocolDetailBytes)
|
||||
}
|
||||
if err := watermarks("flow.raw_output", raw.Flow.RawOutputLowBytes, raw.Flow.RawOutputHighBytes); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw.Flow.UnacknowledgedPerCommandBytes == 0 || raw.Flow.UnacknowledgedPerCommandBytes >= raw.Storage.CommandTotalLimitBytes || raw.Flow.UnacknowledgedPerSessionBytes == 0 || raw.Flow.UnacknowledgedPerSessionBytes >= raw.Storage.ClientTotalLimitBytes || raw.Flow.UnacknowledgedPerCommandBytes > raw.Flow.UnacknowledgedPerSessionBytes {
|
||||
return nil, fmt.Errorf("flow send windows must be positive, ordered, and below their durable tiers")
|
||||
}
|
||||
if heartbeatIdle >= livenessTimeout {
|
||||
return nil, fmt.Errorf("protocol.heartbeat_idle must be below liveness_timeout")
|
||||
}
|
||||
if err := protocolLimits(raw.Protocol.MaxAgentEnvelopeBytes, raw.Protocol.MaxExecutionSpecBytes, raw.Protocol.MaxRawChunkBytes, raw.Protocol.MaxScriptBytes, raw.Protocol.MaxControlRequestBytes, raw.Protocol.MaxJSONRPCBodyBytes); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
observability, err := validateObservability(raw.Observability, PlatformUnix)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
host, _, _ := net.SplitHostPort(raw.JSONRPC.Listen)
|
||||
return &Server{
|
||||
Server: ServerCore{DataDir: dataDir, AgentListen: raw.Server.AgentListen, AgentPath: raw.Server.AgentPath, ControlSocket: controlSocket, ShutdownGrace: shutdownGrace},
|
||||
JSONRPC: JSONRPC{Enabled: raw.JSONRPC.Enabled, Listen: raw.JSONRPC.Listen, NonLoopbackBind: raw.JSONRPC.Enabled && !isLoopbackHost(host)},
|
||||
Queue: ServerQueue{DefaultTTL: defaultTTL, MaxPerClient: raw.Queue.MaxPerClient, MaxServer: raw.Queue.MaxServer, RetryInitial: retryInitial, RetryMax: retryMax},
|
||||
Storage: ServerStorage{
|
||||
CommandOutputLimitBytes: raw.Storage.CommandOutputLimitBytes, CommandTotalLimitBytes: raw.Storage.CommandTotalLimitBytes,
|
||||
ClientTotalLimitBytes: raw.Storage.ClientTotalLimitBytes, ServerTotalLimitBytes: raw.Storage.ServerTotalLimitBytes,
|
||||
TerminalRetention: terminalRetention, AuditLimitBytes: raw.Storage.AuditLimitBytes, AuditRetention: auditRetention,
|
||||
TombstoneMaxEntries: raw.Storage.TombstoneMaxEntries, CommandCloseoutReserveBytes: raw.Storage.CommandCloseoutReserveBytes,
|
||||
FreeSpaceFloorBytes: raw.Storage.FreeSpaceFloorBytes, SegmentTargetBytes: raw.Storage.SegmentTargetBytes,
|
||||
DurabilityInterval: durabilityInterval, SQLiteBusyTimeout: sqliteBusyTimeout,
|
||||
IncidentNoteMaxBytes: raw.Storage.IncidentNoteMaxBytes, ProtocolDetailMaxBytes: raw.Storage.ProtocolDetailMaxBytes,
|
||||
},
|
||||
Flow: ServerFlow{RawOutputHighBytes: raw.Flow.RawOutputHighBytes, RawOutputLowBytes: raw.Flow.RawOutputLowBytes, UnacknowledgedPerCommandBytes: raw.Flow.UnacknowledgedPerCommandBytes, UnacknowledgedPerSessionBytes: raw.Flow.UnacknowledgedPerSessionBytes, WriteDeadline: writeDeadline},
|
||||
Protocol: ServerProtocol{HeartbeatIdle: heartbeatIdle, LivenessTimeout: livenessTimeout, TakeoverTTL: takeoverTTL, MaxAgentEnvelopeBytes: raw.Protocol.MaxAgentEnvelopeBytes, MaxExecutionSpecBytes: raw.Protocol.MaxExecutionSpecBytes, MaxRawChunkBytes: raw.Protocol.MaxRawChunkBytes, MaxScriptBytes: raw.Protocol.MaxScriptBytes, MaxControlRequestBytes: raw.Protocol.MaxControlRequestBytes, MaxJSONRPCBodyBytes: raw.Protocol.MaxJSONRPCBodyBytes},
|
||||
Observability: observability,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func validateClient(raw clientFile, options ClientOptions) (*Client, error) {
|
||||
if options.Platform != PlatformUnix && options.Platform != PlatformWindows {
|
||||
return nil, fmt.Errorf("client target platform is required")
|
||||
}
|
||||
parsedURL, err := url.Parse(raw.Client.ServerURL)
|
||||
if err != nil || parsedURL.Scheme != "wss" || parsedURL.Host == "" || parsedURL.User != nil {
|
||||
return nil, fmt.Errorf("client.server_url must be an absolute wss URL without userinfo")
|
||||
}
|
||||
stateDir, err := validateAbsolutePath("client.state_dir", raw.Client.StateDir, options.Platform, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
daemonCWD, err := validateAbsolutePath("client.daemon_cwd", raw.Client.DaemonCWD, options.Platform, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if options.CheckFilesystem {
|
||||
daemonCWD, err = checkNativeDirectory("client.daemon_cwd", daemonCWD, options.Platform)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if err := validateClientID(raw.Client.ClientID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw.Client.MaxRunningCommands == 0 || raw.Client.MaxQueuedCommands == 0 {
|
||||
return nil, fmt.Errorf("client command limits must be positive")
|
||||
}
|
||||
shutdownGrace, err := parseDuration("client.shutdown_grace", raw.Client.ShutdownGrace, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
caFile, err := validateAbsolutePath("tls.ca_file", raw.TLS.CAFile, options.Platform, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if options.CheckFilesystem && caFile != "" {
|
||||
caFile, err = checkNativeRegularFile("tls.ca_file", caFile, options.Platform)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
shells, err := validateShells(raw.Shells, options)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
network, err := validateClientNetwork(raw.Network)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
durability, err := parseDuration("storage.durability_interval", raw.Storage.DurabilityInterval, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := increasing("client storage tiers", raw.Storage.CommandOutputLimitBytes, raw.Storage.CommandTotalLimitBytes, raw.Storage.ClientTotalLimitBytes); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw.Storage.TombstoneMaxEntries == 0 || raw.Storage.FreeSpaceFloorBytes == 0 || raw.Storage.SegmentTargetBytes == 0 || raw.Storage.CommandCloseoutReserveBytes == 0 || raw.Storage.CommandCloseoutReserveBytes >= raw.Storage.CommandTotalLimitBytes {
|
||||
return nil, fmt.Errorf("client storage counts must be positive and closeout reserve must fit within command quota")
|
||||
}
|
||||
if err := watermarks("flow.raw_output_command", raw.Flow.RawOutputCommandLowBytes, raw.Flow.RawOutputCommandHighBytes); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := watermarks("flow.raw_output_client", raw.Flow.RawOutputClientLowBytes, raw.Flow.RawOutputClientHighBytes); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if raw.Flow.RawOutputCommandHighBytes >= raw.Flow.RawOutputClientHighBytes {
|
||||
return nil, fmt.Errorf("command raw-output high watermark must be below client high watermark")
|
||||
}
|
||||
if raw.Flow.UnacknowledgedPerCommandBytes == 0 || raw.Flow.UnacknowledgedPerCommandBytes >= raw.Storage.CommandTotalLimitBytes || raw.Flow.UnacknowledgedPerSessionBytes == 0 || raw.Flow.UnacknowledgedPerSessionBytes >= raw.Storage.ClientTotalLimitBytes || raw.Flow.UnacknowledgedPerCommandBytes > raw.Flow.UnacknowledgedPerSessionBytes {
|
||||
return nil, fmt.Errorf("client send windows must be positive, ordered, and below durable tiers")
|
||||
}
|
||||
execution, err := validateExecution(raw.Execution)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
observability, err := validateObservability(raw.Observability, options.Platform)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
profiles, err := validateProfiles(raw.Profiles, options.Platform)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Client{
|
||||
Client: ClientCore{ServerURL: raw.Client.ServerURL, StateDir: stateDir, ClientID: raw.Client.ClientID, DaemonCWD: daemonCWD, MaxRunningCommands: raw.Client.MaxRunningCommands, MaxQueuedCommands: raw.Client.MaxQueuedCommands, ShutdownGrace: shutdownGrace},
|
||||
TLS: TLS{CAFile: caFile, ServerName: raw.TLS.ServerName}, Shells: shells, Network: network,
|
||||
Storage: ClientStorage{CommandOutputLimitBytes: raw.Storage.CommandOutputLimitBytes, CommandTotalLimitBytes: raw.Storage.CommandTotalLimitBytes, ClientTotalLimitBytes: raw.Storage.ClientTotalLimitBytes, TombstoneMaxEntries: raw.Storage.TombstoneMaxEntries, CommandCloseoutReserveBytes: raw.Storage.CommandCloseoutReserveBytes, FreeSpaceFloorBytes: raw.Storage.FreeSpaceFloorBytes, SegmentTargetBytes: raw.Storage.SegmentTargetBytes, DurabilityInterval: durability},
|
||||
Flow: ClientFlow{RawOutputCommandHighBytes: raw.Flow.RawOutputCommandHighBytes, RawOutputCommandLowBytes: raw.Flow.RawOutputCommandLowBytes, RawOutputClientHighBytes: raw.Flow.RawOutputClientHighBytes, RawOutputClientLowBytes: raw.Flow.RawOutputClientLowBytes, UnacknowledgedPerCommandBytes: raw.Flow.UnacknowledgedPerCommandBytes, UnacknowledgedPerSessionBytes: raw.Flow.UnacknowledgedPerSessionBytes},
|
||||
Execution: execution, Observability: observability, Profiles: profiles,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func validateClientNetwork(raw clientNetworkFile) (ClientNetwork, error) {
|
||||
values := make([]timePair, 0, 7)
|
||||
for _, item := range []struct{ name, value string }{
|
||||
{"network.heartbeat_idle", raw.HeartbeatIdle}, {"network.liveness_timeout", raw.LivenessTimeout},
|
||||
{"network.reconnect_initial", raw.ReconnectInitial}, {"network.reconnect_max", raw.ReconnectMax},
|
||||
{"network.stable_session_reset", raw.StableSessionReset}, {"network.connect_timeout", raw.ConnectTimeout},
|
||||
{"network.write_deadline", raw.WriteDeadline},
|
||||
} {
|
||||
value, err := parseDuration(item.name, item.value, false)
|
||||
if err != nil {
|
||||
return ClientNetwork{}, err
|
||||
}
|
||||
values = append(values, timePair{item.name, value})
|
||||
}
|
||||
if values[0].value >= values[1].value {
|
||||
return ClientNetwork{}, fmt.Errorf("network.heartbeat_idle must be below liveness_timeout")
|
||||
}
|
||||
if values[2].value > values[3].value {
|
||||
return ClientNetwork{}, fmt.Errorf("network.reconnect_initial must be <= reconnect_max")
|
||||
}
|
||||
return ClientNetwork{HeartbeatIdle: values[0].value, LivenessTimeout: values[1].value, ReconnectInitial: values[2].value, ReconnectMax: values[3].value, StableSessionReset: values[4].value, ConnectTimeout: values[5].value, WriteDeadline: values[6].value}, nil
|
||||
}
|
||||
|
||||
type timePair struct {
|
||||
name string
|
||||
value time.Duration
|
||||
}
|
||||
|
||||
func validateExecution(raw executionFile) (Execution, error) {
|
||||
durations := make([]time.Duration, 4)
|
||||
for index, item := range []struct{ name, value string }{{"execution.descendant_drain_grace", raw.DescendantDrainGrace}, {"execution.windows_term_grace", raw.WindowsTermGrace}, {"execution.hung_threshold", raw.HungThreshold}, {"execution.diagnostic_interval", raw.DiagnosticInterval}} {
|
||||
value, err := parseDuration(item.name, item.value, false)
|
||||
if err != nil {
|
||||
return Execution{}, err
|
||||
}
|
||||
durations[index] = value
|
||||
}
|
||||
if err := protocolLimits(raw.MaxAgentEnvelopeBytes, raw.MaxExecutionSpecBytes, raw.MaxRawChunkBytes, raw.MaxScriptBytes, HardMaxControlRequestBytes, HardMaxJSONRPCBodyBytes); err != nil {
|
||||
return Execution{}, err
|
||||
}
|
||||
if raw.ProtocolDetailMaxBytes == 0 || raw.ProtocolDetailMaxBytes > HardMaxProtocolDetailBytes {
|
||||
return Execution{}, fmt.Errorf("execution.protocol_detail_max_bytes exceeds hard ceiling")
|
||||
}
|
||||
return Execution{DescendantDrainGrace: durations[0], WindowsTermGrace: durations[1], HungThreshold: durations[2], DiagnosticInterval: durations[3], MaxScriptBytes: raw.MaxScriptBytes, MaxExecutionSpecBytes: raw.MaxExecutionSpecBytes, MaxAgentEnvelopeBytes: raw.MaxAgentEnvelopeBytes, MaxRawChunkBytes: raw.MaxRawChunkBytes, ProtocolDetailMaxBytes: raw.ProtocolDetailMaxBytes}, nil
|
||||
}
|
||||
|
||||
func validateShells(raw shellsFile, options ClientOptions) (Shells, error) {
|
||||
if !slices.Contains([]string{"sh", "bash"}, raw.DefaultUnix) {
|
||||
return Shells{}, fmt.Errorf("shells.default_unix must be sh or bash")
|
||||
}
|
||||
if !slices.Contains([]string{"cmd", "powershell"}, raw.DefaultWindows) {
|
||||
return Shells{}, fmt.Errorf("shells.default_windows must be cmd or powershell")
|
||||
}
|
||||
type shellPath struct {
|
||||
name, value string
|
||||
platform Platform
|
||||
}
|
||||
paths := []shellPath{{"sh", raw.SH, PlatformUnix}, {"bash", raw.Bash, PlatformUnix}, {"cmd", raw.CMD, PlatformWindows}, {"powershell", raw.PowerShell, PlatformWindows}}
|
||||
advertised := make(map[string]string)
|
||||
resolved := make(map[string]string)
|
||||
for _, shell := range paths {
|
||||
canonical, err := validateAbsolutePath("shells."+shell.name, shell.value, shell.platform, true)
|
||||
if err != nil {
|
||||
return Shells{}, err
|
||||
}
|
||||
resolved[shell.name] = canonical
|
||||
if shell.platform == options.Platform && canonical != "" {
|
||||
if options.CheckFilesystem {
|
||||
canonical, err = checkNativeExecutable("shells."+shell.name, canonical, options.Platform)
|
||||
if err != nil {
|
||||
return Shells{}, err
|
||||
}
|
||||
}
|
||||
advertised[shell.name] = canonical
|
||||
resolved[shell.name] = canonical
|
||||
}
|
||||
}
|
||||
defaultShell := raw.DefaultUnix
|
||||
if options.Platform == PlatformWindows {
|
||||
defaultShell = raw.DefaultWindows
|
||||
}
|
||||
if _, exists := advertised[defaultShell]; !exists {
|
||||
return Shells{}, fmt.Errorf("current-platform default shell %q is not validated and advertised", defaultShell)
|
||||
}
|
||||
roots := make([]string, len(raw.AllowedCWDRoots))
|
||||
for index, root := range raw.AllowedCWDRoots {
|
||||
canonical, err := validateAbsolutePath(fmt.Sprintf("shells.allowed_cwd_roots[%d]", index), root, options.Platform, false)
|
||||
if err != nil {
|
||||
return Shells{}, err
|
||||
}
|
||||
if options.CheckFilesystem {
|
||||
canonical, err = checkNativeDirectory(fmt.Sprintf("shells.allowed_cwd_roots[%d]", index), canonical, options.Platform)
|
||||
if err != nil {
|
||||
return Shells{}, err
|
||||
}
|
||||
}
|
||||
roots[index] = canonical
|
||||
}
|
||||
return Shells{DefaultUnix: raw.DefaultUnix, DefaultWindows: raw.DefaultWindows, SH: resolved["sh"], Bash: resolved["bash"], CMD: resolved["cmd"], PowerShell: resolved["powershell"], AllowedCWDRoots: roots, Advertised: advertised}, nil
|
||||
}
|
||||
|
||||
func validateProfiles(raw profilesFile, platform Platform) (Profiles, error) {
|
||||
convert := func(name string, source profileFile) (Profile, error) {
|
||||
seen := map[string]bool{}
|
||||
for _, control := range source.RequiredControls {
|
||||
if seen[control] || !slices.Contains([]string{"cpu", "memory", "pids", "io"}, control) {
|
||||
return Profile{}, fmt.Errorf("profiles.%s has duplicate or unknown required control %q", name, control)
|
||||
}
|
||||
seen[control] = true
|
||||
}
|
||||
for key, value := range source.LinuxIOReadBPS {
|
||||
if err := validateDeviceRate(key, value); err != nil {
|
||||
return Profile{}, fmt.Errorf("profiles.%s.linux_io_read_bps: %w", name, err)
|
||||
}
|
||||
}
|
||||
for key, value := range source.LinuxIOWriteBPS {
|
||||
if err := validateDeviceRate(key, value); err != nil {
|
||||
return Profile{}, fmt.Errorf("profiles.%s.linux_io_write_bps: %w", name, err)
|
||||
}
|
||||
}
|
||||
if source.Enabled {
|
||||
if len(source.RequiredControls) == 0 {
|
||||
return Profile{}, fmt.Errorf("profiles.%s is enabled without required_controls", name)
|
||||
}
|
||||
if seen["cpu"] && source.CPUPercent == 0 {
|
||||
return Profile{}, fmt.Errorf("profiles.%s requires nonzero cpu_percent", name)
|
||||
}
|
||||
if seen["memory"] && source.MemoryMaxBytes == 0 {
|
||||
return Profile{}, fmt.Errorf("profiles.%s requires nonzero memory_max_bytes", name)
|
||||
}
|
||||
if seen["pids"] && source.PIDsMax == 0 {
|
||||
return Profile{}, fmt.Errorf("profiles.%s requires nonzero pids_max", name)
|
||||
}
|
||||
if seen["io"] {
|
||||
if platform == PlatformWindows && source.WindowsIOReadBPS == 0 && source.WindowsIOWriteBPS == 0 {
|
||||
return Profile{}, fmt.Errorf("profiles.%s requires a nonzero Windows IO rate", name)
|
||||
}
|
||||
if platform == PlatformUnix && len(source.LinuxIOReadBPS) == 0 && len(source.LinuxIOWriteBPS) == 0 {
|
||||
return Profile{}, fmt.Errorf("profiles.%s requires a nonzero Linux IO rate", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
return Profile{Enabled: source.Enabled, RequiredControls: slices.Clone(source.RequiredControls), CPUPercent: source.CPUPercent, MemoryMaxBytes: source.MemoryMaxBytes, PIDsMax: source.PIDsMax, WindowsIOReadBPS: source.WindowsIOReadBPS, WindowsIOWriteBPS: source.WindowsIOWriteBPS, LinuxIOReadBPS: cloneMap(source.LinuxIOReadBPS), LinuxIOWriteBPS: cloneMap(source.LinuxIOWriteBPS)}, nil
|
||||
}
|
||||
var result Profiles
|
||||
for _, item := range []struct {
|
||||
name string
|
||||
source profileFile
|
||||
target *Profile
|
||||
}{{"light", raw.Light, &result.Light}, {"cpu_medium", raw.CPUMedium, &result.CPUMedium}, {"cpu_heavy", raw.CPUHeavy, &result.CPUHeavy}, {"mem_medium", raw.MemMedium, &result.MemMedium}, {"mem_heavy", raw.MemHeavy, &result.MemHeavy}, {"disk_medium", raw.DiskMedium, &result.DiskMedium}, {"disk_heavy", raw.DiskHeavy, &result.DiskHeavy}} {
|
||||
profile, err := convert(item.name, item.source)
|
||||
if err != nil {
|
||||
return Profiles{}, err
|
||||
}
|
||||
*item.target = profile
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func ValidateProfileCombination(names []string) error {
|
||||
seen := map[string]bool{}
|
||||
dimensions := map[string]string{"cpu_medium": "cpu", "cpu_heavy": "cpu", "mem_medium": "memory", "mem_heavy": "memory", "disk_medium": "disk", "disk_heavy": "disk"}
|
||||
for _, name := range names {
|
||||
name = strings.ToLower(name)
|
||||
if seen[name] {
|
||||
return fmt.Errorf("duplicate execution profile %q", name)
|
||||
}
|
||||
seen[name] = true
|
||||
if name == "light" {
|
||||
if len(names) != 1 {
|
||||
return fmt.Errorf("light execution profile is exclusive")
|
||||
}
|
||||
continue
|
||||
}
|
||||
dimension, exists := dimensions[name]
|
||||
if !exists {
|
||||
return fmt.Errorf("unknown execution profile %q", name)
|
||||
}
|
||||
for prior := range seen {
|
||||
if prior != name && dimensions[prior] == dimension {
|
||||
return fmt.Errorf("multiple %s execution profiles", dimension)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateDeviceRate(key string, value uint64) error {
|
||||
parts := strings.Split(key, ":")
|
||||
if len(parts) != 2 || value == 0 {
|
||||
return fmt.Errorf("device %q must be canonical major:minor with nonzero rate", key)
|
||||
}
|
||||
for _, part := range parts {
|
||||
parsed, err := strconv.ParseUint(part, 10, 32)
|
||||
if err != nil || strconv.FormatUint(parsed, 10) != part {
|
||||
return fmt.Errorf("device %q is not canonical major:minor", key)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateObservability(raw observabilityFile, platform Platform) (Observability, error) {
|
||||
if err := validateListener("observability.listen", raw.Listen); err != nil {
|
||||
return Observability{}, err
|
||||
}
|
||||
for name, value := range map[string]string{"liveness_path": raw.LivenessPath, "readiness_path": raw.ReadinessPath, "metrics_path": raw.MetricsPath} {
|
||||
if err := validateHTTPPath("observability."+name, value); err != nil {
|
||||
return Observability{}, err
|
||||
}
|
||||
}
|
||||
if raw.LivenessPath == raw.ReadinessPath || raw.LivenessPath == raw.MetricsPath || raw.ReadinessPath == raw.MetricsPath {
|
||||
return Observability{}, fmt.Errorf("observability paths must be distinct")
|
||||
}
|
||||
if !slices.Contains([]string{"debug", "info", "warn", "error"}, raw.LogLevel) {
|
||||
return Observability{}, fmt.Errorf("observability.log_level is invalid")
|
||||
}
|
||||
if !slices.Contains([]string{"json", "text"}, raw.LogFormat) {
|
||||
return Observability{}, fmt.Errorf("observability.log_format is invalid")
|
||||
}
|
||||
logFile, err := validateAbsolutePath("observability.log_file", raw.LogFile, platform, true)
|
||||
if err != nil {
|
||||
return Observability{}, err
|
||||
}
|
||||
if (raw.LogMaxBytes == 0) != (raw.LogMaxFiles == 0) {
|
||||
return Observability{}, fmt.Errorf("observability log rotation byte/file limits must both be zero or positive")
|
||||
}
|
||||
return Observability{Listen: raw.Listen, LivenessPath: raw.LivenessPath, ReadinessPath: raw.ReadinessPath, MetricsPath: raw.MetricsPath, LogLevel: raw.LogLevel, LogFormat: raw.LogFormat, LogFile: logFile, LogMaxBytes: raw.LogMaxBytes, LogMaxFiles: raw.LogMaxFiles}, nil
|
||||
}
|
||||
|
||||
func increasing(name string, values ...uint64) error {
|
||||
for index, value := range values {
|
||||
if value == 0 {
|
||||
return fmt.Errorf("%s values must be positive", name)
|
||||
}
|
||||
if index > 0 && value <= values[index-1] {
|
||||
return fmt.Errorf("%s must be strictly increasing", name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func watermarks(name string, low, high uint64) error {
|
||||
if low == 0 || low >= high {
|
||||
return fmt.Errorf("%s requires 0 < low < high", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func protocolLimits(agent, spec, chunk, script, control, jsonRPC uint64) error {
|
||||
values := []struct {
|
||||
name string
|
||||
value, max uint64
|
||||
}{{"max_agent_envelope_bytes", agent, HardMaxAgentEnvelopeBytes}, {"max_execution_spec_bytes", spec, HardMaxExecutionSpecBytes}, {"max_raw_chunk_bytes", chunk, HardMaxRawChunkBytes}, {"max_script_bytes", script, HardMaxScriptBytes}, {"max_control_request_bytes", control, HardMaxControlRequestBytes}, {"max_json_rpc_body_bytes", jsonRPC, HardMaxJSONRPCBodyBytes}}
|
||||
for _, item := range values {
|
||||
if item.value == 0 || item.value > item.max {
|
||||
return fmt.Errorf("protocol.%s must be in 1..%d", item.name, item.max)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func validateListener(name, value string) error {
|
||||
host, port, err := net.SplitHostPort(value)
|
||||
if err != nil || host == "" || port == "" {
|
||||
return fmt.Errorf("%s must be a host:port listener", name)
|
||||
}
|
||||
if _, err := strconv.ParseUint(port, 10, 16); err != nil || port == "0" {
|
||||
return fmt.Errorf("%s has invalid port", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func validateHTTPPath(name, value string) error {
|
||||
if value == "" || !strings.HasPrefix(value, "/") || strings.HasPrefix(value, "//") {
|
||||
return fmt.Errorf("%s must be an absolute HTTP path", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func isLoopbackHost(host string) bool {
|
||||
if strings.EqualFold(host, "localhost") {
|
||||
return true
|
||||
}
|
||||
ip := net.ParseIP(strings.Trim(host, "[]"))
|
||||
return ip != nil && ip.IsLoopback()
|
||||
}
|
||||
func validateClientID(value string) error {
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
if len(value) > 128 {
|
||||
return fmt.Errorf("client.client_id exceeds 128 bytes")
|
||||
}
|
||||
for _, current := range []byte(value) {
|
||||
if current < 0x21 || current > 0x7e {
|
||||
return fmt.Errorf("client.client_id must contain printable ASCII without spaces")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func pathWithin(value, root string) bool {
|
||||
relative, err := filepath.Rel(root, value)
|
||||
return err == nil && relative != "." && relative != ".." && !strings.HasPrefix(relative, ".."+string(filepath.Separator))
|
||||
}
|
||||
func cloneMap(source map[string]uint64) map[string]uint64 {
|
||||
if source == nil {
|
||||
return nil
|
||||
}
|
||||
result := make(map[string]uint64, len(source))
|
||||
for key, value := range source {
|
||||
result[key] = value
|
||||
}
|
||||
return result
|
||||
}
|
||||
Reference in New Issue
Block a user