package admin import ( "bytes" "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" "git.asio.asia/nixevol/NixMsg/internal/httpx" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) type csvLineError struct { Line int `json:"line"` Reason string `json:"reason"` } type importPrepared struct { Insert endpointInsert // PlainLogin 始终填入响应(生成或原文,仅此一次)。 PlainLogin string Name string Line int } func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) { p, _ := principalFrom(r.Context()) ip := httpx.ClientIP(r, h.trusted) raw, err := readImportCSV(r) if err != nil { if httpx.IsBodyTooLarge(err) { h.auditP(p, "endpoint_import", "", "payload_too_large", ip) httpx.WriteError(w, http.StatusRequestEntityTooLarge, "payload_too_large", "请求体过大") return } h.auditP(p, "endpoint_import", "", "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", err.Error()) return } prepared, errs, fatal := h.validateImportCSV(r.Context(), raw) if fatal != nil { h.auditP(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.auditP(p, "endpoint_import", "", "bad_request", ip) writeCSVValidationError(w, errs) return } rows := make([]endpointInsert, len(prepared)) items := make([]map[string]any, 0, len(prepared)) for i, pRow := range prepared { rows[i] = pRow.Insert items = append(items, map[string]any{ "id": pRow.Insert.ID, "login_password": pRow.PlainLogin, "name": pRow.Name, }) } if err := h.insertEndpointsBatch(r.Context(), rows); err != nil { if isUniqueConstraint(err) { h.auditP(p, "endpoint_import", "", "id_taken", ip) writeCSVConflictError(w, h.importUniqueLineErrors(r.Context(), prepared)) return } h.auditP(p, "endpoint_import", "", "error", ip) httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return } ids := make([]string, 0, len(prepared)) for _, pRow := range prepared { ids = append(ids, pRow.Insert.ID) } h.auditPD(p, "endpoint_import", strconv.Itoa(len(items)), "ok", ip, importAuditDetail(ids)) httpx.WriteOK(w, map[string]any{"items": items}) } func readImportCSV(r *http.Request) ([]byte, error) { ct := r.Header.Get("Content-Type") mediaType, params, err := mime.ParseMediaType(ct) if err != nil { mediaType = strings.TrimSpace(strings.Split(ct, ";")[0]) } switch { case strings.HasPrefix(mediaType, "multipart/"): boundary := params["boundary"] if boundary == "" { return nil, errBadRequest("multipart 缺少 boundary") } mr := multipart.NewReader(r.Body, boundary) for { part, pErr := mr.NextPart() if pErr == io.EOF { break } if pErr != nil { return nil, errBadRequest("读取 multipart 失败") } name := part.FormName() if name == "file" || name == "" { b, readErr := io.ReadAll(part) _ = part.Close() if readErr != nil { if httpx.IsBodyTooLarge(readErr) { return nil, readErr } return nil, errBadRequest("读取文件失败") } return b, nil } _ = part.Close() } return nil, errBadRequest("缺少 file 字段") default: // text/csv 或未标明时按原始体;大小由 ServeHTTP 的 MaxBytesReader 限制。 b, readErr := io.ReadAll(r.Body) if readErr != nil { if httpx.IsBodyTooLarge(readErr) { return nil, readErr } return nil, errBadRequest("读取 CSV 失败") } return b, nil } } type badRequestError string func (e badRequestError) Error() string { return string(e) } func errBadRequest(msg string) error { return badRequestError(msg) } 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 header, err := reader.Read() if err != nil { return nil, []csvLineError{{Line: csvErrorLine(err, 1), Reason: "CSV 解析失败"}}, nil } headerLine, _ := reader.FieldPos(0) if headerLine < 1 { headerLine = 1 } norm := normalizeCSVHeader(header) expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"} if len(norm) < len(expected) { return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil } for i, want := range expected { if norm[i] != want { return nil, []csvLineError{{Line: headerLine, Reason: "表头不正确"}}, nil } } errs := make([]csvLineError, 0) seen := make(map[string]int) checkIDs := make([]string, 0) pendings := make([]importPending, 0) 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, "") } id := strings.TrimSpace(cols[0]) name := cols[1] loginPW := cols[2] talkPW := cols[3] delayRaw := strings.TrimSpace(cols[4]) remark := cols[5] delaySec := int64(0) if delayRaw != "" { n, parseErr := strconv.ParseInt(delayRaw, 10, 64) if parseErr != nil || n < 0 { errs = append(errs, csvLineError{Line: line, Reason: "默认延迟无效"}) continue } delaySec = n } if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec, h.maxScheduleSeconds()); msg != "" { errs = append(errs, csvLineError{Line: line, Reason: msg}) continue } if utf8.RuneCountInString(name) > protocol.MaxNameChars { errs = append(errs, csvLineError{Line: line, Reason: "名称不合法"}) continue } p := importPending{ line: line, id: id, name: name, remark: remark, loginPW: loginPW, talkPW: talkPW, delaySec: delaySec, needGenerateID: id == "", needGenerateLogin: loginPW == "", } if !p.needGenerateID { if first, ok := seen[id]; ok { errs = append(errs, csvLineError{Line: line, Reason: "编号与第 " + strconv.Itoa(first) + " 行重复"}) continue } seen[id] = line checkIDs = append(checkIDs, id) } pendings = append(pendings, p) } if len(errs) > 0 { 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, nil, err } for _, p := range pendings { if p.needGenerateID { continue } if _, ok := existing[p.id]; ok { errs = append(errs, csvLineError{Line: p.line, Reason: "编号已占用"}) } } if len(errs) > 0 { return nil, errs, nil } 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, nil, genErr } if _, clash := seen[genID]; clash { continue } if _, clash := existing[genID]; clash { continue } useID = genID seen[useID] = p.line break } if useID == "" { return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}}, nil } p.id = useID } if p.needGenerateLogin { pw, genErr := generateLoginPassword() if genErr != nil { return nil, nil, genErr } p.loginPW = pw } } 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 { out := make([]string, len(cols)) for i, c := range cols { out[i] = strings.TrimSpace(strings.TrimPrefix(c, "\ufeff")) } return out }