384 lines
9.4 KiB
Go
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)
|
|
}
|
|
}
|