feat: 实现 WebSocket、mochi broker 与下行发布
This commit is contained in:
@@ -0,0 +1,196 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user