package broker import ( "bytes" "context" "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.OnQosComplete, mqtt.OnSubscribed, }, []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, } // ClientID、Username 都必须等于端编号 if clientID == "" || endpointID == "" || clientID != endpointID { st.authOK = false h.rememberPending(cl, st) 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 } res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP) if err != nil { st.authErr = err h.rememberPending(cl, st) return err // mochi 不回 CONNACK,直接断开 } st.authOK = res.OK 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.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) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) { h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload)) } func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) { h.b.connsMu.Lock() st := h.b.byClient[cl] if st != nil { h.b.current[st.endpointID] = st } h.b.connsMu.Unlock() if st == nil { return } 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) } func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) { h.b.connsMu.Lock() st := h.b.byClient[cl] delete(h.b.byClient, cl) 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 } 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) return } h.b.uplink.OnDisconnect(context.Background(), info, reason) } func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) { if len(pk.Payload) <= largeFrameBytes { return } h.b.connsMu.RLock() st := h.b.byClient[cl] h.b.connsMu.RUnlock() if st == nil { return } h.b.releaseOneLarge(st) }