376 lines
10 KiB
Go
376 lines
10 KiB
Go
package config
|
|
|
|
import (
|
|
"encoding/pem"
|
|
"fmt"
|
|
"log/slog"
|
|
"net"
|
|
"os"
|
|
"regexp"
|
|
"strconv"
|
|
"strings"
|
|
)
|
|
|
|
const (
|
|
defaultPortFirmware = 8005
|
|
defaultPortLookup = 8007
|
|
defaultPortXMPP = 5223
|
|
defaultHealthPort = 8080
|
|
defaultBindAddress = "0.0.0.0"
|
|
defaultMQTTBase = "ecovacs"
|
|
defaultHADiscovery = "homeassistant"
|
|
defaultControllerJID = "n95bridge@ecouser.net/homeassistant"
|
|
defaultLogLevel = "info"
|
|
defaultMQTTTLSPort = 8883
|
|
defaultMQTTPlainPort = 1883
|
|
)
|
|
|
|
var (
|
|
clientIDRe = regexp.MustCompile(`^[A-Za-z0-9_-]{1,64}$`)
|
|
topicPartRe = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
|
|
)
|
|
|
|
// Config holds every variable from the frozen environment table.
|
|
// MQTTPassword is never emitted by String, Format, GoString, or LogValue.
|
|
type Config struct {
|
|
PortFirmware int
|
|
PortLookup int
|
|
PortXMPP int
|
|
HealthPort int
|
|
BindAddress string
|
|
AdvertiseIP net.IP
|
|
MQTTHost string
|
|
MQTTPort int
|
|
MQTTTLS bool
|
|
MQTTCAFile string
|
|
MQTTUsername string
|
|
MQTTPassword string
|
|
MQTTClientID string
|
|
MQTTBase string
|
|
HADiscoveryPrefix string
|
|
ControllerJID string
|
|
RawCommands bool
|
|
LogLevel slog.Level
|
|
}
|
|
|
|
// Load parses and validates environ, which must be in KEY=value form.
|
|
// Absent keys use defaults; empty values are invalid except for MQTT_PORT and
|
|
// MQTT_CLIENT_ID, which fall back to their defaults.
|
|
func Load(environ []string) (Config, error) {
|
|
env := make(map[string]string, len(environ))
|
|
present := make(map[string]bool, len(environ))
|
|
for _, e := range environ {
|
|
if i := strings.IndexByte(e, '='); i >= 0 {
|
|
key := e[:i]
|
|
env[key] = e[i+1:]
|
|
present[key] = true
|
|
}
|
|
}
|
|
|
|
var cfg Config
|
|
var errs []string
|
|
|
|
parsePort := func(key, raw string, def int, allowEmpty bool) int {
|
|
if raw == "" {
|
|
if allowEmpty {
|
|
return def
|
|
}
|
|
errs = append(errs, fmt.Sprintf("%s must be an integer from 1 to 65535", key))
|
|
return 0
|
|
}
|
|
n, err := strconv.Atoi(raw)
|
|
if err != nil || n < 1 || n > 65535 {
|
|
errs = append(errs, fmt.Sprintf("%s must be an integer from 1 to 65535", key))
|
|
return 0
|
|
}
|
|
return n
|
|
}
|
|
|
|
// MQTT_TLS is parsed before applying the MQTT_PORT default.
|
|
mqttTLS := false
|
|
if !present["MQTT_TLS"] {
|
|
mqttTLS = false
|
|
} else if v := env["MQTT_TLS"]; v == "" {
|
|
errs = append(errs, "MQTT_TLS must be true or false")
|
|
} else {
|
|
b, ok := parseBoolStrict(v)
|
|
if !ok {
|
|
errs = append(errs, "MQTT_TLS must be true or false")
|
|
} else {
|
|
mqttTLS = b
|
|
}
|
|
}
|
|
cfg.MQTTTLS = mqttTLS
|
|
|
|
cfg.PortFirmware = parsePort("PORT_FIRMWARE", env["PORT_FIRMWARE"], defaultPortFirmware, !present["PORT_FIRMWARE"])
|
|
cfg.PortLookup = parsePort("PORT_LOOKUP", env["PORT_LOOKUP"], defaultPortLookup, !present["PORT_LOOKUP"])
|
|
cfg.PortXMPP = parsePort("PORT_XMPP", env["PORT_XMPP"], defaultPortXMPP, !present["PORT_XMPP"])
|
|
cfg.HealthPort = parsePort("HEALTH_PORT", env["HEALTH_PORT"], defaultHealthPort, !present["HEALTH_PORT"])
|
|
|
|
if !present["BIND_ADDRESS"] {
|
|
cfg.BindAddress = defaultBindAddress
|
|
} else if env["BIND_ADDRESS"] == "" || net.ParseIP(env["BIND_ADDRESS"]) == nil {
|
|
errs = append(errs, "BIND_ADDRESS must be an IP address")
|
|
} else {
|
|
cfg.BindAddress = env["BIND_ADDRESS"]
|
|
}
|
|
|
|
advRaw := strings.TrimSpace(env["ADVERTISE_IP"])
|
|
if advRaw == "" {
|
|
errs = append(errs, "ADVERTISE_IP is required")
|
|
} else {
|
|
ip := net.ParseIP(advRaw)
|
|
if ip == nil || ip.To4() == nil || ip.To4().String() != advRaw {
|
|
errs = append(errs, "ADVERTISE_IP must be a literal IPv4 address")
|
|
} else if ip.Equal(net.IPv4zero) {
|
|
errs = append(errs, "ADVERTISE_IP must not be 0.0.0.0")
|
|
} else {
|
|
cfg.AdvertiseIP = ip.To4()
|
|
}
|
|
}
|
|
|
|
hostRaw := env["MQTT_HOST"]
|
|
if hostRaw == "" {
|
|
errs = append(errs, "MQTT_HOST is required")
|
|
} else if strings.Contains(hostRaw, "://") {
|
|
errs = append(errs, "MQTT_HOST must be a host or address without a scheme or port")
|
|
} else if _, _, err := net.SplitHostPort(hostRaw); err == nil {
|
|
errs = append(errs, "MQTT_HOST must be a host or address without a scheme or port")
|
|
} else {
|
|
cfg.MQTTHost = hostRaw
|
|
}
|
|
|
|
mqttPortDefault := defaultMQTTPlainPort
|
|
if mqttTLS {
|
|
mqttPortDefault = defaultMQTTTLSPort
|
|
}
|
|
cfg.MQTTPort = parsePort("MQTT_PORT", env["MQTT_PORT"], mqttPortDefault, true)
|
|
|
|
cfg.MQTTCAFile = env["MQTT_CA_FILE"]
|
|
if cfg.MQTTCAFile != "" {
|
|
if err := validatePEMCertBundle(cfg.MQTTCAFile); err != nil {
|
|
errs = append(errs, "MQTT_CA_FILE must be a PEM certificate bundle")
|
|
}
|
|
}
|
|
|
|
cfg.MQTTUsername = env["MQTT_USERNAME"]
|
|
cfg.MQTTPassword = env["MQTT_PASSWORD"]
|
|
if cfg.MQTTPassword != "" && cfg.MQTTUsername == "" {
|
|
errs = append(errs, "MQTT_PASSWORD requires MQTT_USERNAME")
|
|
}
|
|
|
|
if !present["MQTT_CLIENT_ID"] || env["MQTT_CLIENT_ID"] == "" {
|
|
cfg.MQTTClientID = defaultMQTTClientID()
|
|
} else if !clientIDRe.MatchString(env["MQTT_CLIENT_ID"]) {
|
|
errs = append(errs, "MQTT_CLIENT_ID must match ^[A-Za-z0-9_-]{1,64}$")
|
|
} else {
|
|
cfg.MQTTClientID = env["MQTT_CLIENT_ID"]
|
|
}
|
|
|
|
if !present["MQTT_BASE"] {
|
|
cfg.MQTTBase = defaultMQTTBase
|
|
} else if env["MQTT_BASE"] == "" || !topicPartRe.MatchString(env["MQTT_BASE"]) {
|
|
errs = append(errs, "MQTT_BASE must match ^[A-Za-z0-9_-]+$")
|
|
} else {
|
|
cfg.MQTTBase = env["MQTT_BASE"]
|
|
}
|
|
|
|
if !present["HA_DISCOVERY_PREFIX"] {
|
|
cfg.HADiscoveryPrefix = defaultHADiscovery
|
|
} else if env["HA_DISCOVERY_PREFIX"] == "" || !topicPartRe.MatchString(env["HA_DISCOVERY_PREFIX"]) {
|
|
errs = append(errs, "HA_DISCOVERY_PREFIX must match ^[A-Za-z0-9_-]+$")
|
|
} else {
|
|
cfg.HADiscoveryPrefix = env["HA_DISCOVERY_PREFIX"]
|
|
}
|
|
|
|
if !present["CONTROLLER_JID"] {
|
|
cfg.ControllerJID = defaultControllerJID
|
|
} else if env["CONTROLLER_JID"] == "" || !validControllerJID(env["CONTROLLER_JID"]) {
|
|
errs = append(errs, "CONTROLLER_JID must be local@domain/resource")
|
|
} else {
|
|
cfg.ControllerJID = env["CONTROLLER_JID"]
|
|
}
|
|
|
|
if !present["RAW_COMMANDS"] {
|
|
cfg.RawCommands = false
|
|
} else if env["RAW_COMMANDS"] == "" {
|
|
errs = append(errs, "RAW_COMMANDS must be true or false")
|
|
} else {
|
|
b, ok := parseBoolStrict(env["RAW_COMMANDS"])
|
|
if !ok {
|
|
errs = append(errs, "RAW_COMMANDS must be true or false")
|
|
} else {
|
|
cfg.RawCommands = b
|
|
}
|
|
}
|
|
|
|
levelRaw := env["LOG_LEVEL"]
|
|
if !present["LOG_LEVEL"] {
|
|
levelRaw = defaultLogLevel
|
|
}
|
|
switch strings.ToLower(levelRaw) {
|
|
case "debug":
|
|
cfg.LogLevel = slog.LevelDebug
|
|
case "info":
|
|
cfg.LogLevel = slog.LevelInfo
|
|
case "warn":
|
|
cfg.LogLevel = slog.LevelWarn
|
|
case "error":
|
|
cfg.LogLevel = slog.LevelError
|
|
default:
|
|
errs = append(errs, "LOG_LEVEL must be debug, info, warn, or error")
|
|
}
|
|
|
|
if cfg.PortFirmware != 0 && cfg.PortLookup != 0 && cfg.PortXMPP != 0 && cfg.HealthPort != 0 {
|
|
if cfg.PortFirmware == cfg.PortLookup || cfg.PortFirmware == cfg.PortXMPP || cfg.PortFirmware == cfg.HealthPort ||
|
|
cfg.PortLookup == cfg.PortXMPP || cfg.PortLookup == cfg.HealthPort || cfg.PortXMPP == cfg.HealthPort {
|
|
errs = append(errs, "listener ports must be distinct")
|
|
}
|
|
}
|
|
|
|
if len(errs) > 0 {
|
|
return Config{}, fmt.Errorf("%s", strings.Join(errs, "; "))
|
|
}
|
|
return cfg, nil
|
|
}
|
|
|
|
func parseBoolStrict(s string) (bool, bool) {
|
|
switch s {
|
|
case "true":
|
|
return true, true
|
|
case "false":
|
|
return false, true
|
|
}
|
|
return false, false
|
|
}
|
|
|
|
func validatePEMCertBundle(path string) error {
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
for {
|
|
block, rest := pem.Decode(data)
|
|
if block == nil {
|
|
break
|
|
}
|
|
if block.Type == "CERTIFICATE" {
|
|
return nil
|
|
}
|
|
data = rest
|
|
}
|
|
return fmt.Errorf("no certificate block")
|
|
}
|
|
|
|
func validControllerJID(jid string) bool {
|
|
if strings.Count(jid, "@") != 1 {
|
|
return false
|
|
}
|
|
local, domainResource, _ := strings.Cut(jid, "@")
|
|
if local == "" {
|
|
return false
|
|
}
|
|
if strings.Count(domainResource, "/") != 1 {
|
|
return false
|
|
}
|
|
domain, resource, _ := strings.Cut(domainResource, "/")
|
|
return domain != "" && resource != ""
|
|
}
|
|
|
|
func defaultMQTTClientID() string {
|
|
host, err := os.Hostname()
|
|
if err != nil {
|
|
host = "host"
|
|
}
|
|
if i := strings.IndexByte(host, '.'); i >= 0 {
|
|
host = host[:i]
|
|
}
|
|
host = strings.ToLower(host)
|
|
var b strings.Builder
|
|
for i := 0; i < len(host); i++ {
|
|
c := host[i]
|
|
if (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') {
|
|
b.WriteByte(c)
|
|
} else {
|
|
b.WriteByte('-')
|
|
}
|
|
}
|
|
hostPart := b.String()
|
|
hostPart = collapseTrim(hostPart, '-')
|
|
if hostPart == "" {
|
|
hostPart = "host"
|
|
}
|
|
prefix := "n95bridge-"
|
|
result := prefix + hostPart
|
|
if len(result) > 64 {
|
|
maxHost := 64 - len(prefix)
|
|
if maxHost < 0 {
|
|
maxHost = 0
|
|
}
|
|
if len(hostPart) > maxHost {
|
|
hostPart = hostPart[len(hostPart)-maxHost:]
|
|
}
|
|
hostPart = strings.TrimLeft(hostPart, "-")
|
|
result = prefix + hostPart
|
|
}
|
|
return result
|
|
}
|
|
|
|
func collapseTrim(s string, ch byte) string {
|
|
var b strings.Builder
|
|
prev := byte(0)
|
|
for i := 0; i < len(s); i++ {
|
|
c := s[i]
|
|
if c == ch {
|
|
if prev == ch {
|
|
continue
|
|
}
|
|
}
|
|
b.WriteByte(c)
|
|
prev = c
|
|
}
|
|
return strings.Trim(b.String(), string(ch))
|
|
}
|
|
|
|
// String returns a redacted summary that omits MQTTPassword.
|
|
func (c Config) String() string {
|
|
var b strings.Builder
|
|
fmt.Fprintf(&b, "PortFirmware=%d PortLookup=%d PortXMPP=%d HealthPort=%d ", c.PortFirmware, c.PortLookup, c.PortXMPP, c.HealthPort)
|
|
fmt.Fprintf(&b, "BindAddress=%q AdvertiseIP=%s ", c.BindAddress, c.AdvertiseIP)
|
|
fmt.Fprintf(&b, "MQTTHost=%s MQTTPort=%d MQTTTLS=%t MQTTCAFile=%q MQTTUsername=%q ", c.MQTTHost, c.MQTTPort, c.MQTTTLS, c.MQTTCAFile, c.MQTTUsername)
|
|
fmt.Fprintf(&b, "MQTTClientID=%q MQTTBase=%q HADiscoveryPrefix=%q ControllerJID=%q ", c.MQTTClientID, c.MQTTBase, c.HADiscoveryPrefix, c.ControllerJID)
|
|
fmt.Fprintf(&b, "RawCommands=%t LogLevel=%s", c.RawCommands, c.LogLevel)
|
|
return b.String()
|
|
}
|
|
|
|
// Format redirects all fmt verbs to the redacted String.
|
|
func (c Config) Format(s fmt.State, verb rune) {
|
|
_, _ = fmt.Fprint(s, c.String())
|
|
}
|
|
|
|
// GoString returns the redacted String.
|
|
func (c Config) GoString() string { return c.String() }
|
|
|
|
// LogValue returns a slog group that omits MQTTPassword.
|
|
func (c Config) LogValue() slog.Value {
|
|
return slog.GroupValue(
|
|
slog.Int("PortFirmware", c.PortFirmware),
|
|
slog.Int("PortLookup", c.PortLookup),
|
|
slog.Int("PortXMPP", c.PortXMPP),
|
|
slog.Int("HealthPort", c.HealthPort),
|
|
slog.String("BindAddress", c.BindAddress),
|
|
slog.String("AdvertiseIP", c.AdvertiseIP.String()),
|
|
slog.String("MQTTHost", c.MQTTHost),
|
|
slog.Int("MQTTPort", c.MQTTPort),
|
|
slog.Bool("MQTTTLS", c.MQTTTLS),
|
|
slog.String("MQTTCAFile", c.MQTTCAFile),
|
|
slog.String("MQTTUsername", c.MQTTUsername),
|
|
slog.String("MQTTClientID", c.MQTTClientID),
|
|
slog.String("MQTTBase", c.MQTTBase),
|
|
slog.String("HADiscoveryPrefix", c.HADiscoveryPrefix),
|
|
slog.String("ControllerJID", c.ControllerJID),
|
|
slog.Bool("RawCommands", c.RawCommands),
|
|
slog.String("LogLevel", c.LogLevel.String()),
|
|
)
|
|
}
|