348 lines
8.1 KiB
Go
348 lines
8.1 KiB
Go
// Package mqttbridge owns the per-robot MQTT clients that publish Home
|
|
// Assistant discovery, availability, state, and attribute documents and
|
|
// translate inbound MQTT commands into robot command submissions.
|
|
package mqttbridge
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/sha256"
|
|
"crypto/tls"
|
|
"crypto/x509"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"net"
|
|
"os"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"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"
|
|
)
|
|
|
|
const (
|
|
publishTimeout = 5 * time.Second
|
|
connectTimeout = 10 * time.Second
|
|
)
|
|
|
|
var errPublishTimeout = errors.New("mqttbridge: publish timeout")
|
|
|
|
type SubmitFunc func(context.Context, string, robot.Command) error
|
|
|
|
type mqttClient interface {
|
|
Connect() mqtt.Token
|
|
Disconnect(quiesce uint)
|
|
Subscribe(topic string, qos byte, callback mqtt.MessageHandler) mqtt.Token
|
|
Publish(topic string, qos byte, retained bool, payload interface{}) mqtt.Token
|
|
IsConnected() bool
|
|
}
|
|
|
|
type clientFactory func(*mqtt.ClientOptions) mqttClient
|
|
|
|
type Bridge struct {
|
|
SendCommand func(ctx context.Context, serial string, payload []byte) (robot.Command, error)
|
|
|
|
ctx context.Context
|
|
cfg config.Config
|
|
submit SubmitFunc
|
|
factory clientFactory
|
|
broker string
|
|
tlsCfg *tls.Config
|
|
|
|
mu sync.Mutex
|
|
slots map[string]*slot
|
|
closing bool
|
|
}
|
|
|
|
type slot struct {
|
|
serial string
|
|
ready chan struct{} // closed once client is assigned
|
|
connectMu sync.Mutex
|
|
availabilityMu sync.Mutex
|
|
stateMu sync.Mutex
|
|
client mqttClient
|
|
|
|
mu sync.Mutex
|
|
latest robot.Snapshot
|
|
live bool
|
|
announced uint64
|
|
handled map[dedupKey]struct{}
|
|
}
|
|
|
|
type dedupKey struct {
|
|
topic string
|
|
id uint16
|
|
}
|
|
|
|
func New(ctx context.Context, cfg config.Config, submit SubmitFunc) (*Bridge, error) {
|
|
b := &Bridge{
|
|
ctx: ctx,
|
|
cfg: cfg,
|
|
submit: submit,
|
|
factory: func(o *mqtt.ClientOptions) mqttClient { return mqtt.NewClient(o) },
|
|
slots: map[string]*slot{},
|
|
}
|
|
scheme := "tcp"
|
|
if cfg.MQTTTLS {
|
|
scheme = "tls"
|
|
tlsCfg, err := buildTLS(cfg)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
b.tlsCfg = tlsCfg
|
|
}
|
|
b.broker = scheme + "://" + net.JoinHostPort(cfg.MQTTHost, strconv.Itoa(cfg.MQTTPort))
|
|
return b, nil
|
|
}
|
|
|
|
func buildTLS(cfg config.Config) (*tls.Config, error) {
|
|
pool, err := x509.SystemCertPool()
|
|
if err != nil || pool == nil {
|
|
pool = x509.NewCertPool()
|
|
}
|
|
if cfg.MQTTCAFile != "" {
|
|
pemBytes, err := os.ReadFile(cfg.MQTTCAFile)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("MQTT_CA_FILE: %w", err)
|
|
}
|
|
if !pool.AppendCertsFromPEM(pemBytes) {
|
|
return nil, fmt.Errorf("MQTT_CA_FILE: no certificates found in %s", cfg.MQTTCAFile)
|
|
}
|
|
}
|
|
return &tls.Config{
|
|
ServerName: cfg.MQTTHost,
|
|
RootCAs: pool,
|
|
MinVersion: tls.VersionTLS12,
|
|
}, nil
|
|
}
|
|
|
|
func clientID(prefix, serial string) string {
|
|
if id := prefix + "-" + serial; len(id) <= 128 {
|
|
return id
|
|
}
|
|
sum := sha256.Sum256([]byte(prefix + "\n" + serial))
|
|
return "n95-" + hex.EncodeToString(sum[:])[:36]
|
|
}
|
|
|
|
func (b *Bridge) root(serial string) string {
|
|
return b.cfg.MQTTBase + "/" + serial
|
|
}
|
|
|
|
func (b *Bridge) isClosing() bool {
|
|
b.mu.Lock()
|
|
defer b.mu.Unlock()
|
|
return b.closing
|
|
}
|
|
|
|
func (b *Bridge) slotFor(serial string) *slot {
|
|
b.mu.Lock()
|
|
defer b.mu.Unlock()
|
|
return b.slots[serial]
|
|
}
|
|
|
|
// ensureSlot returns the per-serial slot, creating and connecting its
|
|
// broker client on first use. The slot is installed under b.mu so exactly
|
|
// one client is ever built; the external factory call happens after unlock,
|
|
// and connectMu serializes the closing check with Connect so Shutdown
|
|
// cannot pass between them. Nil is returned once shutdown has begun.
|
|
func (b *Bridge) ensureSlot(serial string) *slot {
|
|
b.mu.Lock()
|
|
if b.closing {
|
|
b.mu.Unlock()
|
|
return nil
|
|
}
|
|
if s := b.slots[serial]; s != nil {
|
|
b.mu.Unlock()
|
|
return s
|
|
}
|
|
s := &slot{
|
|
serial: serial,
|
|
ready: make(chan struct{}),
|
|
latest: robot.NewSnapshot(),
|
|
handled: map[dedupKey]struct{}{},
|
|
}
|
|
b.slots[serial] = s
|
|
b.mu.Unlock()
|
|
|
|
opts := b.options(serial, s)
|
|
s.client = b.factory(opts)
|
|
close(s.ready)
|
|
s.connectMu.Lock()
|
|
defer s.connectMu.Unlock()
|
|
if b.isClosing() {
|
|
return nil
|
|
}
|
|
slog.Info("mqtt connect", "host", b.cfg.MQTTHost, "port", b.cfg.MQTTPort,
|
|
"tls", b.tlsCfg != nil, "client_id", opts.ClientID, "username", b.cfg.MQTTUsername)
|
|
s.client.Connect() // async: ConnectRetry recovers a down broker
|
|
return s
|
|
}
|
|
|
|
func (b *Bridge) options(serial string, s *slot) *mqtt.ClientOptions {
|
|
o := mqtt.NewClientOptions()
|
|
o.AddBroker(b.broker)
|
|
o.SetClientID(clientID(b.cfg.MQTTClientID, serial))
|
|
o.SetCleanSession(true)
|
|
o.SetProtocolVersion(4) // MQTT 3.1.1
|
|
o.SetAutoReconnect(true)
|
|
o.SetConnectRetry(true)
|
|
o.SetConnectTimeout(connectTimeout)
|
|
o.SetOrderMatters(false)
|
|
if b.tlsCfg != nil {
|
|
o.SetTLSConfig(b.tlsCfg)
|
|
}
|
|
if b.cfg.MQTTUsername != "" {
|
|
o.SetUsername(b.cfg.MQTTUsername)
|
|
o.SetPassword(b.cfg.MQTTPassword)
|
|
}
|
|
o.SetWill(b.root(serial)+"/availability", "offline", 0, true)
|
|
o.SetOnConnectHandler(func(mqtt.Client) { b.onConnect(s) })
|
|
o.SetDefaultPublishHandler(func(_ mqtt.Client, m mqtt.Message) { b.onMessage(s, m) })
|
|
return o
|
|
}
|
|
|
|
func (b *Bridge) onConnect(s *slot) {
|
|
s.mu.Lock()
|
|
s.handled = map[dedupKey]struct{}{}
|
|
s.mu.Unlock()
|
|
|
|
root := b.root(s.serial)
|
|
for _, sub := range []struct {
|
|
topic string
|
|
qos byte
|
|
}{
|
|
{root + "/command", 1},
|
|
{root + "/set_fan_speed", 1},
|
|
{root + "/send_command", 1},
|
|
{b.cfg.HADiscoveryPrefix + "/status", 0},
|
|
} {
|
|
tok := s.client.Subscribe(sub.topic, sub.qos, nil)
|
|
if !tok.WaitTimeout(publishTimeout) || tok.Error() != nil {
|
|
slog.Error("mqtt subscribe failed", "topic", sub.topic, "err", tokenErr(tok))
|
|
return
|
|
}
|
|
}
|
|
|
|
_ = b.publish(s, ha.DiscoveryTopic(b.cfg, s.serial), true, ha.Discovery(b.cfg, s.serial))
|
|
|
|
s.availabilityMu.Lock()
|
|
s.mu.Lock()
|
|
live := s.live
|
|
if b.isClosing() {
|
|
s.live = false
|
|
live = false
|
|
}
|
|
s.mu.Unlock()
|
|
|
|
availability := "offline"
|
|
if live {
|
|
availability = "online"
|
|
}
|
|
_ = b.publish(s, root+"/availability", true, []byte(availability))
|
|
s.availabilityMu.Unlock()
|
|
|
|
s.stateMu.Lock()
|
|
s.mu.Lock()
|
|
snap := s.latest
|
|
s.mu.Unlock()
|
|
stateJSON, _ := json.Marshal(snap.State)
|
|
attrJSON, _ := json.Marshal(snap.Attributes)
|
|
_ = b.publish(s, root+"/state", true, stateJSON)
|
|
_ = b.publish(s, root+"/json_attributes", true, attrJSON)
|
|
s.stateMu.Unlock()
|
|
}
|
|
|
|
func (b *Bridge) onMessage(s *slot, m mqtt.Message) {
|
|
if b.isClosing() {
|
|
return
|
|
}
|
|
topic := m.Topic()
|
|
if topic == b.cfg.HADiscoveryPrefix+"/status" {
|
|
if bytes.Equal(m.Payload(), []byte("online")) {
|
|
_ = b.publish(s, ha.DiscoveryTopic(b.cfg, s.serial), true, ha.Discovery(b.cfg, s.serial))
|
|
}
|
|
return
|
|
}
|
|
suffix, ok := strings.CutPrefix(topic, b.root(s.serial)+"/")
|
|
if !ok {
|
|
return
|
|
}
|
|
key := dedupKey{topic: topic, id: m.MessageID()}
|
|
s.mu.Lock()
|
|
if m.Duplicate() {
|
|
if _, seen := s.handled[key]; seen {
|
|
s.mu.Unlock()
|
|
return
|
|
}
|
|
}
|
|
s.handled[key] = struct{}{}
|
|
s.mu.Unlock()
|
|
|
|
cmd := ha.Command(suffix, m.Payload())
|
|
go func() {
|
|
_ = b.submit(b.ctx, s.serial, cmd)
|
|
}()
|
|
}
|
|
|
|
func tokenErr(tok mqtt.Token) error {
|
|
if err := tok.Error(); err != nil {
|
|
return err
|
|
}
|
|
return errPublishTimeout
|
|
}
|
|
|
|
func (b *Bridge) SessionReady(session.ReadyEvent) {}
|
|
|
|
func (b *Bridge) Stanza(session.StanzaEvent) {}
|
|
|
|
func (b *Bridge) AnnounceOK(e session.ReadyEvent) {
|
|
if b.isClosing() {
|
|
return
|
|
}
|
|
s := b.ensureSlot(e.Serial)
|
|
if s == nil {
|
|
return
|
|
}
|
|
s.availabilityMu.Lock()
|
|
defer s.availabilityMu.Unlock()
|
|
if b.isClosing() {
|
|
return
|
|
}
|
|
s.mu.Lock()
|
|
if e.Generation < s.announced {
|
|
s.mu.Unlock()
|
|
return
|
|
}
|
|
s.announced = e.Generation
|
|
s.live = true
|
|
s.mu.Unlock()
|
|
_ = b.publish(s, b.root(s.serial)+"/availability", true, []byte("online"))
|
|
}
|
|
|
|
func (b *Bridge) SessionDown(e session.DownEvent) {
|
|
s := b.slotFor(e.Serial)
|
|
if s == nil {
|
|
return
|
|
}
|
|
s.availabilityMu.Lock()
|
|
defer s.availabilityMu.Unlock()
|
|
s.mu.Lock()
|
|
applicable := s.live && e.Generation == s.announced
|
|
if applicable {
|
|
s.live = false
|
|
}
|
|
s.mu.Unlock()
|
|
if applicable {
|
|
_ = b.publish(s, b.root(s.serial)+"/availability", true, []byte("offline"))
|
|
}
|
|
}
|