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