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 }