package broker import ( "context" "crypto/rand" "encoding/hex" "errors" "log/slog" "net" "sync" "sync/atomic" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/metrics" mqtt "github.com/mochi-mqtt/server/v2" "github.com/mochi-mqtt/server/v2/packets" ) const ( maxClients = 2000 maxPacketSize = 786432 uplinkQueueSize = 256 largeFrameBytes = 64 * 1024 largeFrameSlots = 64 packetOverheadBudget = 128 // 主题与 MQTT 包头预留 keepaliveMin = 10 keepaliveMax = 600 ) // ErrPayloadTooLarge 下行超过客户端 Maximum Packet Size(减包头预留)或 max_receive_bytes。 var ErrPayloadTooLarge = errors.New("broker: payload exceeds client limit") // ErrNoConnection 目标端没有当前连接。 var ErrNoConnection = errors.New("broker: no active connection") // ErrLargeFrameTimeout 全局大帧名额在有界等待内拿不到。 var ErrLargeFrameTimeout = errors.New("broker: large frame quota timeout") // ErrBackpressure 该连接下行队列已满(帧数或字节数)。 var ErrBackpressure = errors.New("broker: downlink backpressure") // ErrNotSubscribed 当前连接尚未订阅下行主题。 var ErrNotSubscribed = errors.New("broker: down topic not subscribed") // ErrSessionWriteConflict 密码登录写令牌时发现库已被并发更新。 var ErrSessionWriteConflict = errors.New("broker: session token write conflict") const ( largeAcquireWait = 5 * time.Second downQueueMax = 256 downQueueBytes = 16 << 20 ) // AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。 type AuthResult struct { OK bool SessionToken string // 密码登录成功时由 N3 填写 } // Authenticator 由 N3 实现;内部故障必须返回 error,不得当成密码错误。 type Authenticator interface { Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error) } // RejectAuthenticator 默认拒绝所有客户端(CONNACK 用户名密码错误)。 type RejectAuthenticator struct{} func (RejectAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) { return AuthResult{OK: false}, nil } // AllowAuthenticator 测试用:允许任意编号。 type AllowAuthenticator struct{} func (AllowAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) { return AuthResult{OK: true}, nil } // PublishDroppedFunc 下行未写入发送队列时回调(对接消息线 OnPublishDropped)。 type PublishDroppedFunc func(ctx context.Context, endpointID string, connID port.ConnID, payload []byte) // Options 装配 broker。 type Options struct { Authenticator Authenticator Uplink port.UplinkHandler Logger *slog.Logger // OnPublishDropped 可选;nil 时仅打 debug 日志。 OnPublishDropped PublishDroppedFunc // Metrics 可选;会话建立/断开时更新 nixmsg_connections。 Metrics *metrics.Registry } // Broker 内置 mochi,不自带监听端口。 type Broker struct { server *mqtt.Server auth Authenticator uplink port.UplinkHandler log *slog.Logger onDrop PublishDroppedFunc metrics *metrics.Registry hook *nixHook connsMu sync.RWMutex current map[string]*connState byClient map[*mqtt.Client]*connState byConnID map[port.ConnID]*connState closedCh chan struct{} queuesMu sync.Mutex queues map[string]*uplinkQueue largeSem chan struct{} closed atomic.Bool lifeMu sync.Mutex lifeLocks map[string]*sync.Mutex } type connState struct { connID port.ConnID endpointID string transport port.Transport remoteIP string client *mqtt.Client maxPacketSize uint32 maxRecvBytes int authOK bool sessionToken string handshook bool subscribedDown bool largePIDs map[uint16]struct{} largePending int metricsCounted bool established bool createdAt time.Time closing bool superseded bool downCh chan downItem downStop chan struct{} downDone chan struct{} downBytes atomic.Int64 sentPub atomic.Int64 mu sync.Mutex handshakeTimer *time.Timer } // New 创建并 Serve mochi(无监听器)。 func New(opts Options) (*Broker, error) { auth := opts.Authenticator if auth == nil { auth = RejectAuthenticator{} } uplink := opts.Uplink if uplink == nil { uplink = port.StubUplinkHandler{} } log := opts.Logger if log == nil { log = slog.Default() } log = slog.New(newRedactHandler(log.Handler())) caps := mqtt.NewDefaultServerCapabilities() caps.MaximumClients = maxClients caps.MaximumQos = 1 caps.MaximumPacketSize = maxPacketSize caps.MaximumSessionExpiryInterval = 0 caps.ReceiveMaximum = 1024 caps.MaximumInflight = 1024 caps.MaximumClientWritesPending = 1024 caps.RetainAvailable = 0 caps.WildcardSubAvailable = 0 caps.SharedSubAvailable = 0 caps.TopicAliasMaximum = 0 caps.Compatibilities.ObscureNotAuthorized = true srv := mqtt.New(&mqtt.Options{ InlineClient: true, Capabilities: caps, Logger: log, }) b := &Broker{ server: srv, auth: auth, uplink: uplink, log: log, onDrop: opts.OnPublishDropped, metrics: opts.Metrics, current: make(map[string]*connState), byClient: make(map[*mqtt.Client]*connState), byConnID: make(map[port.ConnID]*connState), closedCh: make(chan struct{}), queues: make(map[string]*uplinkQueue), largeSem: make(chan struct{}, largeFrameSlots), lifeLocks: make(map[string]*sync.Mutex), } b.hook = &nixHook{b: b} if err := srv.AddHook(b.hook, nil); err != nil { return nil, err } if err := srv.Serve(); err != nil { return nil, err } go b.sweepLoop() return b, nil } // Server 返回底层 mochi(测试用)。 func (b *Broker) Server() *mqtt.Server { return b.server } // Close 关闭 broker。 func (b *Broker) Close() error { if b.closed.Swap(true) { return nil } select { case <-b.closedCh: default: close(b.closedCh) } b.queuesMu.Lock() for _, q := range b.queues { q.close() } b.queues = make(map[string]*uplinkQueue) b.queuesMu.Unlock() return b.server.Close() } // Shutdown 向所有连接发 MQTT 5 0x8B 后关闭。完整 HTTP 停机顺序见 L-03。 func (b *Broker) Shutdown(ctx context.Context) error { if b.closed.Load() { return nil } b.connsMu.RLock() clients := make([]*mqtt.Client, 0, len(b.byClient)) for cl := range b.byClient { if cl != nil { clients = append(clients, cl) } } b.connsMu.RUnlock() for _, cl := range clients { _ = b.server.DisconnectClient(cl, packets.ErrServerShuttingDown) } if ctx != nil { select { case <-ctx.Done(): default: } } return b.Close() } // AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。 func (b *Broker) AttachTCP(conn net.Conn) error { return b.server.EstablishConnection("tcp", conn) } // AttachWS 把 WebSocket NetConn 交给 mochi;阻塞到连接结束。 func (b *Broker) AttachWS(conn net.Conn) error { return b.server.EstablishConnection("ws", conn) } // PublishDown 实现 port.Downlink。 func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error { qos := opts.QoS if qos > 1 { qos = 1 } return b.enqueueDownlink(endpointID, connID, payload, qos, "") } // PublishThenDisconnect 把一帧写入该连接下行队列,写出后再断开(无固定 sleep)。 func (b *Broker) PublishThenDisconnect(_ context.Context, endpointID string, connID port.ConnID, payload []byte, qos byte, reason port.DisconnectReason) error { if qos > 1 { qos = 1 } if reason == "" { reason = port.DisconnectNormal } return b.enqueueDownlink(endpointID, connID, payload, qos, reason) } func (b *Broker) enqueueDownlink(endpointID string, connID port.ConnID, payload []byte, qos byte, disconnect port.DisconnectReason) error { if b.closed.Load() { return errors.New("broker: closed") } st := b.lookupCurrent(endpointID, connID) if st == nil { return ErrNoConnection } st.mu.Lock() maxRecv := st.maxRecvBytes closing := st.closing superseded := st.superseded st.mu.Unlock() if closing || superseded { return ErrNoConnection } if !b.hasDownSub(st) { return ErrNotSubscribed } limit := EffectivePayloadLimit(st.maxPacketSize, maxRecv) if limit > 0 && len(payload) > limit { return ErrPayloadTooLarge } return st.enqueueDown(downItem{ payload: append([]byte(nil), payload...), qos: qos, disconnect: disconnect, }) } func (b *Broker) acquireLarge(ctx context.Context) error { timer := time.NewTimer(largeAcquireWait) defer timer.Stop() select { case b.largeSem <- struct{}{}: return nil case <-ctx.Done(): return ctx.Err() case <-timer.C: return ErrLargeFrameTimeout } } func (b *Broker) releaseLargeSlot() { select { case <-b.largeSem: default: } } func (b *Broker) finishLargePublish(st *connState) { b.reconcileLargeInflight(st) st.mu.Lock() n := st.largePending st.largePending = 0 st.mu.Unlock() for i := 0; i < n; i++ { b.releaseLargeSlot() } } func (b *Broker) releaseLargePID(st *connState, id uint16) { st.mu.Lock() _, ok := st.largePIDs[id] if ok { delete(st.largePIDs, id) } st.mu.Unlock() if ok { b.releaseLargeSlot() } } func (b *Broker) reconcileLargeInflight(st *connState) { if st == nil { return } st.mu.Lock() ids := make([]uint16, 0, len(st.largePIDs)) for id := range st.largePIDs { ids = append(ids, id) } st.mu.Unlock() for _, id := range ids { if st.client != nil { if _, ok := st.client.State.Inflight.Get(id); ok { continue } } b.releaseLargePID(st, id) } } func (b *Broker) releaseAllLarge(st *connState) { st.mu.Lock() n := len(st.largePIDs) + st.largePending st.largePIDs = nil st.largePending = 0 st.mu.Unlock() for i := 0; i < n; i++ { b.releaseLargeSlot() } } // Disconnect 实现 port.ConnControl。 func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.ConnID, reason port.DisconnectReason) error { st := b.lookupConn(endpointID, connID) if st == nil { return ErrNoConnection } code := packets.CodeDisconnect switch reason { case port.DisconnectTakenOver: code = packets.ErrSessionTakenOver case port.DisconnectKicked, port.DisconnectFatal: code = packets.ErrAdministrativeAction } err := b.server.DisconnectClient(st.client, code) // mochi 对错误类原因码会把 Code 当作 error 返回,表示已按该原因断开,不算失败。 if _, ok := err.(packets.Code); ok { return nil } return err } func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState { b.connsMu.RLock() defer b.connsMu.RUnlock() if connID != "" { st := b.byConnID[connID] if st != nil && st.endpointID == endpointID { return st } return nil } return b.current[endpointID] } // lookupCurrent 只返回该端当前连接;connID 非空时必须仍是当前连接。 func (b *Broker) lookupCurrent(endpointID string, connID port.ConnID) *connState { b.connsMu.RLock() defer b.connsMu.RUnlock() cur := b.current[endpointID] if cur == nil { return nil } if connID != "" && cur.connID != connID { return nil } return cur } func (b *Broker) endpointLife(endpointID string) *sync.Mutex { b.lifeMu.Lock() defer b.lifeMu.Unlock() m := b.lifeLocks[endpointID] if m == nil { m = &sync.Mutex{} b.lifeLocks[endpointID] = m } return m } func downTopic(endpointID string) string { return "nix/c/" + endpointID + "/down" } func upTopic(endpointID string) string { return "nix/c/" + endpointID + "/up" } // EffectivePayloadLimit 下行载荷上限:客户端 Maximum Packet Size 减包头预留,再与 max_receive_bytes 取更严者。 func EffectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int { return effectivePayloadLimit(maxPacketSize, maxRecvBytes) } func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int { limit := 0 if maxPacketSize > 0 { if maxPacketSize > packetOverheadBudget { limit = int(maxPacketSize) - packetOverheadBudget } } if maxRecvBytes > 0 { if limit == 0 || maxRecvBytes < limit { limit = maxRecvBytes } } return limit } func randomConnID() port.ConnID { var b [16]byte _, _ = rand.Read(b[:]) return port.ConnID(hex.EncodeToString(b[:])) } func transportOf(cl *mqtt.Client) port.Transport { if cl != nil && cl.Net.Listener == "ws" { return port.TransportWS } return port.TransportTCP } func remoteIPOf(cl *mqtt.Client) string { if cl == nil { return "" } addr := cl.Net.Remote if addr == "" && cl.Net.Conn != nil && cl.Net.Conn.RemoteAddr() != nil { addr = cl.Net.Conn.RemoteAddr().String() } host, _, err := net.SplitHostPort(addr) if err != nil { return addr } return host } // SetMaxReceiveBytes 供 N3 握手后设置;0 表示不限。 func (b *Broker) SetMaxReceiveBytes(endpointID string, connID port.ConnID, n int) { st := b.lookupConn(endpointID, connID) if st == nil { return } st.mu.Lock() st.maxRecvBytes = n st.mu.Unlock() } // ConnInfoOf 返回连接信息(测试/N3)。 func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) { b.connsMu.RLock() st := b.current[endpointID] b.connsMu.RUnlock() if st == nil { return port.ConnInfo{}, false } return port.ConnInfo{ ConnID: st.connID, EndpointID: st.endpointID, Transport: st.transport, RemoteIP: st.remoteIP, SessionToken: st.sessionToken, MaxPacketSize: st.maxPacketSize, }, true } // IsHandshook 当前连接是否已完成握手。 func (b *Broker) IsHandshook(endpointID string) bool { b.connsMu.RLock() st := b.current[endpointID] b.connsMu.RUnlock() if st == nil { return false } st.mu.Lock() defer st.mu.Unlock() return st.handshook } // CurrentConnID 返回端的当前连接代号。 func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) { b.connsMu.RLock() st := b.current[endpointID] b.connsMu.RUnlock() if st == nil { return "", false } return st.connID, true } func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState { b.connsMu.RLock() defer b.connsMu.RUnlock() st := b.byConnID[connID] if st == nil || st.endpointID != endpointID { return nil } return st } func (b *Broker) sweepLoop() { tick := time.NewTicker(time.Minute) defer tick.Stop() for { select { case <-tick.C: b.sweepUnestablished(time.Minute) case <-b.closedCh: return } } } func (b *Broker) sweepUnestablished(minAge time.Duration) { now := time.Now() b.connsMu.Lock() defer b.connsMu.Unlock() for cl, st := range b.byClient { if st.established { continue } if cl != nil && !cl.Closed() { continue } if minAge > 0 && now.Sub(st.createdAt) < minAge { continue } delete(b.byClient, cl) delete(b.byConnID, st.connID) if b.current[st.endpointID] == st { delete(b.current, st.endpointID) } } } func (b *Broker) hasDownSub(st *connState) bool { if st == nil { return false } st.mu.Lock() defer st.mu.Unlock() if st.subscribedDown { return true } // 回退:直接看 mochi 订阅表 if st.client != nil && st.client.State.Subscriptions != nil { _, ok := st.client.State.Subscriptions.Get(downTopic(st.endpointID)) return ok } return false } func (b *Broker) startHandshakeDeadline(endpointID string, connID port.ConnID, d time.Duration) { st := b.connStateOf(endpointID, connID) if st == nil { return } st.mu.Lock() if st.handshook { st.mu.Unlock() return } if st.handshakeTimer != nil { st.handshakeTimer.Stop() } st.handshakeTimer = time.AfterFunc(d, func() { cur := b.connStateOf(endpointID, connID) if cur == nil { return } cur.mu.Lock() done := cur.handshook cur.mu.Unlock() if done { return } _ = b.Disconnect(context.Background(), endpointID, connID, port.DisconnectIdle) }) st.mu.Unlock() } func (b *Broker) cancelHandshakeDeadline(endpointID string, connID port.ConnID) { st := b.connStateOf(endpointID, connID) if st == nil { return } st.mu.Lock() if st.handshakeTimer != nil { st.handshakeTimer.Stop() st.handshakeTimer = nil } st.mu.Unlock() } func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) { b.queuesMu.Lock() q, ok := b.queues[endpointID] if !ok { q = newUplinkQueue(b, endpointID) b.queues[endpointID] = q } b.queuesMu.Unlock() q.push(uplinkItem{conn: conn, payload: payload}) } var ( _ port.Downlink = (*Broker)(nil) _ port.ConnControl = (*Broker)(nil) )