Files

1140 lines
32 KiB
Go

package mqttbridge
import (
"bytes"
"context"
"crypto/ed25519"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/json"
"encoding/pem"
"encoding/xml"
"log/slog"
"math/big"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
mqtt "github.com/eclipse/paho.mqtt.golang"
"git.i3omb.com/gronod/ha-n95-local-control/internal/config"
"git.i3omb.com/gronod/ha-n95-local-control/internal/ha"
"git.i3omb.com/gronod/ha-n95-local-control/internal/robot"
"git.i3omb.com/gronod/ha-n95-local-control/internal/session"
)
type fakeToken struct{ err error }
func (t fakeToken) Wait() bool { return true }
func (t fakeToken) WaitTimeout(time.Duration) bool { return true }
func (t fakeToken) Done() <-chan struct{} {
ch := make(chan struct{})
close(ch)
return ch
}
func (t fakeToken) Error() error { return t.err }
type fakeOp struct {
kind string // "connect", "sub", "pub", "disconnect"
topic string
qos byte
retain bool
payload []byte
}
type fakeClient struct {
mu sync.Mutex
opts *mqtt.ClientOptions
ops []fakeOp
beforePublish func(topic string, payload []byte)
}
func (c *fakeClient) Connect() mqtt.Token {
c.mu.Lock()
c.ops = append(c.ops, fakeOp{kind: "connect"})
c.mu.Unlock()
if h := c.opts.OnConnect; h != nil {
h(nil)
}
return fakeToken{}
}
func (c *fakeClient) Disconnect(uint) {
c.mu.Lock()
defer c.mu.Unlock()
c.ops = append(c.ops, fakeOp{kind: "disconnect"})
}
func (c *fakeClient) Subscribe(topic string, qos byte, _ mqtt.MessageHandler) mqtt.Token {
c.mu.Lock()
defer c.mu.Unlock()
c.ops = append(c.ops, fakeOp{kind: "sub", topic: topic, qos: qos})
return fakeToken{}
}
func (c *fakeClient) Publish(topic string, qos byte, retained bool, payload interface{}) mqtt.Token {
var p []byte
switch v := payload.(type) {
case []byte:
p = append([]byte(nil), v...)
case string:
p = []byte(v)
}
if c.beforePublish != nil {
c.beforePublish(topic, p)
}
c.mu.Lock()
defer c.mu.Unlock()
c.ops = append(c.ops, fakeOp{kind: "pub", topic: topic, qos: qos, retain: retained, payload: p})
return fakeToken{}
}
func (c *fakeClient) IsConnected() bool { return true }
func (c *fakeClient) reconnect() {
if h := c.opts.OnConnect; h != nil {
h(nil)
}
}
func (c *fakeClient) deliver(m fakeMsg) {
if h := c.opts.DefaultPublishHandler; h != nil {
h(nil, m)
}
}
func (c *fakeClient) pubs(topic string) []fakeOp {
c.mu.Lock()
defer c.mu.Unlock()
var out []fakeOp
for _, op := range c.ops {
if op.kind == "pub" && op.topic == topic {
out = append(out, op)
}
}
return out
}
func (c *fakeClient) opsSnapshot() []fakeOp {
c.mu.Lock()
defer c.mu.Unlock()
return append([]fakeOp(nil), c.ops...)
}
type fakeMsg struct {
topic string
payload []byte
dup bool
id uint16
}
func (m fakeMsg) Duplicate() bool { return m.dup }
func (m fakeMsg) Qos() byte { return 1 }
func (m fakeMsg) Retained() bool { return false }
func (m fakeMsg) Topic() string { return m.topic }
func (m fakeMsg) MessageID() uint16 { return m.id }
func (m fakeMsg) Payload() []byte { return m.payload }
func (m fakeMsg) Ack() {}
type harness struct {
b *Bridge
cfg config.Config
mu sync.Mutex
clients []*fakeClient
}
func newHarness(t *testing.T, submit SubmitFunc) *harness {
t.Helper()
cfg := config.Config{
MQTTHost: "mqtt.test",
MQTTPort: 1883,
MQTTBase: "ecovacs",
HADiscoveryPrefix: "homeassistant",
MQTTClientID: "n95-test",
}
b, err := New(context.Background(), cfg, submit)
if err != nil {
t.Fatalf("New: %v", err)
}
h := &harness{b: b, cfg: cfg}
b.factory = func(o *mqtt.ClientOptions) mqttClient {
c := &fakeClient{opts: o}
h.mu.Lock()
h.clients = append(h.clients, c)
h.mu.Unlock()
return c
}
return h
}
func (h *harness) client(t *testing.T) *fakeClient {
t.Helper()
h.mu.Lock()
defer h.mu.Unlock()
if len(h.clients) != 1 {
t.Fatalf("clients = %d, want exactly 1", len(h.clients))
}
return h.clients[0]
}
func noopSubmit(context.Context, string, robot.Command) error { return nil }
func TestMQTTAndMQTTSConfig(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
o := c.opts
if got := o.Servers[0].String(); got != "tcp://mqtt.test:1883" {
t.Fatalf("broker = %q, want tcp://mqtt.test:1883", got)
}
if o.TLSConfig != nil {
t.Fatal("non-TLS options must have nil TLSConfig")
}
if o.ClientID != "n95-test-SER1" {
t.Fatalf("ClientID = %q, want n95-test-SER1", o.ClientID)
}
if !o.CleanSession || !o.AutoReconnect || !o.ConnectRetry || o.Order {
t.Fatalf("session flags wrong: %+v", o)
}
if o.ProtocolVersion != 4 {
t.Fatalf("ProtocolVersion = %d, want 4 (MQTT 3.1.1)", o.ProtocolVersion)
}
if o.ConnectTimeout != connectTimeout {
t.Fatalf("ConnectTimeout = %v, want %v", o.ConnectTimeout, connectTimeout)
}
if o.Username != "" || o.Password != "" {
t.Fatal("credentials must be empty when MQTT_USERNAME is unset")
}
cfg := config.Config{
MQTTHost: "mqtt.test",
MQTTPort: 8883,
MQTTTLS: true,
MQTTCAFile: writeTestCA(t),
MQTTUsername: "u1",
MQTTPassword: "s3cret",
MQTTBase: "ecovacs",
HADiscoveryPrefix: "homeassistant",
MQTTClientID: "n95-test",
}
b2, err := New(context.Background(), cfg, noopSubmit)
if err != nil {
t.Fatalf("New TLS: %v", err)
}
var tlsClients []*fakeClient
b2.factory = func(o *mqtt.ClientOptions) mqttClient {
fc := &fakeClient{opts: o}
tlsClients = append(tlsClients, fc)
return fc
}
var logBuf bytes.Buffer
prevLogger := slog.Default()
slog.SetDefault(slog.New(slog.NewTextHandler(&logBuf, nil)))
defer slog.SetDefault(prevLogger)
b2.Republisher("SER1")
to := tlsClients[0].opts
connectLog := logBuf.String()
for _, want := range []string{"mqtt connect", "host=mqtt.test", "port=8883", "tls=true", "client_id=n95-test-SER1", "username=u1"} {
if !strings.Contains(connectLog, want) {
t.Fatalf("connect log %q missing %q", connectLog, want)
}
}
if strings.Contains(connectLog, "s3cret") || strings.Contains(connectLog, "ClientOptions") {
t.Fatalf("connect log leaks credentials or options struct: %q", connectLog)
}
if got := to.Servers[0].String(); got != "tls://mqtt.test:8883" {
t.Fatalf("TLS broker = %q, want tls://mqtt.test:8883", got)
}
tc := to.TLSConfig
if tc == nil {
t.Fatal("TLS options must have TLSConfig")
}
if tc.ServerName != "mqtt.test" {
t.Fatalf("ServerName = %q, want mqtt.test", tc.ServerName)
}
if tc.MinVersion != tls.VersionTLS12 {
t.Fatalf("MinVersion = %x, want TLS1.2", tc.MinVersion)
}
if tc.InsecureSkipVerify {
t.Fatal("InsecureSkipVerify must never be set")
}
if tc.RootCAs == nil || len(tc.RootCAs.Subjects()) == 0 {
t.Fatal("RootCAs must include the configured CA")
}
if to.Username != "u1" || to.Password != "s3cret" {
t.Fatal("username/password not applied")
}
v6 := config.Config{MQTTHost: "2001:db8::1", MQTTPort: 1883}
bv6, err := New(context.Background(), v6, noopSubmit)
if err != nil {
t.Fatalf("New IPv6: %v", err)
}
if bv6.broker != "tcp://[2001:db8::1]:1883" {
t.Fatalf("IPv6 broker = %q, want tcp://[2001:db8::1]:1883", bv6.broker)
}
longPrefix := strings.Repeat("a", 130)
id := clientID(longPrefix, "SER1")
if !strings.HasPrefix(id, "n95-") || len(id) != 40 {
t.Fatalf("hashed client id = %q (len %d)", id, len(id))
}
if clientID(longPrefix, "SER1") != id {
t.Fatal("hashed client id must be deterministic")
}
if clientID(longPrefix, "SER2") == id {
t.Fatal("hashed client id must differ per serial")
}
badCfg := cfg
badCfg.MQTTCAFile = filepath.Join(t.TempDir(), "missing.pem")
if _, err := New(context.Background(), badCfg, noopSubmit); err == nil {
t.Fatal("New must fail when MQTT_CA_FILE cannot be read")
}
}
func writeTestCA(t *testing.T) string {
t.Helper()
pub, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatal(err)
}
tmpl := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "test-ca"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
IsCA: true,
KeyUsage: x509.KeyUsageCertSign,
BasicConstraintsValid: true,
}
der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, pub, priv)
if err != nil {
t.Fatal(err)
}
path := filepath.Join(t.TempDir(), "ca.pem")
if err := os.WriteFile(path, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o600); err != nil {
t.Fatal(err)
}
return path
}
func TestSubscribeBeforeDiscovery(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
ops := c.opsSnapshot()
var kinds []string
for _, op := range ops {
label := op.kind
if op.topic != "" {
label += ":" + op.topic
}
kinds = append(kinds, label)
}
want := []string{
"connect",
"sub:ecovacs/SER1/command",
"sub:ecovacs/SER1/set_fan_speed",
"sub:ecovacs/SER1/send_command",
"sub:homeassistant/status",
"pub:homeassistant/vacuum/ecovacs_SER1/config",
"pub:ecovacs/SER1/availability",
"pub:ecovacs/SER1/state",
"pub:ecovacs/SER1/json_attributes",
}
if len(kinds) != len(want) {
t.Fatalf("ops = %v, want %v", kinds, want)
}
for i := range want {
if kinds[i] != want[i] {
t.Fatalf("ops[%d] = %q, want %q (full: %v)", i, kinds[i], want[i], kinds)
}
}
for _, q := range []byte{ops[1].qos, ops[2].qos, ops[3].qos} {
if q != 1 {
t.Fatalf("command subscription qos = %d, want 1", q)
}
}
if ops[4].qos != 0 {
t.Fatalf("status subscription qos = %d, want 0", ops[4].qos)
}
if string(ops[6].payload) != "offline" {
t.Fatalf("initial availability = %q, want offline", ops[6].payload)
}
}
func TestDiscoveryRetainedPrefix(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
pubs := c.pubs("homeassistant/vacuum/ecovacs_SER1/config")
if len(pubs) != 1 {
t.Fatalf("discovery publishes = %d, want 1", len(pubs))
}
if !pubs[0].retain || pubs[0].qos != 0 {
t.Fatalf("discovery publish retain=%v qos=%d, want retained qos0", pubs[0].retain, pubs[0].qos)
}
var doc map[string]any
if err := json.Unmarshal(pubs[0].payload, &doc); err != nil {
t.Fatalf("discovery payload invalid: %v", err)
}
if doc["unique_id"] != "ecovacs_SER1" {
t.Fatalf("unique_id = %v", doc["unique_id"])
}
}
func statePub(t *testing.T, c *fakeClient, serial string) map[string]string {
t.Helper()
pubs := c.pubs("ecovacs/" + serial + "/state")
if len(pubs) == 0 {
t.Fatal("no state publish")
}
var doc map[string]string
if err := json.Unmarshal(pubs[len(pubs)-1].payload, &doc); err != nil {
t.Fatalf("state JSON invalid: %v", err)
}
return doc
}
func TestStateJSONAlwaysValid(t *testing.T) {
h := newHarness(t, noopSubmit)
repub := h.b.Republisher("SER1")
c := h.client(t)
states := []string{"idle", "cleaning", "returning", "docked", "error"}
for _, st := range states {
snap := robot.NewSnapshot()
snap.State.State = st
snap.State.FanSpeed = "standard"
repub(context.Background(), snap)
doc := statePub(t, c, "SER1")
if doc["state"] != st {
t.Fatalf("published state = %q, want %q", doc["state"], st)
}
if doc["fan_speed"] != "standard" {
t.Fatalf("published fan_speed = %q", doc["fan_speed"])
}
}
pubs := c.pubs("ecovacs/SER1/state")
last := pubs[len(pubs)-1]
if string(last.payload) != `{"state":"error","fan_speed":"standard"}` {
t.Fatalf("state payload = %s", last.payload)
}
if !last.retain || last.qos != 0 {
t.Fatal("state publish must be retained qos0")
}
for _, charge := range []string{"", "Idle", "going", "SlotCharging"} {
for _, clean := range []string{"", "auto", "spot", "stop"} {
for _, latched := range []bool{false, true} {
got := robot.Derive(robot.Facts{Charge: charge, CleanType: clean, ErrorLatched: latched})
if got == "paused" {
t.Fatalf("Derive produced paused for %+v", charge)
}
}
}
}
}
func TestFanSpeedListMatchesState(t *testing.T) {
h := newHarness(t, noopSubmit)
repub := h.b.Republisher("SER1")
c := h.client(t)
pubs := c.pubs("homeassistant/vacuum/ecovacs_SER1/config")
if len(pubs) != 1 {
t.Fatalf("discovery publishes = %d", len(pubs))
}
var doc map[string]any
if err := json.Unmarshal(pubs[0].payload, &doc); err != nil {
t.Fatal(err)
}
list, ok := doc["fan_speed_list"].([]any)
if !ok || len(list) != 2 || list[0] != "standard" || list[1] != "strong" {
t.Fatalf("fan_speed_list = %v, want [standard strong]", doc["fan_speed_list"])
}
snap := robot.NewSnapshot()
snap.State.State = "cleaning"
snap.State.FanSpeed = "strong"
repub(context.Background(), snap)
statePubs := c.pubs("ecovacs/SER1/state")
last := statePubs[len(statePubs)-1]
if !last.retain {
t.Fatal("state publish must be retained")
}
var state map[string]string
if err := json.Unmarshal(last.payload, &state); err != nil {
t.Fatal(err)
}
if state["fan_speed"] != "strong" {
t.Fatalf("republished fan_speed = %q, want strong", state["fan_speed"])
}
}
func TestRetainAvailabilityAndLWT(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
o := c.opts
if !o.WillEnabled || !o.WillRetained || o.WillQos != 0 {
t.Fatalf("LWT flags wrong: enabled=%v retained=%v qos=%d", o.WillEnabled, o.WillRetained, o.WillQos)
}
if o.WillTopic != "ecovacs/SER1/availability" || string(o.WillPayload) != "offline" {
t.Fatalf("LWT topic/payload = %q %q", o.WillTopic, o.WillPayload)
}
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 1})
pubs := c.pubs("ecovacs/SER1/availability")
if len(pubs) != 2 || string(pubs[1].payload) != "online" || !pubs[1].retain {
t.Fatalf("availability pubs after AnnounceOK = %+v", pubs)
}
h.b.SessionDown(session.DownEvent{JID: "b@x/r", Serial: "SER1", Generation: 1})
pubs = c.pubs("ecovacs/SER1/availability")
if len(pubs) != 3 || string(pubs[2].payload) != "offline" || !pubs[2].retain {
t.Fatalf("availability pubs after SessionDown = %+v", pubs)
}
}
func TestRetainAvailabilityAndNotCommands(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
for _, suffix := range []string{"command", "set_fan_speed", "send_command"} {
if pubs := c.pubs("ecovacs/SER1/" + suffix); len(pubs) != 0 {
t.Fatalf("bridge must never publish to command topic %s", suffix)
}
}
for _, op := range c.opsSnapshot() {
if op.kind == "sub" && strings.HasPrefix(op.topic, "ecovacs/") && op.qos != 1 {
t.Fatalf("command subscription %s qos %d, want 1", op.topic, op.qos)
}
}
}
func TestBirthRepublish(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
disc := "homeassistant/vacuum/ecovacs_SER1/config"
c.deliver(fakeMsg{topic: "homeassistant/status", payload: []byte("offline")})
if n := len(c.pubs(disc)); n != 1 {
t.Fatalf("offline status republished discovery (n=%d)", n)
}
c.deliver(fakeMsg{topic: "homeassistant/status", payload: []byte("online")})
pubs := c.pubs(disc)
if len(pubs) != 2 || !pubs[1].retain {
t.Fatalf("birth message must republish retained discovery, pubs=%d", len(pubs))
}
}
func TestReconnectRepublish(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 1})
c.reconnect()
var subs, pubs int
var avail []fakeOp
for _, op := range c.opsSnapshot() {
switch op.kind {
case "sub":
subs++
case "pub":
pubs++
if op.topic == "ecovacs/SER1/availability" {
avail = append(avail, op)
}
}
}
if subs != 8 {
t.Fatalf("subscriptions after reconnect = %d, want 8 (two rounds of 4)", subs)
}
if pubs != 4+1+4 { // first connect 4, announce online, reconnect 4
t.Fatalf("publishes = %d, want 9", pubs)
}
last := avail[len(avail)-1]
if string(last.payload) != "online" || !last.retain {
t.Fatalf("reconnect availability = %q retain=%v, want retained online", last.payload, last.retain)
}
}
type outboundIQ struct {
From string `xml:"from,attr"`
To string `xml:"to,attr"`
Type string `xml:"type,attr"`
ID string `xml:"id,attr"`
Ctl struct {
TD string `xml:"td,attr"`
ID string `xml:"id,attr"`
Sid string `xml:"sid,attr"`
Speed string `xml:"speed,attr"`
Inner string `xml:",innerxml"`
} `xml:"query>ctl"`
}
const controllerJID = "app@im.ecovacs.local/app"
const botJID = "ser1@im.ecovacs.local/resource"
func TestSection11CommandMapping(t *testing.T) {
var mu sync.Mutex
var sent [][]byte
send := func(b []byte) error {
mu.Lock()
defer mu.Unlock()
sent = append(sent, append([]byte(nil), b...))
return nil
}
lastSent := func() []byte {
mu.Lock()
defer mu.Unlock()
return sent[len(sent)-1]
}
sentCount := func() int {
mu.Lock()
defer mu.Unlock()
return len(sent)
}
fleet := robot.NewFleet(context.Background(), controllerJID, nil, nil, nil, nil)
fleet.SessionReady(session.ReadyEvent{
JID: botJID, Serial: "SER1",
Generation: 1, Send: send,
})
ctx := context.Background()
for _, tc := range []struct {
suffix, payload string
wantTD string
wantInner string
wantCtlSid string
wantCtlSpeed string
}{
{"command", "start", "Clean", `<clean type="auto" speed="standard" act="s"/>`, "", ""},
{"command", "stop", "Clean", `<clean type="stop" speed="standard" act="h"/>`, "", ""},
{"command", "return_to_base", "Charge", `<charge type="go"/>`, "", ""},
{"command", "clean_spot", "Clean", `<clean type="spot" speed="standard" act="s"/>`, "", ""},
{"command", "locate", "PlaySound", "", "0", ""},
{"set_fan_speed", "standard", "SetCleanSpeed", "", "", "standard"},
{"set_fan_speed", "strong", "SetCleanSpeed", "", "", "strong"},
} {
before := sentCount()
cmd := ha.Command(tc.suffix, []byte(tc.payload))
if err := fleet.Submit(ctx, "SER1", cmd); err != nil {
t.Fatalf("Submit(%q, %q): %v", tc.suffix, tc.payload, err)
}
if sentCount() != before+1 {
t.Fatalf("Submit(%q, %q) produced %d sends", tc.suffix, tc.payload, sentCount()-before)
}
var iq outboundIQ
if err := xml.Unmarshal(lastSent(), &iq); err != nil {
t.Fatalf("Submit(%q, %q) produced malformed stanza: %v", tc.suffix, tc.payload, err)
}
if iq.From != controllerJID {
t.Fatalf("iq from = %q, want %q", iq.From, controllerJID)
}
if iq.To != botJID {
t.Fatalf("iq to = %q, want %q", iq.To, botJID)
}
if iq.Type != "set" {
t.Fatalf("iq type = %q, want set", iq.Type)
}
if iq.ID == "" {
t.Fatal("iq stanza id empty")
}
if iq.Ctl.ID == "" {
t.Fatal("ctl id empty")
}
if iq.Ctl.TD != tc.wantTD {
t.Fatalf("ctl td = %q, want %q", iq.Ctl.TD, tc.wantTD)
}
if iq.Ctl.Inner != tc.wantInner {
t.Fatalf("ctl inner = %q, want %q", iq.Ctl.Inner, tc.wantInner)
}
if tc.wantCtlSid != "" && iq.Ctl.Sid != tc.wantCtlSid {
t.Fatalf("ctl sid = %q, want %q", iq.Ctl.Sid, tc.wantCtlSid)
}
if tc.wantCtlSpeed != "" && iq.Ctl.Speed != tc.wantCtlSpeed {
t.Fatalf("ctl speed = %q, want %q", iq.Ctl.Speed, tc.wantCtlSpeed)
}
}
for _, tc := range []struct {
suffix, payload, wantErr string
}{
{"command", "pause", "rejected:pause"},
{"command", "junk", "rejected:junk"},
{"set_fan_speed", "turbo", "rejected:turbo"},
{"weird", "x", "rejected:unsupported"},
} {
before := sentCount()
err := fleet.Submit(ctx, "SER1", ha.Command(tc.suffix, []byte(tc.payload)))
if err == nil || err.Error() != tc.wantErr {
t.Fatalf("Submit(%q, %q) err = %v, want %q", tc.suffix, tc.payload, err, tc.wantErr)
}
if sentCount() != before {
t.Fatalf("rejected command %q wrote XMPP", tc.payload)
}
}
before := sentCount()
err := fleet.Submit(ctx, "SER1", ha.Command("send_command", []byte(`{"command":"x"}`)))
if err == nil || err.Error() != "rejected:unsupported" {
t.Fatalf("send_command without hook err = %v", err)
}
var hookPayload []byte
fleet.SetSendCommandFunc(func(_ context.Context, serial string, payload []byte) (robot.Command, error) {
if serial != "SER1" {
t.Errorf("hook serial = %q", serial)
}
hookPayload = append([]byte(nil), payload...)
return robot.Command{Name: "locate"}, nil
})
if err := fleet.Submit(ctx, "SER1", ha.Command("send_command", []byte(`{"command":"move"}`))); err != nil {
t.Fatalf("send_command with hook: %v", err)
}
if string(hookPayload) != `{"command":"move"}` {
t.Fatalf("hook payload = %q", hookPayload)
}
if sentCount() <= before {
t.Fatal("send_command MUST produce actor XMPP output now")
}
}
func TestOfflineCommandDroppedNotReplayed(t *testing.T) {
var mu sync.Mutex
var sent [][]byte
var snaps []robot.Snapshot
send := func(b []byte) error {
mu.Lock()
defer mu.Unlock()
sent = append(sent, append([]byte(nil), b...))
return nil
}
republish := func(_ context.Context, s robot.Snapshot) {
mu.Lock()
defer mu.Unlock()
snaps = append(snaps, s)
}
fleet := robot.NewFleet(context.Background(), "app@im.ecovacs.local/app", republish, nil, nil, nil)
fleet.SessionReady(session.ReadyEvent{
JID: "ser1@im.ecovacs.local/resource", Serial: "SER1",
Generation: 1, Send: send,
})
fleet.SessionDown(session.DownEvent{JID: "ser1@im.ecovacs.local/resource", Serial: "SER1", Generation: 1, Reason: session.ReasonPingTimeout})
_ = fleet.Submit(context.Background(), "SER1", ha.Command("command", []byte("stop")))
err := fleet.Submit(context.Background(), "SER1", ha.Command("command", []byte("start")))
if err == nil || err.Error() != "offline" {
t.Fatalf("offline submit err = %v, want offline", err)
}
mu.Lock()
lastSnap := snaps[len(snaps)-1]
mu.Unlock()
if lastSnap.Attributes.LastCommandError == nil || *lastSnap.Attributes.LastCommandError != "offline" {
t.Fatalf("last_command_error = %v, want offline", lastSnap.Attributes.LastCommandError)
}
mu.Lock()
sent = nil
mu.Unlock()
fleet.SessionReady(session.ReadyEvent{
JID: "ser1@im.ecovacs.local/resource", Serial: "SER1",
Generation: 2, Send: send,
})
fleet.AnnounceOK(session.ReadyEvent{
JID: "ser1@im.ecovacs.local/resource", Serial: "SER1",
Generation: 2, Send: send,
})
deadline := time.Now().Add(2 * time.Second)
for {
mu.Lock()
n := len(sent)
mu.Unlock()
if n >= 9 || time.Now().After(deadline) {
break
}
time.Sleep(time.Millisecond)
}
mu.Lock()
defer mu.Unlock()
if len(sent) != 9 {
t.Fatalf("announce fan-out = %d stanzas, want 9", len(sent))
}
for _, s := range sent {
if strings.Contains(string(s), `act="s"`) || strings.Contains(string(s), `type="auto"`) {
t.Fatalf("dropped command replayed: %s", s)
}
}
}
func TestHalfOpenOfflineThenOnline(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
avail := "ecovacs/SER1/availability"
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 1})
h.b.SessionDown(session.DownEvent{JID: "b@x/r", Serial: "SER1", Generation: 1, Reason: session.ReasonPingTimeout})
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 2})
h.b.SessionDown(session.DownEvent{JID: "b@x/r", Serial: "SER1", Generation: 1, Reason: session.ReasonTCPClose})
pubs := c.pubs(avail)
var got []string
for _, p := range pubs {
got = append(got, string(p.payload))
}
want := []string{"offline", "online", "offline", "online"} // connect, announce, down, announce
if len(got) != len(want) {
t.Fatalf("availability sequence = %v, want %v", got, want)
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("availability sequence = %v, want %v", got, want)
}
}
}
func TestRedeliveryDedup(t *testing.T) {
submitted := make(chan robot.Command, 8)
h := newHarness(t, func(_ context.Context, _ string, c robot.Command) error {
submitted <- c
return nil
})
h.b.Republisher("SER1")
c := h.client(t)
topic := "ecovacs/SER1/command"
c.deliver(fakeMsg{topic: topic, payload: []byte("start"), id: 7})
select {
case got := <-submitted:
if got.Name != "start" {
t.Fatalf("submitted command = %+v", got)
}
case <-time.After(2 * time.Second):
t.Fatal("first delivery not submitted")
}
c.deliver(fakeMsg{topic: topic, payload: []byte("start"), dup: true, id: 7})
c.deliver(fakeMsg{topic: topic, payload: []byte("stop"), id: 7})
select {
case got := <-submitted:
if got.Name != "stop" {
t.Fatalf("non-dup resend submitted = %+v", got)
}
case <-time.After(2 * time.Second):
t.Fatal("non-dup resend not submitted")
}
deadline := time.Now().Add(50 * time.Millisecond)
for time.Now().Before(deadline) {
select {
case extra := <-submitted:
t.Fatalf("duplicate delivery submitted %+v", extra)
default:
time.Sleep(time.Millisecond)
}
}
}
func TestPublishUnknownSerial(t *testing.T) {
h := newHarness(t, noopSubmit)
if err := h.b.Publish("NOPE", "state", true, []byte("{}")); err == nil {
t.Fatal("Publish for unknown serial must fail")
}
h.b.Republisher("SER1")
c := h.client(t)
if err := h.b.Publish("SER1", "state", false, []byte(`{"state":"idle","fan_speed":"standard"}`)); err != nil {
t.Fatalf("Publish: %v", err)
}
pubs := c.pubs("ecovacs/SER1/state")
last := pubs[len(pubs)-1]
if last.retain {
t.Fatal("Publish must honour the supplied retain flag (false)")
}
}
func TestShutdownPublishesOfflineAndDisconnects(t *testing.T) {
submitted := make(chan robot.Command, 8)
h := newHarness(t, func(_ context.Context, _ string, c robot.Command) error {
submitted <- c
return nil
})
h.b.Republisher("SER1")
c := h.client(t)
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 1})
if err := h.b.Shutdown(context.Background()); err != nil {
t.Fatalf("Shutdown: %v", err)
}
pubs := c.pubs("ecovacs/SER1/availability")
last := pubs[len(pubs)-1]
if string(last.payload) != "offline" || !last.retain || last.qos != 0 {
t.Fatalf("shutdown publish = %+v, want retained qos0 offline", last)
}
ops := c.opsSnapshot()
if ops[len(ops)-1].kind != "disconnect" {
t.Fatalf("last op = %+v, want disconnect", ops[len(ops)-1])
}
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 2})
pubs = c.pubs("ecovacs/SER1/availability")
if string(pubs[len(pubs)-1].payload) != "offline" {
t.Fatalf("AnnounceOK after shutdown published %q", pubs[len(pubs)-1].payload)
}
c.deliver(fakeMsg{topic: "ecovacs/SER1/command", payload: []byte("start"), id: 9})
select {
case cmd := <-submitted:
t.Fatalf("command dispatched after shutdown: %+v", cmd)
case <-time.After(50 * time.Millisecond):
}
c.reconnect()
pubs = c.pubs("ecovacs/SER1/availability")
if got := string(pubs[len(pubs)-1].payload); got != "offline" {
t.Fatalf("post-shutdown availability = %q, want offline", got)
}
}
func TestShutdownDuringSlotCreationDoesNotConnect(t *testing.T) {
cfg := config.Config{
MQTTHost: "mqtt.test",
MQTTPort: 1883,
MQTTBase: "ecovacs",
HADiscoveryPrefix: "homeassistant",
MQTTClientID: "n95-test",
}
b, err := New(context.Background(), cfg, noopSubmit)
if err != nil {
t.Fatal(err)
}
entered := make(chan struct{})
release := make(chan struct{})
fc := &fakeClient{}
b.factory = func(o *mqtt.ClientOptions) mqttClient {
fc.opts = o
close(entered)
<-release
return fc
}
repubDone := make(chan robot.RepublishFunc, 1)
go func() { repubDone <- b.Republisher("SER1") }()
select {
case <-entered:
case <-time.After(2 * time.Second):
t.Fatal("factory never invoked")
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
shutDone := make(chan error, 1)
go func() { shutDone <- b.Shutdown(ctx) }()
deadline := time.Now().Add(2 * time.Second)
for !b.isClosing() {
if time.Now().After(deadline) {
t.Fatal("closing never set")
}
time.Sleep(time.Millisecond)
}
close(release)
var repub robot.RepublishFunc
select {
case repub = <-repubDone:
case <-time.After(2 * time.Second):
t.Fatal("Republisher never returned")
}
select {
case err := <-shutDone:
if err != nil {
t.Fatalf("Shutdown: %v", err)
}
case <-time.After(2 * time.Second):
t.Fatal("Shutdown never returned")
}
ops := fc.opsSnapshot()
for _, op := range ops {
if op.kind == "connect" {
t.Fatalf("client connected after shutdown began: %+v", ops)
}
}
repub(context.Background(), robot.NewSnapshot())
if n := len(fc.opsSnapshot()); n != len(ops) {
t.Fatal("republisher published after shutdown")
}
late := b.Republisher("SER1")
late(context.Background(), robot.NewSnapshot())
if n := len(fc.opsSnapshot()); n != len(ops) {
t.Fatal("post-shutdown Republisher created activity")
}
}
func TestShutdownWinsOverReconnectOnline(t *testing.T) {
h := newHarness(t, noopSubmit)
h.b.Republisher("SER1")
c := h.client(t)
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 1})
availTopic := "ecovacs/SER1/availability"
entered := make(chan struct{})
release := make(chan struct{})
var once sync.Once
c.beforePublish = func(topic string, payload []byte) {
if topic == availTopic && string(payload) == "online" {
once.Do(func() { close(entered) })
<-release
}
}
reconnectDone := make(chan struct{})
go func() {
c.reconnect()
close(reconnectDone)
}()
select {
case <-entered:
case <-time.After(2 * time.Second):
t.Fatal("reconnect online publish never reached hook")
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
shutDone := make(chan error, 1)
go func() { shutDone <- h.b.Shutdown(ctx) }()
deadline := time.Now().Add(2 * time.Second)
for !h.b.isClosing() {
if time.Now().After(deadline) {
t.Fatal("closing never set")
}
time.Sleep(time.Millisecond)
}
close(release)
select {
case <-reconnectDone:
case <-time.After(2 * time.Second):
t.Fatal("reconnect never finished")
}
select {
case err := <-shutDone:
if err != nil {
t.Fatalf("Shutdown: %v", err)
}
case <-time.After(2 * time.Second):
t.Fatal("Shutdown never returned")
}
pubs := c.pubs(availTopic)
if last := pubs[len(pubs)-1]; string(last.payload) != "offline" || !last.retain {
t.Fatalf("final availability = %+v, want retained offline", last)
}
ops := c.opsSnapshot()
disconnectAt, lastAvailAt := -1, -1
for i, op := range ops {
if op.kind == "disconnect" {
disconnectAt = i
}
if op.kind == "pub" && op.topic == availTopic {
lastAvailAt = i
}
}
if disconnectAt < 0 || lastAvailAt < 0 || disconnectAt < lastAvailAt {
t.Fatalf("disconnect at %d must follow last availability publish at %d: %+v", disconnectAt, lastAvailAt, ops)
}
}
func TestReconnectDoesNotOverwriteNewerSnapshot(t *testing.T) {
h := newHarness(t, noopSubmit)
repub := h.b.Republisher("SER1")
c := h.client(t)
h.b.AnnounceOK(session.ReadyEvent{JID: "b@x/r", Serial: "SER1", Generation: 1})
stateTopic := "ecovacs/SER1/state"
attrsTopic := "ecovacs/SER1/json_attributes"
entered := make(chan struct{})
release := make(chan struct{})
var once sync.Once
c.beforePublish = func(topic string, payload []byte) {
if topic == stateTopic && string(payload) == `{"state":"idle","fan_speed":"standard"}` {
once.Do(func() { close(entered) })
<-release
}
}
reconnectDone := make(chan struct{})
go func() {
c.reconnect()
close(reconnectDone)
}()
select {
case <-entered:
case <-time.After(2 * time.Second):
t.Fatal("reconnect never reached state publish")
}
newer := robot.NewSnapshot()
newer.State.State = "cleaning"
newer.State.FanSpeed = "strong"
battery := 87
newer.Attributes.BatteryLevel = &battery
repubDone := make(chan struct{})
go func() {
repub(context.Background(), newer)
close(repubDone)
}()
close(release)
select {
case <-repubDone:
case <-time.After(2 * time.Second):
t.Fatal("republish never finished")
}
select {
case <-reconnectDone:
case <-time.After(2 * time.Second):
t.Fatal("reconnect never finished")
}
c.beforePublish = nil
statePubs := c.pubs(stateTopic)
if got := string(statePubs[len(statePubs)-1].payload); got != `{"state":"cleaning","fan_speed":"strong"}` {
t.Fatalf("final retained state = %s, want newer snapshot", got)
}
attrPubs := c.pubs(attrsTopic)
var attrs map[string]any
if err := json.Unmarshal(attrPubs[len(attrPubs)-1].payload, &attrs); err != nil {
t.Fatal(err)
}
if attrs["battery_level"] != float64(87) {
t.Fatalf("final retained attributes = %s, want battery_level 87", attrPubs[len(attrPubs)-1].payload)
}
}