268 lines
6.4 KiB
Go
268 lines
6.4 KiB
Go
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)
|
||
}
|