package broker import ( "bytes" "context" "time" "git.asio.asia/nixevol/NixMsg/internal/app/port" mqtt "github.com/mochi-mqtt/server/v2" "github.com/mochi-mqtt/server/v2/packets" ) type nixHook struct { mqtt.HookBase b *Broker } func (h *nixHook) ID() string { return "nixmsg" } func (h *nixHook) Provides(b byte) bool { return bytes.Contains([]byte{ mqtt.OnConnect, mqtt.OnConnectAuthenticate, mqtt.OnACLCheck, mqtt.OnPublish, mqtt.OnPublishDropped, mqtt.OnSessionEstablished, mqtt.OnDisconnect, mqtt.OnQosPublish, mqtt.OnQosComplete, mqtt.OnQosDropped, mqtt.OnSubscribed, mqtt.OnPacketSent, }, []byte{b}) } func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { endpointID := string(pk.Connect.Username) clientID := pk.Connect.ClientIdentifier if endpointID == "" { endpointID = clientID } remoteIP := remoteIPOf(cl) st := &connState{ connID: randomConnID(), endpointID: endpointID, transport: transportOf(cl), remoteIP: remoteIP, client: cl, maxPacketSize: pk.Properties.MaximumPacketSize, createdAt: time.Now(), } // ClientID、Username 都必须等于端编号 if clientID == "" || endpointID == "" || clientID != endpointID { return nil } // 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive ka := pk.Connect.Keepalive if ka < keepaliveMin || ka > keepaliveMax { if ka < keepaliveMin { ka = keepaliveMin } if ka > keepaliveMax { ka = keepaliveMax } cl.State.Keepalive = ka cl.State.ServerKeepalive = true } // B-01:绕开 mochi 发送配额路径(NextImmediate 递归读锁 + PUBACK 配额泄漏)。 // ParseConnect 已按客户端 Receive Maximum 设过 sendQuota;此处一律置 0。 if cl.State.Inflight != nil { cl.State.Inflight.ResetSendQuota(0) } if rm := pk.Properties.ReceiveMaximum; rm > 0 && rm < 256 { h.b.log.Warn("client receive maximum below 256; server ignores MQTT send quota", "endpoint", endpointID, "receive_maximum", rm) } authCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() res, err := h.b.auth.Authenticate(authCtx, endpointID, pk.Connect.Password, remoteIP) if err != nil { return err // mochi 不回 CONNACK,直接断开;不登记连接表 } if !res.OK { return nil } st.authOK = true st.sessionToken = res.SessionToken h.rememberPending(cl, st) return nil } func (h *nixHook) rememberPending(cl *mqtt.Client, st *connState) { h.b.connsMu.Lock() h.b.byClient[cl] = st h.b.byConnID[st.connID] = st h.b.connsMu.Unlock() } func (h *nixHook) OnConnectAuthenticate(cl *mqtt.Client, _ packets.Packet) bool { h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st == nil { return false } // 内部故障已在 OnConnect 返回 error;此处只反映业务上的拒绝 return st.authOK } func (h *nixHook) OnACLCheck(cl *mqtt.Client, topic string, write bool) bool { h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st == nil || st.endpointID == "" { return false } up := upTopic(st.endpointID) down := downTopic(st.endpointID) if write { return topic == up } return topic == down } func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet, error) { // InlineClient 的 PublishDown 走 InjectPacket → OnPublish;必须放行才能分发给订阅者。 if cl != nil && cl.Net.Inline { return pk, nil } h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st == nil { return pk, packets.CodeSuccessIgnore } payload := append([]byte(nil), pk.Payload...) info := port.ConnInfo{ ConnID: st.connID, EndpointID: st.endpointID, Transport: st.transport, RemoteIP: st.remoteIP, SessionToken: st.sessionToken, MaxPacketSize: st.maxPacketSize, } h.b.enqueueUplink(st.endpointID, info, payload) return pk, packets.CodeSuccessIgnore } func (h *nixHook) OnSubscribed(cl *mqtt.Client, pk packets.Packet, reasonCodes []byte) { h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st == nil { return } down := downTopic(st.endpointID) for i, sub := range pk.Filters { if sub.Filter != down { continue } if i < len(reasonCodes) && reasonCodes[i] >= 0x80 { continue } st.mu.Lock() st.subscribedDown = true st.mu.Unlock() return } } func (h *nixHook) OnPacketSent(cl *mqtt.Client, pk packets.Packet, _ []byte) { if pk.FixedHeader.Type != packets.Publish { return } h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st != nil { st.sentPub.Add(1) } } func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload)) h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st != nil { h.b.reconcileLargeInflight(st) } if h.b.onDrop == nil || st == nil { return } h.b.onDrop(context.Background(), st.endpointID, st.connID, append([]byte(nil), pk.Payload...)) } func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) { h.b.connsMu.Lock() st := h.b.byClient[cl] var old *connState if st != nil { old = h.b.current[st.endpointID] h.b.current[st.endpointID] = st st.established = true } h.b.connsMu.Unlock() if st == nil { return } lk := h.b.endpointLife(st.endpointID) lk.Lock() if old != nil && old != st { old.mu.Lock() old.superseded = true old.mu.Unlock() } lk.Unlock() st.startDownLoop(h.b) info := port.ConnInfo{ ConnID: st.connID, EndpointID: st.endpointID, Transport: st.transport, RemoteIP: st.remoteIP, SessionToken: st.sessionToken, MaxPacketSize: st.maxPacketSize, } _ = h.b.uplink.OnSessionEstablished(context.Background(), info) h.noteConnectionOpen(st) } func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { h.b.connsMu.Lock() st := h.b.byClient[cl] delete(h.b.byClient, cl) if st != nil { delete(h.b.byConnID, st.connID) } isCurrent := false if st != nil && h.b.current[st.endpointID] == st { delete(h.b.current, st.endpointID) isCurrent = true } h.b.connsMu.Unlock() if st == nil { return } lk := h.b.endpointLife(st.endpointID) lk.Lock() st.stopDownLoop() lk.Unlock() h.b.releaseAllLarge(st) h.b.cancelHandshakeDeadline(st.endpointID, st.connID) reason := port.DisconnectNormal if err != nil { if code, ok := err.(packets.Code); ok { switch code.Code { case packets.ErrSessionTakenOver.Code: reason = port.DisconnectTakenOver case packets.ErrAdministrativeAction.Code: reason = port.DisconnectKicked } } } info := port.ConnInfo{ ConnID: st.connID, EndpointID: st.endpointID, Transport: st.transport, RemoteIP: st.remoteIP, SessionToken: st.sessionToken, MaxPacketSize: st.maxPacketSize, } if sess, ok := h.b.uplink.(*Session); ok { sess.HandleDisconnect(context.Background(), info, reason, isCurrent) h.noteConnectionClose(st) return } h.b.uplink.OnDisconnect(context.Background(), info, reason) h.noteConnectionClose(st) } func (h *nixHook) noteConnectionOpen(st *connState) { if h.b.metrics == nil || st == nil || st.metricsCounted { return } h.b.metrics.Connections.WithLabelValues(string(st.transport)).Inc() st.metricsCounted = true } func (h *nixHook) noteConnectionClose(st *connState) { if h.b.metrics == nil || st == nil || !st.metricsCounted { return } h.b.metrics.Connections.WithLabelValues(string(st.transport)).Dec() st.metricsCounted = false } func (h *nixHook) OnQosPublish(cl *mqtt.Client, pk packets.Packet, _ int64, _ int) { if len(pk.Payload) <= largeFrameBytes { return } h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st == nil { return } st.mu.Lock() if st.largePending > 0 { st.largePending-- } if st.largePIDs == nil { st.largePIDs = make(map[uint16]struct{}) } st.largePIDs[pk.PacketID] = struct{}{} st.mu.Unlock() } func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) { h.releaseLargeByPacketID(cl, pk.PacketID) } func (h *nixHook) OnQosDropped(cl *mqtt.Client, pk packets.Packet) { h.releaseLargeByPacketID(cl, pk.PacketID) } func (h *nixHook) releaseLargeByPacketID(cl *mqtt.Client, id uint16) { h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st == nil { return } h.b.releaseLargePID(st, id) }