feat: 接线 serve 真实 auth/注册/管理/MQTT 与消息循环

This commit is contained in:
Nixevol
2026-09-30 07:56:25 +08:00
parent 550ed534e4
commit 0efcfbb8b6
9 changed files with 699 additions and 80 deletions
+313
View File
@@ -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, &regEnv); 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 ""
}