feat: 实现管理后台端管理接口(A2)
This commit is contained in:
@@ -0,0 +1,310 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user