Files
NixMsg/internal/admin/endpoints_test.go
T

365 lines
10 KiB
Go

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