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