375 lines
11 KiB
Go
375 lines
11 KiB
Go
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)
|
||
}
|
||
}
|