feat: 接线 serve 真实 auth/注册/管理/MQTT 与消息循环
This commit is contained in:
@@ -0,0 +1,313 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
"git.asio.asia/nixevol/NixMsg/test/harness"
|
||||
"github.com/mochi-mqtt/server/v2/packets"
|
||||
)
|
||||
|
||||
func TestWireAdminLoginRegisterMQTTHandshake(t *testing.T) {
|
||||
dataDir := t.TempDir()
|
||||
cfgPath := writeTestConfig(t, dataDir)
|
||||
initAdminForTest(t, dataDir)
|
||||
enableRegistration(t, dataDir, "wire-code-99")
|
||||
|
||||
cfg, err := config.Load(cfgPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if vErr := cfg.Validate(); vErr != nil {
|
||||
t.Fatal(vErr)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- runServe(ctx, cfg) }()
|
||||
|
||||
addr := waitListenAddr(t, dataDir, 15*time.Second)
|
||||
base := "http://" + addr
|
||||
|
||||
// 1) admin init 后真实进程可管理登录
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar, Timeout: 10 * time.Second}
|
||||
loginBody, _ := json.Marshal(map[string]string{
|
||||
"username": "admin",
|
||||
"password": "test-admin-password-xx",
|
||||
})
|
||||
resp, err := client.Post(base+"/api/admin/login", "application/json", bytes.NewReader(loginBody))
|
||||
if err != nil {
|
||||
t.Fatalf("admin login: %v", err)
|
||||
}
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("admin login status=%d body=%s", resp.StatusCode, body)
|
||||
}
|
||||
var loginEnv struct {
|
||||
OK bool `json:"ok"`
|
||||
}
|
||||
if uErr := json.Unmarshal(body, &loginEnv); uErr != nil || !loginEnv.OK {
|
||||
t.Fatalf("admin login resp=%s", body)
|
||||
}
|
||||
|
||||
// 2) 已写入注册开关与安全码后可注册
|
||||
regBody := `{"registration_code":"wire-code-99","id":"ep_wire1","login_password":"password12","name":"接线端"}`
|
||||
regResp, err := http.Post(base+"/api/client/register", "application/json", strings.NewReader(regBody))
|
||||
if err != nil {
|
||||
t.Fatalf("register: %v", err)
|
||||
}
|
||||
regBytes, _ := io.ReadAll(regResp.Body)
|
||||
_ = regResp.Body.Close()
|
||||
if regResp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("register status=%d body=%s", regResp.StatusCode, regBytes)
|
||||
}
|
||||
var regEnv struct {
|
||||
OK bool `json:"ok"`
|
||||
Data struct {
|
||||
ID string `json:"id"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if uErr := json.Unmarshal(regBytes, ®Env); uErr != nil || !regEnv.OK || regEnv.Data.ID != "ep_wire1" {
|
||||
t.Fatalf("register resp=%s", regBytes)
|
||||
}
|
||||
|
||||
// 3) 注册出的端用密码完成 MQTT 握手并拿到 session_token
|
||||
tok := mqttPasswordHandshake(t, base, "ep_wire1", "password12")
|
||||
if tok == "" || !strings.HasPrefix(tok, protocol.SessionTokenPrefix) {
|
||||
t.Fatalf("session_token=%q", tok)
|
||||
}
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case err := <-errCh:
|
||||
if err != nil {
|
||||
t.Fatalf("serve exit: %v", err)
|
||||
}
|
||||
case <-time.After(15 * time.Second):
|
||||
t.Fatal("serve did not stop")
|
||||
}
|
||||
}
|
||||
|
||||
func mqttPasswordHandshake(t *testing.T, httpBase, endpointID, password string) string {
|
||||
t.Helper()
|
||||
mc, err := harness.DialMQTTWebSocket(httpBase, 10*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("dial mqtt ws: %v", err)
|
||||
}
|
||||
defer func() { _ = mc.Close() }()
|
||||
|
||||
pk := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
||||
ProtocolVersion: 5,
|
||||
Connect: packets.ConnectParams{
|
||||
ProtocolName: []byte("MQTT"),
|
||||
Clean: true,
|
||||
ClientIdentifier: endpointID,
|
||||
Keepalive: 30,
|
||||
UsernameFlag: true,
|
||||
Username: []byte(endpointID),
|
||||
PasswordFlag: true,
|
||||
Password: []byte(password),
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if encErr := pk.ConnectEncode(&buf); encErr != nil {
|
||||
t.Fatal(encErr)
|
||||
}
|
||||
if sendErr := mc.Send(buf.Bytes()); sendErr != nil {
|
||||
t.Fatal(sendErr)
|
||||
}
|
||||
ack, err := mc.Recv()
|
||||
if err != nil {
|
||||
t.Fatalf("connack: %v", err)
|
||||
}
|
||||
if len(ack) < 2 || ack[0]>>4 != packets.Connack {
|
||||
t.Fatalf("want CONNACK, got %x", ack)
|
||||
}
|
||||
// MQTT5 CONNACK: remaining length, flags, reason code
|
||||
reason := byte(0)
|
||||
if len(ack) >= 4 {
|
||||
reason = ack[3]
|
||||
}
|
||||
if reason != 0 {
|
||||
t.Fatalf("connack reason=%d raw=%x", reason, ack)
|
||||
}
|
||||
|
||||
sub := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1},
|
||||
ProtocolVersion: 5,
|
||||
PacketID: 1,
|
||||
Filters: packets.Subscriptions{
|
||||
{Filter: "nix/c/" + endpointID + "/down", Qos: 1},
|
||||
},
|
||||
}
|
||||
buf.Reset()
|
||||
if err := sub.SubscribeEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := mc.Send(buf.Bytes()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := mc.Recv(); err != nil { // SUBACK
|
||||
t.Fatalf("suback: %v", err)
|
||||
}
|
||||
|
||||
hello, _ := protocol.Marshal(protocol.Hello{
|
||||
V: protocol.Version, Type: protocol.TypeHello, RID: "h1",
|
||||
})
|
||||
pub := packets.Packet{
|
||||
FixedHeader: packets.FixedHeader{Type: packets.Publish, Qos: 1},
|
||||
ProtocolVersion: 5,
|
||||
TopicName: "nix/c/" + endpointID + "/up",
|
||||
PacketID: 2,
|
||||
Payload: hello,
|
||||
}
|
||||
buf.Reset()
|
||||
if err := pub.PublishEncode(&buf); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := mc.Send(buf.Bytes()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(10 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
raw, err := mc.Recv()
|
||||
if err != nil {
|
||||
t.Fatalf("recv down: %v", err)
|
||||
}
|
||||
if len(raw) < 2 {
|
||||
continue
|
||||
}
|
||||
typ := raw[0] >> 4
|
||||
if typ == packets.Puback || typ == packets.Pingresp {
|
||||
continue
|
||||
}
|
||||
if typ != packets.Publish {
|
||||
continue
|
||||
}
|
||||
payload, err := decodePublishPayload(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("publish decode: %v raw=%x", err, raw)
|
||||
}
|
||||
var m map[string]any
|
||||
if err := json.Unmarshal(payload, &m); err != nil {
|
||||
t.Fatalf("json: %v payload=%s", err, payload)
|
||||
}
|
||||
if m["type"] == "resp" {
|
||||
if tok := extractSessionToken(m); tok != "" {
|
||||
return tok
|
||||
}
|
||||
t.Fatalf("hello resp without token: %v", m)
|
||||
}
|
||||
}
|
||||
t.Fatal("timeout waiting hello resp")
|
||||
return ""
|
||||
}
|
||||
|
||||
func decodePublishPayload(raw []byte) ([]byte, error) {
|
||||
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
|
||||
}
|
||||
|
||||
func extractSessionToken(m map[string]any) string {
|
||||
if m["ok"] != true {
|
||||
return ""
|
||||
}
|
||||
data, _ := m["data"].(map[string]any)
|
||||
tok, _ := data["session_token"].(string)
|
||||
return tok
|
||||
}
|
||||
|
||||
func enableRegistration(t *testing.T, dataDir, code string) {
|
||||
t.Helper()
|
||||
db, err := store.Open(dataDir, "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = db.Close() }()
|
||||
now := time.Now().UnixMilli()
|
||||
err = db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`INSERT INTO settings(key, value, updated_at) VALUES(?,?,?)
|
||||
ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at`,
|
||||
"registration_enabled", "1", now); e != nil {
|
||||
return e
|
||||
}
|
||||
_, e := tx.Exec(`INSERT INTO settings(key, value, updated_at) VALUES(?,?,?)
|
||||
ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at`,
|
||||
"registration_code", code, now)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func waitListenAddr(t *testing.T, dataDir string, timeout time.Duration) string {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(timeout)
|
||||
path := filepath.Join(dataDir, "listen.addr")
|
||||
for time.Now().Before(deadline) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err == nil {
|
||||
addr := strings.TrimSpace(string(b))
|
||||
if addr != "" {
|
||||
return addr
|
||||
}
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
t.Fatal("listen.addr not written")
|
||||
return ""
|
||||
}
|
||||
Reference in New Issue
Block a user