Files
NixMsg/internal/broker/hooks.go

197 lines
4.6 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)
}