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, }, []byte{b}) } func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error { endpointID := string(pk.Connect.Username) if endpointID == "" { endpointID = pk.Connect.ClientIdentifier } remoteIP := remoteIPOf(cl) st := &connState{ connID: randomConnID(), endpointID: endpointID, transport: transportOf(cl), remoteIP: remoteIP, client: cl, maxPacketSize: pk.Properties.MaximumPacketSize, } // 心跳校正:超出 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) { 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) 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) if st != nil && h.b.current[st.endpointID] == st { delete(h.b.current, st.endpointID) } h.b.connsMu.Unlock() if st == nil { return } h.b.releaseAllLarge(st) 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, } 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) }