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

384 lines
9.4 KiB
Go

package config
import (
"bytes"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"fmt"
"log/slog"
"math/big"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
func validEnv() []string {
return []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
}
}
func TestLoadValidMinimal(t *testing.T) {
cfg, err := Load(validEnv())
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if cfg.PortFirmware != defaultPortFirmware {
t.Errorf("PortFirmware = %d, want %d", cfg.PortFirmware, defaultPortFirmware)
}
if cfg.PortLookup != defaultPortLookup {
t.Errorf("PortLookup = %d, want %d", cfg.PortLookup, defaultPortLookup)
}
if cfg.PortXMPP != defaultPortXMPP {
t.Errorf("PortXMPP = %d, want %d", cfg.PortXMPP, defaultPortXMPP)
}
if cfg.HealthPort != defaultHealthPort {
t.Errorf("HealthPort = %d, want %d", cfg.HealthPort, defaultHealthPort)
}
if cfg.BindAddress != defaultBindAddress {
t.Errorf("BindAddress = %q, want %q", cfg.BindAddress, defaultBindAddress)
}
if cfg.MQTTPort != defaultMQTTPlainPort {
t.Errorf("MQTTPort = %d, want %d", cfg.MQTTPort, defaultMQTTPlainPort)
}
if cfg.MQTTTLS {
t.Error("MQTTTLS should be false")
}
if cfg.MQTTBase != defaultMQTTBase {
t.Errorf("MQTTBase = %q, want %q", cfg.MQTTBase, defaultMQTTBase)
}
if cfg.HADiscoveryPrefix != defaultHADiscovery {
t.Errorf("HADiscoveryPrefix = %q, want %q", cfg.HADiscoveryPrefix, defaultHADiscovery)
}
if cfg.ControllerJID != defaultControllerJID {
t.Errorf("ControllerJID = %q, want %q", cfg.ControllerJID, defaultControllerJID)
}
if cfg.RawCommands {
t.Error("RawCommands should be false")
}
if cfg.LogLevel != slog.LevelInfo {
t.Errorf("LogLevel = %v, want info", cfg.LogLevel)
}
if cfg.MQTTClientID == "" || !clientIDRe.MatchString(cfg.MQTTClientID) {
t.Errorf("MQTTClientID = %q, want hostname-derived id matching regex", cfg.MQTTClientID)
}
if !strings.HasPrefix(cfg.MQTTClientID, "n95bridge-") {
t.Errorf("MQTTClientID = %q, want prefix n95bridge-", cfg.MQTTClientID)
}
}
func TestLoadPortsAndTLS(t *testing.T) {
cfg, err := Load([]string{
"PORT_FIRMWARE=9005",
"PORT_LOOKUP=9007",
"PORT_XMPP=9223",
"HEALTH_PORT=9080",
"BIND_ADDRESS=127.0.0.1",
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_TLS=true",
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if cfg.PortFirmware != 9005 || cfg.PortLookup != 9007 || cfg.PortXMPP != 9223 || cfg.HealthPort != 9080 {
t.Errorf("ports mismatch: %+v", cfg)
}
if cfg.BindAddress != "127.0.0.1" {
t.Errorf("BindAddress = %q", cfg.BindAddress)
}
if cfg.MQTTPort != defaultMQTTTLSPort {
t.Errorf("MQTTPort with TLS = %d, want %d", cfg.MQTTPort, defaultMQTTTLSPort)
}
if !cfg.MQTTTLS {
t.Error("MQTTTLS should be true")
}
}
func TestLoadPortDefaultFollowsTLS(t *testing.T) {
cfg, err := Load([]string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_TLS=true",
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if cfg.MQTTPort != 8883 {
t.Errorf("MQTTPort = %d, want 8883", cfg.MQTTPort)
}
cfg, err = Load([]string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_PORT=9999",
"MQTT_TLS=true",
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if cfg.MQTTPort != 9999 {
t.Errorf("MQTTPort = %d, want 9999", cfg.MQTTPort)
}
}
func TestLoadValidationMessages(t *testing.T) {
tmp := t.TempDir()
badPEM := filepath.Join(tmp, "bad.pem")
if err := os.WriteFile(badPEM, []byte("not a certificate"), 0600); err != nil {
t.Fatal(err)
}
goodPEM := filepath.Join(tmp, "good.pem")
writeTestCertificate(t, goodPEM)
cases := []struct {
name string
env []string
wantAll []string
}{
{
name: "ports invalid",
env: []string{
"PORT_FIRMWARE=abc",
"PORT_LOOKUP=0",
"PORT_XMPP=99999",
"HEALTH_PORT=-1",
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
},
wantAll: []string{
"PORT_FIRMWARE must be an integer from 1 to 65535",
"PORT_LOOKUP must be an integer from 1 to 65535",
"PORT_XMPP must be an integer from 1 to 65535",
"HEALTH_PORT must be an integer from 1 to 65535",
},
},
{
name: "advertise",
env: []string{
"ADVERTISE_IP=0.0.0.0",
"MQTT_HOST=mqtt.example.invalid",
},
wantAll: []string{
"ADVERTISE_IP must not be 0.0.0.0",
},
},
{
name: "mqtt host",
env: []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=tcp://mqtt.example.invalid",
},
wantAll: []string{
"MQTT_HOST must be a host or address without a scheme or port",
},
},
{
name: "mqtt ca",
env: []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_CA_FILE=" + badPEM,
},
wantAll: []string{
"MQTT_CA_FILE must be a PEM certificate bundle",
},
},
{
name: "password without username",
env: []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_PASSWORD=secret",
},
wantAll: []string{
"MQTT_PASSWORD requires MQTT_USERNAME",
},
},
{
name: "client id",
env: []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_CLIENT_ID=bad id!",
},
wantAll: []string{
"MQTT_CLIENT_ID must match ^[A-Za-z0-9_-]{1,64}$",
},
},
{
name: "topics jid",
env: []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_BASE=",
"HA_DISCOVERY_PREFIX=home assistant",
"CONTROLLER_JID=nobody",
},
wantAll: []string{
"MQTT_BASE must match ^[A-Za-z0-9_-]+$",
"HA_DISCOVERY_PREFIX must match ^[A-Za-z0-9_-]+$",
"CONTROLLER_JID must be local@domain/resource",
},
},
{
name: "log and raw",
env: []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"RAW_COMMANDS=maybe",
"LOG_LEVEL=verbose",
},
wantAll: []string{
"RAW_COMMANDS must be true or false",
"LOG_LEVEL must be debug, info, warn, or error",
},
},
{
name: "duplicate ports",
env: []string{
"PORT_FIRMWARE=8005",
"PORT_LOOKUP=8005",
"PORT_XMPP=5223",
"HEALTH_PORT=8080",
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
},
wantAll: []string{
"listener ports must be distinct",
},
},
{
name: "ca ok",
env: []string{
"ADVERTISE_IP=192.0.2.10",
"MQTT_HOST=mqtt.example.invalid",
"MQTT_CA_FILE=" + goodPEM,
},
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
_, err := Load(tc.env)
if len(tc.wantAll) == 0 {
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
return
}
if err == nil {
t.Fatalf("expected error containing %q, got nil", tc.wantAll)
}
got := err.Error()
for _, want := range tc.wantAll {
if !strings.Contains(got, want) {
t.Errorf("error %q does not contain %q", got, want)
}
}
})
}
}
func TestLoadAdvertiseIPValidation(t *testing.T) {
cases := []struct {
name string
value string
wantErr string
}{
{"missing", "", "ADVERTISE_IP is required"},
{"not v4", "::1", "ADVERTISE_IP must be a literal IPv4 address"},
{"with port", "192.0.2.10:5223", "ADVERTISE_IP must be a literal IPv4 address"},
{"zero", "0.0.0.0", "ADVERTISE_IP must not be 0.0.0.0"},
{"loopback ok", "127.0.0.1", ""},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
env := []string{"MQTT_HOST=mqtt.example.invalid"}
if tc.value != "" {
env = append(env, "ADVERTISE_IP="+tc.value)
}
_, err := Load(env)
if tc.wantErr == "" {
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
return
}
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
t.Fatalf("expected error containing %q, got %v", tc.wantErr, err)
}
})
}
}
func TestConfigPasswordNotLogged(t *testing.T) {
env := append(validEnv(), "MQTT_USERNAME=user", "MQTT_PASSWORD=top-secret-123")
cfg, err := Load(env)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if cfg.MQTTPassword != "top-secret-123" {
t.Fatalf("password not stored")
}
for _, s := range []string{cfg.String(), fmt.Sprintf("%v", cfg), fmt.Sprintf("%+v", cfg), fmt.Sprintf("%#v", cfg)} {
if strings.Contains(s, "top-secret") {
t.Errorf("redacted output contains password: %q", s)
}
}
var buf bytes.Buffer
handler := slog.NewTextHandler(&buf, nil)
logger := slog.New(handler)
logger.Info("test", "config", cfg)
if strings.Contains(buf.String(), "top-secret") {
t.Errorf("log output contains password: %s", buf.String())
}
}
func TestDefaultMQTTClientID(t *testing.T) {
id := defaultMQTTClientID()
if !strings.HasPrefix(id, "n95bridge-") {
t.Errorf("defaultMQTTClientID() = %q, want prefix n95bridge-", id)
}
if !clientIDRe.MatchString(id) {
t.Errorf("defaultMQTTClientID() = %q, does not match regex", id)
}
if len(id) > 64 {
t.Errorf("defaultMQTTClientID() length %d > 64", len(id))
}
}
func writeTestCertificate(t *testing.T, path string) {
t.Helper()
key, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
tmpl := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "test"},
NotBefore: time.Now(),
NotAfter: time.Now().Add(time.Hour),
}
der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
if err != nil {
t.Fatal(err)
}
f, err := os.Create(path)
if err != nil {
t.Fatal(err)
}
defer f.Close()
if err := pem.Encode(f, &pem.Block{Type: "CERTIFICATE", Bytes: der}); err != nil {
t.Fatal(err)
}
}