fix: 批量导入并发哈希并校正 CSV 行号
This commit is contained in:
@@ -908,6 +908,15 @@
|
|||||||
- 备选方案:同一事务里写 enabled 与作废;Identity 接口当前不暴露事务。
|
- 备选方案:同一事务里写 enabled 与作废;Identity 接口当前不暴露事务。
|
||||||
- 影响:Identity 失败时端仍保持原启用状态。
|
- 影响:Identity 失败时端仍保持原启用状态。
|
||||||
|
|
||||||
|
### 复审修复 H-04
|
||||||
|
|
||||||
|
- 日期:2026-09-30
|
||||||
|
- 原条款:DEVELOPMENT §8 批量开通哈希走并发池;PRD F01 指出行号;审查 #47。
|
||||||
|
- 实际做法:导入按 `max(1, NumCPU-1)` 并发算哈希并按行下标收集;UNIQUE 冲突返回 409 附行号;哈希失败 500/503。CSV 逐条 `Read()`,用 `FieldPos` / `csv.ParseError` 取物理行号。界面进度与错误表在 W-05。
|
||||||
|
- 原因:串行 argon2 过慢,且 `i+2` 行号会被空行和跨行字段带偏。
|
||||||
|
- 备选方案:占满哈希池;未采用,给 MQTT 登录留槽位。
|
||||||
|
- 影响:导入吞吐上升;冲突不再表现为无行号的 500。
|
||||||
|
|
||||||
## 后台网页 W
|
## 后台网页 W
|
||||||
|
|
||||||
1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`**
|
1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`**
|
||||||
|
|||||||
@@ -649,3 +649,14 @@ func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) {
|
|||||||
Data: map[string]any{"errors": errs},
|
Data: map[string]any{"errors": errs},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func writeCSVConflictError(w http.ResponseWriter, errs []csvLineError) {
|
||||||
|
httpx.WriteJSON(w, http.StatusConflict, httpx.Envelope{
|
||||||
|
OK: false,
|
||||||
|
Error: &httpx.ErrorBody{
|
||||||
|
Code: "id_taken",
|
||||||
|
Message: "编号已占用",
|
||||||
|
},
|
||||||
|
Data: map[string]any{"errors": errs},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
+187
-77
@@ -5,12 +5,15 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"encoding/csv"
|
"encoding/csv"
|
||||||
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"mime"
|
"mime"
|
||||||
"mime/multipart"
|
"mime/multipart"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"runtime"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
@@ -47,7 +50,16 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
prepared, errs := h.validateImportCSV(r.Context(), raw)
|
prepared, errs, fatal := h.validateImportCSV(r.Context(), raw)
|
||||||
|
if fatal != nil {
|
||||||
|
h.audit(actorString(p), "endpoint_import", "", "error", ip)
|
||||||
|
if errors.Is(fatal, context.Canceled) || errors.Is(fatal, context.DeadlineExceeded) {
|
||||||
|
httpx.WriteError(w, http.StatusServiceUnavailable, "busy", "哈希繁忙,请稍后再试")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||||
|
return
|
||||||
|
}
|
||||||
if len(errs) > 0 {
|
if len(errs) > 0 {
|
||||||
h.audit(actorString(p), "endpoint_import", "", "bad_request", ip)
|
h.audit(actorString(p), "endpoint_import", "", "bad_request", ip)
|
||||||
writeCSVValidationError(w, errs)
|
writeCSVValidationError(w, errs)
|
||||||
@@ -65,6 +77,11 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
if err := h.insertEndpointsBatch(r.Context(), rows); err != nil {
|
if err := h.insertEndpointsBatch(r.Context(), rows); err != nil {
|
||||||
|
if isUniqueConstraint(err) {
|
||||||
|
h.audit(actorString(p), "endpoint_import", "", "id_taken", ip)
|
||||||
|
writeCSVConflictError(w, h.importUniqueLineErrors(r.Context(), prepared))
|
||||||
|
return
|
||||||
|
}
|
||||||
h.audit(actorString(p), "endpoint_import", "", "error", ip)
|
h.audit(actorString(p), "endpoint_import", "", "error", ip)
|
||||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||||
return
|
return
|
||||||
@@ -128,59 +145,86 @@ func (e badRequestError) Error() string { return string(e) }
|
|||||||
|
|
||||||
func errBadRequest(msg string) error { return badRequestError(msg) }
|
func errBadRequest(msg string) error { return badRequestError(msg) }
|
||||||
|
|
||||||
func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPrepared, []csvLineError) {
|
type importPending struct {
|
||||||
|
line int
|
||||||
|
id string
|
||||||
|
name string
|
||||||
|
remark string
|
||||||
|
loginPW string
|
||||||
|
talkPW string
|
||||||
|
delaySec int64
|
||||||
|
needGenerateID bool
|
||||||
|
needGenerateLogin bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func csvErrorLine(err error, fallback int) int {
|
||||||
|
var pe *csv.ParseError
|
||||||
|
if errors.As(err, &pe) {
|
||||||
|
if pe.StartLine > 0 {
|
||||||
|
return pe.StartLine
|
||||||
|
}
|
||||||
|
if pe.Line > 0 {
|
||||||
|
return pe.Line
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if fallback > 0 {
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
|
func importHashConcurrency() int {
|
||||||
|
n := runtime.NumCPU() - 1
|
||||||
|
if n < 1 {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPrepared, []csvLineError, error) {
|
||||||
raw = bytes.TrimPrefix(raw, []byte{0xEF, 0xBB, 0xBF})
|
raw = bytes.TrimPrefix(raw, []byte{0xEF, 0xBB, 0xBF})
|
||||||
reader := csv.NewReader(bytes.NewReader(raw))
|
reader := csv.NewReader(bytes.NewReader(raw))
|
||||||
reader.FieldsPerRecord = -1
|
reader.FieldsPerRecord = -1
|
||||||
reader.TrimLeadingSpace = true
|
reader.TrimLeadingSpace = true
|
||||||
|
reader.ReuseRecord = false
|
||||||
|
|
||||||
records, err := reader.ReadAll()
|
header, err := reader.Read()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, []csvLineError{{Line: 1, Reason: "CSV 解析失败"}}
|
return nil, []csvLineError{{Line: csvErrorLine(err, 1), Reason: "CSV 解析失败"}}, nil
|
||||||
}
|
}
|
||||||
if len(records) < 1 {
|
headerLine, _ := reader.FieldPos(0)
|
||||||
return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}}
|
if headerLine < 1 {
|
||||||
|
headerLine = 1
|
||||||
}
|
}
|
||||||
|
norm := normalizeCSVHeader(header)
|
||||||
header := normalizeCSVHeader(records[0])
|
|
||||||
expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"}
|
expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"}
|
||||||
if len(header) < len(expected) {
|
if len(norm) < len(expected) {
|
||||||
return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}}
|
return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil
|
||||||
}
|
}
|
||||||
for i, want := range expected {
|
for i, want := range expected {
|
||||||
if header[i] != want {
|
if norm[i] != want {
|
||||||
return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}}
|
return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
dataRows := records[1:]
|
|
||||||
if len(dataRows) == 0 {
|
|
||||||
return nil, []csvLineError{{Line: 2, Reason: "没有数据行"}}
|
|
||||||
}
|
|
||||||
if len(dataRows) > maxImportRows {
|
|
||||||
return nil, []csvLineError{{Line: 1, Reason: "最多 1000 行"}}
|
|
||||||
}
|
|
||||||
|
|
||||||
errs := make([]csvLineError, 0)
|
errs := make([]csvLineError, 0)
|
||||||
prepared := make([]importPrepared, 0, len(dataRows))
|
seen := make(map[string]int)
|
||||||
seen := make(map[string]int) // id -> first line
|
checkIDs := make([]string, 0)
|
||||||
checkIDs := make([]string, 0, len(dataRows))
|
pendings := make([]importPending, 0)
|
||||||
|
|
||||||
type pending struct {
|
for {
|
||||||
line int
|
cols, readErr := reader.Read()
|
||||||
id string
|
if errors.Is(readErr, io.EOF) {
|
||||||
name string
|
break
|
||||||
remark string
|
}
|
||||||
loginPW string
|
if readErr != nil {
|
||||||
talkPW string
|
errs = append(errs, csvLineError{Line: csvErrorLine(readErr, 1), Reason: "CSV 解析失败"})
|
||||||
delaySec int64
|
continue
|
||||||
needGenerateID bool
|
}
|
||||||
needGenerateLogin bool
|
line, _ := reader.FieldPos(0)
|
||||||
}
|
if line < 1 {
|
||||||
pendings := make([]pending, 0, len(dataRows))
|
line = headerLine + 1
|
||||||
|
}
|
||||||
for i, cols := range dataRows {
|
|
||||||
line := i + 2 // 表头为 1
|
|
||||||
for len(cols) < 6 {
|
for len(cols) < 6 {
|
||||||
cols = append(cols, "")
|
cols = append(cols, "")
|
||||||
}
|
}
|
||||||
@@ -209,7 +253,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
p := pending{
|
p := importPending{
|
||||||
line: line,
|
line: line,
|
||||||
id: id,
|
id: id,
|
||||||
name: name,
|
name: name,
|
||||||
@@ -232,12 +276,18 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
|
|||||||
}
|
}
|
||||||
|
|
||||||
if len(errs) > 0 {
|
if len(errs) > 0 {
|
||||||
return nil, errs
|
return nil, errs, nil
|
||||||
|
}
|
||||||
|
if len(pendings) == 0 {
|
||||||
|
return nil, []csvLineError{{Line: headerLine + 1, Reason: "没有数据行"}}, nil
|
||||||
|
}
|
||||||
|
if len(pendings) > maxImportRows {
|
||||||
|
return nil, []csvLineError{{Line: headerLine, Reason: "最多 1000 行"}}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
existing, err := h.existingEndpointIDs(ctx, checkIDs)
|
existing, err := h.existingEndpointIDs(ctx, checkIDs)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, []csvLineError{{Line: 1, Reason: "校验失败"}}
|
return nil, nil, err
|
||||||
}
|
}
|
||||||
for _, p := range pendings {
|
for _, p := range pendings {
|
||||||
if p.needGenerateID {
|
if p.needGenerateID {
|
||||||
@@ -248,16 +298,19 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if len(errs) > 0 {
|
if len(errs) > 0 {
|
||||||
return nil, errs
|
return nil, errs, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, p := range pendings {
|
hashed := make([]importPending, len(pendings))
|
||||||
useID := p.id
|
copy(hashed, pendings)
|
||||||
|
for i := range hashed {
|
||||||
|
p := &hashed[i]
|
||||||
if p.needGenerateID {
|
if p.needGenerateID {
|
||||||
|
useID := ""
|
||||||
for attempt := 0; attempt < 16; attempt++ {
|
for attempt := 0; attempt < 16; attempt++ {
|
||||||
genID, genErr := generateEndpointID()
|
genID, genErr := generateEndpointID()
|
||||||
if genErr != nil {
|
if genErr != nil {
|
||||||
return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}}
|
return nil, nil, genErr
|
||||||
}
|
}
|
||||||
if _, clash := seen[genID]; clash {
|
if _, clash := seen[genID]; clash {
|
||||||
continue
|
continue
|
||||||
@@ -270,46 +323,103 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
if useID == "" {
|
if useID == "" {
|
||||||
return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}}
|
return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}}, nil
|
||||||
}
|
}
|
||||||
|
p.id = useID
|
||||||
}
|
}
|
||||||
|
|
||||||
loginPW := p.loginPW
|
|
||||||
if p.needGenerateLogin {
|
if p.needGenerateLogin {
|
||||||
pw, genErr := generateLoginPassword()
|
pw, genErr := generateLoginPassword()
|
||||||
if genErr != nil {
|
if genErr != nil {
|
||||||
return nil, []csvLineError{{Line: p.line, Reason: "生成密码失败"}}
|
return nil, nil, genErr
|
||||||
}
|
}
|
||||||
loginPW = pw
|
p.loginPW = pw
|
||||||
}
|
}
|
||||||
loginHash, hashErr := h.hash.Hash(ctx, auth.PasswordLogin, loginPW)
|
|
||||||
if hashErr != nil {
|
|
||||||
return nil, []csvLineError{{Line: p.line, Reason: "哈希失败"}}
|
|
||||||
}
|
|
||||||
var talkHash sql.NullString
|
|
||||||
if p.talkPW != "" {
|
|
||||||
th, thErr := h.hash.Hash(ctx, auth.PasswordTalk, p.talkPW)
|
|
||||||
if thErr != nil {
|
|
||||||
return nil, []csvLineError{{Line: p.line, Reason: "哈希失败"}}
|
|
||||||
}
|
|
||||||
talkHash = sql.NullString{String: th, Valid: true}
|
|
||||||
}
|
|
||||||
prepared = append(prepared, importPrepared{
|
|
||||||
Insert: endpointInsert{
|
|
||||||
ID: useID,
|
|
||||||
Name: p.name,
|
|
||||||
Remark: p.remark,
|
|
||||||
Source: sourceAdmin,
|
|
||||||
LoginHash: loginHash,
|
|
||||||
TalkHash: talkHash,
|
|
||||||
DefaultDelayMs: p.delaySec * 1000,
|
|
||||||
},
|
|
||||||
PlainLogin: loginPW,
|
|
||||||
Name: p.name,
|
|
||||||
Line: p.line,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
return prepared, nil
|
|
||||||
|
prepared, hashErr := h.hashImportPendings(ctx, hashed)
|
||||||
|
if hashErr != nil {
|
||||||
|
return nil, nil, hashErr
|
||||||
|
}
|
||||||
|
return prepared, nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) hashImportPendings(ctx context.Context, pendings []importPending) ([]importPrepared, error) {
|
||||||
|
out := make([]importPrepared, len(pendings))
|
||||||
|
errCh := make([]error, len(pendings))
|
||||||
|
workers := importHashConcurrency()
|
||||||
|
sem := make(chan struct{}, workers)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := range pendings {
|
||||||
|
i := i
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
select {
|
||||||
|
case sem <- struct{}{}:
|
||||||
|
defer func() { <-sem }()
|
||||||
|
case <-ctx.Done():
|
||||||
|
errCh[i] = ctx.Err()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
p := pendings[i]
|
||||||
|
loginHash, hashErr := h.hash.Hash(ctx, auth.PasswordLogin, p.loginPW)
|
||||||
|
if hashErr != nil {
|
||||||
|
errCh[i] = hashErr
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var talkHash sql.NullString
|
||||||
|
if p.talkPW != "" {
|
||||||
|
th, thErr := h.hash.Hash(ctx, auth.PasswordTalk, p.talkPW)
|
||||||
|
if thErr != nil {
|
||||||
|
errCh[i] = thErr
|
||||||
|
return
|
||||||
|
}
|
||||||
|
talkHash = sql.NullString{String: th, Valid: true}
|
||||||
|
}
|
||||||
|
out[i] = importPrepared{
|
||||||
|
Insert: endpointInsert{
|
||||||
|
ID: p.id,
|
||||||
|
Name: p.name,
|
||||||
|
Remark: p.remark,
|
||||||
|
Source: sourceAdmin,
|
||||||
|
LoginHash: loginHash,
|
||||||
|
TalkHash: talkHash,
|
||||||
|
DefaultDelayMs: p.delaySec * 1000,
|
||||||
|
},
|
||||||
|
PlainLogin: p.loginPW,
|
||||||
|
Name: p.name,
|
||||||
|
Line: p.line,
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
for _, e := range errCh {
|
||||||
|
if e != nil {
|
||||||
|
return nil, e
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) importUniqueLineErrors(ctx context.Context, prepared []importPrepared) []csvLineError {
|
||||||
|
ids := make([]string, 0, len(prepared))
|
||||||
|
for _, p := range prepared {
|
||||||
|
ids = append(ids, p.Insert.ID)
|
||||||
|
}
|
||||||
|
existing, err := h.existingEndpointIDs(ctx, ids)
|
||||||
|
if err != nil {
|
||||||
|
return []csvLineError{{Line: 1, Reason: "编号已占用"}}
|
||||||
|
}
|
||||||
|
errs := make([]csvLineError, 0)
|
||||||
|
for _, p := range prepared {
|
||||||
|
if _, ok := existing[p.Insert.ID]; ok {
|
||||||
|
errs = append(errs, csvLineError{Line: p.Line, Reason: "编号已占用"})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(errs) == 0 {
|
||||||
|
return []csvLineError{{Line: 1, Reason: "编号已占用"}}
|
||||||
|
}
|
||||||
|
return errs
|
||||||
}
|
}
|
||||||
|
|
||||||
func normalizeCSVHeader(cols []string) []string {
|
func normalizeCSVHeader(cols []string) []string {
|
||||||
|
|||||||
@@ -0,0 +1,276 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user