Files
NixMsg/sdk/go/transport_mqtt.go

259 lines
5.8 KiB
Go

package nixmsg
import (
"context"
"errors"
"fmt"
"net/url"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/eclipse/paho.golang/autopaho"
"github.com/eclipse/paho.golang/paho"
)
type mqttTransport struct {
mu sync.Mutex
cm *autopaho.ConnectionManager
cancel context.CancelFunc
cfg transportConfig
cred atomic.Value // string
upTopic string
downTopic string
stopped atomic.Bool
ready chan struct{}
}
func newMQTTTransport() *mqttTransport {
t := &mqttTransport{ready: make(chan struct{})}
t.cred.Store("")
return t
}
func (t *mqttTransport) SetCredential(passwordOrToken string) {
t.cred.Store(passwordOrToken)
}
func (t *mqttTransport) Start(ctx context.Context, cfg transportConfig) error {
t.mu.Lock()
defer t.mu.Unlock()
if t.cm != nil {
return errors.New("transport already started")
}
t.cfg = cfg
t.upTopic = fmt.Sprintf("nix/c/%s/up", cfg.EndpointID)
t.downTopic = fmt.Sprintf("nix/c/%s/down", cfg.EndpointID)
u, err := normalizeMQTTURL(cfg.URL, cfg.AllowTCP)
if err != nil {
return err
}
innerCtx, cancel := context.WithCancel(ctx)
t.cancel = cancel
var sessionExpiry uint32 // 0;由 ConnectPacketBuilder 显式写入 Properties
cliCfg := autopaho.ClientConfig{
ServerUrls: []*url.URL{u},
KeepAlive: 30,
ConnectTimeout: cfg.ConnectTimeout,
CleanStartOnInitialConnection: false, // 不要只靠这个;每次用 ConnectPacketBuilder
SessionExpiryInterval: sessionExpiry,
ConnectUsername: cfg.EndpointID,
ReconnectBackoff: cfg.Backoff.Func,
OnConnectError: func(err error) {
var ce *autopaho.ConnackError
if errors.As(err, &ce) {
if isAuthCONNACK(ce.ReasonCode) {
reason := AuthBadCredentials
if tok, _ := t.cred.Load().(string); strings.HasPrefix(tok, "nst_") {
reason = AuthSessionInvalid
}
if cfg.OnAuthFailed != nil {
cfg.OnAuthFailed(reason)
}
cancel()
}
}
},
OnConnectionDown: func() bool {
if t.stopped.Load() {
return false
}
if cfg.Backoff != nil {
cfg.Backoff.MarkOffline()
}
if cfg.OnOffline != nil {
cfg.OnOffline()
}
return !t.stopped.Load()
},
OnConnectionUp: func(cm *autopaho.ConnectionManager, _ *paho.Connack) {
if cfg.Backoff != nil {
cfg.Backoff.MarkOnline()
}
go func() {
_, err := cm.Subscribe(innerCtx, &paho.Subscribe{
Subscriptions: []paho.SubscribeOptions{{
Topic: t.downTopic,
QoS: 1,
}},
})
if err != nil {
return
}
if cfg.MQTTReady != nil {
_ = cfg.MQTTReady(innerCtx)
}
if cfg.OnOnline != nil {
cfg.OnOnline()
}
}()
},
ClientConfig: paho.ClientConfig{
ClientID: cfg.EndpointID,
OnServerDisconnect: func(d *paho.Disconnect) {
if d != nil && d.ReasonCode == 0x8E {
if cfg.OnKicked != nil {
cfg.OnKicked()
}
cancel()
}
},
OnPublishReceived: []func(paho.PublishReceived) (bool, error){
func(pr paho.PublishReceived) (bool, error) {
if cfg.OnDown != nil && pr.Packet != nil {
cfg.OnDown(pr.Packet.Payload)
}
return true, nil
},
},
},
}
cliCfg.ConnectPacketBuilder = func(c *paho.Connect, _ *url.URL) (*paho.Connect, error) {
c.CleanStart = true
zero := uint32(0)
if c.Properties == nil {
c.Properties = &paho.ConnectProperties{}
}
c.Properties.SessionExpiryInterval = &zero
pass, _ := t.cred.Load().(string)
c.UsernameFlag = true
c.Username = cfg.EndpointID
c.PasswordFlag = true
c.Password = []byte(pass)
if cfg.OnConnectPacket != nil {
cfg.OnConnectPacket(c.CleanStart, zero)
}
return c, nil
}
cm, err := autopaho.NewConnection(innerCtx, cliCfg)
if err != nil {
cancel()
return err
}
t.cm = cm
return nil
}
func (t *mqttTransport) PublishUp(payload []byte) error {
t.mu.Lock()
cm := t.cm
topic := t.upTopic
t.mu.Unlock()
if cm == nil {
return apiErr(CodeNotConnected, "未连接")
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
_, err := cm.Publish(ctx, &paho.Publish{
Topic: topic,
QoS: 1,
Payload: payload,
})
return err
}
func (t *mqttTransport) Stop(ctx context.Context) error {
t.stopped.Store(true)
t.mu.Lock()
cm := t.cm
cancel := t.cancel
t.mu.Unlock()
if cancel != nil {
cancel()
}
if cm != nil {
return cm.Disconnect(ctx)
}
return nil
}
func isAuthCONNACK(code byte) bool {
switch code {
case 0x86, 0x87, 0x8A, // MQTT 5
4, 5: // MQTT 3.1.1 bad user/pass, not authorized
return true
default:
return false
}
}
func normalizeMQTTURL(raw string, allowTCP bool) (*url.URL, error) {
u, err := url.Parse(raw)
if err != nil {
return nil, err
}
switch strings.ToLower(u.Scheme) {
case "ws", "wss":
if u.Path == "" || u.Path == "/" {
u.Path = "/mqtt"
}
return u, nil
case "http":
u.Scheme = "ws"
if u.Path == "" || u.Path == "/" {
u.Path = "/mqtt"
}
return u, nil
case "https":
u.Scheme = "wss"
if u.Path == "" || u.Path == "/" {
u.Path = "/mqtt"
}
return u, nil
case "mqtt", "tcp", "mqtts", "ssl", "tls":
if !allowTCP {
return nil, apiErr(CodeBadRequest, "裸 TCP 需在选项中显式打开 AllowTCP")
}
return u, nil
default:
return nil, fmt.Errorf("不支持的 URL scheme: %s", u.Scheme)
}
}
// RegisterURLFromConnect 从连接地址推出注册 HTTP 地址(第 6.9 节)。
func RegisterURLFromConnect(connectURL string) (string, error) {
u, err := url.Parse(connectURL)
if err != nil {
return "", err
}
out := *u
switch strings.ToLower(u.Scheme) {
case "wss", "https", "mqtts", "ssl", "tls":
out.Scheme = "https"
case "ws", "http", "mqtt", "tcp":
out.Scheme = "http"
default:
return "", fmt.Errorf("无法从 %s 推出注册地址", u.Scheme)
}
out.Path = "/api/client/register"
out.RawQuery = ""
out.Fragment = ""
return out.String(), nil
}