merge: admin-web
This commit is contained in:
@@ -272,6 +272,33 @@ func TestGroupsCRUD(t *testing.T) {
|
||||
t.Fatalf("created=%+v", created)
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodGet, base+"/api/admin/groups/"+created.ID+"?limit=1", "", nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("get page1: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var page1 struct {
|
||||
MemberTotal float64 `json:"member_total"`
|
||||
Members []any `json:"members"`
|
||||
NextCursor string `json:"next_cursor"`
|
||||
}
|
||||
_ = json.Unmarshal(env.Data, &page1)
|
||||
if page1.MemberTotal != 3 || len(page1.Members) != 1 || page1.NextCursor == "" {
|
||||
t.Fatalf("page1=%+v raw=%s", page1, env.Data)
|
||||
}
|
||||
res = doReq(t, client, http.MethodGet, base+"/api/admin/groups/"+created.ID+"?limit=1&cursor="+page1.NextCursor, "", nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("get page2: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var page2 struct {
|
||||
Members []any `json:"members"`
|
||||
}
|
||||
_ = json.Unmarshal(env.Data, &page2)
|
||||
if len(page2.Members) != 1 {
|
||||
t.Fatalf("page2 members=%d", len(page2.Members))
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodPatch, base+"/api/admin/groups/"+created.ID,
|
||||
`{"name":"新名"}`, csrf())
|
||||
env = decodeEnv(t, res)
|
||||
|
||||
@@ -233,7 +233,7 @@ func (h *Handler) handleEndpointCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if req.DefaultDelaySeconds != nil {
|
||||
delaySec = *req.DefaultDelaySeconds
|
||||
}
|
||||
if errMsg := validateEndpointFields(req.ID, req.Name, req.Remark, req.LoginPassword, req.TalkPassword, delaySec); errMsg != "" {
|
||||
if errMsg := validateEndpointFields(req.ID, req.Name, req.Remark, req.LoginPassword, req.TalkPassword, delaySec, h.maxScheduleSeconds()); errMsg != "" {
|
||||
h.audit(actorString(p), "endpoint_create", req.ID, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", errMsg)
|
||||
return
|
||||
@@ -341,17 +341,24 @@ func (h *Handler) handleEndpointPatch(w http.ResponseWriter, r *http.Request) {
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "备注过长")
|
||||
return
|
||||
}
|
||||
if req.DefaultDelaySeconds != nil && *req.DefaultDelaySeconds < 0 {
|
||||
h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "默认延迟无效")
|
||||
return
|
||||
if req.DefaultDelaySeconds != nil {
|
||||
if msg := validateDelaySeconds(*req.DefaultDelaySeconds, h.maxScheduleSeconds()); msg != "" {
|
||||
h.audit(actorString(p), "endpoint_patch", id, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", msg)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
hasMeta := req.Name != nil || req.Remark != nil || req.DefaultDelaySeconds != nil
|
||||
var wasEnabled bool
|
||||
var err error
|
||||
// 注入 Identity 时启停只走 identity,避免先写 enabled 再级联失败造成半生效。
|
||||
patchEnabled := req.Enabled
|
||||
if h.identity != nil {
|
||||
patchEnabled = nil
|
||||
}
|
||||
if hasMeta || (req.Enabled != nil && h.identity == nil) {
|
||||
wasEnabled, err = h.patchEndpoint(r.Context(), id, req.Name, req.Remark, req.DefaultDelaySeconds, req.Enabled)
|
||||
wasEnabled, err = h.patchEndpoint(r.Context(), id, req.Name, req.Remark, req.DefaultDelaySeconds, patchEnabled)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
h.audit(actorString(p), "endpoint_patch", id, "not_found", ip)
|
||||
@@ -609,7 +616,7 @@ func (h *Handler) handleEndpointUnlock(w http.ResponseWriter, r *http.Request) {
|
||||
httpx.WriteOK(w, map[string]any{})
|
||||
}
|
||||
|
||||
func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec int64) string {
|
||||
func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec, maxDelaySec int64) string {
|
||||
if id != "" && !protocol.ValidEndpointID(id) {
|
||||
return "编号不合法"
|
||||
}
|
||||
@@ -628,12 +635,26 @@ func validateEndpointFields(id, name, remark, loginPW, talkPW string, delaySec i
|
||||
if !protocol.ValidTalkPassword(talkPW) {
|
||||
return "对话密码不合法"
|
||||
}
|
||||
if msg := validateDelaySeconds(delaySec, maxDelaySec); msg != "" {
|
||||
return msg
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func validateDelaySeconds(delaySec, maxDelaySec int64) string {
|
||||
if delaySec < 0 {
|
||||
return "默认延迟无效"
|
||||
}
|
||||
if maxDelaySec > 0 && delaySec > maxDelaySec {
|
||||
return "默认延迟无效"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *Handler) maxScheduleSeconds() int64 {
|
||||
return int64(h.cfg.Limits.MaxScheduleSeconds)
|
||||
}
|
||||
|
||||
func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) {
|
||||
httpx.WriteJSON(w, http.StatusBadRequest, httpx.Envelope{
|
||||
OK: false,
|
||||
@@ -644,3 +665,14 @@ func writeCSVValidationError(w http.ResponseWriter, errs []csvLineError) {
|
||||
Data: map[string]any{"errors": errs},
|
||||
})
|
||||
}
|
||||
|
||||
func writeCSVConflictError(w http.ResponseWriter, errs []csvLineError) {
|
||||
httpx.WriteJSON(w, http.StatusConflict, httpx.Envelope{
|
||||
OK: false,
|
||||
Error: &httpx.ErrorBody{
|
||||
Code: "id_taken",
|
||||
Message: "编号已占用",
|
||||
},
|
||||
Data: map[string]any{"errors": errs},
|
||||
})
|
||||
}
|
||||
|
||||
+202
-81
@@ -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"
|
||||
@@ -37,12 +40,26 @@ func (h *Handler) handleEndpointImport(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
raw, err := readImportCSV(r)
|
||||
if err != nil {
|
||||
if httpx.IsBodyTooLarge(err) {
|
||||
h.audit(actorString(p), "endpoint_import", "", "payload_too_large", ip)
|
||||
httpx.WriteError(w, http.StatusRequestEntityTooLarge, "payload_too_large", "请求体过大")
|
||||
return
|
||||
}
|
||||
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)
|
||||
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)
|
||||
@@ -60,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
|
||||
@@ -91,9 +113,12 @@ func readImportCSV(r *http.Request) ([]byte, error) {
|
||||
}
|
||||
name := part.FormName()
|
||||
if name == "file" || name == "" {
|
||||
b, readErr := io.ReadAll(io.LimitReader(part, 8<<20))
|
||||
b, readErr := io.ReadAll(part)
|
||||
_ = part.Close()
|
||||
if readErr != nil {
|
||||
if httpx.IsBodyTooLarge(readErr) {
|
||||
return nil, readErr
|
||||
}
|
||||
return nil, errBadRequest("读取文件失败")
|
||||
}
|
||||
return b, nil
|
||||
@@ -102,9 +127,12 @@ func readImportCSV(r *http.Request) ([]byte, error) {
|
||||
}
|
||||
return nil, errBadRequest("缺少 file 字段")
|
||||
default:
|
||||
// text/csv 或未标明时按原始体
|
||||
b, readErr := io.ReadAll(io.LimitReader(r.Body, 8<<20))
|
||||
// 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
|
||||
@@ -117,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, "")
|
||||
}
|
||||
@@ -189,7 +244,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
|
||||
}
|
||||
delaySec = n
|
||||
}
|
||||
if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec); msg != "" {
|
||||
if msg := validateEndpointFields(id, name, remark, loginPW, talkPW, delaySec, h.maxScheduleSeconds()); msg != "" {
|
||||
errs = append(errs, csvLineError{Line: line, Reason: msg})
|
||||
continue
|
||||
}
|
||||
@@ -198,7 +253,7 @@ func (h *Handler) validateImportCSV(ctx context.Context, raw []byte) ([]importPr
|
||||
continue
|
||||
}
|
||||
|
||||
p := pending{
|
||||
p := importPending{
|
||||
line: line,
|
||||
id: id,
|
||||
name: name,
|
||||
@@ -221,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 {
|
||||
@@ -237,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
|
||||
@@ -259,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 {
|
||||
|
||||
@@ -209,6 +209,7 @@ LIMIT ? OFFSET ?`, id, limit, offset)
|
||||
"name": name,
|
||||
"owner_id": owner,
|
||||
"created_at_ms": created,
|
||||
"member_total": memberTotal,
|
||||
"members": members,
|
||||
"next_cursor": next,
|
||||
})
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestLoginRejectsOversizeBody(t *testing.T) {
|
||||
_, srv, client, _ := setup(t)
|
||||
body := `{"username":"admin","password":"` + strings.Repeat("a", 1<<20) + `"}`
|
||||
req, err := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/login", strings.NewReader(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusRequestEntityTooLarge {
|
||||
t.Fatalf("want 413 got %d env=%+v", res.StatusCode, env)
|
||||
}
|
||||
if env.Error == nil || env.Error.Code != "payload_too_large" {
|
||||
t.Fatalf("want payload_too_large got %+v", env.Error)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportRejectsOversizeCSV(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
payload := bytes.Repeat([]byte("x"), 8<<20+1)
|
||||
req, err := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints/import", bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "text/csv")
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, _ := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
if res.StatusCode != http.StatusRequestEntityTooLarge {
|
||||
t.Fatalf("want 413 got %d body=%s", res.StatusCode, raw)
|
||||
}
|
||||
if !strings.Contains(string(raw), "payload_too_large") {
|
||||
t.Fatalf("want payload_too_large in body: %s", raw)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,146 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/admin"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
"net/http/cookiejar"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
type failDisableIdentity struct {
|
||||
identity.Stub
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *failDisableIdentity) Disable(context.Context, string) error { return f.err }
|
||||
func (f *failDisableIdentity) Enable(context.Context, string) error { return nil }
|
||||
|
||||
type trackIdentity struct {
|
||||
identity.Stub
|
||||
mu sync.Mutex
|
||||
disableN int
|
||||
enableN int
|
||||
}
|
||||
|
||||
func (t *trackIdentity) Disable(context.Context, string) error {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
t.disableN++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *trackIdentity) Enable(context.Context, string) error {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
t.enableN++
|
||||
return nil
|
||||
}
|
||||
|
||||
func setupEndpointsWithIdentity(t *testing.T, ident identity.Service) (*store.DB, *httptest.Server, *http.Client) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
|
||||
hash := auth.NewStubHashPool()
|
||||
if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil {
|
||||
t.Fatal(seedErr)
|
||||
}
|
||||
h := admin.New(admin.Deps{
|
||||
DB: db,
|
||||
Hash: hash,
|
||||
Tokens: admin.NewRandomAPITokens(),
|
||||
Locks: admin.NewMemoryLoginLocks(),
|
||||
Identity: ident,
|
||||
})
|
||||
srv := httptest.NewServer(h)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar}
|
||||
res := postJSON(t, client, srv.URL+"/api/admin/login",
|
||||
`{"username":"admin","password":"`+testPassword+`"}`, nil)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusOK || !env.OK {
|
||||
t.Fatalf("login failed: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
return db, srv, client
|
||||
}
|
||||
|
||||
func TestPatchEnabledFalseKeepsEnabledWhenIdentityFails(t *testing.T) {
|
||||
ident := &failDisableIdentity{err: errors.New("disable failed")}
|
||||
db, srv, client := setupEndpointsWithIdentity(t, ident)
|
||||
base := srv.URL
|
||||
|
||||
res := postJSON(t, client, base+"/api/admin/endpoints",
|
||||
`{"id":"keep-on","name":"仍启用","login_password":"password1"}`,
|
||||
csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodPatch, base+"/api/admin/endpoints/keep-on",
|
||||
`{"enabled":false}`, csrfHeaders())
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusInternalServerError {
|
||||
t.Fatalf("want 500 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var enabled int
|
||||
if err := db.Read.QueryRow(`SELECT enabled FROM endpoints WHERE id='keep-on'`).Scan(&enabled); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if enabled != 1 {
|
||||
t.Fatalf("want still enabled, got %d", enabled)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPatchNameOnlyDoesNotTouchEnabled(t *testing.T) {
|
||||
ident := &trackIdentity{}
|
||||
db, srv, client := setupEndpointsWithIdentity(t, ident)
|
||||
base := srv.URL
|
||||
|
||||
res := postJSON(t, client, base+"/api/admin/endpoints",
|
||||
`{"id":"name-only","name":"旧名","login_password":"password1"}`,
|
||||
csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodPatch, base+"/api/admin/endpoints/name-only",
|
||||
`{"name":"新名"}`, csrfHeaders())
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("patch: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
ident.mu.Lock()
|
||||
d, e := ident.disableN, ident.enableN
|
||||
ident.mu.Unlock()
|
||||
if d != 0 || e != 0 {
|
||||
t.Fatalf("identity enable/disable should not run, disable=%d enable=%d", d, e)
|
||||
}
|
||||
var enabled int
|
||||
var name string
|
||||
if err := db.Read.QueryRow(`SELECT enabled, name FROM endpoints WHERE id='name-only'`).Scan(&enabled, &name); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if enabled != 1 || name != "新名" {
|
||||
t.Fatalf("enabled=%d name=%q", enabled, name)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,276 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/admin"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
type concHashPool struct {
|
||||
inner *auth.StubHashPool
|
||||
mu sync.Mutex
|
||||
active, max int
|
||||
}
|
||||
|
||||
func (p *concHashPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) {
|
||||
p.mu.Lock()
|
||||
p.active++
|
||||
if p.active > p.max {
|
||||
p.max = p.active
|
||||
}
|
||||
p.mu.Unlock()
|
||||
time.Sleep(40 * time.Millisecond)
|
||||
defer func() {
|
||||
p.mu.Lock()
|
||||
p.active--
|
||||
p.mu.Unlock()
|
||||
}()
|
||||
return p.inner.Hash(ctx, kind, password)
|
||||
}
|
||||
|
||||
func (p *concHashPool) Verify(ctx context.Context, kind auth.PasswordKind, password, phc string) (bool, error) {
|
||||
return p.inner.Verify(ctx, kind, password, phc)
|
||||
}
|
||||
|
||||
func (p *concHashPool) QueueLen() int { return 0 }
|
||||
|
||||
func (p *concHashPool) maxActive() int {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
return p.max
|
||||
}
|
||||
|
||||
type gateHashPool struct {
|
||||
inner *auth.StubHashPool
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
startOnce sync.Once
|
||||
}
|
||||
|
||||
func (p *gateHashPool) Hash(ctx context.Context, kind auth.PasswordKind, password string) (string, error) {
|
||||
p.startOnce.Do(func() { close(p.started) })
|
||||
select {
|
||||
case <-p.release:
|
||||
case <-ctx.Done():
|
||||
return "", ctx.Err()
|
||||
}
|
||||
return p.inner.Hash(ctx, kind, password)
|
||||
}
|
||||
|
||||
func (p *gateHashPool) Verify(ctx context.Context, kind auth.PasswordKind, password, phc string) (bool, error) {
|
||||
return p.inner.Verify(ctx, kind, password, phc)
|
||||
}
|
||||
|
||||
func (p *gateHashPool) QueueLen() int { return 0 }
|
||||
|
||||
func setupEndpointsHash(t *testing.T, hash auth.HashPool) (*store.DB, *httptest.Server, *http.Client) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
// 管理员口令用假哈希写入,避免导入用的 Hash 闸门卡住 seed。
|
||||
if seedErr := admin.SeedAdminPassword(context.Background(), db, auth.NewStubHashPool(), testPassword); seedErr != nil {
|
||||
t.Fatal(seedErr)
|
||||
}
|
||||
h := admin.New(admin.Deps{
|
||||
DB: db,
|
||||
Hash: hash,
|
||||
Tokens: admin.NewRandomAPITokens(),
|
||||
Locks: admin.NewMemoryLoginLocks(),
|
||||
})
|
||||
srv := httptest.NewServer(h)
|
||||
t.Cleanup(srv.Close)
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar}
|
||||
res := postJSON(t, client, srv.URL+"/api/admin/login",
|
||||
`{"username":"admin","password":"`+testPassword+`"}`, nil)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusOK || !env.OK {
|
||||
t.Fatalf("login failed: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
return db, srv, client
|
||||
}
|
||||
|
||||
func postImportCSV(t *testing.T, client *http.Client, base, csvBody string) *http.Response {
|
||||
t.Helper()
|
||||
req, err := http.NewRequest(http.MethodPost, base+"/api/admin/endpoints/import", strings.NewReader(csvBody))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "text/csv")
|
||||
req.Header.Set("X-Nixmsg-Request", "1")
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func decodeCSVErrors(t *testing.T, res *http.Response) (status int, code string, lines []int) {
|
||||
t.Helper()
|
||||
raw, _ := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
var env struct {
|
||||
OK bool `json:"ok"`
|
||||
Error *struct {
|
||||
Code string `json:"code"`
|
||||
} `json:"error"`
|
||||
Data *struct {
|
||||
Errors []struct {
|
||||
Line int `json:"line"`
|
||||
Reason string `json:"reason"`
|
||||
} `json:"errors"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &env); err != nil {
|
||||
t.Fatalf("json: %v body=%s", err, raw)
|
||||
}
|
||||
code = ""
|
||||
if env.Error != nil {
|
||||
code = env.Error.Code
|
||||
}
|
||||
if env.Data != nil {
|
||||
for _, e := range env.Data.Errors {
|
||||
lines = append(lines, e.Line)
|
||||
}
|
||||
}
|
||||
return res.StatusCode, code, lines
|
||||
}
|
||||
|
||||
func TestImportHashConcurrencyGreaterThanOne(t *testing.T) {
|
||||
if runtime.NumCPU() < 2 {
|
||||
t.Skip("need at least 2 CPUs to observe concurrent hashing")
|
||||
}
|
||||
pool := &concHashPool{inner: auth.NewStubHashPool()}
|
||||
_, srv, client := setupEndpointsHash(t, pool)
|
||||
var b strings.Builder
|
||||
b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n")
|
||||
for i := 0; i < 8; i++ {
|
||||
fmt.Fprintf(&b, "conc-%d,名,password1,,0,\n", i)
|
||||
}
|
||||
res := postImportCSV(t, client, srv.URL, b.String())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("import: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
if pool.maxActive() < 2 {
|
||||
t.Fatalf("want concurrent hash > 1, got %d", pool.maxActive())
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportThousandRows(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
var b strings.Builder
|
||||
b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n")
|
||||
for i := 0; i < 1000; i++ {
|
||||
fmt.Fprintf(&b, "r%04d,名,password1,,0,\n", i)
|
||||
}
|
||||
res := postImportCSV(t, client, srv.URL, b.String())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("import 1000: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportUniqueRaceReturns409WithLine(t *testing.T) {
|
||||
stub := auth.NewStubHashPool()
|
||||
gate := &gateHashPool{inner: stub, started: make(chan struct{}), release: make(chan struct{})}
|
||||
db, srv, client := setupEndpointsHash(t, gate)
|
||||
csvBody := "" +
|
||||
"id,name,login_password,talk_password,default_delay_seconds,remark\n" +
|
||||
"race-1,甲,password1,,0,\n"
|
||||
|
||||
done := make(chan *http.Response, 1)
|
||||
go func() {
|
||||
done <- postImportCSV(t, client, srv.URL, csvBody)
|
||||
}()
|
||||
select {
|
||||
case <-gate.started:
|
||||
case <-time.After(5 * time.Second):
|
||||
close(gate.release)
|
||||
t.Fatal("hash did not start")
|
||||
}
|
||||
err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`
|
||||
INSERT INTO endpoints(id, name, remark, source, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
|
||||
VALUES ('race-1', '占', '', 'admin', 'stub$x', NULL, 0, 0, 1, 1)`)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
close(gate.release)
|
||||
t.Fatal(err)
|
||||
}
|
||||
close(gate.release)
|
||||
res := <-done
|
||||
status, code, lines := decodeCSVErrors(t, res)
|
||||
if status != http.StatusConflict || code != "id_taken" {
|
||||
t.Fatalf("want 409 id_taken got %d %s lines=%v", status, code, lines)
|
||||
}
|
||||
found := false
|
||||
for _, ln := range lines {
|
||||
if ln == 2 {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("want line 2 in conflict errors, got %v", lines)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportErrorLineSkipsEmptyRows(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
csvBody := "" +
|
||||
"id,name,login_password,talk_password,default_delay_seconds,remark\n" +
|
||||
"\n" +
|
||||
"bad-1,甲,password1,,-1,\n"
|
||||
res := postImportCSV(t, client, srv.URL, csvBody)
|
||||
status, _, lines := decodeCSVErrors(t, res)
|
||||
if status != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d lines=%v", status, lines)
|
||||
}
|
||||
found := false
|
||||
for _, ln := range lines {
|
||||
if ln == 3 {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("want physical line 3, got %v", lines)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportUnclosedQuoteLine(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
csvBody := "" +
|
||||
"id,name,login_password,talk_password,default_delay_seconds,remark\n" +
|
||||
"\"not-closed\n"
|
||||
res := postImportCSV(t, client, srv.URL, csvBody)
|
||||
status, _, lines := decodeCSVErrors(t, res)
|
||||
if status != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d lines=%v", status, lines)
|
||||
}
|
||||
if len(lines) == 0 || lines[0] < 2 {
|
||||
t.Fatalf("want parse error line >= 2, got %v", lines)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCreateRejectsDefaultDelayOverMax(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
body := `{"id":"dly-c","name":"延迟","login_password":"password1","default_delay_seconds":31536001}`
|
||||
res := postJSON(t, client, srv.URL+"/api/admin/endpoints", body, csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPatchRejectsDefaultDelayOverMax(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
res := postJSON(t, client, srv.URL+"/api/admin/endpoints",
|
||||
`{"id":"dly-p","name":"延迟","login_password":"password1"}`, csrfHeaders())
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
res = doReq(t, client, http.MethodPatch, srv.URL+"/api/admin/endpoints/dly-p",
|
||||
`{"default_delay_seconds":31536001}`, csrfHeaders())
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportRejectsDefaultDelayOverMax(t *testing.T) {
|
||||
_, srv, client, _, _ := setupEndpoints(t)
|
||||
csvBody := "" +
|
||||
"id,name,login_password,talk_password,default_delay_seconds,remark\n" +
|
||||
"dly-i,甲,password1,,31536001,\n"
|
||||
res := postImportCSV(t, client, srv.URL, csvBody)
|
||||
raw, _ := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
if res.StatusCode != http.StatusBadRequest {
|
||||
t.Fatalf("want 400 got %d body=%s", res.StatusCode, raw)
|
||||
}
|
||||
var parsed struct {
|
||||
Data *struct {
|
||||
Errors []struct {
|
||||
Line int `json:"line"`
|
||||
Reason string `json:"reason"`
|
||||
} `json:"errors"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &parsed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if parsed.Data == nil || len(parsed.Data.Errors) == 0 {
|
||||
t.Fatalf("want row error, body=%s", raw)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/admin"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
func TestPasswordWrongOldLocksAfterTen(t *testing.T) {
|
||||
_, srv, client, _ := setup(t)
|
||||
login(t, client, srv.URL)
|
||||
base := srv.URL
|
||||
var last *http.Response
|
||||
var env envelope
|
||||
for i := 0; i < 10; i++ {
|
||||
last = postJSON(t, client, base+"/api/admin/password",
|
||||
`{"old_password":"not-the-password","new_password":"new-password-12"}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1"})
|
||||
env = decodeEnv(t, last)
|
||||
}
|
||||
if last.StatusCode != http.StatusTooManyRequests {
|
||||
t.Fatalf("want 429 after 10 wrong old passwords, got %d %+v", last.StatusCode, env)
|
||||
}
|
||||
if env.Error == nil || env.Error.Code != "rate_limited" {
|
||||
t.Fatalf("want rate_limited got %+v", env.Error)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginDeletesExpiredAdminSessions(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
hash := auth.NewStubHashPool()
|
||||
if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil {
|
||||
t.Fatal(seedErr)
|
||||
}
|
||||
h := admin.New(admin.Deps{
|
||||
DB: db,
|
||||
Hash: hash,
|
||||
Tokens: admin.NewRandomAPITokens(),
|
||||
Locks: admin.NewMemoryLoginLocks(),
|
||||
})
|
||||
srv := httptest.NewServer(h)
|
||||
t.Cleanup(srv.Close)
|
||||
if err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`INSERT INTO admin_sessions(token_hash, created_at, expires_at) VALUES ('expired-hash', 1, 1)`)
|
||||
return e
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar}
|
||||
login(t, client, srv.URL)
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM admin_sessions WHERE token_hash = 'expired-hash'`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("expired session still present, count=%d", n)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
@@ -11,6 +12,7 @@ import (
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
@@ -23,6 +25,13 @@ const (
|
||||
defaultSessionTTL = 12 * time.Hour
|
||||
minPasswordLen = 12
|
||||
lastUsedMinGap = time.Minute
|
||||
|
||||
maxLoginBodyBytes = 8 << 10
|
||||
maxJSONBodyBytes = 1 << 20
|
||||
maxImportBodyBytes = 8 << 20
|
||||
defaultReadFor = 15 * time.Second
|
||||
loginReadFor = 10 * time.Second
|
||||
importReadFor = 2 * time.Minute
|
||||
)
|
||||
|
||||
// Deps 是管理 Handler 的依赖。
|
||||
@@ -129,11 +138,36 @@ func New(d Deps) *Handler {
|
||||
return h
|
||||
}
|
||||
|
||||
// ServeHTTP 实现 http.Handler。
|
||||
// ServeHTTP 实现 http.Handler。按路由限制请求体大小与读截止时间。
|
||||
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
limit, readFor := requestBodyBudget(r)
|
||||
if err := http.NewResponseController(w).SetReadDeadline(time.Now().Add(readFor)); err != nil && !errors.Is(err, http.ErrNotSupported) {
|
||||
// 测试用 ResponseRecorder 或不支持截止时间的封装:忽略。
|
||||
}
|
||||
if r.Body != nil {
|
||||
r.Body = http.MaxBytesReader(w, r.Body, limit)
|
||||
}
|
||||
h.mux.ServeHTTP(w, r)
|
||||
}
|
||||
|
||||
func requestBodyBudget(r *http.Request) (int64, time.Duration) {
|
||||
if r.Method == http.MethodPost && r.URL.Path == "/api/admin/endpoints/import" {
|
||||
return maxImportBodyBytes, importReadFor
|
||||
}
|
||||
if r.Method == http.MethodPost && r.URL.Path == "/api/admin/login" {
|
||||
return maxLoginBodyBytes, loginReadFor
|
||||
}
|
||||
return maxJSONBodyBytes, defaultReadFor
|
||||
}
|
||||
|
||||
func writeDecodeError(w http.ResponseWriter, err error) {
|
||||
if httpx.IsBodyTooLarge(err) {
|
||||
httpx.WriteError(w, http.StatusRequestEntityTooLarge, "payload_too_large", "请求体过大")
|
||||
return
|
||||
}
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
|
||||
}
|
||||
|
||||
func (h *Handler) routes() {
|
||||
// 公开
|
||||
h.mux.HandleFunc("POST /api/admin/login", h.handleLogin)
|
||||
|
||||
+25
-5
@@ -4,6 +4,7 @@ import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"net/http"
|
||||
"unicode/utf8"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
@@ -35,7 +36,7 @@ func (h *Handler) handleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
Password string `json:"password"`
|
||||
}
|
||||
if err := httpx.DecodeJSON(r, &req); err != nil {
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
|
||||
writeDecodeError(w, err)
|
||||
return
|
||||
}
|
||||
if req.Username != adminUsername {
|
||||
@@ -111,16 +112,23 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) {
|
||||
p, _ := principalFrom(r.Context())
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
|
||||
if locked, retry := h.locks.Check(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}); locked {
|
||||
w.Header().Set("Retry-After", formatRetryAfter(retry))
|
||||
h.audit(actorString(p), "password_change", "", "rate_limited", ip)
|
||||
httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
OldPassword string `json:"old_password"`
|
||||
NewPassword string `json:"new_password"`
|
||||
}
|
||||
if err := httpx.DecodeJSON(r, &req); err != nil {
|
||||
h.audit(actorString(p), "password_change", "", "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
|
||||
writeDecodeError(w, err)
|
||||
return
|
||||
}
|
||||
if len(req.NewPassword) < minPasswordLen {
|
||||
if utf8.RuneCountInString(req.NewPassword) < minPasswordLen {
|
||||
h.audit(actorString(p), "password_change", "", "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "新密码至少 12 位")
|
||||
return
|
||||
@@ -134,10 +142,18 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
ok, err := h.hash.Verify(r.Context(), auth.PasswordAdmin, req.OldPassword, phc)
|
||||
if err != nil || !ok {
|
||||
locked, retry := h.locks.Fail(auth.LockKey{Kind: auth.LockAdminIP, IP: ip})
|
||||
if locked {
|
||||
w.Header().Set("Retry-After", formatRetryAfter(retry))
|
||||
h.audit(actorString(p), "password_change", "", "rate_limited", ip)
|
||||
httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试")
|
||||
return
|
||||
}
|
||||
h.audit(actorString(p), "password_change", "", "unauthorized", ip)
|
||||
httpx.WriteError(w, http.StatusUnauthorized, "unauthorized", "旧密码错误")
|
||||
return
|
||||
}
|
||||
h.locks.Clear(auth.LockKey{Kind: auth.LockAdminIP, IP: ip})
|
||||
newPHC, err := h.hash.Hash(r.Context(), auth.PasswordAdmin, req.NewPassword)
|
||||
if err != nil {
|
||||
h.audit(actorString(p), "password_change", "", "error", ip)
|
||||
@@ -149,9 +165,13 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) {
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
// 保留当前会话,作废其它会话
|
||||
// 保留当前会话,作废其它会话;失败则返回 500,避免其它会话继续有效。
|
||||
if p.Session != "" {
|
||||
_ = h.deleteOtherSessions(r.Context(), hashSessionHex(p.Session))
|
||||
if err := h.deleteOtherSessions(r.Context(), hashSessionHex(p.Session)); err != nil {
|
||||
h.audit(actorString(p), "password_change", "", "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
}
|
||||
h.audit(actorString(p), "password_change", "", "ok", ip)
|
||||
httpx.WriteOK(w, map[string]any{})
|
||||
|
||||
@@ -42,6 +42,9 @@ func (h *Handler) setAdminPasswordHash(ctx context.Context, phc string) error {
|
||||
func (h *Handler) createSession(ctx context.Context, hashHex string, ttl time.Duration) error {
|
||||
now := time.Now()
|
||||
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, err := tx.Exec(`DELETE FROM admin_sessions WHERE expires_at <= ?`, now.UnixMilli()); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := tx.Exec(
|
||||
`INSERT INTO admin_sessions(token_hash, created_at, expires_at) VALUES(?, ?, ?)`,
|
||||
hashHex, now.UnixMilli(), now.Add(ttl).UnixMilli(),
|
||||
|
||||
Reference in New Issue
Block a user