Files
NixMsg/test/harness/mqtt.go
T

322 lines
7.1 KiB
Go

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.Now().Add(timeout))
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)
}
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
}