324 lines
7.2 KiB
Go
324 lines
7.2 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.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
|
|
}
|