Files
NixMsg/internal/admin/endpoints_csv.go
T

436 lines
11 KiB
Go

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
}