95 lines
5.1 KiB
Go
95 lines
5.1 KiB
Go
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
|
|
}
|