259 lines
5.8 KiB
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
|
|
}
|