package harness import ( "bufio" "crypto/rand" "crypto/sha1" "encoding/base64" "encoding/binary" "fmt" "io" "net" "net/http" "net/url" "strings" "time" ) // MQTTClient 测试用 MQTT 传输层:能连上 WebSocket /mqtt 或裸 TCP,收发原始 MQTT 控制包字节。 // 不实现业务握手(hello)与主题约定;完整帧协议留给各线集成测试自行组合。 type MQTTClient interface { Send(packet []byte) error Recv() ([]byte, error) Close() error LocalAddr() net.Addr RemoteAddr() net.Addr } // DialMQTTTCP 连接裸 MQTT TCP(与 listen 同一地址)。 func DialMQTTTCP(addr string, timeout time.Duration) (MQTTClient, error) { if timeout <= 0 { timeout = 5 * time.Second } conn, err := net.DialTimeout("tcp", addr, timeout) if err != nil { return nil, err } _ = conn.SetDeadline(time.Time{}) return &tcpMQTT{conn: conn, r: bufio.NewReader(conn)}, nil } // DialMQTTWebSocket 连接 ws(s)://host/mqtt,子协议 mqtt。 func DialMQTTWebSocket(httpBase string, timeout time.Duration) (MQTTClient, error) { if timeout <= 0 { timeout = 5 * time.Second } base := strings.TrimRight(httpBase, "/") u, err := url.Parse(base) if err != nil { return nil, err } switch u.Scheme { case "http": u.Scheme = "ws" case "https": u.Scheme = "wss" case "ws", "wss": case "": u.Scheme = "ws" default: return nil, fmt.Errorf("unsupported scheme %q", u.Scheme) } u.Path = "/mqtt" u.RawQuery = "" u.Fragment = "" key := make([]byte, 16) if _, err = rand.Read(key); err != nil { return nil, err } secKey := base64.StdEncoding.EncodeToString(key) httpURL := *u if u.Scheme == "ws" { httpURL.Scheme = "http" } else { httpURL.Scheme = "https" } req, err := http.NewRequest(http.MethodGet, httpURL.String(), nil) if err != nil { return nil, err } req.Header.Set("Connection", "Upgrade") req.Header.Set("Upgrade", "websocket") req.Header.Set("Sec-WebSocket-Version", "13") req.Header.Set("Sec-WebSocket-Key", secKey) req.Header.Set("Sec-WebSocket-Protocol", "mqtt") host := u.Hostname() port := u.Port() if port == "" { if u.Scheme == "wss" { port = "443" } else { port = "80" } } dialer := net.Dialer{Timeout: timeout} raw, err := dialer.Dial("tcp", net.JoinHostPort(host, port)) if err != nil { return nil, err } _ = raw.SetDeadline(time.Now().Add(timeout)) if err = req.Write(raw); err != nil { _ = raw.Close() return nil, err } br := bufio.NewReader(raw) resp, err := http.ReadResponse(br, req) if err != nil { _ = raw.Close() return nil, err } if resp.StatusCode != http.StatusSwitchingProtocols { body, _ := io.ReadAll(io.LimitReader(resp.Body, 512)) _ = resp.Body.Close() _ = raw.Close() return nil, fmt.Errorf("websocket upgrade status %d: %s", resp.StatusCode, body) } if resp.Header.Get("Sec-WebSocket-Accept") != wsAcceptKey(secKey) { _ = raw.Close() return nil, fmt.Errorf("bad Sec-WebSocket-Accept") } if proto := resp.Header.Get("Sec-WebSocket-Protocol"); proto != "" && proto != "mqtt" { _ = raw.Close() return nil, fmt.Errorf("unexpected subprotocol %q", proto) } // 握手完成后清掉超时,否则长会话后续读写会在 dial timeout 到期后全部失败。 _ = raw.SetDeadline(time.Time{}) return &wsMQTT{conn: raw, r: br}, nil } func wsAcceptKey(secKey string) string { const guid = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11" sum := sha1.Sum([]byte(secKey + guid)) return base64.StdEncoding.EncodeToString(sum[:]) } type tcpMQTT struct { conn net.Conn r *bufio.Reader } func (c *tcpMQTT) Send(packet []byte) error { _, err := c.conn.Write(packet) return err } func (c *tcpMQTT) Recv() ([]byte, error) { return readMQTTPacket(c.r) } func (c *tcpMQTT) Close() error { return c.conn.Close() } func (c *tcpMQTT) LocalAddr() net.Addr { return c.conn.LocalAddr() } func (c *tcpMQTT) RemoteAddr() net.Addr { return c.conn.RemoteAddr() } type wsMQTT struct { conn net.Conn r *bufio.Reader } func (c *wsMQTT) Send(packet []byte) error { return writeWSClientBinary(c.conn, packet) } func (c *wsMQTT) Recv() ([]byte, error) { for { payload, opcode, err := readWSFrame(c.r) if err != nil { return nil, err } switch opcode { case 0x2: return payload, nil case 0x8: return nil, io.EOF case 0x9: _ = writeWSClientControl(c.conn, 0xA, payload) case 0xA: continue default: continue } } } func (c *wsMQTT) Close() error { return c.conn.Close() } func (c *wsMQTT) LocalAddr() net.Addr { return c.conn.LocalAddr() } func (c *wsMQTT) RemoteAddr() net.Addr { return c.conn.RemoteAddr() } func readMQTTPacket(r *bufio.Reader) ([]byte, error) { first, err := r.ReadByte() if err != nil { return nil, err } remaining, remBytes, err := readMQTTRemaining(r) if err != nil { return nil, err } buf := make([]byte, 1+len(remBytes)+remaining) buf[0] = first copy(buf[1:], remBytes) if remaining > 0 { if _, err := io.ReadFull(r, buf[1+len(remBytes):]); err != nil { return nil, err } } return buf, nil } func readMQTTRemaining(r *bufio.Reader) (value int, raw []byte, err error) { multiplier := 1 for i := 0; i < 4; i++ { encoded, readErr := r.ReadByte() if readErr != nil { return 0, nil, readErr } raw = append(raw, encoded) value += int(encoded&127) * multiplier if encoded&128 == 0 { return value, raw, nil } multiplier *= 128 } return 0, nil, fmt.Errorf("mqtt remaining length overflow") } func writeWSClientBinary(w io.Writer, payload []byte) error { return writeWSClientFrame(w, 0x2, payload) } func writeWSClientControl(w io.Writer, opcode byte, payload []byte) error { return writeWSClientFrame(w, opcode, payload) } func writeWSClientFrame(w io.Writer, opcode byte, payload []byte) error { mask := make([]byte, 4) if _, err := rand.Read(mask); err != nil { return err } header := []byte{0x80 | (opcode & 0x0f)} n := len(payload) switch { case n < 126: header = append(header, 0x80|byte(n)) case n <= 65535: header = append(header, 0x80|126, byte(n>>8), byte(n)) default: var ext [8]byte binary.BigEndian.PutUint64(ext[:], uint64(n)) header = append(header, 0x80|127) header = append(header, ext[:]...) } header = append(header, mask...) masked := make([]byte, n) for i := 0; i < n; i++ { masked[i] = payload[i] ^ mask[i%4] } if _, err := w.Write(header); err != nil { return err } _, err := w.Write(masked) return err } func readWSFrame(r *bufio.Reader) (payload []byte, opcode byte, err error) { b0, err := r.ReadByte() if err != nil { return nil, 0, err } opcode = b0 & 0x0f b1, err := r.ReadByte() if err != nil { return nil, 0, err } masked := b1&0x80 != 0 n := int(b1 & 0x7f) switch n { case 126: var ext [2]byte if _, err := io.ReadFull(r, ext[:]); err != nil { return nil, 0, err } n = int(binary.BigEndian.Uint16(ext[:])) case 127: var ext [8]byte if _, err := io.ReadFull(r, ext[:]); err != nil { return nil, 0, err } n = int(binary.BigEndian.Uint64(ext[:])) } var maskKey [4]byte if masked { if _, err := io.ReadFull(r, maskKey[:]); err != nil { return nil, 0, err } } payload = make([]byte, n) if n > 0 { if _, err := io.ReadFull(r, payload); err != nil { return nil, 0, err } if masked { for i := 0; i < n; i++ { payload[i] ^= maskKey[i%4] } } } return payload, opcode, nil }