267 lines
6.1 KiB
Go
267 lines
6.1 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: 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
|
||
}
|