Files
NixMsg/internal/broker/broker.go
T

491 lines
11 KiB
Go

package broker
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"log/slog"
"net"
"sync"
"sync/atomic"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
mqtt "github.com/mochi-mqtt/server/v2"
"github.com/mochi-mqtt/server/v2/packets"
)
const (
maxClients = 2000
maxPacketSize = 786432
uplinkQueueSize = 256
largeFrameBytes = 64 * 1024
largeFrameSlots = 64
packetOverheadBudget = 128 // 主题与 MQTT 包头预留
keepaliveMin = 10
keepaliveMax = 600
)
// ErrPayloadTooLarge 下行超过客户端 Maximum Packet Size(减包头预留)或 max_receive_bytes。
var ErrPayloadTooLarge = errors.New("broker: payload exceeds client limit")
// ErrNoConnection 目标端没有当前连接。
var ErrNoConnection = errors.New("broker: no active connection")
// AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。
type AuthResult struct {
OK bool
SessionToken string // 密码登录成功时由 N3 填写
}
// Authenticator 由 N3 实现;内部故障必须返回 error,不得当成密码错误。
type Authenticator interface {
Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error)
}
// RejectAuthenticator 默认拒绝所有客户端(CONNACK 用户名密码错误)。
type RejectAuthenticator struct{}
func (RejectAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
return AuthResult{OK: false}, nil
}
// AllowAuthenticator 测试用:允许任意编号。
type AllowAuthenticator struct{}
func (AllowAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
return AuthResult{OK: true}, nil
}
// PublishDroppedFunc 下行未写入发送队列时回调(对接消息线 OnPublishDropped)。
type PublishDroppedFunc func(ctx context.Context, endpointID string, connID port.ConnID, payload []byte)
// Options 装配 broker。
type Options struct {
Authenticator Authenticator
Uplink port.UplinkHandler
Logger *slog.Logger
// OnPublishDropped 可选;nil 时仅打 debug 日志。
OnPublishDropped PublishDroppedFunc
}
// Broker 内置 mochi,不自带监听端口。
type Broker struct {
server *mqtt.Server
auth Authenticator
uplink port.UplinkHandler
log *slog.Logger
onDrop PublishDroppedFunc
hook *nixHook
connsMu sync.RWMutex
current map[string]*connState
byClient map[*mqtt.Client]*connState
queuesMu sync.Mutex
queues map[string]*uplinkQueue
largeSem chan struct{}
closed atomic.Bool
}
type connState struct {
connID port.ConnID
endpointID string
transport port.Transport
remoteIP string
client *mqtt.Client
maxPacketSize uint32
maxRecvBytes int
authOK bool
authErr error
sessionToken string
handshook bool
subscribedDown bool
largeHeld int
mu sync.Mutex
handshakeTimer *time.Timer
}
// New 创建并 Serve mochi(无监听器)。
func New(opts Options) (*Broker, error) {
auth := opts.Authenticator
if auth == nil {
auth = RejectAuthenticator{}
}
uplink := opts.Uplink
if uplink == nil {
uplink = port.StubUplinkHandler{}
}
log := opts.Logger
if log == nil {
log = slog.Default()
}
caps := mqtt.NewDefaultServerCapabilities()
caps.MaximumClients = maxClients
caps.MaximumQos = 1
caps.MaximumPacketSize = maxPacketSize
caps.MaximumSessionExpiryInterval = 0
caps.ReceiveMaximum = 1024
caps.MaximumInflight = 1024
caps.MaximumClientWritesPending = 1024
caps.RetainAvailable = 0
caps.WildcardSubAvailable = 0
caps.SharedSubAvailable = 0
caps.TopicAliasMaximum = 0
caps.Compatibilities.ObscureNotAuthorized = true
srv := mqtt.New(&mqtt.Options{
InlineClient: true,
Capabilities: caps,
Logger: log,
})
b := &Broker{
server: srv,
auth: auth,
uplink: uplink,
log: log,
onDrop: opts.OnPublishDropped,
current: make(map[string]*connState),
byClient: make(map[*mqtt.Client]*connState),
queues: make(map[string]*uplinkQueue),
largeSem: make(chan struct{}, largeFrameSlots),
}
b.hook = &nixHook{b: b}
if err := srv.AddHook(b.hook, nil); err != nil {
return nil, err
}
if err := srv.Serve(); err != nil {
return nil, err
}
return b, nil
}
// Server 返回底层 mochi(测试用)。
func (b *Broker) Server() *mqtt.Server { return b.server }
// Close 关闭 broker。
func (b *Broker) Close() error {
if b.closed.Swap(true) {
return nil
}
b.queuesMu.Lock()
for _, q := range b.queues {
q.close()
}
b.queuesMu.Unlock()
return b.server.Close()
}
// AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。
func (b *Broker) AttachTCP(conn net.Conn) error {
return b.server.EstablishConnection("tcp", conn)
}
// AttachWS 把 WebSocket NetConn 交给 mochi;阻塞到连接结束。
func (b *Broker) AttachWS(conn net.Conn) error {
return b.server.EstablishConnection("ws", conn)
}
// PublishDown 实现 port.Downlink。
func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error {
if b.closed.Load() {
return errors.New("broker: closed")
}
st := b.lookupConn(endpointID, connID)
if st == nil {
return ErrNoConnection
}
limit := effectivePayloadLimit(st.maxPacketSize, st.maxRecvBytes)
if limit > 0 && len(payload) > limit {
return ErrPayloadTooLarge
}
qos := opts.QoS
if qos > 1 {
qos = 1
}
topic := downTopic(endpointID)
large := len(payload) > largeFrameBytes
if large {
select {
case b.largeSem <- struct{}{}:
case <-ctx.Done():
return ctx.Err()
}
st.mu.Lock()
st.largeHeld++
st.mu.Unlock()
}
if err := b.server.Publish(topic, payload, false, qos); err != nil {
if large {
b.releaseOneLarge(st)
}
return err
}
if large && qos == 0 {
b.releaseOneLarge(st)
}
return nil
}
func (b *Broker) releaseOneLarge(st *connState) {
st.mu.Lock()
if st.largeHeld > 0 {
st.largeHeld--
st.mu.Unlock()
select {
case <-b.largeSem:
default:
}
return
}
st.mu.Unlock()
}
func (b *Broker) releaseAllLarge(st *connState) {
st.mu.Lock()
n := st.largeHeld
st.largeHeld = 0
st.mu.Unlock()
for i := 0; i < n; i++ {
select {
case <-b.largeSem:
default:
}
}
}
// Disconnect 实现 port.ConnControl。
func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.ConnID, reason port.DisconnectReason) error {
st := b.lookupConn(endpointID, connID)
if st == nil {
return ErrNoConnection
}
code := packets.CodeDisconnect
switch reason {
case port.DisconnectTakenOver:
code = packets.ErrSessionTakenOver
case port.DisconnectKicked, port.DisconnectFatal:
code = packets.ErrAdministrativeAction
}
err := b.server.DisconnectClient(st.client, code)
// mochi 对错误类原因码会把 Code 当作 error 返回,表示已按该原因断开,不算失败。
if _, ok := err.(packets.Code); ok {
return nil
}
return err
}
func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
b.connsMu.RLock()
defer b.connsMu.RUnlock()
if connID != "" {
for _, st := range b.byClient {
if st.endpointID == endpointID && st.connID == connID {
return st
}
}
return nil
}
return b.current[endpointID]
}
func downTopic(endpointID string) string {
return "nix/c/" + endpointID + "/down"
}
func upTopic(endpointID string) string {
return "nix/c/" + endpointID + "/up"
}
func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int {
limit := 0
if maxPacketSize > 0 {
if maxPacketSize > packetOverheadBudget {
limit = int(maxPacketSize) - packetOverheadBudget
}
}
if maxRecvBytes > 0 {
if limit == 0 || maxRecvBytes < limit {
limit = maxRecvBytes
}
}
return limit
}
func randomConnID() port.ConnID {
var b [16]byte
_, _ = rand.Read(b[:])
return port.ConnID(hex.EncodeToString(b[:]))
}
func transportOf(cl *mqtt.Client) port.Transport {
if cl != nil && cl.Net.Listener == "ws" {
return port.TransportWS
}
return port.TransportTCP
}
func remoteIPOf(cl *mqtt.Client) string {
if cl == nil {
return ""
}
addr := cl.Net.Remote
if addr == "" && cl.Net.Conn != nil && cl.Net.Conn.RemoteAddr() != nil {
addr = cl.Net.Conn.RemoteAddr().String()
}
host, _, err := net.SplitHostPort(addr)
if err != nil {
return addr
}
return host
}
// SetMaxReceiveBytes 供 N3 握手后设置;0 表示不限。
func (b *Broker) SetMaxReceiveBytes(endpointID string, connID port.ConnID, n int) {
st := b.lookupConn(endpointID, connID)
if st == nil {
return
}
st.mu.Lock()
st.maxRecvBytes = n
st.mu.Unlock()
}
// ConnInfoOf 返回连接信息(测试/N3)。
func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) {
b.connsMu.RLock()
st := b.current[endpointID]
b.connsMu.RUnlock()
if st == nil {
return port.ConnInfo{}, false
}
return port.ConnInfo{
ConnID: st.connID,
EndpointID: st.endpointID,
Transport: st.transport,
RemoteIP: st.remoteIP,
SessionToken: st.sessionToken,
MaxPacketSize: st.maxPacketSize,
}, true
}
// IsHandshook 当前连接是否已完成握手。
func (b *Broker) IsHandshook(endpointID string) bool {
b.connsMu.RLock()
st := b.current[endpointID]
b.connsMu.RUnlock()
if st == nil {
return false
}
st.mu.Lock()
defer st.mu.Unlock()
return st.handshook
}
// CurrentConnID 返回端的当前连接代号。
func (b *Broker) CurrentConnID(endpointID string) (port.ConnID, bool) {
b.connsMu.RLock()
st := b.current[endpointID]
b.connsMu.RUnlock()
if st == nil {
return "", false
}
return st.connID, true
}
func (b *Broker) connStateOf(endpointID string, connID port.ConnID) *connState {
b.connsMu.RLock()
defer b.connsMu.RUnlock()
for _, st := range b.byClient {
if st.endpointID == endpointID && st.connID == connID {
return st
}
}
return nil
}
func (b *Broker) hasDownSub(st *connState) bool {
if st == nil {
return false
}
st.mu.Lock()
defer st.mu.Unlock()
if st.subscribedDown {
return true
}
// 回退:直接看 mochi 订阅表
if st.client != nil && st.client.State.Subscriptions != nil {
_, ok := st.client.State.Subscriptions.Get(downTopic(st.endpointID))
return ok
}
return false
}
func (b *Broker) startHandshakeDeadline(endpointID string, connID port.ConnID, d time.Duration) {
st := b.connStateOf(endpointID, connID)
if st == nil {
return
}
st.mu.Lock()
if st.handshook {
st.mu.Unlock()
return
}
if st.handshakeTimer != nil {
st.handshakeTimer.Stop()
}
st.handshakeTimer = time.AfterFunc(d, func() {
cur := b.connStateOf(endpointID, connID)
if cur == nil {
return
}
cur.mu.Lock()
done := cur.handshook
cur.mu.Unlock()
if done {
return
}
_ = b.Disconnect(context.Background(), endpointID, connID, port.DisconnectIdle)
})
st.mu.Unlock()
}
func (b *Broker) cancelHandshakeDeadline(endpointID string, connID port.ConnID) {
st := b.connStateOf(endpointID, connID)
if st == nil {
return
}
st.mu.Lock()
if st.handshakeTimer != nil {
st.handshakeTimer.Stop()
st.handshakeTimer = nil
}
st.mu.Unlock()
}
func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) {
b.queuesMu.Lock()
q, ok := b.queues[endpointID]
if !ok {
q = newUplinkQueue(b, endpointID)
b.queues[endpointID] = q
}
b.queuesMu.Unlock()
q.push(uplinkItem{conn: conn, payload: payload})
}
var (
_ port.Downlink = (*Broker)(nil)
_ port.ConnControl = (*Broker)(nil)
)