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 }