package broker import ( "context" "crypto/rand" "encoding/hex" "errors" "log/slog" "net" "sync" "sync/atomic" "git.asio.asia/nixevol/NixMsg/internal/app/port" 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") // 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 } // Options 装配 broker。 type Options struct { Authenticator Authenticator Uplink port.UplinkHandler Logger *slog.Logger } // Broker 内置 mochi,不自带监听端口。 type Broker struct { server *mqtt.Server auth Authenticator uplink port.UplinkHandler log *slog.Logger hook *nixHook connsMu sync.RWMutex current map[string]*connState byClient map[*mqtt.Client]*connState queuesMu sync.Mutex queues map[string]*uplinkQueue largeSem chan struct{} closed atomic.Bool } type connState struct { connID port.ConnID endpointID string transport port.Transport remoteIP string client *mqtt.Client maxPacketSize uint32 maxRecvBytes int authOK bool authErr error sessionToken string largeHeld int mu sync.Mutex } // 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() } 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, current: make(map[string]*connState), byClient: make(map[*mqtt.Client]*connState), queues: make(map[string]*uplinkQueue), largeSem: make(chan struct{}, largeFrameSlots), } 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 } 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 } b.queuesMu.Lock() for _, q := range b.queues { q.close() } b.queuesMu.Unlock() return b.server.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 { if b.closed.Load() { return errors.New("broker: closed") } st := b.lookupConn(endpointID, connID) if st == nil { return ErrNoConnection } limit := effectivePayloadLimit(st.maxPacketSize, st.maxRecvBytes) if limit > 0 && len(payload) > limit { return ErrPayloadTooLarge } qos := opts.QoS if qos > 1 { qos = 1 } topic := downTopic(endpointID) large := len(payload) > largeFrameBytes if large { select { case b.largeSem <- struct{}{}: case <-ctx.Done(): return ctx.Err() } st.mu.Lock() st.largeHeld++ st.mu.Unlock() } if err := b.server.Publish(topic, payload, false, qos); err != nil { if large { b.releaseOneLarge(st) } return err } if large && qos == 0 { b.releaseOneLarge(st) } return nil } func (b *Broker) releaseOneLarge(st *connState) { st.mu.Lock() if st.largeHeld > 0 { st.largeHeld-- st.mu.Unlock() select { case <-b.largeSem: default: } return } st.mu.Unlock() } func (b *Broker) releaseAllLarge(st *connState) { st.mu.Lock() n := st.largeHeld st.largeHeld = 0 st.mu.Unlock() for i := 0; i < n; i++ { select { case <-b.largeSem: default: } } } // 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 } return b.server.DisconnectClient(st.client, code) } func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState { b.connsMu.RLock() defer b.connsMu.RUnlock() if connID != "" { for _, st := range b.byClient { if st.endpointID == endpointID && st.connID == connID { return st } } return nil } return b.current[endpointID] } func downTopic(endpointID string) string { return "nix/c/" + endpointID + "/down" } func upTopic(endpointID string) string { return "nix/c/" + endpointID + "/up" } 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 } 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) )