260 lines
6.5 KiB
Go
260 lines
6.5 KiB
Go
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)
|
|
}
|
|
}
|