diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 8e1c47e..ceff40f 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -760,6 +760,51 @@ - 备选:无。 - 影响:无。 +### 复审修复 H-01 + +- 日期:2026-09-30 +- 原条款:PRD §8 管理员登录防暴力;审查 #44。 +- 实际做法:`httpx.DecodeJSON` 内部 `MaxBytesReader` 1 MiB;`admin.Handler.ServeHTTP` 按路由限制(登录 8 KiB、导入 8 MiB、其余 1 MiB)并设读截止时间;超限 413 JSON `payload_too_large`。CSV 导入不再在 8 MiB 处静默截断。不在 listener 加全局 `ReadTimeout`。 +- 原因:公开登录接口与 JSON 解码原先不限大小。 +- 备选方案:登录 4 KiB(与注册一致);未采用,8 KiB 对口令字段更宽裕。 +- 影响:超大请求快速失败;合法批量 JSON(约 200 KiB)仍低于 1 MiB。 + +### 复审修复 H-03 + +- 日期:2026-09-30 +- 原条款:PRD F01 停用作废不可恢复;审查 #46。 +- 实际做法:注入 Identity 时 `handleEndpointPatch` 给 `patchEndpoint` 传 nil 的 enabled,启停只由 `identity.Disable`/`Enable` 执行。界面部分在 W-05。 +- 原因:先写 enabled 再级联失败会留下半生效状态。 +- 备选方案:同一事务里写 enabled 与作废;Identity 接口当前不暴露事务。 +- 影响:Identity 失败时端仍保持原启用状态。 + +### 复审修复 H-04 + +- 日期:2026-09-30 +- 原条款:DEVELOPMENT §8 批量开通哈希走并发池;PRD F01 指出行号;审查 #47。 +- 实际做法:导入按 `max(1, NumCPU-1)` 并发算哈希并按行下标收集;UNIQUE 冲突返回 409 附行号;哈希失败 500/503。CSV 逐条 `Read()`,用 `FieldPos` / `csv.ParseError` 取物理行号。界面进度与错误表在 W-05。 +- 原因:串行 argon2 过慢,且 `i+2` 行号会被空行和跨行字段带偏。 +- 备选方案:占满哈希池;未采用,给 MQTT 登录留槽位。 +- 影响:导入吞吐上升;冲突不再表现为无行号的 500。 + +### 复审修复 H-05 + +- 日期:2026-09-30 +- 原条款:PRD F11 最远可定到 365 天后(可配置);审查 #48。 +- 实际做法:`validateEndpointFields` 增加 `MaxScheduleSeconds`,开通/修改/导入与端侧改默认延迟同一上限(0 ≤ delay ≤ max)。界面 `:max` 在 W-05。 +- 原因:后台原先只检查 `>= 0`,可设出永远发不出消息的端,极大值还会溢出。 +- 备选方案:导入超限静默改成上限;未采用,改为行错误让管理员改表。 +- 影响:超上限返回 400 或 CSV 行错误。 + +### 复审修复 H-06 + +- 日期:2026-09-30 +- 原条款:PRD §8 管理员登录防暴力;admin-api 2.4 改密作废其它会话;审查 #49。 +- 实际做法:改管理员密码时旧密码校验计入 `LockAdminIP`,锁定返回 429,成功清零。登录写入会话的同一事务删除过期 `admin_sessions`。作废其它会话失败返回 500 并写审计。 +- 原因:持有 Cookie 可无限猜旧密码;过期会话从不清理;作废失败被忽略。 +- 备选方案:改密失败不返回 401(与契约不符);未采用。W-01 排除改密 401 自动登出。 +- 影响:连续 10 次旧密码错误后锁定;过期会话行在下次登录清除。 + ## 后台网页 W 1. **W1–W3 阶段使用内存假数据,不请求真实 `/api/admin`** @@ -818,6 +863,61 @@ - 备选:仅文档要求手跑 e2e;未采用。 - 影响:`task itest` 变长;需本机 Go/Node/Chromium。 +### 复审修复 W-01 + +- 日期:2026-09-30 +- 原条款:TASKS W1 请求封装;审查 #51;PRD F01 名称可选。 +- 实际做法:`requestAdmin` 对 401(排除 login/password)清空会话并跳转登录页,提示一次「登录已过期」;`run()` 不再对 `ApiError` 重复 toast;网络错误「无法连接服务器」、429「请稍后再试」。概览/系统/端/群/投递/令牌加载失败显示 `LoadFailed` 与重试。开通名称改为可选;改密用 n-form rules 按字符数校验,后端 `login.go` 同步 `utf8.RuneCountInString`;清除对话密码加确认。 +- 未改:`RegistrationView.vue` / `registration.go`(U-01 归属;本波指令禁止改注册页)。生成新安全码确认留给身份线。 +- 原因:会话过期后页面不跳转且错误弹两次;加载失败空白;帮助文字误称 nst_。 +- 备选方案:改密 401 也自动登出;未采用,与契约「旧密码错误 401」冲突。 +- 影响:改密失败不会踢当前会话;注册页错误提示仍走原逻辑。 + +### 复审修复 W-04 + +- 日期:2026-09-30 +- 原条款:PRD F01 密码只显示一次并可下载;审查 #54。 +- 实际做法:新增 `utils/csv.ts` 按 RFC 4180 引号转义、CRLF、BOM、表头;名称列公式前缀 `'`,密码列不加前缀。`SecretOnceAlert` 增加 `mime`、复制成功/失败提示与 clipboard 回退。开通手填密码时不返回 `login_password`,只提示「已开通 {id}」。导入结果弹窗增加下载。 +- 原因:拼接 CSV 会错列;手填密码显示 undefined;明文 HTTP 下 clipboard 不可用。 +- 备选方案:密码列也加公式前缀;未采用,避免改写密码。 +- 影响:导入下载文件可被 Excel 正确打开。 + +### 复审修复 W-05 + +- 日期:2026-09-30 +- 原条款:PRD F01/F17 解锁与批量清理;审查 #55;H-03/H-04/H-05 界面部分。 +- 实际做法:解锁始终可点,锁定列 HelpTip 说明只反映编号锁。筛选/翻页清空选择;批量确认列出前 20 个编号,删除说明级联,结果列出失败项。编辑只提交变更字段;启用改停用弹确认。导入中禁用按钮并显示预估秒数,校验错误用可滚动表格。默认延迟 `:max` 取自 `limits.max_schedule_seconds`。 +- 原因:仅 IP 锁定时无法解锁;跨页多选会误删;PATCH 总带 enabled 可能误停用。 +- 备选方案:扩展 LoginLocks 查询任一锁定;留给 P 线。 +- 影响:管理员可解除 IP 锁定;停用需二次确认。 + +### 复审修复 W-02 + +- 日期:2026-09-30 +- 原条款:PRD F17 按时间筛选并查看原因/完整统计;审查 #52。 +- 实际做法:投递记录加 `n-date-picker datetimerange` 传 `from_ms`/`to_ms`;列表与详情增加原因;统计六项;详情接收端 `limit=200` 并以「加载更多」翻页。 +- 原因:接口已提供字段,页面未接线。 +- 备选方案:详情用页码分页;未采用,加载更多更贴合游标。 +- 影响:大群投递可看完全部接收端。 + +### 复审修复 W-03 + +- 日期:2026-09-30 +- 原条款:PRD F17 群成员增减;审查 #53;DEVIATIONS A3.3。 +- 实际做法:群详情返回 `member_total`(只增字段)并同步 admin-api 7.8;成员表远程分页。去掉建群编号输入,占位改为「群名称」。移除/转让加确认;加人失败原因译成中文。 +- 原因:默认每页 50,千人群看不到后半;前端仍按可自定 id 设计。 +- 备选方案:后端按编号前缀搜索;未采用,与现接口「名称包含」一致。 +- 影响:网页线改了 `internal/admin/groups.go` 与 admin-api 7.8(issue 允许)。 + +### 复审修复 W-06 + +- 日期:2026-09-30 +- 原条款:用户界面规则长列表内部滚动;PRD F17 界面中文与自助注册数;审查 #56。 +- 实际做法:列表页 `.page-body` 改为 flex 列,筛选不伸缩,表格 `flex-height`。新增 `utils/labels.ts` 译消息/投递状态与原因(原值 title)。运行参数中文标签、键名放 HelpTip。端列表拆成最近上线/离线。概览展示「其中自助注册」,运行时长按天时分。e2e 断言「定时中」。Playwright 双分辨率截图未跑(本波以 Vitest 验收布局类改动)。 +- 原因:写死 480 高度双滚动;英文枚举;漏自助数。 +- 备选方案:一次性密码改弹窗;仍放顶部不参与伸缩。 +- 影响:投递页不再显示 `scheduled` 英文单元格。 + ### S1.1 传输层可注入假实现(Go / JS) - 相关文档:DEVELOPMENT 第 9 节单元测试要求「用假的 MQTT/HTTP,不要起真实服务器」。 diff --git a/docs/api/admin-api.md b/docs/api/admin-api.md index 877d68c..55e9938 100644 --- a/docs/api/admin-api.md +++ b/docs/api/admin-api.md @@ -559,6 +559,7 @@ Authorization: Bearer nxm_... "name": "一组", "owner_id": "a", "created_at_ms": 1750000000000, + "member_total": 1, "members": [ {"id": "a", "name": "", "online": true, "joined_at_ms": 1750000000000} ], diff --git a/internal/admin/a3_test.go b/internal/admin/a3_test.go index 953e3bb..32aa4b3 100644 --- a/internal/admin/a3_test.go +++ b/internal/admin/a3_test.go @@ -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) diff --git a/internal/admin/endpoints.go b/internal/admin/endpoints.go index d6ef5bc..5e466cd 100644 --- a/internal/admin/endpoints.go +++ b/internal/admin/endpoints.go @@ -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}, + }) +} diff --git a/internal/admin/endpoints_csv.go b/internal/admin/endpoints_csv.go index 11b3fa5..bf9844f 100644 --- a/internal/admin/endpoints_csv.go +++ b/internal/admin/endpoints_csv.go @@ -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 { diff --git a/internal/admin/groups.go b/internal/admin/groups.go index 14c122d..c8eb577 100644 --- a/internal/admin/groups.go +++ b/internal/admin/groups.go @@ -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, }) diff --git a/internal/admin/h01_test.go b/internal/admin/h01_test.go new file mode 100644 index 0000000..f1e11d4 --- /dev/null +++ b/internal/admin/h01_test.go @@ -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) + } +} diff --git a/internal/admin/h03_test.go b/internal/admin/h03_test.go new file mode 100644 index 0000000..d4ebd92 --- /dev/null +++ b/internal/admin/h03_test.go @@ -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) + } +} diff --git a/internal/admin/h04_test.go b/internal/admin/h04_test.go new file mode 100644 index 0000000..a72e853 --- /dev/null +++ b/internal/admin/h04_test.go @@ -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) + } +} diff --git a/internal/admin/h05_test.go b/internal/admin/h05_test.go new file mode 100644 index 0000000..1af4142 --- /dev/null +++ b/internal/admin/h05_test.go @@ -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) + } +} diff --git a/internal/admin/h06_test.go b/internal/admin/h06_test.go new file mode 100644 index 0000000..2d16503 --- /dev/null +++ b/internal/admin/h06_test.go @@ -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) + } +} diff --git a/internal/admin/handler.go b/internal/admin/handler.go index e44b0a8..6787226 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -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) diff --git a/internal/admin/login.go b/internal/admin/login.go index 306184e..6f8140f 100644 --- a/internal/admin/login.go +++ b/internal/admin/login.go @@ -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{}) diff --git a/internal/admin/store.go b/internal/admin/store.go index d4e52b9..b736999 100644 --- a/internal/admin/store.go +++ b/internal/admin/store.go @@ -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(), diff --git a/internal/httpx/json.go b/internal/httpx/json.go index ecfac96..93175df 100644 --- a/internal/httpx/json.go +++ b/internal/httpx/json.go @@ -8,6 +8,17 @@ import ( "net/http" ) +const maxJSONBodyBytes = 1 << 20 + +// IsBodyTooLarge 判断是否因请求体超过 MaxBytesReader 上限而失败。 +func IsBodyTooLarge(err error) bool { + if err == nil { + return false + } + var maxErr *http.MaxBytesError + return errors.As(err, &maxErr) +} + // ErrorBody 是失败响应里的 error 对象。 type ErrorBody struct { Code string `json:"code"` @@ -46,10 +57,15 @@ func WriteError(w http.ResponseWriter, status int, code, message string) { }) } -// DecodeJSON 解码请求 JSON 体;空体对 dst 保持零值。 +// DecodeJSON 解码请求 JSON 体;空体对 dst 保持零值。内部把请求体限制在 1 MiB。 func DecodeJSON(r *http.Request, dst any) error { defer func() { _ = r.Body.Close() }() - dec := json.NewDecoder(r.Body) + body := r.Body + if body == nil { + return nil + } + body = http.MaxBytesReader(nil, body, maxJSONBodyBytes) + dec := json.NewDecoder(body) dec.DisallowUnknownFields() if err := dec.Decode(dst); err != nil { if errors.Is(err, io.EOF) { diff --git a/internal/httpx/json_test.go b/internal/httpx/json_test.go new file mode 100644 index 0000000..9307cc4 --- /dev/null +++ b/internal/httpx/json_test.go @@ -0,0 +1,39 @@ +package httpx + +import ( + "bytes" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestDecodeJSONRejectsOversizeBody(t *testing.T) { + t.Parallel() + payload := `{"password":"` + strings.Repeat("a", 2<<20) + `"}` + req := httptest.NewRequest(http.MethodPost, "/x", strings.NewReader(payload)) + var dst struct { + Password string `json:"password"` + } + err := DecodeJSON(req, &dst) + if err == nil { + t.Fatal("want error for 2 MiB JSON body") + } + if !IsBodyTooLarge(err) { + t.Fatalf("want IsBodyTooLarge, got %v", err) + } +} + +func TestDecodeJSONAcceptsSmallBody(t *testing.T) { + t.Parallel() + req := httptest.NewRequest(http.MethodPost, "/x", bytes.NewReader([]byte(`{"password":"ok"}`))) + var dst struct { + Password string `json:"password"` + } + if err := DecodeJSON(req, &dst); err != nil { + t.Fatal(err) + } + if dst.Password != "ok" { + t.Fatalf("password=%q", dst.Password) + } +} diff --git a/web/e2e/admin-main.spec.ts b/web/e2e/admin-main.spec.ts index d862b77..152f868 100644 --- a/web/e2e/admin-main.spec.ts +++ b/web/e2e/admin-main.spec.ts @@ -127,7 +127,7 @@ test.describe("后台主路径 W4", () => { await nav(page, "投递记录"); await expect(page).toHaveURL(/\/messages/); await expect(page.getByText("w4-e2e-scheduled")).toBeVisible(); - await expect(page.getByRole("cell", { name: "scheduled", exact: true })).toBeVisible(); + await expect(page.getByRole("cell", { name: "定时中", exact: true })).toBeVisible(); await nav(page, "API 令牌"); await expect(page).toHaveURL(/\/tokens/); diff --git a/web/src/api/admin-mock.ts b/web/src/api/admin-mock.ts index 53a1f60..ac6c9d6 100644 --- a/web/src/api/admin-mock.ts +++ b/web/src/api/admin-mock.ts @@ -22,9 +22,7 @@ async function run(fn: () => Promise, silent = false): Promise { try { return await fn(); } catch (e) { - if (!silent && e instanceof ApiError) { - message.error(e.message); - } else if (!silent && e instanceof Error) { + if (!silent && e instanceof Error && !(e instanceof ApiError)) { message.error(e.message); } throw e; @@ -163,12 +161,14 @@ export function listMessages(q: { endpoint_id?: string; group_id?: string; state?: string; + from_ms?: number; + to_ms?: number; }) { return run(() => mockApi.listMessages(q)); } -export function getMessage(seq: number) { - return run(() => mockApi.getMessage(seq)); +export function getMessage(seq: number, cursor?: string, limit?: number) { + return run(() => mockApi.getMessage(seq, cursor, limit)); } export function getSettings() { diff --git a/web/src/api/admin.ts b/web/src/api/admin.ts index 7944f69..444114e 100644 --- a/web/src/api/admin.ts +++ b/web/src/api/admin.ts @@ -33,9 +33,7 @@ async function run(fn: () => Promise, silent = false): Promise { try { return await fn(); } catch (e) { - if (!silent && e instanceof ApiError) { - message.error(e.message); - } else if (!silent && e instanceof Error) { + if (!silent && e instanceof Error && !(e instanceof ApiError)) { message.error(e.message); } throw e; @@ -333,6 +331,8 @@ export function listMessages(q: { endpoint_id?: string; group_id?: string; state?: string; + from_ms?: number; + to_ms?: number; }): Promise> { return run(() => requestAdmin>( @@ -343,13 +343,17 @@ export function listMessages(q: { endpoint_id: q.endpoint_id, group_id: q.group_id, state: q.state, + from_ms: q.from_ms, + to_ms: q.to_ms, })}`, ), ); } -export function getMessage(seq: number): Promise { - return run(() => requestAdmin(`/api/admin/messages/${seq}`)); +export function getMessage(seq: number, cursor?: string, limit?: number): Promise { + return run(() => + requestAdmin(`/api/admin/messages/${seq}${buildQuery({ cursor, limit })}`), + ); } export function getSettings(): Promise { diff --git a/web/src/api/http.spec.ts b/web/src/api/http.spec.ts new file mode 100644 index 0000000..afe7b5a --- /dev/null +++ b/web/src/api/http.spec.ts @@ -0,0 +1,78 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { ApiError, requestAdmin, setUnauthorizedHandler } from "./http"; + +vi.mock("@/utils/notify", () => ({ + message: { + error: vi.fn(), + success: vi.fn(), + warning: vi.fn(), + }, +})); + +import { message } from "@/utils/notify"; + +describe("requestAdmin", () => { + beforeEach(() => { + vi.stubGlobal("fetch", vi.fn()); + setUnauthorizedHandler(null); + vi.mocked(message.error).mockClear(); + }); + + afterEach(() => { + vi.unstubAllGlobals(); + setUnauthorizedHandler(null); + }); + + it("401 跳转登录页且只提示一次", async () => { + const fetchMock = vi.mocked(fetch); + fetchMock.mockResolvedValue( + new Response(JSON.stringify({ ok: false, error: { code: "unauthorized", message: "未登录" } }), { + status: 401, + headers: { "Content-Type": "application/json" }, + }), + ); + const handler = vi.fn(); + setUnauthorizedHandler(handler); + + await expect(requestAdmin("/api/admin/overview")).rejects.toBeInstanceOf(ApiError); + await expect(requestAdmin("/api/admin/endpoints")).rejects.toBeInstanceOf(ApiError); + + expect(handler).toHaveBeenCalled(); + expect(vi.mocked(message.error).mock.calls.filter((c) => c[0] === "登录已过期")).toHaveLength(1); + }); + + it("改密 401 不自动登出", async () => { + const fetchMock = vi.mocked(fetch); + fetchMock.mockResolvedValue( + new Response(JSON.stringify({ ok: false, error: { code: "unauthorized", message: "旧密码错误" } }), { + status: 401, + headers: { "Content-Type": "application/json" }, + }), + ); + const handler = vi.fn(); + setUnauthorizedHandler(handler); + await expect( + requestAdmin("/api/admin/password", { method: "POST", body: { old_password: "x", new_password: "yyyyyyyyyyyy" } }), + ).rejects.toMatchObject({ message: "旧密码错误" }); + expect(handler).not.toHaveBeenCalled(); + }); + + it("网络错误显示中文", async () => { + vi.mocked(fetch).mockRejectedValue(new TypeError("Failed to fetch")); + await expect(requestAdmin("/api/admin/me")).rejects.toMatchObject({ message: "无法连接服务器" }); + expect(message.error).toHaveBeenCalledWith("无法连接服务器"); + }); + + it("429 提示稍后再试", async () => { + vi.mocked(fetch).mockResolvedValue( + new Response(JSON.stringify({ ok: false, error: { code: "rate_limited", message: "locked" } }), { + status: 429, + headers: { "Content-Type": "application/json" }, + }), + ); + await expect(requestAdmin("/api/admin/login", { method: "POST", body: {} })).rejects.toMatchObject({ + message: "请求过于频繁,请稍后再试", + }); + expect(message.error).toHaveBeenCalledWith("请求过于频繁,请稍后再试"); + }); +}); diff --git a/web/src/api/http.ts b/web/src/api/http.ts index b92c1a4..6db2679 100644 --- a/web/src/api/http.ts +++ b/web/src/api/http.ts @@ -27,9 +27,52 @@ export interface RequestOptions { contentType?: string | null; } +const noAutoLogout = new Set(["/api/admin/login", "/api/admin/password"]); + +type UnauthorizedHandler = (redirectPath: string) => void; +let unauthorizedHandler: UnauthorizedHandler | null = null; +let expiredNotified = false; + +export function setUnauthorizedHandler(handler: UnauthorizedHandler | null) { + unauthorizedHandler = handler; +} + +function requestPath(path: string): string { + const q = path.indexOf("?"); + return q >= 0 ? path.slice(0, q) : path; +} + +function currentRedirect(): string { + if (typeof window === "undefined") { + return "/"; + } + return `${window.location.pathname}${window.location.search}` || "/"; +} + +function notifyExpiredOnce() { + if (expiredNotified) { + return; + } + expiredNotified = true; + message.error("登录已过期"); + window.setTimeout(() => { + expiredNotified = false; + }, 2000); +} + +function handleUnauthorized(path: string, silent: boolean) { + if (noAutoLogout.has(requestPath(path))) { + return; + } + if (!silent) { + notifyExpiredOnce(); + } + unauthorizedHandler?.(currentRedirect()); +} + /** - * 管理接口请求封装:自动带 credentials 与 X-Nixmsg-Request: 1,错误直接 message 显示。 - * W4 接真实后端时页面不必改,只需让 admin 模块走本函数。 + * 管理接口请求封装:自动带 credentials 与 X-Nixmsg-Request: 1。 + * 错误只在这里提示一次;401(登录/改密除外)清空会话并跳转登录页。 */ export async function requestAdmin(path: string, opts: RequestOptions = {}): Promise { const method = opts.method ?? (opts.body != null || opts.rawBody != null ? "POST" : "GET"); @@ -46,19 +89,32 @@ export async function requestAdmin(path: string, opts: RequestOptions = {}): headers["Content-Type"] = opts.contentType; } - const res = await fetch(path, { - method, - credentials: "include", - headers, - body: body ?? undefined, - }); + let res: Response; + try { + res = await fetch(path, { + method, + credentials: "include", + headers, + body: body ?? undefined, + }); + } catch { + const err = new ApiError("network", "无法连接服务器", 0); + if (!opts.silent) { + message.error(err.message); + } + throw err; + } + + if (res.status === 401) { + handleUnauthorized(path, Boolean(opts.silent)); + } let envelope: ApiEnvelope; try { envelope = (await res.json()) as ApiEnvelope; } catch { const err = new ApiError("internal", `响应不是 JSON(HTTP ${res.status})`, res.status); - if (!opts.silent) { + if (!opts.silent && res.status !== 401) { message.error(err.message); } throw err; @@ -66,7 +122,7 @@ export async function requestAdmin(path: string, opts: RequestOptions = {}): if (typeof envelope !== "object" || envelope === null) { const err = new ApiError("internal", "响应格式无效", res.status); - if (!opts.silent) { + if (!opts.silent && res.status !== 401) { message.error(err.message); } throw err; @@ -77,8 +133,14 @@ export async function requestAdmin(path: string, opts: RequestOptions = {}): } const errBody: ApiErrorBody = envelope.error ?? { code: "internal", message: "未知错误" }; - const err = new ApiError(errBody.code, errBody.message, res.status, "data" in envelope ? envelope.data : undefined); - if (!opts.silent) { + let display = errBody.message; + if (res.status === 429) { + display = "请求过于频繁,请稍后再试"; + } else if (res.status === 401 && !noAutoLogout.has(requestPath(path))) { + display = "登录已过期"; + } + const err = new ApiError(errBody.code, display, res.status, "data" in envelope ? envelope.data : undefined); + if (!opts.silent && res.status !== 401) { message.error(err.message); } throw err; diff --git a/web/src/api/mock.ts b/web/src/api/mock.ts index 7ce9b35..c62de3e 100644 --- a/web/src/api/mock.ts +++ b/web/src/api/mock.ts @@ -151,6 +151,14 @@ const messages: MessageDetail[] = [ pushed_at_ms: now() - 3_599_000, updated_at_ms: now() - 3_598_000, }, + { + endpoint_id: "ops-bot", + state: "rejected", + reason: "disabled", + attempts: 1, + pushed_at_ms: now() - 3_599_000, + updated_at_ms: now() - 3_598_000, + }, ], next_cursor: "", }, @@ -320,6 +328,7 @@ export const mockApi = { endpoints_total: endpoints.length, endpoints_online: endpoints.filter((e) => e.online).length, endpoints_disabled: endpoints.filter((e) => !e.enabled).length, + endpoints_self: endpoints.filter((e) => e.source === "self").length, groups_total: groups.length, messages_pending: messages.filter((m) => m.state === "dispatched").length, messages_scheduled: messages.filter((m) => m.state === "scheduled").length, @@ -354,7 +363,8 @@ export const mockApi = { if (endpoints.some((e) => e.id === id)) { throw new ApiError("id_taken", "编号已占用", 409); } - const loginPassword = (body.login_password || "").trim() || randomPassword(); + const provided = (body.login_password || "").trim(); + const loginPassword = provided || randomPassword(); endpoints.unshift({ id, name: body.name, @@ -369,7 +379,7 @@ export const mockApi = { created_at_ms: now(), login_locked: false, }); - return { id, login_password: loginPassword }; + return provided ? { id } : { id, login_password: loginPassword }; }, async importEndpoints(csvText: string): Promise<{ items: ImportItem[] }> { @@ -598,7 +608,7 @@ export const mockApi = { let list = groups.map(toGroupSummary); if (query) { const n = query.toLowerCase(); - list = list.filter((g) => g.name.toLowerCase().includes(n) || g.id.toLowerCase().includes(n)); + list = list.filter((g) => g.name.toLowerCase().includes(n)); } return paginate(list, cursor, limit); }, @@ -655,6 +665,7 @@ export const mockApi = { name: g.name, owner_id: g.owner_id, created_at_ms: g.created_at_ms, + member_total: all.length, members: page.items, next_cursor: page.next_cursor, }; @@ -729,6 +740,8 @@ export const mockApi = { endpoint_id?: string; group_id?: string; state?: string; + from_ms?: number; + to_ms?: number; }): Promise> { requireSession(); let list = messages.map(toMessageSummary); @@ -741,16 +754,23 @@ export const mockApi = { }); } if (q.state) list = list.filter((m) => m.state === q.state); + if (q.from_ms != null) list = list.filter((m) => m.created_at_ms >= q.from_ms!); + if (q.to_ms != null) list = list.filter((m) => m.created_at_ms <= q.to_ms!); return paginate(list, q.cursor, q.limit); }, - async getMessage(seq: number): Promise { + async getMessage(seq: number, cursor?: string, limit?: number): Promise { requireSession(); const m = messages.find((x) => x.seq === seq); if (!m) { throw new ApiError("not_found", "消息不存在", 404); } - return clone(m); + const page = paginate(m.deliveries, cursor, limit ?? 200); + return { + ...clone(m), + deliveries: page.items, + next_cursor: page.next_cursor, + }; }, async getSettings(): Promise { diff --git a/web/src/api/types.ts b/web/src/api/types.ts index 3501fa9..a892f1e 100644 --- a/web/src/api/types.ts +++ b/web/src/api/types.ts @@ -34,8 +34,7 @@ export interface Overview { endpoints_total: number; endpoints_online: number; endpoints_disabled: number; - /** A3 补充:自助注册端数量;契约示例未列,前端可选展示 */ - endpoints_self?: number; + endpoints_self: number; groups_total: number; messages_pending: number; messages_scheduled: number; @@ -72,7 +71,7 @@ export interface EndpointCreateRequest { export interface EndpointCreateResult { id: string; - login_password: string; + login_password?: string; } export interface EndpointPatchRequest { @@ -98,7 +97,7 @@ export interface ImportError { export interface ImportItem { id: string; - login_password: string; + login_password?: string; name: string; } @@ -155,6 +154,7 @@ export interface GroupDetail { name: string; owner_id: string; created_at_ms: number; + member_total?: number; members: GroupMember[]; next_cursor: string; } diff --git a/web/src/components/LoadFailed.vue b/web/src/components/LoadFailed.vue new file mode 100644 index 0000000..fd8b4e6 --- /dev/null +++ b/web/src/components/LoadFailed.vue @@ -0,0 +1,19 @@ + + + diff --git a/web/src/components/SecretOnceAlert.spec.ts b/web/src/components/SecretOnceAlert.spec.ts new file mode 100644 index 0000000..f08a27f --- /dev/null +++ b/web/src/components/SecretOnceAlert.spec.ts @@ -0,0 +1,73 @@ +import { config, mount, flushPromises } from "@vue/test-utils"; +import { describe, expect, it, vi, afterEach } from "vitest"; +import { NConfigProvider, NMessageProvider, zhCN, dateZhCN } from "naive-ui"; +import { defineComponent, h } from "vue"; +import SecretOnceAlert from "./SecretOnceAlert.vue"; + +vi.mock("@/utils/notify", () => ({ + message: { + success: vi.fn(), + error: vi.fn(), + warning: vi.fn(), + }, +})); + +import { message } from "@/utils/notify"; + +config.global.stubs = { teleport: true }; + +function wrap() { + return defineComponent({ + setup() { + return () => + h(NConfigProvider, { locale: zhCN, dateLocale: dateZhCN, size: "small" }, { + default: () => + h(NMessageProvider, null, { + default: () => + h(SecretOnceAlert, { + title: "令牌只显示一次", + secret: "nxm_abc", + filename: "api-token.txt", + }), + }), + }); + }, + }); +} + +describe("SecretOnceAlert 复制", () => { + afterEach(() => { + vi.mocked(message.success).mockClear(); + vi.mocked(message.error).mockClear(); + }); + + it("不支持 clipboard 时回退 execCommand 并提示成功", async () => { + Object.defineProperty(navigator, "clipboard", { value: undefined, configurable: true }); + Object.defineProperty(document, "execCommand", { value: vi.fn(() => true), configurable: true }); + + const w = mount(wrap(), { attachTo: document.body }); + await flushPromises(); + await w.find('[data-testid="secret-copy"]').trigger("click"); + await flushPromises(); + + expect(document.execCommand).toHaveBeenCalledWith("copy"); + expect(message.success).toHaveBeenCalledWith("已复制"); + w.unmount(); + }); + + it("writeText 失败时提示错误", async () => { + Object.defineProperty(navigator, "clipboard", { + value: { writeText: vi.fn().mockRejectedValue(new Error("denied")) }, + configurable: true, + }); + Object.defineProperty(document, "execCommand", { value: vi.fn(() => false), configurable: true }); + + const w = mount(wrap(), { attachTo: document.body }); + await flushPromises(); + await w.find('[data-testid="secret-copy"]').trigger("click"); + await flushPromises(); + + expect(message.error).toHaveBeenCalledWith("复制失败,请手动选择文本"); + w.unmount(); + }); +}); diff --git a/web/src/components/SecretOnceAlert.vue b/web/src/components/SecretOnceAlert.vue index 437a726..c932dc4 100644 --- a/web/src/components/SecretOnceAlert.vue +++ b/web/src/components/SecretOnceAlert.vue @@ -1,40 +1,71 @@
+ +
- - + + + 创建后自动分配 @@ -315,7 +376,23 @@ async function onTransfer(endpointId: string) { - +