Files
NixMsg/sdk/go/transport_mqtt.go

267 lines
6.1 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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: uint16(DefaultKeepAliveSeconds),
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()
return
}
}
if cfg.Backoff != nil {
cfg.Backoff.MarkOffline()
}
},
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,
PacketTimeout: cfg.ConnectTimeout,
OnServerDisconnect: func(d *paho.Disconnect) {
if d != nil && d.ReasonCode == 0x8E {
if cfg.OnKicked != nil {
cfg.OnKicked()
}
cancel()
return
}
// 0x8B 等其它原因按可重试处理,继续重连。
},
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
// 不设置 ReceiveMaximum(B-01 / K-00)。
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
}