feat: 实现 WebSocket、mochi broker 与下行发布

This commit is contained in:
Nixevol
2026-09-30 06:56:51 +08:00
parent 407a023a68
commit 532ee44da3
8 changed files with 970 additions and 1 deletions
+259
View File
@@ -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)
}
}