Files
NixMsg/internal/broker/hooks.go
T

268 lines
6.4 KiB
Go
Raw 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,
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))
if h.b.onDrop == nil {
return
}
h.b.connsMu.RLock()
st := h.b.byClient[cl]
h.b.connsMu.RUnlock()
if 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]
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)
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)
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)
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) 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)
}