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 }