feat: 实现 test/harness 集成测试启动器
This commit is contained in:
@@ -0,0 +1,161 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrAdminInitUnsupported 表示当前二进制还没有 admin init 子命令。
|
||||
var ErrAdminInitUnsupported = errors.New("admin init unsupported")
|
||||
|
||||
// AdminInitializer 在启动 serve 前初始化管理员密码。
|
||||
// P1 实现 admin init 后,默认实现会解析终端输出中的密码。
|
||||
type AdminInitializer interface {
|
||||
Init(binPath, configPath string) (password string, err error)
|
||||
}
|
||||
|
||||
// CLIAdminInit 调用 `nixmsg admin init`;若命令不存在则返回 ErrAdminInitUnsupported。
|
||||
type CLIAdminInit struct{}
|
||||
|
||||
// Init 执行 admin init。
|
||||
func (CLIAdminInit) Init(binPath, configPath string) (string, error) {
|
||||
cmd := exec.Command(binPath, "admin", "init")
|
||||
cmd.Env = append(os.Environ(), "NIXMSG_CONFIG="+configPath)
|
||||
out, err := cmd.CombinedOutput()
|
||||
text := string(out)
|
||||
if err != nil {
|
||||
if isAdminInitUnsupported(text) {
|
||||
return "", ErrAdminInitUnsupported
|
||||
}
|
||||
return "", fmt.Errorf("admin init: %w\n%s", err, text)
|
||||
}
|
||||
pass := parseAdminPassword(text)
|
||||
if pass == "" {
|
||||
return "", fmt.Errorf("admin init succeeded but password not found in output:\n%s", text)
|
||||
}
|
||||
return pass, nil
|
||||
}
|
||||
|
||||
func isAdminInitUnsupported(text string) bool {
|
||||
lower := strings.ToLower(text)
|
||||
return strings.Contains(lower, "unknown command: admin") ||
|
||||
strings.Contains(lower, "unknown command") && strings.Contains(lower, "admin")
|
||||
}
|
||||
|
||||
func parseAdminPassword(out string) string {
|
||||
// 约定:P1 实现后密码单独占一行,或出现在 "password:" 之后。
|
||||
lines := strings.Split(out, "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
lower := strings.ToLower(line)
|
||||
switch {
|
||||
case strings.HasPrefix(lower, "password:"):
|
||||
return strings.TrimSpace(line[len("password:"):])
|
||||
case strings.HasPrefix(lower, "admin password:"):
|
||||
return strings.TrimSpace(line[len("admin password:"):])
|
||||
}
|
||||
}
|
||||
for i := len(lines) - 1; i >= 0; i-- {
|
||||
line := strings.TrimSpace(lines[i])
|
||||
if len(line) >= 12 && !strings.Contains(strings.ToLower(line), "error") {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) == 1 {
|
||||
return fields[0]
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// AdminClient 管理接口 HTTP 客户端。
|
||||
// Cookie 会话模式下,改变状态的方法会自动加 X-Nixmsg-Request: 1。
|
||||
// 设置 APIToken 后走 Bearer,不再要求该头。
|
||||
type AdminClient struct {
|
||||
BaseURL string
|
||||
HTTP *http.Client
|
||||
APIToken string
|
||||
}
|
||||
|
||||
// NewAdminClient 创建带 CookieJar 的管理客户端。baseURL 形如 http://127.0.0.1:12345。
|
||||
func NewAdminClient(baseURL string) (*AdminClient, error) {
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &AdminClient{
|
||||
BaseURL: strings.TrimRight(baseURL, "/"),
|
||||
HTTP: &http.Client{
|
||||
Timeout: 30 * time.Second,
|
||||
Jar: jar,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// SetSessionCookie 手动写入管理员会话 Cookie(测试辅助)。
|
||||
func (c *AdminClient) SetSessionCookie(value string) error {
|
||||
u, err := url.Parse(c.BaseURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.HTTP.Jar.SetCookies(u, []*http.Cookie{{
|
||||
Name: "nixmsg_admin",
|
||||
Value: value,
|
||||
Path: "/",
|
||||
}})
|
||||
return nil
|
||||
}
|
||||
|
||||
// Do 发送请求。method 为 POST/PUT/PATCH/DELETE 且未使用 API 令牌时,自动加 X-Nixmsg-Request: 1。
|
||||
func (c *AdminClient) Do(method, path string, body []byte, contentType string) (*http.Response, error) {
|
||||
if !strings.HasPrefix(path, "/") {
|
||||
path = "/" + path
|
||||
}
|
||||
var rdr io.Reader
|
||||
if body != nil {
|
||||
rdr = bytes.NewReader(body)
|
||||
}
|
||||
req, err := http.NewRequest(method, c.BaseURL+path, rdr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if contentType != "" {
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
}
|
||||
if c.APIToken != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+c.APIToken)
|
||||
} else if isMutatingMethod(method) {
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
}
|
||||
return c.HTTP.Do(req)
|
||||
}
|
||||
|
||||
// Get JSON GET。
|
||||
func (c *AdminClient) Get(path string) (*http.Response, error) {
|
||||
return c.Do(http.MethodGet, path, nil, "")
|
||||
}
|
||||
|
||||
// PostJSON POST application/json。
|
||||
func (c *AdminClient) PostJSON(path string, body []byte) (*http.Response, error) {
|
||||
return c.Do(http.MethodPost, path, body, "application/json")
|
||||
}
|
||||
|
||||
func isMutatingMethod(method string) bool {
|
||||
switch strings.ToUpper(method) {
|
||||
case http.MethodPost, http.MethodPut, http.MethodPatch, http.MethodDelete:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package harness_test
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/test/harness"
|
||||
)
|
||||
|
||||
func TestAdminClientMutatingHeader(t *testing.T) {
|
||||
t.Parallel()
|
||||
var sawRequestHdr, sawAuth bool
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("X-Nixmsg-Request") == "1" {
|
||||
sawRequestHdr = true
|
||||
}
|
||||
if r.Header.Get("Authorization") != "" {
|
||||
sawAuth = true
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
c, err := harness.NewAdminClient(srv.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp, err := c.PostJSON("/api/admin/login", []byte(`{"password":"x"}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if !sawRequestHdr {
|
||||
t.Fatal("expected X-Nixmsg-Request on POST")
|
||||
}
|
||||
|
||||
sawRequestHdr = false
|
||||
c.APIToken = "nxm_test"
|
||||
resp, err = c.PostJSON("/api/admin/endpoints", []byte(`{}`))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _ = io.Copy(io.Discard, resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if sawRequestHdr {
|
||||
t.Fatal("API token mode should not set X-Nixmsg-Request")
|
||||
}
|
||||
if !sawAuth {
|
||||
t.Fatal("expected Authorization bearer")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var (
|
||||
binOnce sync.Once
|
||||
binPath string
|
||||
binErr error
|
||||
)
|
||||
|
||||
// Binary 返回已编译好的 nixmsg 可执行文件路径(整个进程只编译一次)。
|
||||
func Binary() (string, error) {
|
||||
binOnce.Do(func() {
|
||||
root, err := moduleRoot()
|
||||
if err != nil {
|
||||
binErr = err
|
||||
return
|
||||
}
|
||||
dir, err := os.MkdirTemp("", "nixmsg-harness-bin-*")
|
||||
if err != nil {
|
||||
binErr = err
|
||||
return
|
||||
}
|
||||
name := "nixmsg"
|
||||
if runtime.GOOS == "windows" {
|
||||
name += ".exe"
|
||||
}
|
||||
out := filepath.Join(dir, name)
|
||||
cmd := exec.Command("go", "build", "-o", out, "./cmd/nixmsg")
|
||||
cmd.Dir = root
|
||||
cmd.Env = append(os.Environ(), "CGO_ENABLED=0")
|
||||
if outBytes, runErr := cmd.CombinedOutput(); runErr != nil {
|
||||
binErr = fmt.Errorf("build nixmsg: %w\n%s", runErr, outBytes)
|
||||
return
|
||||
}
|
||||
binPath = out
|
||||
})
|
||||
return binPath, binErr
|
||||
}
|
||||
|
||||
func moduleRoot() (string, error) {
|
||||
dir, err := os.Getwd()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
for {
|
||||
if _, statErr := os.Stat(filepath.Join(dir, "go.mod")); statErr == nil {
|
||||
return dir, nil
|
||||
}
|
||||
parent := filepath.Dir(dir)
|
||||
if parent == dir {
|
||||
return "", fmt.Errorf("go.mod not found from %s", dir)
|
||||
}
|
||||
dir = parent
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/test/harness"
|
||||
)
|
||||
|
||||
func main() {
|
||||
srv, err := harness.Start(harness.Options{KeepDir: true})
|
||||
if err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
info := srv.ConnectionInfo()
|
||||
enc := json.NewEncoder(os.Stdout)
|
||||
enc.SetEscapeHTML(false)
|
||||
if err := enc.Encode(info); err != nil {
|
||||
_ = srv.Stop()
|
||||
_ = os.RemoveAll(srv.DataDir)
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
sigCh := make(chan os.Signal, 1)
|
||||
signal.Notify(sigCh, os.Interrupt, syscall.SIGTERM)
|
||||
<-sigCh
|
||||
|
||||
_ = srv.Stop()
|
||||
_ = os.RemoveAll(srv.DataDir)
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package harness_test
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/test/harness"
|
||||
)
|
||||
|
||||
func TestHealthzAndCleanup(t *testing.T) {
|
||||
srv, err := harness.Start(harness.Options{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dataDir := srv.DataDir
|
||||
httpBase := srv.HTTPBase
|
||||
|
||||
resp, err := http.Get(httpBase + "/healthz")
|
||||
if err != nil {
|
||||
_ = srv.Stop()
|
||||
t.Fatalf("healthz: %v", err)
|
||||
}
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
_ = srv.Stop()
|
||||
t.Fatalf("status=%d body=%s", resp.StatusCode, body)
|
||||
}
|
||||
if string(body) != "ok" {
|
||||
_ = srv.Stop()
|
||||
t.Fatalf("body=%q", body)
|
||||
}
|
||||
|
||||
if err = srv.Stop(); err != nil {
|
||||
t.Fatalf("stop: %v", err)
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if _, statErr := os.Stat(dataDir); os.IsNotExist(statErr) {
|
||||
break
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
if _, statErr := os.Stat(dataDir); !os.IsNotExist(statErr) {
|
||||
t.Fatalf("data dir still exists: %s", dataDir)
|
||||
}
|
||||
|
||||
_, err = http.Get(httpBase + "/healthz")
|
||||
if err == nil {
|
||||
t.Fatal("healthz still reachable after stop")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParallelIsolation(t *testing.T) {
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
addrs []string
|
||||
dataDirs []string
|
||||
)
|
||||
runOne := func(i int) {
|
||||
defer wg.Done()
|
||||
srv, err := harness.Start(harness.Options{})
|
||||
if err != nil {
|
||||
t.Errorf("start %d: %v", i, err)
|
||||
return
|
||||
}
|
||||
defer func() { _ = srv.Stop() }()
|
||||
|
||||
resp, err := http.Get(srv.HTTPBase + "/healthz")
|
||||
if err != nil {
|
||||
t.Errorf("healthz %d: %v", i, err)
|
||||
return
|
||||
}
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Errorf("healthz %d status=%d", i, resp.StatusCode)
|
||||
return
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
addrs = append(addrs, srv.Addr)
|
||||
dataDirs = append(dataDirs, srv.DataDir)
|
||||
mu.Unlock()
|
||||
}
|
||||
|
||||
const n = 2
|
||||
wg.Add(n)
|
||||
for i := 0; i < n; i++ {
|
||||
go runOne(i)
|
||||
}
|
||||
wg.Wait()
|
||||
if t.Failed() {
|
||||
return
|
||||
}
|
||||
if len(addrs) != n || len(dataDirs) != n {
|
||||
t.Fatalf("got %d addrs %d dirs", len(addrs), len(dataDirs))
|
||||
}
|
||||
if addrs[0] == addrs[1] {
|
||||
t.Fatalf("same listen addr: %s", addrs[0])
|
||||
}
|
||||
if dataDirs[0] == dataDirs[1] {
|
||||
t.Fatalf("same data dir: %s", dataDirs[0])
|
||||
}
|
||||
for _, d := range dataDirs {
|
||||
if filepath.Clean(d) == "" {
|
||||
t.Fatal("empty data dir")
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// Package harness 提供集成测试启动器。
|
||||
package harness
|
||||
|
||||
// Info 是命令行启动器打印到 stdout 的连接信息。
|
||||
type Info struct {
|
||||
Listen string `json:"listen"`
|
||||
HTTPBase string `json:"http_base"`
|
||||
AdminListen string `json:"admin_listen,omitempty"`
|
||||
AdminHTTPBase string `json:"admin_http_base"`
|
||||
MQTTWS string `json:"mqtt_ws"`
|
||||
MQTTTCP string `json:"mqtt_tcp"`
|
||||
AdminPassword string `json:"admin_password"`
|
||||
DataDir string `json:"data_dir"`
|
||||
ConfigPath string `json:"config_path"`
|
||||
}
|
||||
|
||||
// ConnectionInfo 组装给外部语言测试用的 JSON 结构。
|
||||
func (s *Server) ConnectionInfo() Info {
|
||||
adminBase := s.AdminHTTPBase
|
||||
if adminBase == "" {
|
||||
adminBase = s.HTTPBase
|
||||
}
|
||||
return Info{
|
||||
Listen: s.Addr,
|
||||
HTTPBase: s.HTTPBase,
|
||||
AdminListen: s.AdminAddr,
|
||||
AdminHTTPBase: adminBase,
|
||||
MQTTWS: "ws://" + s.Addr + "/mqtt",
|
||||
MQTTTCP: s.Addr,
|
||||
AdminPassword: s.AdminPassword,
|
||||
DataDir: s.DataDir,
|
||||
ConfigPath: s.ConfigPath,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,321 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,162 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Options 控制测试进程启动。
|
||||
type Options struct {
|
||||
// Listen 写入配置的 listen,默认 127.0.0.1:0。
|
||||
Listen string
|
||||
// AdminListen 可选;非空时写入 admin_listen,并读取 admin.addr。
|
||||
AdminListen string
|
||||
// AdminInit 可选;nil 时用 CLIAdminInit。不支持时忽略并继续启动。
|
||||
AdminInit AdminInitializer
|
||||
// KeepDir 为 true 时 Stop 不删除数据目录(命令行启动器在清理前可读)。
|
||||
KeepDir bool
|
||||
}
|
||||
|
||||
// Server 是一次集成测试用的真实 nixmsg 进程。
|
||||
type Server struct {
|
||||
BinPath string
|
||||
ConfigPath string
|
||||
DataDir string
|
||||
Addr string // 来自 listen.addr,形如 127.0.0.1:12345
|
||||
AdminAddr string // 来自 admin.addr(若有)
|
||||
AdminPassword string
|
||||
HTTPBase string
|
||||
AdminHTTPBase string
|
||||
|
||||
cmd *exec.Cmd
|
||||
keepDir bool
|
||||
stopped bool
|
||||
}
|
||||
|
||||
// Start 编译(如需)并启动服务:临时目录、listen 端口 0、读 listen.addr。
|
||||
func Start(opts Options) (*Server, error) {
|
||||
bin, err := Binary()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dataDir, err := os.MkdirTemp("", "nixmsg-harness-*")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
listen := opts.Listen
|
||||
if listen == "" {
|
||||
listen = "127.0.0.1:0"
|
||||
}
|
||||
cfgPath := filepath.Join(dataDir, "config.yaml")
|
||||
cfg := fmt.Sprintf("listen: %q\ndata_dir: %q\n", listen, filepath.ToSlash(dataDir))
|
||||
if opts.AdminListen != "" {
|
||||
cfg += fmt.Sprintf("admin_listen: %q\n", opts.AdminListen)
|
||||
}
|
||||
if err = os.WriteFile(cfgPath, []byte(cfg), 0o644); err != nil {
|
||||
_ = os.RemoveAll(dataDir)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
adminInit := opts.AdminInit
|
||||
if adminInit == nil {
|
||||
adminInit = CLIAdminInit{}
|
||||
}
|
||||
password, initErr := adminInit.Init(bin, cfgPath)
|
||||
if initErr != nil && !errors.Is(initErr, ErrAdminInitUnsupported) {
|
||||
_ = os.RemoveAll(dataDir)
|
||||
return nil, initErr
|
||||
}
|
||||
if errors.Is(initErr, ErrAdminInitUnsupported) {
|
||||
password = ""
|
||||
}
|
||||
|
||||
cmd := exec.Command(bin, "serve")
|
||||
cmd.Env = append(os.Environ(), "NIXMSG_CONFIG="+cfgPath)
|
||||
cmd.Stdout = os.Stderr
|
||||
cmd.Stderr = os.Stderr
|
||||
if err = cmd.Start(); err != nil {
|
||||
_ = os.RemoveAll(dataDir)
|
||||
return nil, fmt.Errorf("start serve: %w", err)
|
||||
}
|
||||
|
||||
s := &Server{
|
||||
BinPath: bin,
|
||||
ConfigPath: cfgPath,
|
||||
DataDir: dataDir,
|
||||
AdminPassword: password,
|
||||
cmd: cmd,
|
||||
keepDir: opts.KeepDir,
|
||||
}
|
||||
|
||||
addr, err := waitAddrFile(filepath.Join(dataDir, "listen.addr"), 10*time.Second)
|
||||
if err != nil {
|
||||
_ = s.Stop()
|
||||
return nil, fmt.Errorf("wait listen.addr: %w", err)
|
||||
}
|
||||
s.Addr = addr
|
||||
s.HTTPBase = "http://" + addr
|
||||
|
||||
if opts.AdminListen != "" {
|
||||
adminAddr, adminErr := waitAddrFile(filepath.Join(dataDir, "admin.addr"), 10*time.Second)
|
||||
if adminErr == nil {
|
||||
s.AdminAddr = adminAddr
|
||||
s.AdminHTTPBase = "http://" + adminAddr
|
||||
}
|
||||
}
|
||||
if s.AdminHTTPBase == "" {
|
||||
s.AdminHTTPBase = s.HTTPBase
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// AdminClient 返回指向管理接口的 HTTP 客户端。
|
||||
func (s *Server) AdminClient() (*AdminClient, error) {
|
||||
base := s.AdminHTTPBase
|
||||
if base == "" {
|
||||
base = s.HTTPBase
|
||||
}
|
||||
return NewAdminClient(base)
|
||||
}
|
||||
|
||||
// Stop 结束进程并删除临时目录(除非 KeepDir)。
|
||||
func (s *Server) Stop() error {
|
||||
if s == nil || s.stopped {
|
||||
return nil
|
||||
}
|
||||
s.stopped = true
|
||||
var stopErr error
|
||||
if s.cmd != nil && s.cmd.Process != nil {
|
||||
_ = s.cmd.Process.Kill()
|
||||
_, _ = s.cmd.Process.Wait()
|
||||
}
|
||||
if !s.keepDir && s.DataDir != "" {
|
||||
if err := os.RemoveAll(s.DataDir); err != nil {
|
||||
stopErr = err
|
||||
}
|
||||
}
|
||||
return stopErr
|
||||
}
|
||||
|
||||
func waitAddrFile(path string, timeout time.Duration) (string, error) {
|
||||
deadline := time.Now().Add(timeout)
|
||||
var lastErr error
|
||||
for time.Now().Before(deadline) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
addr := strings.TrimSpace(string(b))
|
||||
if addr != "" {
|
||||
return addr, nil
|
||||
}
|
||||
lastErr = fmt.Errorf("empty addr file")
|
||||
} else {
|
||||
lastErr = err
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
return "", lastErr
|
||||
}
|
||||
Reference in New Issue
Block a user