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 "" }