Files
2026-09-24 09:52:57 +01:00

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()),
)
}