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) } }