diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 1a77095..d0f634c 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -908,6 +908,15 @@ - 备选方案:同一事务里写 enabled 与作废;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 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** diff --git a/internal/admin/endpoints.go b/internal/admin/endpoints.go index 93c433e..ba21f22 100644 --- a/internal/admin/endpoints.go +++ b/internal/admin/endpoints.go @@ -649,3 +649,14 @@ func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) { 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}, + }) +} diff --git a/internal/admin/endpoints_csv.go b/internal/admin/endpoints_csv.go index eff12ec..ee56630 100644 --- a/internal/admin/endpoints_csv.go +++ b/internal/admin/endpoints_csv.go @@ -5,12 +5,15 @@ import ( "context" "database/sql" "encoding/csv" + "errors" "io" "mime" "mime/multipart" "net/http" + "runtime" "strconv" "strings" + "sync" "unicode/utf8" "git.asio.asia/nixevol/NixMsg/internal/auth" @@ -47,7 +50,16 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) { 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 { h.audit(actorString(p), "endpoint_import", "", "bad_request", ip) 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 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) httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return @@ -128,59 +145,86 @@ func (e badRequestError) Error() string { return string(e) } 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}) reader := csv.NewReader(bytes.NewReader(raw)) reader.FieldsPerRecord = -1 reader.TrimLeadingSpace = true + reader.ReuseRecord = false - records, err := reader.ReadAll() + header, err := reader.Read() if err != nil { - return nil, []csvLineError{{Line: 1, Reason: "CSV 解析失败"}} + return nil, []csvLineError{{Line: csvErrorLine(err, 1), Reason: "CSV 解析失败"}}, nil } - if len(records) < 1 { - return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + headerLine, _ := reader.FieldPos(0) + if headerLine < 1 { + headerLine = 1 } - - header := normalizeCSVHeader(records[0]) + norm := normalizeCSVHeader(header) expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"} - if len(header) < len(expected) { - return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + if len(norm) < len(expected) { + return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil } for i, want := range expected { - if header[i] != want { - return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} + if norm[i] != want { + 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) - prepared := make([]importPrepared, 0, len(dataRows)) - seen := make(map[string]int) // id -> first line - checkIDs := make([]string, 0, len(dataRows)) + seen := make(map[string]int) + checkIDs := make([]string, 0) + pendings := make([]importPending, 0) - type pending struct { - line int - id string - name string - remark string - loginPW string - talkPW string - delaySec int64 - needGenerateID bool - needGenerateLogin bool - } - pendings := make([]pending, 0, len(dataRows)) - - for i, cols := range dataRows { - line := i + 2 // 表头为 1 + for { + cols, readErr := reader.Read() + if errors.Is(readErr, io.EOF) { + break + } + if readErr != nil { + errs = append(errs, csvLineError{Line: csvErrorLine(readErr, 1), Reason: "CSV 解析失败"}) + continue + } + line, _ := reader.FieldPos(0) + if line < 1 { + line = headerLine + 1 + } for len(cols) < 6 { cols = append(cols, "") } @@ -209,7 +253,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr continue } - p := pending{ + p := importPending{ line: line, id: id, name: name, @@ -232,12 +276,18 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr } 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) if err != nil { - return nil, []csvLineError{{Line: 1, Reason: "校验失败"}} + return nil, nil, err } for _, p := range pendings { if p.needGenerateID { @@ -248,16 +298,19 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr } } if len(errs) > 0 { - return nil, errs + return nil, errs, nil } - for _, p := range pendings { - useID := p.id + hashed := make([]importPending, len(pendings)) + copy(hashed, pendings) + for i := range hashed { + p := &hashed[i] if p.needGenerateID { + useID := "" for attempt := 0; attempt < 16; attempt++ { genID, genErr := generateEndpointID() if genErr != nil { - return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}} + return nil, nil, genErr } if _, clash := seen[genID]; clash { continue @@ -270,46 +323,103 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr break } 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 { pw, genErr := generateLoginPassword() 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 { diff --git a/internal/admin/h04_test.go b/internal/admin/h04_test.go new file mode 100644 index 0000000..a72e853 --- /dev/null +++ b/internal/admin/h04_test.go @@ -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) + } +}