Files
NixMsg/internal/admin/a3_test.go
T

375 lines
11 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)
}
}