package admin_test import ( "context" "database/sql" "encoding/json" "net/http" "net/http/cookiejar" "net/http/httptest" "path/filepath" "strings" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/admin" "git.asio.asia/nixevol/NixMsg/internal/app/group" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/config" "git.asio.asia/nixevol/NixMsg/internal/store" ) func setupA3(t *testing.T) (*store.DB, *httptest.Server, *http.Client) { t.Helper() dir := t.TempDir() db, err := store.Open(filepath.Join(dir, "data"), "FULL") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = db.Close() }) hash := auth.NewStubHashPool() if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil { t.Fatal(seedErr) } gApp := group.New(group.Config{DB: db, MaxGroupMembers: 100}) h := admin.New(admin.Deps{ DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(), Locks: admin.NewMemoryLoginLocks(), Groups: gApp, Config: config.Default(), Version: "0.1.0-test", }) srv := httptest.NewServer(h) t.Cleanup(srv.Close) jar, err := cookiejar.New(nil) if err != nil { t.Fatal(err) } client := &http.Client{Jar: jar} login(t, client, srv.URL) return db, srv, client } func csrf() map[string]string { return map[string]string{"X-Nixmsg-Request": "1"} } func insertEndpoint(t *testing.T, db *store.DB, id, source string, enabled bool, online bool) { t.Helper() now := time.Now().UnixMilli() en := 1 if !enabled { en = 0 } var onlineSince, offlineSince any if online { onlineSince = now } else { offlineSince = now } err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { _, e := tx.Exec(`INSERT INTO endpoints( id, name, remark, source, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at, online_since, offline_since ) VALUES(?,?,?,?,?,?,?,?,?,?,?,?)`, id, id, "", source, "stub$hash", nil, 0, 0, en, now, onlineSince, offlineSince) return e }) if err != nil { t.Fatal(err) } } func TestRegistrationToggleAndGenerate(t *testing.T) { _, srv, client := setupA3(t) base := srv.URL res := doReq(t, client, http.MethodGet, base+"/api/admin/registration", "", nil) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("get: %d %+v", res.StatusCode, env) } var got map[string]any _ = json.Unmarshal(env.Data, &got) if got["enabled"] != false { t.Fatalf("default enabled=%v", got["enabled"]) } res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", `{"enabled":true,"code":"abcdefgh"}`, csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("put enable: %d %+v", res.StatusCode, env) } _ = json.Unmarshal(env.Data, &got) if got["enabled"] != true || got["code"] != "abcdefgh" { t.Fatalf("after put: %v", got) } res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", `{"generate":true}`, csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("generate: %d %+v", res.StatusCode, env) } _ = json.Unmarshal(env.Data, &got) code, _ := got["code"].(string) if len(code) != 16 { t.Fatalf("generated len=%d code=%q", len(code), code) } res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", `{"enabled":false}`, csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("disable: %d %+v", res.StatusCode, env) } _ = json.Unmarshal(env.Data, &got) if got["enabled"] != false { t.Fatalf("disabled=%v", got["enabled"]) } } func TestRegistrationEnableRequiresCode(t *testing.T) { db, srv, client := setupA3(t) base := srv.URL res := doReq(t, client, http.MethodPut, base+"/api/admin/registration", `{"enabled":true}`, csrf()) env := decodeEnv(t, res) if res.StatusCode != 400 || env.OK { t.Fatalf("enable without code: %d %+v", res.StatusCode, env) } if env.Error == nil || !strings.Contains(env.Error.Message, "8–64") { t.Fatalf("want 8–64 message, got %+v", env.Error) } res = doReq(t, client, http.MethodGet, base+"/api/admin/registration", "", nil) env = decodeEnv(t, res) var got map[string]any _ = json.Unmarshal(env.Data, &got) if got["enabled"] != false { t.Fatalf("switch must stay off after 400: %v", got) } var n int if err := db.Read.QueryRow(`SELECT COUNT(*) FROM settings WHERE key = ? AND value = '1'`, "registration_enabled").Scan(&n); err != nil { t.Fatal(err) } if n != 0 { t.Fatalf("enabled setting rolled back, count=%d", n) } res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", `{"enabled":true,"generate":true}`, csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("enable+generate: %d %+v", res.StatusCode, env) } _ = json.Unmarshal(env.Data, &got) code, _ := got["code"].(string) if got["enabled"] != true || len(code) != 16 { t.Fatalf("after generate: %v", got) } res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", `{"enabled":false}`, csrf()) if res.StatusCode != 200 { t.Fatalf("disable: %d", res.StatusCode) } _ = res.Body.Close() res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", `{"enabled":true,"code":"abcdefgh"}`, csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("enable+code: %d %+v", res.StatusCode, env) } _ = json.Unmarshal(env.Data, &got) if got["enabled"] != true || got["code"] != "abcdefgh" { t.Fatalf("after enable+code: %v", got) } } func TestMessageDetailHasNoBody(t *testing.T) { db, srv, client := setupA3(t) base := srv.URL now := time.Now().UnixMilli() err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { if _, e := tx.Exec(`INSERT INTO messages( seq, id, sender_id, dest_kind, dest_id, meta, content_type, body_enc, send_at, keep, ttl_seconds, receipt, state, reason, created_at ) VALUES(1,'m1','a','endpoint','b','{}','text/plain','enc',?,?,0,1,'dispatched','',?)`, now, 1, now); e != nil { return e } if _, e := tx.Exec(`INSERT INTO message_bodies(seq, body) VALUES(1, ?)`, []byte("SECRET_BODY_SHOULD_NOT_APPEAR")); e != nil { return e } _, e := tx.Exec(`INSERT INTO deliveries( seq, endpoint_id, send_at, keep, state, reason, attempts, updated_at, pushed_at ) VALUES(1,'b',?,1,'pending','',2,?,?)`, now, now, now) return e }) if err != nil { t.Fatal(err) } res := doReq(t, client, http.MethodGet, base+"/api/admin/messages/1", "", nil) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("detail: %d %+v", res.StatusCode, env) } raw := string(env.Data) if strings.Contains(raw, "body") || strings.Contains(raw, "SECRET_BODY") { t.Fatalf("response contains body/secret: %s", raw) } var detail map[string]any _ = json.Unmarshal(env.Data, &detail) dels, _ := detail["deliveries"].([]any) if len(dels) != 1 { t.Fatalf("deliveries=%v", detail["deliveries"]) } d0 := dels[0].(map[string]any) if d0["attempts"].(float64) != 2 { t.Fatalf("attempts=%v", d0["attempts"]) } res = doReq(t, client, http.MethodGet, base+"/api/admin/messages", "", nil) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("list: %d %+v", res.StatusCode, env) } if strings.Contains(string(env.Data), "SECRET_BODY") || strings.Contains(string(env.Data), `"body"`) { t.Fatalf("list leaked body: %s", env.Data) } } func TestGroupsCRUD(t *testing.T) { db, srv, client := setupA3(t) base := srv.URL insertEndpoint(t, db, "alice", "admin", true, false) insertEndpoint(t, db, "bob", "admin", true, true) insertEndpoint(t, db, "carol", "self", true, false) res := doReq(t, client, http.MethodPost, base+"/api/admin/groups", `{"name":"一组","owner_id":"alice","member_ids":["bob","carol"]}`, csrf()) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("create: %d %+v", res.StatusCode, env) } var created struct { ID string `json:"id"` OwnerID string `json:"owner_id"` } _ = json.Unmarshal(env.Data, &created) if created.ID == "" || created.OwnerID != "alice" { t.Fatalf("created=%+v", created) } res = doReq(t, client, http.MethodGet, base+"/api/admin/groups/"+created.ID+"?limit=1", "", nil) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("get page1: %d %+v", res.StatusCode, env) } var page1 struct { MemberTotal float64 `json:"member_total"` Members []any `json:"members"` NextCursor string `json:"next_cursor"` } _ = json.Unmarshal(env.Data, &page1) if page1.MemberTotal != 3 || len(page1.Members) != 1 || page1.NextCursor == "" { t.Fatalf("page1=%+v raw=%s", page1, env.Data) } res = doReq(t, client, http.MethodGet, base+"/api/admin/groups/"+created.ID+"?limit=1&cursor="+page1.NextCursor, "", nil) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("get page2: %d %+v", res.StatusCode, env) } var page2 struct { Members []any `json:"members"` } _ = json.Unmarshal(env.Data, &page2) if len(page2.Members) != 1 { t.Fatalf("page2 members=%d", len(page2.Members)) } res = doReq(t, client, http.MethodPatch, base+"/api/admin/groups/"+created.ID, `{"name":"新名"}`, csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("rename: %d %+v", res.StatusCode, env) } res = doReq(t, client, http.MethodPost, base+"/api/admin/groups/"+created.ID+"/transfer", `{"endpoint_id":"bob"}`, csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("transfer: %d %+v", res.StatusCode, env) } res = doReq(t, client, http.MethodDelete, base+"/api/admin/groups/"+created.ID+"/members/carol", "", csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("remove: %d %+v", res.StatusCode, env) } res = doReq(t, client, http.MethodDelete, base+"/api/admin/groups/"+created.ID, "", csrf()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("dissolve: %d %+v", res.StatusCode, env) } } func TestOverviewAndSettings(t *testing.T) { db, srv, client := setupA3(t) base := srv.URL insertEndpoint(t, db, "a1", "admin", true, true) insertEndpoint(t, db, "s1", "self", true, false) insertEndpoint(t, db, "d1", "admin", false, false) res := doReq(t, client, http.MethodGet, base+"/api/admin/overview", "", nil) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("overview: %d %+v", res.StatusCode, env) } var ov map[string]any _ = json.Unmarshal(env.Data, &ov) if ov["version"] != "0.1.0-test" { t.Fatalf("version=%v", ov["version"]) } if ov["endpoints_total"].(float64) != 3 { t.Fatalf("total=%v", ov["endpoints_total"]) } if ov["endpoints_self"].(float64) != 1 { t.Fatalf("self=%v", ov["endpoints_self"]) } if ov["endpoints_online"].(float64) != 1 { t.Fatalf("online=%v", ov["endpoints_online"]) } if ov["endpoints_disabled"].(float64) != 1 { t.Fatalf("disabled=%v", ov["endpoints_disabled"]) } res = doReq(t, client, http.MethodGet, base+"/api/admin/settings", "", nil) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("settings: %d %+v", res.StatusCode, env) } raw := string(env.Data) if strings.Contains(raw, "token") || strings.Contains(raw, "password") { t.Fatalf("settings leaked secret: %s", raw) } var st map[string]any _ = json.Unmarshal(env.Data, &st) if st["listen"] == nil || st["limits"] == nil { t.Fatalf("settings=%v", st) } }