fix: 批量导入并发哈希并校正 CSV 行号
This commit is contained in:
+187
-77
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user