343 lines
7.8 KiB
Go
343 lines
7.8 KiB
Go
package accept
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/json"
|
|
"io"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
|
"git.asio.asia/nixevol/NixMsg/test/harness"
|
|
"github.com/mochi-mqtt/server/v2/packets"
|
|
)
|
|
|
|
// MQTTSession 是验收/弱网用的端侧 MQTT 会话(WebSocket + hello + 应用帧)。
|
|
type MQTTSession struct {
|
|
t *testing.T
|
|
mc harness.MQTTClient
|
|
EndpointID string
|
|
pktID uint16
|
|
mu sync.Mutex
|
|
inbox []map[string]any
|
|
closed bool
|
|
done chan struct{}
|
|
}
|
|
|
|
// AppResp 是 type=resp 的解析结果。
|
|
type AppResp struct {
|
|
OK bool
|
|
Error map[string]any
|
|
Data any
|
|
Raw map[string]any
|
|
}
|
|
|
|
// MQTTLogin 用密码连上 /mqtt、订阅 down、完成 hello。
|
|
func MQTTLogin(t *testing.T, httpBase, endpointID, password string) *MQTTSession {
|
|
t.Helper()
|
|
mc, err := harness.DialMQTTWebSocket(httpBase, 10*time.Second)
|
|
if err != nil {
|
|
t.Fatalf("dial mqtt: %v", err)
|
|
}
|
|
s := &MQTTSession{t: t, mc: mc, EndpointID: endpointID, pktID: 10, done: make(chan struct{})}
|
|
s.connectSubscribeHello(password)
|
|
go s.readLoop()
|
|
return s
|
|
}
|
|
|
|
// Close 关闭底层连接。
|
|
func (s *MQTTSession) Close() {
|
|
s.mu.Lock()
|
|
if s.closed {
|
|
s.mu.Unlock()
|
|
return
|
|
}
|
|
s.closed = true
|
|
s.mu.Unlock()
|
|
_ = s.mc.Close()
|
|
select {
|
|
case <-s.done:
|
|
case <-time.After(3 * time.Second):
|
|
}
|
|
}
|
|
|
|
func (s *MQTTSession) nextPkt() uint16 {
|
|
s.pktID++
|
|
if s.pktID == 0 {
|
|
s.pktID = 1
|
|
}
|
|
return s.pktID
|
|
}
|
|
|
|
func (s *MQTTSession) connectSubscribeHello(password string) {
|
|
t := s.t
|
|
pk := packets.Packet{
|
|
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
|
ProtocolVersion: 5,
|
|
Connect: packets.ConnectParams{
|
|
ProtocolName: []byte("MQTT"),
|
|
Clean: true,
|
|
ClientIdentifier: s.EndpointID,
|
|
Keepalive: 30,
|
|
UsernameFlag: true,
|
|
Username: []byte(s.EndpointID),
|
|
PasswordFlag: true,
|
|
Password: []byte(password),
|
|
},
|
|
}
|
|
var buf bytes.Buffer
|
|
if err := pk.ConnectEncode(&buf); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := s.mc.Send(buf.Bytes()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ack, err := s.mc.Recv()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(ack) < 4 || ack[0]>>4 != packets.Connack || ack[3] != 0 {
|
|
t.Fatalf("connack %x", ack)
|
|
}
|
|
|
|
sub := packets.Packet{
|
|
FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1},
|
|
ProtocolVersion: 5,
|
|
PacketID: s.nextPkt(),
|
|
Filters: packets.Subscriptions{
|
|
{Filter: "nix/c/" + s.EndpointID + "/down", Qos: 1},
|
|
},
|
|
}
|
|
buf.Reset()
|
|
if err := sub.SubscribeEncode(&buf); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := s.mc.Send(buf.Bytes()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := s.mc.Recv(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
hello, _ := protocol.Marshal(protocol.Hello{V: protocol.Version, Type: protocol.TypeHello, RID: "h0"})
|
|
s.publishRaw(hello)
|
|
deadline := time.Now().Add(10 * time.Second)
|
|
for time.Now().Before(deadline) {
|
|
raw, err := s.mc.Recv()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
m := s.handlePacket(raw)
|
|
if m == nil {
|
|
continue
|
|
}
|
|
if m["type"] == "resp" && m["ok"] == true {
|
|
return
|
|
}
|
|
if m["type"] == "resp" {
|
|
t.Fatalf("hello failed: %v", m)
|
|
}
|
|
s.push(m)
|
|
}
|
|
t.Fatal("hello timeout")
|
|
}
|
|
|
|
func (s *MQTTSession) publishRaw(payload []byte) {
|
|
t := s.t
|
|
pub := packets.Packet{
|
|
FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1},
|
|
ProtocolVersion: 5,
|
|
TopicName: "nix/c/" + s.EndpointID + "/up",
|
|
PacketID: s.nextPkt(),
|
|
Payload: payload,
|
|
}
|
|
var buf bytes.Buffer
|
|
if err := pub.PublishEncode(&buf); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := s.mc.Send(buf.Bytes()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
func (s *MQTTSession) readLoop() {
|
|
defer close(s.done)
|
|
for {
|
|
raw, err := s.mc.Recv()
|
|
if err != nil {
|
|
return
|
|
}
|
|
m := s.handlePacket(raw)
|
|
if m != nil {
|
|
s.push(m)
|
|
}
|
|
}
|
|
}
|
|
|
|
func (s *MQTTSession) handlePacket(raw []byte) map[string]any {
|
|
if len(raw) < 2 {
|
|
return nil
|
|
}
|
|
typ := raw[0] >> 4
|
|
qos := (raw[0] >> 1) & 0x3
|
|
switch typ {
|
|
case packets.Puback, packets.Pingresp, packets.Suback:
|
|
return nil
|
|
case packets.Publish:
|
|
payload, err := decodePublishPayload(raw)
|
|
if err != nil {
|
|
return nil
|
|
}
|
|
if qos == 1 {
|
|
rem, n, _ := decodeRemainingLength(raw[1:])
|
|
body := raw[1+n:]
|
|
pk := packets.Packet{ProtocolVersion: 5, FixedHeader: packets.FixedHeader{Type: packets.Publish, Remaining: rem, Qos: qos}}
|
|
if decErr := pk.PublishDecode(body); decErr == nil && pk.PacketID != 0 {
|
|
ack := packets.Packet{
|
|
FixedHeader: packets.FixedHeader{Type: packets.Puback},
|
|
ProtocolVersion: 5,
|
|
PacketID: pk.PacketID,
|
|
}
|
|
var buf bytes.Buffer
|
|
if encErr := ack.PubackEncode(&buf); encErr == nil {
|
|
_ = s.mc.Send(buf.Bytes())
|
|
}
|
|
}
|
|
}
|
|
var m map[string]any
|
|
if json.Unmarshal(payload, &m) != nil {
|
|
return nil
|
|
}
|
|
return m
|
|
default:
|
|
return nil
|
|
}
|
|
}
|
|
|
|
func (s *MQTTSession) push(m map[string]any) {
|
|
s.mu.Lock()
|
|
s.inbox = append(s.inbox, m)
|
|
s.mu.Unlock()
|
|
}
|
|
|
|
// Request 发上行帧并等同 rid 的 resp。
|
|
func (s *MQTTSession) Request(t *testing.T, frame map[string]any) AppResp {
|
|
t.Helper()
|
|
rid, _ := frame["rid"].(string)
|
|
payload, err := protocol.Marshal(frame)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
s.publishRaw(payload)
|
|
deadline := time.Now().Add(15 * time.Second)
|
|
for time.Now().Before(deadline) {
|
|
m := s.takeMatching(func(x map[string]any) bool {
|
|
return x["type"] == "resp" && x["rid"] == rid
|
|
})
|
|
if m != nil {
|
|
r := AppResp{OK: m["ok"] == true, Raw: m, Data: m["data"]}
|
|
if e, ok := m["error"].(map[string]any); ok {
|
|
r.Error = e
|
|
}
|
|
return r
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
t.Fatalf("timeout waiting resp rid=%s", rid)
|
|
return AppResp{}
|
|
}
|
|
|
|
// WaitType 等到指定 type 的下行帧。
|
|
func (s *MQTTSession) WaitType(t *testing.T, typ string, timeout time.Duration) map[string]any {
|
|
t.Helper()
|
|
deadline := time.Now().Add(timeout)
|
|
for time.Now().Before(deadline) {
|
|
m := s.takeMatching(func(x map[string]any) bool { return x["type"] == typ })
|
|
if m != nil {
|
|
return m
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
t.Fatalf("timeout waiting type=%s", typ)
|
|
return nil
|
|
}
|
|
|
|
// TryType 在超时内尝试取指定 type;超时返回 nil。
|
|
func (s *MQTTSession) TryType(typ string, timeout time.Duration) map[string]any {
|
|
deadline := time.Now().Add(timeout)
|
|
for time.Now().Before(deadline) {
|
|
m := s.takeMatching(func(x map[string]any) bool { return x["type"] == typ })
|
|
if m != nil {
|
|
return m
|
|
}
|
|
time.Sleep(10 * time.Millisecond)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *MQTTSession) takeMatching(pred func(map[string]any) bool) map[string]any {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
for i, m := range s.inbox {
|
|
if pred(m) {
|
|
s.inbox = append(s.inbox[:i], s.inbox[i+1:]...)
|
|
return m
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// DrainEvents 排空 group_event / presence,避免干扰断言。
|
|
func DrainEvents(t *testing.T, s *MQTTSession, d time.Duration) {
|
|
t.Helper()
|
|
deadline := time.Now().Add(d)
|
|
for time.Now().Before(deadline) {
|
|
_ = s.takeMatching(func(x map[string]any) bool {
|
|
typ, _ := x["type"].(string)
|
|
return typ == "group_event" || typ == "presence"
|
|
})
|
|
time.Sleep(20 * time.Millisecond)
|
|
}
|
|
}
|
|
|
|
func decodePublishPayload(raw []byte) ([]byte, error) {
|
|
if len(raw) < 2 {
|
|
return nil, io.ErrUnexpectedEOF
|
|
}
|
|
rem, n, err := decodeRemainingLength(raw[1:])
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
body := raw[1+n:]
|
|
if len(body) != rem {
|
|
return nil, io.ErrUnexpectedEOF
|
|
}
|
|
pk := packets.Packet{
|
|
ProtocolVersion: 5,
|
|
FixedHeader: packets.FixedHeader{
|
|
Type: packets.Publish,
|
|
Remaining: rem,
|
|
Qos: (raw[0] >> 1) & 0x3,
|
|
},
|
|
}
|
|
if err := pk.PublishDecode(body); err != nil {
|
|
return nil, err
|
|
}
|
|
return pk.Payload, nil
|
|
}
|
|
|
|
func decodeRemainingLength(b []byte) (value int, n int, err error) {
|
|
var mul uint32 = 1
|
|
var v uint32
|
|
for i := 0; i < len(b) && i < 4; i++ {
|
|
v += uint32(b[i]&127) * mul
|
|
n++
|
|
if b[i]&128 == 0 {
|
|
return int(v), n, nil
|
|
}
|
|
mul *= 128
|
|
}
|
|
return 0, 0, io.ErrUnexpectedEOF
|
|
}
|