feat: 实现 WebSocket、mochi broker 与下行发布
This commit is contained in:
@@ -0,0 +1,379 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"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
|
||||
}
|
||||
|
||||
// Options 装配 broker。
|
||||
type Options struct {
|
||||
Authenticator Authenticator
|
||||
Uplink port.UplinkHandler
|
||||
Logger *slog.Logger
|
||||
}
|
||||
|
||||
// Broker 内置 mochi,不自带监听端口。
|
||||
type Broker struct {
|
||||
server *mqtt.Server
|
||||
auth Authenticator
|
||||
uplink port.UplinkHandler
|
||||
log *slog.Logger
|
||||
|
||||
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
|
||||
largeHeld int
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// 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,
|
||||
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
|
||||
}
|
||||
return b.server.DisconnectClient(st.client, code)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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)
|
||||
)
|
||||
@@ -0,0 +1,259 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"github.com/coder/websocket"
|
||||
"github.com/mochi-mqtt/server/v2/packets"
|
||||
)
|
||||
|
||||
func TestWSCrossOriginAllowed(t *testing.T) {
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle("/mqtt", b.WSHandler(nil))
|
||||
srv := httptest.NewServer(mux)
|
||||
defer srv.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
c, _, err := websocket.Dial(ctx, "ws"+srv.URL[len("http"):]+"/mqtt", &websocket.DialOptions{
|
||||
HTTPHeader: http.Header{"Origin": []string{"https://other.example"}},
|
||||
Subprotocols: []string{"mqtt"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("cross-origin dial: %v", err)
|
||||
}
|
||||
defer func() { _ = c.Close(websocket.StatusNormalClosure, "") }()
|
||||
if c.Subprotocol() != "mqtt" {
|
||||
t.Fatalf("subprotocol=%q", c.Subprotocol())
|
||||
}
|
||||
}
|
||||
|
||||
func TestWSWrongSubprotocolClosed(t *testing.T) {
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle("/mqtt", b.WSHandler(nil))
|
||||
srv := httptest.NewServer(mux)
|
||||
defer srv.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
c, _, err := websocket.Dial(ctx, "ws"+srv.URL[len("http"):]+"/mqtt", &websocket.DialOptions{
|
||||
HTTPHeader: http.Header{"Origin": []string{"https://other.example"}},
|
||||
Subprotocols: []string{"not-mqtt"},
|
||||
})
|
||||
if err != nil {
|
||||
// 有的实现在握手阶段就失败;也算关闭
|
||||
return
|
||||
}
|
||||
defer func() { _ = c.Close(websocket.StatusNormalClosure, "") }()
|
||||
|
||||
// 服务端应立刻关掉;后续读写会失败
|
||||
c.SetReadLimit(16)
|
||||
_, _, readErr := c.Read(ctx)
|
||||
if readErr == nil {
|
||||
t.Fatal("expected connection closed for wrong subprotocol")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishDownExceedsClientMax(t *testing.T) {
|
||||
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
clientDone := make(chan struct{})
|
||||
r, w := net.Pipe()
|
||||
go func() {
|
||||
defer close(clientDone)
|
||||
_ = b.AttachTCP(r)
|
||||
}()
|
||||
|
||||
endpoint := "ep-limit"
|
||||
connectAndSubscribe(t, w, endpoint, 200) // MaximumPacketSize=200 → payload limit 72
|
||||
|
||||
// 等会话建立
|
||||
deadline := time.Now().Add(3 * time.Second)
|
||||
for {
|
||||
if _, ok := b.ConnInfoOf(endpoint); ok {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatal("session not established")
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
big := bytes.Repeat([]byte("x"), 100) // > 200-128
|
||||
pubErr := b.PublishDown(context.Background(), endpoint, "", big, port.PublishOpts{QoS: 1})
|
||||
if !errors.Is(pubErr, ErrPayloadTooLarge) {
|
||||
t.Fatalf("PublishDown err=%v want ErrPayloadTooLarge", pubErr)
|
||||
}
|
||||
|
||||
// 合法大小应成功
|
||||
small := []byte(`{"v":1,"type":"resp"}`)
|
||||
if err := b.PublishDown(context.Background(), endpoint, "", small, port.PublishOpts{QoS: 0}); err != nil {
|
||||
t.Fatalf("small publish: %v", err)
|
||||
}
|
||||
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-clientDone:
|
||||
case <-time.After(3 * time.Second):
|
||||
}
|
||||
}
|
||||
|
||||
func TestInternalAuthErrorDoesNotReturnBadPassword(t *testing.T) {
|
||||
auth := &errAuthenticator{err: context.DeadlineExceeded}
|
||||
b, err := New(Options{Authenticator: auth})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
r, w := net.Pipe()
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- b.AttachTCP(r) }()
|
||||
|
||||
writeConnect(t, w, "ep-err", 30, 0)
|
||||
// 不应收到 CONNACK(内部错误直接断开)
|
||||
_ = w.SetReadDeadline(time.Now().Add(500 * time.Millisecond))
|
||||
buf := make([]byte, 64)
|
||||
n, readErr := w.Read(buf)
|
||||
if readErr == nil && n > 0 {
|
||||
// 若收到包,不能是 bad username/password CONNACK (reason 0x86)
|
||||
if n >= 2 && buf[0]>>4 == packets.Connack {
|
||||
t.Fatalf("unexpected connack on internal error: %x", buf[:n])
|
||||
}
|
||||
}
|
||||
_ = w.Close()
|
||||
select {
|
||||
case <-errCh:
|
||||
case <-time.After(2 * time.Second):
|
||||
}
|
||||
}
|
||||
|
||||
func TestRejectUnknownByDefault(t *testing.T) {
|
||||
b, err := New(Options{}) // RejectAuthenticator
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = b.Close() }()
|
||||
|
||||
r, w := net.Pipe()
|
||||
go func() { _ = b.AttachTCP(r) }()
|
||||
writeConnect(t, w, "ep-unknown", 30, 0)
|
||||
_ = w.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
buf := make([]byte, 128)
|
||||
n, err := io.ReadAtLeast(w, buf, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if buf[0]>>4 != packets.Connack {
|
||||
t.Fatalf("want connack, got %x", buf[:n])
|
||||
}
|
||||
_ = w.Close()
|
||||
}
|
||||
|
||||
type errAuthenticator struct {
|
||||
err error
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func (a *errAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
return AuthResult{}, a.err
|
||||
}
|
||||
|
||||
func connectAndSubscribe(t *testing.T, w net.Conn, endpoint string, maxPacket uint32) {
|
||||
t.Helper()
|
||||
writeConnect(t, w, endpoint, 30, maxPacket)
|
||||
// read CONNACK
|
||||
_ = w.SetReadDeadline(time.Now().Add(3 * time.Second))
|
||||
buf := make([]byte, 256)
|
||||
n, err := io.ReadAtLeast(w, buf, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if buf[0]>>4 != packets.Connack {
|
||||
t.Fatalf("want connack got %x", buf[:n])
|
||||
}
|
||||
writeSubscribe(t, w, downTopic(endpoint))
|
||||
// read SUBACK
|
||||
n, err = io.ReadAtLeast(w, buf, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if buf[0]>>4 != packets.Suback {
|
||||
t.Fatalf("want suback got %x", buf[:n])
|
||||
}
|
||||
}
|
||||
|
||||
func writeConnect(t *testing.T, w net.Conn, endpoint string, keepalive uint16, maxPacket uint32) {
|
||||
t.Helper()
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
||||
ProtocolVersion: 5,
|
||||
Connect: packets.ConnectParams{
|
||||
ProtocolName: []byte("MQTT"),
|
||||
Clean: true,
|
||||
ClientIdentifier: endpoint,
|
||||
Keepalive: keepalive,
|
||||
UsernameFlag: true,
|
||||
Username: []byte(endpoint),
|
||||
PasswordFlag: true,
|
||||
Password: []byte("test"),
|
||||
},
|
||||
Properties: packets.Properties{
|
||||
MaximumPacketSize: maxPacket,
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := pk.ConnectEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := w.Write(buf.Bytes()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func writeSubscribe(t *testing.T, w net.Conn, topic string) {
|
||||
t.Helper()
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1},
|
||||
ProtocolVersion: 5,
|
||||
PacketID: 1,
|
||||
Filters: packets.Subscriptions{
|
||||
{Filter: topic, Qos: 1},
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := pk.SubscribeEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := w.Write(buf.Bytes()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
)
|
||||
|
||||
type uplinkItem struct {
|
||||
conn port.ConnInfo
|
||||
payload []byte
|
||||
}
|
||||
|
||||
// uplinkQueue 每端串行队列,长度 256,满了堵住 OnPublish(背压)。
|
||||
type uplinkQueue struct {
|
||||
b *Broker
|
||||
endpointID string
|
||||
ch chan uplinkItem
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func newUplinkQueue(b *Broker, endpointID string) *uplinkQueue {
|
||||
q := &uplinkQueue{
|
||||
b: b,
|
||||
endpointID: endpointID,
|
||||
ch: make(chan uplinkItem, uplinkQueueSize),
|
||||
}
|
||||
go q.loop()
|
||||
return q
|
||||
}
|
||||
|
||||
func (q *uplinkQueue) push(item uplinkItem) {
|
||||
q.ch <- item // 满则阻塞读循环,形成背压
|
||||
}
|
||||
|
||||
func (q *uplinkQueue) close() {
|
||||
q.once.Do(func() { close(q.ch) })
|
||||
}
|
||||
|
||||
func (q *uplinkQueue) loop() {
|
||||
for item := range q.ch {
|
||||
_ = q.b.uplink.HandleUplink(context.Background(), item.conn, item.payload)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package broker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/listener"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
|
||||
// WSHandler 返回 /mqtt 的 WebSocket 升级处理。
|
||||
// Accept 时 InsecureSkipVerify=true;之后检查 Subprotocol==mqtt。
|
||||
// NetConn 使用 Background 派生的 context,不用请求 Context。
|
||||
func (b *Broker) WSHandler(proxies *listener.ProxySet) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
c, err := websocket.Accept(w, r, &websocket.AcceptOptions{
|
||||
Subprotocols: []string{"mqtt"},
|
||||
InsecureSkipVerify: true,
|
||||
})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if c.Subprotocol() != "mqtt" {
|
||||
_ = c.Close(websocket.StatusPolicyViolation, "subprotocol must be mqtt")
|
||||
return
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
nc := websocket.NetConn(ctx, c, websocket.MessageBinary)
|
||||
if proxies != nil {
|
||||
ip := proxies.ClientIP(r)
|
||||
if ip != "" {
|
||||
nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)})
|
||||
}
|
||||
}
|
||||
_ = b.AttachWS(nc)
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user