277 lines
7.3 KiB
Go
277 lines
7.3 KiB
Go
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)
|
|
}
|
|
}
|