package admin import ( "bytes" "context" "database/sql" "encoding/csv" "io" "mime" "mime/multipart" "net/http" "strconv" "strings" "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 { h.audit(actorString(p), "endpoint_import", "", "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", err.Error()) return } prepared, errs := h.validateImportCSV(r.Context(), raw) if len(errs) > 0 { h.audit(actorString(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 { h.audit(actorString(p), "endpoint_import", "", "error", ip) httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return } h.audit(actorString(p), "endpoint_import", strconv.Itoa(len(items)), "ok", ip) 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(io.LimitReader(part, 8<<20)) _ = part.Close() if readErr != nil { return nil, errBadRequest("读取文件失败") } return b, nil } _ = part.Close() } return nil, errBadRequest("缺少 file 字段") default: // text/csv 或未标明时按原始体 b, readErr := io.ReadAll(io.LimitReader(r.Body, 8<<20)) if readErr != nil { 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) } func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPrepared, []csvLineError) { raw = bytes.TrimPrefix(raw, []byte{0xEF, 0xBB, 0xBF}) reader := csv.NewReader(bytes.NewReader(raw)) reader.FieldsPerRecord = -1 reader.TrimLeadingSpace = true records, err := reader.ReadAll() if err != nil { return nil, []csvLineError{{Line: 1, Reason: "CSV 解析失败"}} } if len(records) < 1 { return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} } header := normalizeCSVHeader(records[0]) expected := []string{"id", "name", "login_password", "talk_password", "default_delay_seconds", "remark"} if len(header) < len(expected) { return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} } for i, want := range expected { if header[i] != want { return nil, []csvLineError{{Line: 1, Reason: "表头不正确"}} } } 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)) 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 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); 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 := pending{ 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 } existing, err := h.existingEndpointIDs(ctx, checkIDs) if err != nil { return nil, []csvLineError{{Line: 1, Reason: "校验失败"}} } 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 } for _, p := range pendings { useID := p.id if p.needGenerateID { for attempt := 0; attempt < 16; attempt++ { genID, genErr := generateEndpointID() if genErr != nil { return nil, []csvLineError{{Line: p.line, Reason: "生成编号失败"}} } 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: "生成编号失败"}} } } loginPW := p.loginPW if p.needGenerateLogin { pw, genErr := generateLoginPassword() if genErr != nil { return nil, []csvLineError{{Line: p.line, Reason: "生成密码失败"}} } 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 } 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 }