package admin_test import ( "context" "database/sql" "encoding/json" "fmt" "io" "net/http" "net/http/cookiejar" "net/http/httptest" "path/filepath" "runtime" "strings" "sync" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/admin" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/store" ) type concHashPool struct { inner *auth.StubHashPool mu sync.Mutex active, max int } func (p *concHashPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) { p.mu.Lock() p.active++ if p.active > p.max { p.max = p.active } p.mu.Unlock() time.Sleep(40 * time.Millisecond) defer func() { p.mu.Lock() p.active-- p.mu.Unlock() }() return p.inner.Hash(ctx, kind, password) } func (p *concHashPool) Verify(ctx context.Context, kind auth.PasswordKind, password, phc string) (bool, error) { return p.inner.Verify(ctx, kind, password, phc) } func (p *concHashPool) QueueLen() int { return 0 } func (p *concHashPool) maxActive() int { p.mu.Lock() defer p.mu.Unlock() return p.max } type gateHashPool struct { inner *auth.StubHashPool started chan struct{} release chan struct{} startOnce sync.Once } func (p *gateHashPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) { p.startOnce.Do(func() { close(p.started) }) select { case <-p.release: case <-ctx.Done(): return "", ctx.Err() } return p.inner.Hash(ctx, kind, password) } func (p *gateHashPool) Verify(ctx context.Context, kind auth.PasswordKind, password, phc string) (bool, error) { return p.inner.Verify(ctx, kind, password, phc) } func (p *gateHashPool) QueueLen() int { return 0 } func setupEndpointsHash(t *testing.T, hash auth.HashPool) (*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 闸门卡住 seed。 if seedErr := admin.SeedAdminPassword(context.Background(), db, auth.NewStubHashPool(), testPassword); seedErr != nil { t.Fatal(seedErr) } h := admin.New(admin.Deps{ DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(), Locks: admin.NewMemoryLoginLocks(), }) 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 } func postImportCSV(t *testing.T, client *http.Client, base, csvBody string) *http.Response { t.Helper() 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) } return res } func decodeCSVErrors(t *testing.T, res *http.Response) (status int, code string, lines []int) { t.Helper() raw, _ := io.ReadAll(res.Body) _ = res.Body.Close() var env struct { OK bool `json:"ok"` Error *struct { Code string `json:"code"` } `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.Fatalf("json: %v body=%s", err, raw) } code = "" if env.Error != nil { code = env.Error.Code } if env.Data != nil { for _, e := range env.Data.Errors { lines = append(lines, e.Line) } } return res.StatusCode, code, lines } func TestImportHashConcurrencyGreaterThanOne(t *testing.T) { if runtime.NumCPU() < 2 { t.Skip("need at least 2 CPUs to observe concurrent hashing") } pool := &concHashPool{inner: auth.NewStubHashPool()} _, srv, client := setupEndpointsHash(t, pool) var b strings.Builder b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n") for i := 0; i < 8; i++ { fmt.Fprintf(&b, "conc-%d,名,password1,,0,\n", i) } res := postImportCSV(t, client, srv.URL, b.String()) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("import: %d %+v", res.StatusCode, env) } if pool.maxActive() < 2 { t.Fatalf("want concurrent hash > 1, got %d", pool.maxActive()) } } func TestImportThousandRows(t *testing.T) { _, srv, client, _, _ := setupEndpoints(t) var b strings.Builder b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n") for i := 0; i < 1000; i++ { fmt.Fprintf(&b, "r%04d,名,password1,,0,\n", i) } res := postImportCSV(t, client, srv.URL, b.String()) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("import 1000: %d %+v", res.StatusCode, env) } } func TestImportUniqueRaceReturns409WithLine(t *testing.T) { stub := auth.NewStubHashPool() gate := &gateHashPool{inner: stub, started: make(chan struct{}), release: make(chan struct{})} db, srv, client := setupEndpointsHash(t, gate) csvBody := "" + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + "race-1,甲,password1,,0,\n" done := make(chan *http.Response, 1) go func() { done <- postImportCSV(t, client, srv.URL, csvBody) }() select { case <-gate.started: case <-time.After(5 * time.Second): close(gate.release) t.Fatal("hash did not start") } 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) VALUES ('race-1', '占', '', 'admin', 'stub$x', NULL, 0, 0, 1, 1)`) return e }) if err != nil { close(gate.release) t.Fatal(err) } close(gate.release) res := <-done status, code, lines := decodeCSVErrors(t, res) if status != http.StatusConflict || code != "id_taken" { t.Fatalf("want 409 id_taken got %d %s lines=%v", status, code, lines) } found := false for _, ln := range lines { if ln == 2 { found = true } } if !found { t.Fatalf("want line 2 in conflict errors, got %v", lines) } } func TestImportErrorLineSkipsEmptyRows(t *testing.T) { _, srv, client, _, _ := setupEndpoints(t) csvBody := "" + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + "\n" + "bad-1,甲,password1,,-1,\n" res := postImportCSV(t, client, srv.URL, csvBody) status, _, lines := decodeCSVErrors(t, res) if status != http.StatusBadRequest { t.Fatalf("want 400 got %d lines=%v", status, lines) } found := false for _, ln := range lines { if ln == 3 { found = true } } if !found { t.Fatalf("want physical line 3, got %v", lines) } } func TestImportUnclosedQuoteLine(t *testing.T) { _, srv, client, _, _ := setupEndpoints(t) csvBody := "" + "id,name,login_password,talk_password,default_delay_seconds,remark\n" + "\"not-closed\n" res := postImportCSV(t, client, srv.URL, csvBody) status, _, lines := decodeCSVErrors(t, res) if status != http.StatusBadRequest { t.Fatalf("want 400 got %d lines=%v", status, lines) } if len(lines) == 0 || lines[0] < 2 { t.Fatalf("want parse error line >= 2, got %v", lines) } }