feat: 实现管理后台端管理接口(A2)
This commit is contained in:
@@ -0,0 +1,364 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/admin"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
type kickRecorder struct {
|
||||
mu sync.Mutex
|
||||
Calls []string
|
||||
}
|
||||
|
||||
func (k *kickRecorder) Kick(_ context.Context, endpointID string) (bool, error) {
|
||||
k.mu.Lock()
|
||||
defer k.mu.Unlock()
|
||||
k.Calls = append(k.Calls, endpointID)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (k *kickRecorder) count() int {
|
||||
k.mu.Lock()
|
||||
defer k.mu.Unlock()
|
||||
return len(k.Calls)
|
||||
}
|
||||
|
||||
func setupEndpoints(t *testing.T) (*store.DB, *httptest.Server, *http.Client, *kickRecorder, *admin.MemoryLoginLocks) {
|
||||
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)
|
||||
}
|
||||
kick := &kickRecorder{}
|
||||
locks := admin.NewMemoryLoginLocks()
|
||||
h := admin.New(admin.Deps{
|
||||
DB: db,
|
||||
Hash: hash,
|
||||
Tokens: admin.NewRandomAPITokens(),
|
||||
Locks: locks,
|
||||
KickEndpoint: kick.Kick,
|
||||
})
|
||||
srv := httptest.NewServer(h)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar}
|
||||
res := postJSON(t, client, srv.URL+"/api/admin/login",
|
||||
`{"username":"admin","password":"`+testPassword+`"}`, nil)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusOK || !env.OK {
|
||||
t.Fatalf("login failed: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
return db, srv, client, kick, locks
|
||||
}
|
||||
|
||||
func TestEndpointImportDuplicateRejectsAll(t *testing.T) {
|
||||
db, srv, client, _, _ := setupEndpoints(t)
|
||||
base := srv.URL
|
||||
|
||||
csvBody := "" +
|
||||
"id,name,login_password,talk_password,default_delay_seconds,remark\n" +
|
||||
"ep-a,甲,password1,,0,\n" +
|
||||
"ep-b,乙,password2,,0,\n" +
|
||||
"ep-a,丙,password3,,0,\n"
|
||||
|
||||
req, err := http.NewRequest(http.MethodPost, base+"/api/admin/endpoints/import", strings.NewReader(csvBody))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "text/csv")
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _ := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
if res.StatusCode != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d body=%s", res.StatusCode, raw)
|
||||
}
|
||||
var env struct {
|
||||
OK bool `json:"ok"`
|
||||
Error *struct {
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
Data *struct {
|
||||
Errors []struct {
|
||||
Line int `json:"line"`
|
||||
Reason string `json:"reason"`
|
||||
} `json:"errors"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &env); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if env.OK || env.Error == nil || env.Error.Code != "bad_request" {
|
||||
t.Fatalf("env=%+v", env)
|
||||
}
|
||||
if env.Data == nil || len(env.Data.Errors) == 0 {
|
||||
t.Fatalf("missing errors: %s", raw)
|
||||
}
|
||||
foundLine := false
|
||||
for _, e := range env.Data.Errors {
|
||||
if e.Line == 4 {
|
||||
foundLine = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundLine {
|
||||
t.Fatalf("want line 4 in errors: %+v", env.Data.Errors)
|
||||
}
|
||||
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM endpoints`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("want 0 endpoints created, got %d", n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEndpointResetPasswordAppearsOnce(t *testing.T) {
|
||||
db, srv, client, kick, _ := setupEndpoints(t)
|
||||
base := srv.URL
|
||||
|
||||
res := postJSON(t, client, base+"/api/admin/endpoints",
|
||||
`{"id":"dev-1","name":"门口","login_password":"oldpass12"}`,
|
||||
csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
// 写入假会话令牌,确认重置会清空
|
||||
err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE endpoints SET session_hash='abc', session_issued_at=1, session_used_at=1 WHERE id='dev-1'`)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
res = postJSON(t, client, base+"/api/admin/endpoints/dev-1/reset-login-password",
|
||||
`{}`, csrfHeaders())
|
||||
body, _ := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
if res.StatusCode != 200 {
|
||||
t.Fatalf("reset status=%d body=%s", res.StatusCode, body)
|
||||
}
|
||||
var resetEnv struct {
|
||||
OK bool `json:"ok"`
|
||||
Data struct {
|
||||
LoginPassword string `json:"login_password"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &resetEnv); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !resetEnv.OK || resetEnv.Data.LoginPassword == "" {
|
||||
t.Fatalf("want one-time password, got %s", body)
|
||||
}
|
||||
pw := resetEnv.Data.LoginPassword
|
||||
if strings.Count(string(body), pw) != 1 {
|
||||
t.Fatalf("password should appear exactly once in response: %s", body)
|
||||
}
|
||||
if strings.HasPrefix(pw, "nst_") {
|
||||
t.Fatalf("password must not start with nst_: %q", pw)
|
||||
}
|
||||
|
||||
var session sql.NullString
|
||||
if err := db.Read.QueryRow(`SELECT session_hash FROM endpoints WHERE id='dev-1'`).Scan(&session); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if session.Valid {
|
||||
t.Fatal("session_hash should be cleared")
|
||||
}
|
||||
if kick.count() < 1 {
|
||||
t.Fatal("expected kick hook after reset")
|
||||
}
|
||||
|
||||
// 详情中不得再出现明文密码
|
||||
res = doReq(t, client, http.MethodGet, base+"/api/admin/endpoints/dev-1", "", nil)
|
||||
detailBody, _ := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
if strings.Contains(string(detailBody), pw) {
|
||||
t.Fatalf("password leaked in detail: %s", detailBody)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEndpointMutatingRequiresCSRF(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
base := srv.URL
|
||||
|
||||
res := postJSON(t, client, base+"/api/admin/endpoints",
|
||||
`{"id":"no-csrf","name":"x","login_password":"password1"}`,
|
||||
nil) // 无 CSRF
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusForbidden || env.Error == nil || env.Error.Code != "forbidden" {
|
||||
t.Fatalf("want 403 forbidden got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEndpointDisableKickAndUnlock(t *testing.T) {
|
||||
db, srv, client, kick, locks := setupEndpoints(t)
|
||||
base := srv.URL
|
||||
|
||||
res := postJSON(t, client, base+"/api/admin/endpoints",
|
||||
`{"id":"lock-1","name":"锁","login_password":"password1"}`,
|
||||
csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
res = postJSON(t, client, base+"/api/admin/endpoints/batch",
|
||||
`{"ids":["lock-1"],"action":"disable"}`, csrfHeaders())
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("disable: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var enabled int
|
||||
if err := db.Read.QueryRow(`SELECT enabled FROM endpoints WHERE id='lock-1'`).Scan(&enabled); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if enabled != 0 {
|
||||
t.Fatalf("want enabled=0 got %d", enabled)
|
||||
}
|
||||
if kick.count() < 1 {
|
||||
t.Fatal("disable should kick")
|
||||
}
|
||||
|
||||
for i := 0; i < 50; i++ {
|
||||
locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "lock-1"})
|
||||
}
|
||||
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "lock-1"}); !locked {
|
||||
t.Fatal("expected endpoint locked before unlock")
|
||||
}
|
||||
|
||||
res = postJSON(t, client, base+"/api/admin/endpoints/lock-1/unlock", `{}`, csrfHeaders())
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("unlock: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "lock-1"}); locked {
|
||||
t.Fatal("expected unlocked")
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password",
|
||||
`{"talk_password":"talk"}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"})
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("talk-password: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var talkSet struct {
|
||||
TalkPasswordSet bool `json:"talk_password_set"`
|
||||
}
|
||||
_ = json.Unmarshal(env.Data, &talkSet)
|
||||
if !talkSet.TalkPasswordSet {
|
||||
t.Fatal("want talk_password_set true")
|
||||
}
|
||||
|
||||
var ver int
|
||||
if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if ver != 1 {
|
||||
t.Fatalf("talk_version want 1 got %d", ver)
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password",
|
||||
`{"talk_password":""}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"})
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("clear talk: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if ver != 2 {
|
||||
t.Fatalf("talk_version want 2 got %d", ver)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEndpointImportBOMAndCreate(t *testing.T) {
|
||||
db, srv, client, _, _ := setupEndpoints(t)
|
||||
base := srv.URL
|
||||
|
||||
var buf bytes.Buffer
|
||||
buf.Write([]byte{0xEF, 0xBB, 0xBF})
|
||||
buf.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n")
|
||||
buf.WriteString("bom-1,门,,secret,5,备注\n")
|
||||
|
||||
req, err := http.NewRequest(http.MethodPost, base+"/api/admin/endpoints/import", &buf)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "text/csv; charset=utf-8")
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("import: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var data struct {
|
||||
Items []struct {
|
||||
ID string `json:"id"`
|
||||
LoginPassword string `json:"login_password"`
|
||||
Name string `json:"name"`
|
||||
} `json:"items"`
|
||||
}
|
||||
if err := json.Unmarshal(env.Data, &data); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(data.Items) != 1 || data.Items[0].ID != "bom-1" || data.Items[0].LoginPassword == "" {
|
||||
t.Fatalf("items=%+v", data.Items)
|
||||
}
|
||||
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM endpoints WHERE id='bom-1'`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("want 1 row got %d", n)
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodGet, base+"/api/admin/endpoints?source=admin", "", nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("list: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func csrfHeaders() map[string]string {
|
||||
return map[string]string{"X-Nixmsg-Request": "1"}
|
||||
}
|
||||
Reference in New Issue
Block a user